diff --git a/CMakeLists.txt b/CMakeLists.txt index 409d1a716a8a..cb8039a26414 100644 --- a/CMakeLists.txt +++ b/CMakeLists.txt @@ -57,20 +57,13 @@ tvm_option(USE_THREADS "Build with thread support" ON) tvm_option(USE_LLVM "Build with LLVM, can be set to specific llvm-config path" OFF) tvm_option(USE_MLIR "Build with MLIR support" OFF) tvm_option(USE_STACKVM_RUNTIME "Include stackvm into the runtime" OFF) -tvm_option(USE_GRAPH_EXECUTOR "Build with tiny graph executor" ON) -tvm_option(USE_GRAPH_EXECUTOR_CUDA_GRAPH "Build with tiny graph executor with CUDA Graph for GPUs" OFF) -tvm_option(USE_AOT_EXECUTOR "Build with AOT executor" ON) -tvm_option(USE_PROFILER "Build profiler for the VM and graph executor" ON) tvm_option(USE_OPENMP "Build with OpenMP thread pool implementation" OFF) -tvm_option(USE_RELAY_DEBUG "Building Relay in debug mode..." OFF) tvm_option(TVM_DEBUG_WITH_ABI_CHANGE "Enable debug code that may cause ABI changes" OFF) tvm_option(TVM_LOG_BEFORE_THROW "Whether log before throw, for debugging purposes" OFF) tvm_option(USE_RTTI "Build with RTTI" ON) tvm_option(USE_MSVC_MT "Build with MT" 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(INDEX_DEFAULT_I64 "Defaults the index datatype to int64" ON) tvm_option(USE_LIBBACKTRACE "Use libbacktrace to supply linenumbers on stack traces" AUTO) @@ -335,34 +328,6 @@ tvm_file_glob(GLOB CODEGEN_SRCS list(APPEND COMPILER_SRCS ${CODEGEN_SRCS}) -tvm_file_glob(GLOB_RECURSE RELAY_OP_SRCS - src/relay/op/*.cc - ) -tvm_file_glob(GLOB_RECURSE RELAY_PASS_SRCS - src/relay/analysis/*.cc - src/relay/collage/*.cc - src/relay/transforms/*.cc - src/relay/quantize/*.cc - ) -tvm_file_glob(GLOB RELAY_BACKEND_SRCS - src/relay/backend/*.cc - src/relay/backend/vm/*.cc - src/relay/backend/aot/*.cc - ) -tvm_file_glob(GLOB_RECURSE RELAY_IR_SRCS - src/relay/ir/*.cc - src/relay/printer/*.cc - src/relay/parser/*.cc - ) -tvm_file_glob(GLOB_RECURSE RELAY_QNN_SRCS - src/relay/qnn/*.cc -) -list(APPEND COMPILER_SRCS ${RELAY_OP_SRCS}) -list(APPEND COMPILER_SRCS ${RELAY_PASS_SRCS}) -list(APPEND COMPILER_SRCS ${RELAY_BACKEND_SRCS}) -list(APPEND COMPILER_SRCS ${RELAY_IR_SRCS}) -list(APPEND COMPILER_SRCS ${RELAY_QNN_SRCS}) - tvm_file_glob(GLOB DATATYPE_SRCS src/target/datatype/*.cc) list(APPEND COMPILER_SRCS ${DATATYPE_SRCS}) list(APPEND COMPILER_SRCS "src/target/datatype/myfloat/myfloat.cc") @@ -418,51 +383,6 @@ else() list(APPEND COMPILER_SRCS ${STACKVM_RUNTIME_SRCS}) endif(USE_STACKVM_RUNTIME) -# NOTE(areusch): USE_GRAPH_RUNTIME will be deleted in a future release -if(USE_GRAPH_RUNTIME AND NOT DEFINED USE_GRAPH_EXECUTOR) - message(WARNING "USE_GRAPH_RUNTIME renamed to USE_GRAPH_EXECUTOR. Please update your config.cmake") - set(USE_GRAPH_EXECUTOR ${USE_GRAPH_RUNTIME}) - unset(USE_GRAPH_RUNTIME CACHE) -endif(USE_GRAPH_RUNTIME AND NOT DEFINED USE_GRAPH_EXECUTOR) - -# NOTE(areusch): USE_GRAPH_RUNTIME_DEBUG will be deleted in a future release -if(USE_GRAPH_RUNTIME_DEBUG AND NOT DEFINED USE_PROFILER) - message(WARNING "USE_GRAPH_RUNTIME_DEBUG renamed to USE_PROFILER. Please update your config.cmake") - set(USE_PROFILER ${USE_GRAPH_RUNTIME_DEBUG}) - unset(USE_GRAPH_RUNTIME_DEBUG CACHE) -endif(USE_GRAPH_RUNTIME_DEBUG AND NOT DEFINED USE_PROFILER) - -if(USE_GRAPH_EXECUTOR) - message(STATUS "Build with Graph Executor support...") - tvm_file_glob(GLOB RUNTIME_GRAPH_EXECUTOR_SRCS src/runtime/graph_executor/*.cc) - list(APPEND RUNTIME_SRCS ${RUNTIME_GRAPH_EXECUTOR_SRCS}) - -endif(USE_GRAPH_EXECUTOR) - -# convert old options for profiler -if(USE_GRAPH_EXECUTOR_DEBUG) - message(WARNING "USE_GRAPH_EXECUTOR_DEBUG renamed to USE_PROFILER. Please update your config.cmake") - unset(USE_GRAPH_EXECUTOR_DEBUG CACHE) - set(USE_PROFILER ON) -endif() -if(USE_VM_PROFILER) - message(WARNING "USE_VM_PROFILER renamed to USE_PROFILER. Please update your config.cmake") - unset(USE_VM_PROFILER CACHE) - set(USE_PROFILER ON) -endif() - -if(USE_PROFILER) - message(STATUS "Build with profiler...") - - tvm_file_glob(GLOB RUNTIME_GRAPH_EXECUTOR_DEBUG_SRCS src/runtime/graph_executor/debug/*.cc) - list(APPEND RUNTIME_SRCS ${RUNTIME_GRAPH_EXECUTOR_DEBUG_SRCS}) - set_source_files_properties(${RUNTIME_GRAPH_EXECUTOR_SRCS} - PROPERTIES COMPILE_DEFINITIONS "TVM_GRAPH_EXECUTOR_DEBUG") - - tvm_file_glob(GLOB RUNTIME_VM_PROFILER_SRCS src/runtime/vm/profiler/*.cc) - list(APPEND RUNTIME_SRCS ${RUNTIME_VM_PROFILER_SRCS}) -endif(USE_PROFILER) - if(USE_CUDA AND USE_NCCL) message(STATUS "Build with NCCL...") find_nccl(${USE_NCCL}) @@ -493,13 +413,6 @@ if(USE_ROCM AND USE_RCCL) list(APPEND RUNTIME_SRCS ${RUNTIME_RCCL_SRC}) endif() -if(USE_AOT_EXECUTOR) - message(STATUS "Build with AOT Executor support...") - file(GLOB RUNTIME_AOT_EXECUTOR_SRCS src/runtime/aot_executor/*.cc) - list(APPEND RUNTIME_SRCS ${RUNTIME_AOT_EXECUTOR_SRCS}) - -endif(USE_AOT_EXECUTOR) - # Enable ctest if gtest is available if(USE_GTEST) # Check env var for backward compatibility. A better way to specify package @@ -538,12 +451,6 @@ if(USE_GTEST) endif() endif() -if(USE_PIPELINE_EXECUTOR) - message(STATUS "Build with Pipeline Executor support...") - tvm_file_glob(GLOB RUNTIME_PIPELINE_SRCS src/runtime/pipeline/*.cc) - list(APPEND RUNTIME_SRCS ${RUNTIME_PIPELINE_SRCS}) -endif(USE_PIPELINE_EXECUTOR) - if(USE_KALLOC_ALIGNMENT) message(STATUS "Build Alloc alignment set to ${USE_KALLOC_ALIGNMENT}") add_definitions(-DTVM_KALLOC_ALIGNMENT=${USE_KALLOC_ALIGNMENT}) @@ -576,7 +483,6 @@ include(cmake/modules/Metal.cmake) include(cmake/modules/ROCM.cmake) include(cmake/modules/LLVM.cmake) include(cmake/modules/contrib/BLAS.cmake) -include(cmake/modules/contrib/CODEGENC.cmake) include(cmake/modules/contrib/DNNL.cmake) include(cmake/modules/contrib/AMX.cmake) include(cmake/modules/contrib/CUTLASS.cmake) @@ -587,10 +493,7 @@ include(cmake/modules/contrib/MSCCLPP.cmake) include(cmake/modules/contrib/Sort.cmake) include(cmake/modules/contrib/NNPack.cmake) include(cmake/modules/contrib/LibTorch.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) @@ -604,7 +507,6 @@ include(cmake/modules/contrib/MSC.cmake) include(cmake/modules/contrib/vllm.cmake) include(cmake/modules/Git.cmake) include(cmake/modules/LibInfo.cmake) -include(cmake/modules/RustExt.cmake) include(cmake/modules/contrib/Mrvl.cmake) set(LIBINFO_FILE ${CMAKE_CURRENT_LIST_DIR}/src/support/libinfo.cc) diff --git a/Makefile b/Makefile index 5134c949ea48..deb5c94f3800 100644 --- a/Makefile +++ b/Makefile @@ -127,7 +127,7 @@ webclean: # JVM build rules INCLUDE_FLAGS = -Iinclude -I$(DLPACK_PATH)/include -I$(DMLC_CORE_PATH)/include -PKG_CFLAGS = -std=c++11 -Wall -O2 $(INCLUDE_FLAGS) -fPIC +PKG_CFLAGS = -Wall -O3 $(INCLUDE_FLAGS) -fPIC PKG_LDFLAGS = ifeq ($(OS),Windows_NT) diff --git a/apps/android_rpc/app/src/main/jni/tvm_runtime.h b/apps/android_rpc/app/src/main/jni/tvm_runtime.h index 7b4ced7c9c0d..fa9fde747892 100644 --- a/apps/android_rpc/app/src/main/jni/tvm_runtime.h +++ b/apps/android_rpc/app/src/main/jni/tvm_runtime.h @@ -38,8 +38,6 @@ #include "../src/runtime/cpu_device_api.cc" #include "../src/runtime/dso_library.cc" #include "../src/runtime/file_utils.cc" -#include "../src/runtime/graph_executor/graph_executor.cc" -#include "../src/runtime/graph_executor/graph_executor_factory.cc" #include "../src/runtime/library_module.cc" #include "../src/runtime/logging.cc" #include "../src/runtime/memory/memory_manager.cc" diff --git a/apps/cpp_rpc/main.cc b/apps/cpp_rpc/main.cc index a94c45fc94ce..17b79141b15b 100644 --- a/apps/cpp_rpc/main.cc +++ b/apps/cpp_rpc/main.cc @@ -27,7 +27,7 @@ #if defined(__linux__) || defined(__ANDROID__) #include #endif -#include +#include #include #include diff --git a/apps/cpp_rpc/win32_process.cc b/apps/cpp_rpc/win32_process.cc index cd24e34b55e7..610e38adf268 100644 --- a/apps/cpp_rpc/win32_process.cc +++ b/apps/cpp_rpc/win32_process.cc @@ -23,7 +23,7 @@ #include "win32_process.h" #include -#include +#include #include #include diff --git a/apps/cpp_rtvm/CMakeLists.txt b/apps/cpp_rtvm/CMakeLists.txt deleted file mode 100644 index b8f2d50e47aa..000000000000 --- a/apps/cpp_rtvm/CMakeLists.txt +++ /dev/null @@ -1,99 +0,0 @@ -cmake_policy(SET CMP0069 NEW) # suppress cmake warning about IPO - -set(RTVM_SOURCES - main.cc - tvm_runner.cc - ../../3rdparty/cnpy/cnpy.cpp -) -set(TVM_RUNNER_SOURCES - tvm_runner.cc - ../../3rdparty/cnpy/cnpy.cpp -) - -set(RTVM_LINKER_LIBS "") - -if(WIN32) - file(GLOB ZLIB_SRC - "../../3rdparty/zlib/*.c" - ) - list(APPEND RTVM_SOURCES ${ZLIB_SRC}) - list(APPEND TVM_RUNNER_SOURCES ${ZLIB_SRC}) -endif() - -# Set output to same directory as the other TVM libs -set(CMAKE_RUNTIME_OUTPUT_DIRECTORY ${CMAKE_BINARY_DIR}) -add_executable(rtvm ${RTVM_SOURCES}) -add_library(tvm_runner_objs OBJECT ${TVM_RUNNER_SOURCES}) -add_library(tvm_runner SHARED $) - -include(CheckIPOSupported) -check_ipo_supported(RESULT result OUTPUT output) -if(result) - set_property(TARGET rtvm PROPERTY INTERPROCEDURAL_OPTIMIZATION_RELEASE TRUE) -endif() - -if(WIN32) - target_compile_definitions(rtvm PUBLIC -DNOMINMAX) -endif() - -if (OS) - if (OS STREQUAL "Linux") - set_property(TARGET rtvm PROPERTY LINK_FLAGS -lpthread) - set_property(TARGET tvm_runner PROPERTY LINK_FLAGS -lpthread) - endif() -endif() - -if(USE_OPENCL) - if (ANDROID_ABI) - if(DEFINED ENV{ANDROID_NDK_MAJOR}) - if($ENV{ANDROID_NDK_MAJOR} VERSION_LESS "23") - set_property(TARGET rtvm PROPERTY LINK_FLAGS -fuse-ld=gold) - set_property(TARGET tvm_runner PROPERTY LINK_FLAGS -fuse-ld=gold) - endif() - endif() - endif() -endif() - -target_include_directories( - rtvm - PUBLIC "../../include" - PUBLIC "../../3rdparty/cnpy" - PUBLIC "../../3rdparty/zlib" - PUBLIC DLPACK_PATH - PUBLIC DMLC_PATH -) - -if (BUILD_FOR_ANDROID AND USE_HEXAGON) - get_hexagon_sdk_property("${USE_HEXAGON_SDK}" "${USE_HEXAGON_ARCH}" - DSPRPC_LIB DSPRPC_LIB_DIRS - ) - if(DSPRPC_LIB_DIRS) - link_directories(${DSPRPC_LIB_DIRS}) - else() - message(WARNING "Could not locate some Hexagon SDK components") - endif() - list(APPEND RTVM_LINKER_LIBS cdsprpc log) -endif() - -if(BUILD_STATIC_RUNTIME) - list(APPEND RTVM_LINKER_LIBS -Wl,--whole-archive tvm_runtime -Wl,--no-whole-archive) -else() - list(APPEND RTVM_LINKER_LIBS tvm_runtime) -endif() - -if(NOT WIN32) - list(APPEND RTVM_LINKER_LIBS z) -endif() - -target_link_libraries(rtvm ${RTVM_LINKER_LIBS}) - -# Build tvm_runner as a exportable lib -target_include_directories( - tvm_runner_objs - PUBLIC "../../include" - PUBLIC "../../3rdparty/cnpy" - PUBLIC "../../3rdparty/zlib" - PUBLIC DLPACK_PATH - PUBLIC DMLC_PATH -) -target_link_libraries(tvm_runner ${RTVM_LINKER_LIBS}) diff --git a/apps/cpp_rtvm/README.md b/apps/cpp_rtvm/README.md deleted file mode 100644 index 652d46eb5821..000000000000 --- a/apps/cpp_rtvm/README.md +++ /dev/null @@ -1,390 +0,0 @@ - - - - - - - - - - - - - - - - - - -# Native Inference application for CPP Native - -Native inference tool ```rtvm``` helps in deploying TVM compiled models from a standalone cpp environment. -Overall process starts from getting a model from a framework all the way up to running on target device using `rtvm` tool. - -### Models - -Models can be downloaded from well known frameworks like Tensorflow, PyTorch, TFLite, Onnx ..etc. -scripts/download_models.py has a reference to prepare sample network ```resnet50``` from keras framework. - -```bash -python3 scripts/download_models.py -``` - -### Auto Tuning -Auto tuning process tunes various operatrors the given model for respective target. Auto tuning for remote devices use ```tvm_rpc``` and we need to setup the rpc environment before we invoke tuning. -Please refer below section [RPC setup](#rpc-setup) for the same. - -Auto tunng is necessary to obtain best performaning kernels. We can skip this step if we have tuning log already or the tuning cache is available from tophub (implicite by TVM compilation process). -Below message indicate that there exists some kernels not optimized for the selected target. In this case we can proceed with tuning to best performance. -```One or more operators have not been tuned. Please tune your model for better performance. Use DEBUG logging level to see more details.``` - -with below environment from [RPC setup](#rpc-setup) -``` bash -tvm tracker running on ```TVM_TRACKER_HOST``` -tracker port being ```TVM_TRACKER_PORT``` -rpc device access key being ```TVM_RPC_KEY``` -the model to be tuned being ```./model_data/keras-resnet50/resnet50.h5``` -``` - -the below command we can generate the tuning cache to file ```./model_data/keras-resnet50/keras-resnet50.log``` - -```bash -python3 -m tvm.driver.tvmc tune --target="opencl" --target-host="llvm -mtriple=aarch64-linux-gnu" \ -./model_data/keras-resnet50/resnet50.h5 -o ./model_data/keras-resnet50/keras-resnet50.log \ ---early-stopping 0 --repeat 30 --rpc-key ${TVM_RPC_KEY} --rpc-tracker ${TVM_TRACKER_HOST}:${TVM_TRACKER_PORT} --trials 1024 \ ---tuning-records ./model_data/keras-resnet50/keras-resnet50-records.log --tuner xgb -``` - -where -```bash ---target="opencl" refers to opencl device on Android device ---target-host="llvm -mtriple=aarch64-linux-gnu" refers to target_host being an ARM64 CPU -Options --early-stopping, --repeat, --trials, --tuner are Auto TVM specific options. -``` -Please refer to AutoTVM documentation for more details [here](https://tvm.apache.org/docs/how_to/tune_with_autotvm/index.html?highlight=autotvm). - -### Compile the model - -Compilation step generates TVM compiler output artifacts which need to be taken to target device for deployment. -These artifacts is a compressed archive with kernel shared lib, json with graph description and params binary. - -Below command will generate the same - - -```bash -python3 -m tvm.driver.tvmc compile --cross-compiler ${ANDROID_NDK_HOME}/toolchains/llvm/prebuilt/linux-x86_64/bin/aarch64-linux-android28-clang \ ---target="opencl, llvm" --target-llvm-mtriple aarch64-linux-gnu -o keras-resnet50.tar ./model_data/keras-resnet50/resnet50.h5 -``` - -where -``` ---cross-compiler : Indicates the cross compiler path for kernel library generation ---target="opencl, llvm" indicates target and host devices -``` - -### Test Run via RPC - -At this stage we can verify the generated compiler output for execution correctness over the RPC setup interface. -Below command can run the compiled output on remote target device. - -with - -``` bash -tvm tracker running on ```TVM_TRACKER_HOST``` -tracker port being ```TVM_TRACKER_PORT``` -rpc device access key being ```TVM_RPC_KEY``` -compilation out being keras-resnet50.tar -``` - -```bash -python3 -m tvm.driver.tvmc run --device="cl" keras-resnet50.tar --rpc-key ${TVM_RPC_KEY} --rpc-tracker ${TVM_TRACKER_HOST}:${TVM_TRACKER_PORT} --print-time -``` - -This inputs random inputs and validates the execution correctness of the compiled model. - -```tvmc``` tool has various options to input custom data, profile the model and benchmark the execution. - - -### Deployment Run - -Now we will verify the deployment run of the compiled model using ```rtvm``` tool on target device without any RPC or host based execution. - -We need to extract the tar achive on target device. We can copy the extracted contents of ```keras-resnet50.tar``` under Android temp folder at ```/data/local/tmp/keras-resnet50/``` - -Also copy the cross compiled tool ```rtvm``` and ```libtvm_runtime.so``` to ```data/local/tmp/``` - -```rtvm``` usage can be quired as below -```bash -Android:/data/local/tmp $ LD_LIBRARY_PATH=./ ./rtvm -Command line usage ---model - The folder containing tvm artifacts(mod.so, mod.param, mod.json) ---device - The target device to use {llvm, opencl, cpu, cuda, metal, rocm, vpi, oneapi} ---input - Numpy file for the model input (optional and we use random of not given) ---output - Numpy file name to dump the model output as numpy ---dump-meta - Dump model meta information ---pre-compiled - The file name of a file where pre-compiled programs should be stored ---profile - Profile over all execution ---dry-run - Profile after given dry runs, default 10 ---run-count - Profile for given runs, default 50 ---zero-copy - Profile with zero copy api - - Example - ./rtvm --model=keras-resnet50 --device="opencl" --dump-meta - ./rtvm --model=keras-resnet50 --device="opencl" --input input.npz --output=output.npz -``` - -```rtvm``` can run the model using no inputs (just a dry run without any valid inputs) and also with specific input supplied as a numpy npz format file. - -We can create npz dump for all inputs by saving the dict object as shown below. - -With ```keras-resnet50``` having one input ```input_1``` with shape ```[1, 224, 224, 3]``` and dtype ```float32``` - -``` -# Random initilization -input1 = np.random.uniform(low=-1, high=1, size=(1, 224, 224, 3)).astype("float32") -dataset = {"input_1": input1} -np.savez("input.npz", **dataset) -``` - -Copy ```input.npz``` also to the target device as ```/data/local/tmp/input.npz``` - - -Now, on Android shell we can do a dry run as well as with specific input as shown below. -```bash -# Query meta data information -Android:/data/local/tmp/ $ LD_LIBRARY_PATH=./ ./rtvm --model=keras-resnet50 --device=opencl --dump-meta -. . . . . . -Meta Information:keras-resnet50 - Number of Inputs:183 - Number of Outputs:1 - Input MetaInfo: - Input:input_1 - DType:float32 - Shape:[1, 224, 224, 3] - Output MetaInfo: - Output:tvmgen_default_fused_nn_softmax - DType:float32 - Shape:[1, 1000] -. . . . . . - -# Dry run with out any inputs -Android:/data/local/tmp/ $ LD_LIBRARY_PATH=./ ./rtvm --model=keras-resnet50 --device=opencl -Model = keras-resnet50 -Device = opencl -Input = -Output = -Dump Metadata = False -TVMRunner Constructor:keras-resnet50 Devices:opencl -TVMRunner Load:keras-resnet50 -TVMRunner::GetMetaInfo -Executing dry run ... -Set Random Input for :input_1 -TVMRunner::GetInputMemSize:input_1 -Random Input Size:602112 bytes -TVMRunner::SetInput (Raw) -TVMRunner::Run -Get Output for :tvmgen_default_fused_nn_softmax -TVMRunner::GetOutputMemSize:tvmgen_default_fused_nn_softmax -TVMRunner::GetOutput (Raw) -Output Size:4000 bytes - - -# Run with input and dump output as npz file -Android:/data/local/tmp/ $ LD_LIBRARY_PATH=./ ./rtvm --model=keras-resnet50 --device=opencl --input=input.npz --output=output.npz -Model = keras-resnet50 -Device = opencl -Input = input.npz -Output = output.npz -Dump Metadata = False -TVMRunner Constructor:keras-resnet50 Devices:opencl -TVMRunner Load:keras-resnet50 -TVMRunner::GetMetaInfo -Executing with Input:input.npz Output:output.npz -TVMRunner::SetInput (Numpy):input.npz -Set Numpy Input for :input_1 -TVMRunner::Run -TVMRunner::GetOutput (Numpy):output.npz -Get Output for :tvmgen_default_fused_nn_softmax -Output Size:4000 bytes -``` - -output.npz contains the modle outputs. Below is a quick look of its contents. -```bash -tvm-host:~$ unzip -l output.npz -Archive: output.npz - Length Date Time Name ---------- ---------- ----- ---- - 4080 1980-00-00 00:00 tvmgen_default_fused_nn_softmax.npy ---------- ------- - 4080 1 file - -``` - -Building ```cpp_rtvm``` produces ```libtvm_runner.so```, a simplified interface that rtvm use internally for loading and executing tvm compiled models from C/C++ environments. -```tvm_runner.h``` describes the interface definition here. Alternatively pro users can use TVM's [c_native_api](https://github.com/apache/tvm/blob/main/include/tvm/runtime/c_runtime_api.h) interface for more access to TVM features. - - -# RPC Setup - -For Android devices we require cross compilation of tvm_rpc (also libtvm_runtime.so which is a dependency) for remote device. -RPC setup involves running tracker on host device and running tvm_rpc on target device. - -### Tracker - -Below command runs the tracker on host over port ```9100``` - -```bash -python3 -m tvm.exec.rpc_tracker --host 127.0.0.1 --port 9100" -``` -### RPC on Target - -With ```abcd1234ef``` being adb device id and tvm_rpc (and libtvm_runtime.so) is pushed to target device at ```/data/local/tmp/tvm_rpc/``` - -```bash -export ANDROID_SERIAL=abcd1234ef -# Below settings will reroute networking tcm connections on devices to host device via adb interface -adb reverse tcp:9100 tcp:9100 -adb forward tcp:5000 tcp:5000 -# Run the tvm_rpc on device -env adb shell "cd /data/local/tmp/tvm_rpc; killall -9 tvm_rpc; \ -LD_LIBRARY_PATH=/data/local/tmp/tvm_rpc/ ./tvm_rpc server --host=0.0.0.0 --port=5000 --port-end=5010 --tracker=127.0.0.1:9100 --key=android -``` - -Now we have the rpc setup with ```TVM_TRACKER_HOST=127.0.0.1```, ```TVM_TRACKER_PORT=9100``` and ```TVM_RPC_KEY=android```. - -We can also check connected and available devices on tracker as shown below. - -```bash -python3 -m tvm.exec.query_rpc_tracker --port ${TVM_TRACKER_PORT} -Tracker address 127.0.0.1:9100 - -Server List ------------------------------- -server-address key ------------------------------- - 127.0.0.1:5000 server:android ------------------------------- - -Queue Status -------------------------------- -key total free pending -------------------------------- -android 1 1 0 -------------------------------- -``` - - -# Target Specific Configuration - -Below sections describe device/target specific settings to be used with ```tvmc``` tool. - -### Adreno GPU - -Adreno GPU has a docker definition that helps to ease the development environment. - -We can build the docker image by using below command from TVM repo. - -```bash -./docker/build.sh ci_adreno -docker tag tvm.ci_adreno ci_adreno -``` - -Below command builds host and target rpc components for Adreno and drops into an interactive shell. - -```bash -./tests/scripts/ci.py adreno -i -``` - -Also, one can build with Adreno OpenCLML SDK support - -```bash -export ADRENO_OPENCL= -./tests/scripts/ci.py adreno -i -``` - -Above command produces -```build-adreno``` which is host build -```build-adreno-target``` which contains cross compiled tvm_rpc and libtvm_runtime.so - - -Below options to be used for Adreno GPU while working with tvmc - -* Tuning - - ``` - --target="opencl -device=adreno" - --target-host="llvm -mtriple=aarch64-linux-gnu" - ``` - -* Compilation - - ``` - --cross-compiler ${ANDROID_NDK_HOME}/toolchains/llvm/prebuilt/linux-x86_64/bin/aarch64-linux-android28-clang - --target="opencl, llvm" - --target-opencl-device adreno - --target-llvm-mtriple aarch64-linux-gnu - ``` - - While enabling CLML just need to specify below target option for compilation. - ```--target="opencl, clml, llvm"``` - - -* Running - - ```--device="cl"``` - - -For example with a model from keras ```./model_data/keras-resnet50/resnet50.h5``` - - -```bash -# Tuning -python3 -m tvm.driver.tvmc tune --desired-layout NCHW --target="opencl -device=adreno" --target-host="llvm -mtriple=aarch64-linux-gnu" \ -./model_data/keras-resnet50/resnet50.h5 -o ./model_data/keras-resnet50/keras-resnet50.log --early-stopping 0 --repeat 30 \ ---rpc-key ${TVM_RPC_KEY} --rpc-tracker {TVM_TRACKER_HOST}:{TVM_TRACKER_PORT} --trials 1024 --tuning-records ./model_data/keras-resnet50/keras-resnet50-records.log --tuner xgb - -# Tuning produces tuning log ./model_data/keras-resnet50/keras-resnet50.log - - -# Compilation -python3 -m tvm.driver.tvmc compile --cross-compiler ${ANDROID_NDK_HOME}/toolchains/llvm/prebuilt/linux-x86_64/bin/aarch64-linux-android28-clang \ ---desired-layout NCHW --target="opencl, llvm" --target-opencl-device adreno --target-llvm-mtriple aarch64-linux-gnu \ -./model_data/keras-resnet50/resnet50.h5 -o keras-resnet50.tar - -# Compilation produces target artifacts keras-resnet50.tar - -# Run on adreno device via RPC -python3 -m tvm.driver.tvmc run --device="cl" keras-resnet50.tar --rpc-key ${TVM_RPC_KEY} --rpc-tracker {TVM_TRACKER_HOST}:{TVM_TRACKER_PORT} --print-time - -``` - -# Use pre-compiled OpenCL kernels -Using pre-compiled programs might significantly improve inference time of the -first run. E.g. for topology with ~300 kernels compilation time on Adreno was -about 26 seconds. But after dumping compiled programs to binary files and reuse -them on the next runs, the compilation time was significantly decreased (more -than 1000 times) and starts to be around 25 ms. - -To use such functionality, the developer have to pass parameter `--pre-compiled` -to the `rtvm` and specify the file name where pre-compiled programs will be -stored. If the pre-compiled file name was passed to the `rtvm` then After method -`Load`, method `UsePreCompiledProgram` is called. This method loads pre-compiled -programs if the file exists. In opposite case the file will be created and -pre-compiled programs will be saved to this file. - -# Performnace Profiling Options -The tool has added few options to measure wall clock performance of the given model on Target natively. ---profile : Can turn on the profiling ---dry-run : The number of times dry run the model before mearuring the performance. Default value os 10 ---run-count : The number times to run the model and take an average. Default value is 50. ---zero-copy: This option enables graph runtime zero copy to be used for input and output than byte copy to DLTensor. - -Performance profile options dumps information summary as given below. - Module Load :27 ms - Graph Runtime Create :11 ms - Params Read :15 ms - Params Set :41 ms - Pre Compiled Progs Load :24 ms -Total Load Time :118 ms -Average ExecTime :27 ms -Unload Time :35.9236 ms diff --git a/apps/cpp_rtvm/main.cc b/apps/cpp_rtvm/main.cc deleted file mode 100644 index ee3d4d2583c8..000000000000 --- a/apps/cpp_rtvm/main.cc +++ /dev/null @@ -1,418 +0,0 @@ -/* - * 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 main.cc - * \brief TVM runtime utility for TVM. - */ -#include -#include -#include -#if defined(__linux__) || defined(__ANDROID__) -#include -#endif -#include - -#include -#include -#include -#include -#include - -#include "../../src/support/socket.h" -#include "../../src/support/utils.h" -#include "tvm_runner.h" - -using namespace std; -using namespace tvm::runtime; -using namespace tvm::support; - -static const string kUsage = - "Command line usage\n" - "--model - The folder containing tvm artifacts(mod.so, mod.param, mod.json) \n" - "--device - The target device to use {llvm, opencl, cpu, cuda, metal, rocm, vpi, " - "oneapi}\n" - "--input - Numpy file for the model input (optional and we use random of not given)\n" - "--output - Numpy file name to dump the model output as numpy\n" - "--dump-meta - Dump model meta information\n" - "--pre-compiled - The file name of a file where pre-compiled programs should be stored\n" - "--profile - Profile over all execution\n" - "--dry-run - Profile after given dry runs, default 10\n" - "--run-count - Profile for given runs, default 50\n" - "--zero-copy - Profile with zero copy api\n" - "\n" - " Example\n" - " ./rtvm --model=keras-resnet50 --device=\"opencl\" --dump-meta\n" - " ./rtvm --model=keras-resnet50 --device=\"opencl\" --input input.npz --output=output.npz\n" - "\n"; - -/*! - * \brief Tool Arguments. - * \arg model The tvm artifact to load & run - * \arg device The target device to use {llvm, cl, ...etc.} - * \arg input Numpy file for the model input - * \arg output Numpy file name to dump the model output as numpy - * \arg pre_compiled File name where pre-compiled programs should be stored - * \arg profile Do we profile overall execution - */ -struct ToolArgs { - string model; - string device; - string input; - string output; - string pre_compiled; - bool dump_meta{false}; - bool profile{false}; - int dry_run{10}; - int run_count{50}; - bool zero_copy{false}; -}; - -/*! - * \brief PrintArgs print the contents of ToolArgs - * \param args ToolArgs structure - */ -void PrintArgs(const ToolArgs& args) { - LOG(INFO) << "Model = " << args.model; - LOG(INFO) << "Device = " << args.device; - LOG(INFO) << "Input = " << args.input; - LOG(INFO) << "Output = " << args.output; - LOG(INFO) << "Pre-compiled = " << args.pre_compiled; - LOG(INFO) << "Dump Metadata = " << ((args.dump_meta) ? ("True") : ("False")); - LOG(INFO) << "Profile = " << ((args.profile) ? ("True") : ("False")); - LOG(INFO) << "Dry Run = " << args.dry_run; - LOG(INFO) << "Run Count = " << args.run_count; - LOG(INFO) << "Zero Copy = " << ((args.zero_copy) ? ("True") : ("False")); -} - -#if defined(__linux__) || defined(__ANDROID__) -/*! - * \brief CtrlCHandler, exits if Ctrl+C is pressed - * \param s signal - */ -void CtrlCHandler(int s) { - LOG(INFO) << "\nUser pressed Ctrl+C, Exiting"; - exit(1); -} - -/*! - * \brief HandleCtrlC Register for handling Ctrl+C event. - */ -void HandleCtrlC() { - // Ctrl+C handler - struct sigaction sigIntHandler; - sigIntHandler.sa_handler = CtrlCHandler; - sigemptyset(&sigIntHandler.sa_mask); - sigIntHandler.sa_flags = 0; - sigaction(SIGINT, &sigIntHandler, nullptr); -} -#endif -/*! - * \brief GetCmdOption Parse and find the command option. - * \param argc arg counter - * \param argv arg values - * \param option command line option to search for. - * \param key whether the option itself is key - * \return value corresponding to option. - */ -string GetCmdOption(int argc, char* argv[], string option, bool key = false) { - string cmd; - for (int i = 1; i < argc; ++i) { - string arg = argv[i]; - if (arg.find(option) == 0) { - if (key) { - cmd = argv[i]; - return cmd; - } - // We assume "=" is the end of option. - ICHECK_EQ(*option.rbegin(), '='); - cmd = arg.substr(arg.find('=') + 1); - return cmd; - } - } - return cmd; -} - -/*! - * \brief ParseCmdArgs parses the command line arguments. - * \param argc arg counter - * \param argv arg values - * \param args the output structure which holds the parsed values - */ -void ParseCmdArgs(int argc, char* argv[], struct ToolArgs& args) { - const string model = GetCmdOption(argc, argv, "--model="); - if (!model.empty()) { - args.model = model; - } else { - LOG(INFO) << kUsage; - exit(0); - } - - const string device = GetCmdOption(argc, argv, "--device="); - if (!device.empty()) { - args.device = device; - } else { - LOG(INFO) << kUsage; - exit(0); - } - - const string input = GetCmdOption(argc, argv, "--input="); - if (!input.empty()) { - args.input = input; - } - - const string output = GetCmdOption(argc, argv, "--output="); - if (!output.empty()) { - args.output = output; - } - - const string pmeta = GetCmdOption(argc, argv, "--dump-meta", true); - if (!pmeta.empty()) { - args.dump_meta = true; - } - - args.pre_compiled = GetCmdOption(argc, argv, "--pre-compiled="); - - const string pprofile = GetCmdOption(argc, argv, "--profile", true); - if (!pprofile.empty()) { - args.profile = true; - } - - const string pdry_run = GetCmdOption(argc, argv, "--dry-run="); - if (!pdry_run.empty()) { - args.dry_run = stoi(pdry_run); - } - - const string prun = GetCmdOption(argc, argv, "--run-count="); - if (!prun.empty()) { - args.run_count = stoi(prun); - } - - const string pzcopy = GetCmdOption(argc, argv, "--zero-copy", true); - if (!pzcopy.empty()) { - args.zero_copy = true; - } -} - -/*! - * \brief Loads and Executes the model on given Target. - * \param args tool arguments - * \return result of operation. - */ -int ExecuteModel(ToolArgs& args) { -#if defined(__linux__) || defined(__ANDROID__) - // Ctrl+C handler - HandleCtrlC(); -#endif - - // Initialize TVM Runner - auto runner = new TVMRunner(args.model, args.device); - - // Load the model - runner->Load(); - if (!args.pre_compiled.empty()) { - runner->UsePreCompiledPrograms(args.pre_compiled); - } - - // Query Model meta Information - TVMMetaInfo mInfo = runner->GetMetaInfo(); - - // Print Meta Information - if (args.dump_meta) runner->PrintMetaInfo(); - - int total_exec_time = 0; - - if (args.profile) { - if (args.dry_run) { - for (int ii = 0; ii < args.dry_run; ++ii) { - runner->Run(); - } - TVMSynchronize(GetTVMDevice(args.device), 0, nullptr); - } - int total_time = 0; - std::map input_data_even, input_data_odd; - std::map output_data_even, output_data_odd; - - std::map input_data; - std::map output_data; - - // Alloc / populate and keep input data ready - for (auto& elem : mInfo.input_info) { - if (args.zero_copy) { - auto ndarr = - NDArray::Empty(elem.second.first, tvm::runtime::String2DLDataType(elem.second.second), - DLDevice{GetTVMDevice(args.device), 0}); - input_data_even.insert({elem.first, ndarr}); - - ndarr = - NDArray::Empty(elem.second.first, tvm::runtime::String2DLDataType(elem.second.second), - DLDevice{GetTVMDevice(args.device), 0}); - input_data_odd.insert({elem.first, ndarr}); - } else { - char* data = (char*)malloc(runner->GetInputMemSize(elem.first)); - input_data.insert({elem.first, data}); - } - } - - // Alloc and keep output bufers ready - for (auto& elem : mInfo.output_info) { - if (args.zero_copy) { - auto ndarr = - NDArray::Empty(elem.second.first, tvm::runtime::String2DLDataType(elem.second.second), - DLDevice{GetTVMDevice(args.device), 0}); - output_data_even.insert({elem.first, ndarr}); - - ndarr = - NDArray::Empty(elem.second.first, tvm::runtime::String2DLDataType(elem.second.second), - DLDevice{GetTVMDevice(args.device), 0}); - output_data_odd.insert({elem.first, ndarr}); - } else { - char* data = (char*)malloc(runner->GetOutputMemSize(elem.first)); - output_data.insert({elem.first, data}); - } - } - - for (int ii = 0; ii < args.run_count; ++ii) { - // Timer start - auto tstart = std::chrono::high_resolution_clock::now(); - // Set random input for all input - for (auto& elem : mInfo.input_info) { - if (args.zero_copy) { - if (ii % 2) { - runner->SetInput(elem.first, input_data_even[elem.first]); - } else { - runner->SetInput(elem.first, input_data_odd[elem.first]); - } - } else { - runner->SetInput(elem.first, input_data[elem.first]); - } - } - - if (args.zero_copy) { - // With zero copy set the result NDArray up front - for (auto& elem : mInfo.output_info) { - if (ii % 2) { - runner->SetOutput(elem.first, output_data_even[elem.first]); - } else { - runner->SetOutput(elem.first, output_data_odd[elem.first]); - } - } - } - - // Run the model - runner->Run(); - - if (!args.zero_copy) { - // W/o zero copy we need to invoke explicite data copy - for (auto& elem : mInfo.output_info) { - runner->GetOutput(elem.first, output_data[elem.first]); - } - } else { - // Just wait for the run to complete. - TVMSynchronize(GetTVMDevice(args.device), 0, nullptr); - } - - // Timer end - auto tend = std::chrono::high_resolution_clock::now(); - LOG(INFO) << "Exec Time:" << static_cast((tend - tstart).count()) / 1e6; - total_exec_time += static_cast((tend - tstart).count()) / 1e6; - } - - // Free input bufers - for (auto& elem : mInfo.input_info) { - free(input_data[elem.first]); - } - - // Free output bufers - for (auto& elem : mInfo.output_info) { - free(output_data[elem.first]); - } - } else if (!args.input.empty() && !args.output.empty()) { - LOG(INFO) << "Executing with Input:" << args.input << " Output:" << args.output; - // Set Input from Numpy Input - runner->SetInput(args.input); - // Run the model - runner->Run(); - // Get Output as Numpy dump - runner->GetOutput(args.output); - } else { - LOG(INFO) << "Executing dry run ... "; - // Set random input for all inputs - for (auto& elem : mInfo.input_info) { - LOG(INFO) << "Set Random Input for :" << elem.first; - auto shape = elem.second.first; - size_t ssize = runner->GetInputMemSize(elem.first); - char* data = (char*)malloc(ssize); - LOG(INFO) << "Random Input Size:" << ssize << " bytes"; - runner->SetInput(elem.first, data); - free(data); - } - // Run the model - runner->Run(); - // Get Output and dump few values - for (auto& elem : mInfo.output_info) { - LOG(INFO) << "Get Output for :" << elem.first; - auto shape = elem.second.first; - size_t ssize = runner->GetOutputMemSize(elem.first); - char* data = (char*)malloc(ssize); - runner->GetOutput(elem.first, data); - LOG(INFO) << "Output Size:" << ssize << " bytes"; - free(data); - } - } - - if (args.profile) { - // Print Stats - runner->PrintStats(); - } - auto tstart = std::chrono::high_resolution_clock::now(); - delete runner; - auto tend = std::chrono::high_resolution_clock::now(); - - if (args.profile) { - LOG(INFO) << "Average ExecTime :" << total_exec_time / args.run_count << " ms"; - LOG(INFO) << "Unload Time :" << static_cast((tend - tstart).count()) / 1e6 - << " ms"; - } - return 0; -} - -/*! - * \brief main The main function. - * \param argc arg counter - * \param argv arg values - * \return result of operation. - */ -int main(int argc, char* argv[]) { - if (argc <= 1) { - LOG(INFO) << kUsage; - return 0; - } - - ToolArgs args; - ParseCmdArgs(argc, argv, args); - PrintArgs(args); - - if (ExecuteModel(args)) { - PrintArgs(args); - LOG(INFO) << kUsage; - return -1; - } - return 0; -} diff --git a/apps/cpp_rtvm/scripts/download_models.py b/apps/cpp_rtvm/scripts/download_models.py deleted file mode 100644 index ef330cf765d8..000000000000 --- a/apps/cpp_rtvm/scripts/download_models.py +++ /dev/null @@ -1,36 +0,0 @@ -#!/usr/bin/env python3 -# -*- coding: utf-8 -*- - -# 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. - -tmp_dir = "./model_data/" -dload_models = [] - -# Keras : Resnet50 -try: - from tensorflow.keras.applications.resnet50 import ResNet50 - - model_file_name = "{}/{}".format(tmp_dir + "keras-resnet50", "resnet50.h5") - model = ResNet50(include_top=True, weights="imagenet", input_shape=(224, 224, 3), classes=1000) - model.save(model_file_name) - dload_models.append(model_file_name) -except ImportError: - LOG.warning("Keras is not installed, skipping Keras models") - - -print("Models:", dload_models) diff --git a/apps/cpp_rtvm/tvm_runner.cc b/apps/cpp_rtvm/tvm_runner.cc deleted file mode 100644 index 945b541cfa03..000000000000 --- a/apps/cpp_rtvm/tvm_runner.cc +++ /dev/null @@ -1,415 +0,0 @@ -/* - * 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 tvm_runner.cc - * \brief TVM model runner implementation. - */ - -#include "tvm_runner.h" - -#include - -#include -#include -#include -#include -#include -#include - -namespace tvm { -namespace runtime { - -/*! - * \brief Get the TVM device id corresponding to device string. - * \param device the target device in string format. - * \return dl_device corresponding to the device string. - */ -DLDeviceType GetTVMDevice(std::string device) { - if (!device.compare("cpu")) { - return kDLCPU; - } else if (!device.compare("llvm")) { - return kDLCPU; - } else if (!device.compare("cuda")) { - return kDLCUDA; - } else if (!device.compare("opencl")) { - return kDLOpenCL; - } else if (!device.compare("vulkan")) { - return kDLVulkan; - } else if (!device.compare("metal")) { - return kDLMetal; - } else if (!device.compare("vpi")) { - return kDLVPI; - } else if (!device.compare("rocm")) { - return kDLROCM; - } else if (!device.compare("oneapi")) { - return kDLOneAPI; - } else { - LOG(FATAL) << "TVMRunner : Unsupported device :" << device; - } -} - -/*! - * \brief Constructor for TVMRunner. - * \param path where the tfm compiler artifacts present. - * \param device the target device where we need to load the compiled model. - */ -TVMRunner::TVMRunner(std::string path, std::string device) - : r_model_path(path), r_device(device), r_run_was_called(false) { - LOG(INFO) << "TVMRunner Constructor:" << r_model_path << " Devices:" << r_device; -} - -/*! - * \brief Load Setup TVM graph runtime for given model. - * \param 0 on success else error code. - */ -int TVMRunner::Load(void) { - LOG(INFO) << "TVMRunner Load:" << r_model_path; - // Load the lib file - auto tstart = std::chrono::high_resolution_clock::now(); - - r_mod_handle = Module::LoadFromFile((r_model_path + "/mod.so").c_str(), "so"); - auto tend = std::chrono::high_resolution_clock::now(); - r_module_load_ms = static_cast((tend - tstart).count()) / 1e6; - - tstart = std::chrono::high_resolution_clock::now(); - // Read model json file - std::ifstream json_reader((r_model_path + "/mod.json").c_str()); - CHECK(!json_reader.fail()) << "Failed to open json file:" << (r_model_path + "/mod.json").c_str(); - json_reader.seekg(0, std::ios_base::end); - std::size_t json_size = json_reader.tellg(); - json_reader.seekg(0, std::ios_base::beg); - std::string json_data; - json_data.reserve(json_size); - json_reader.read((char*)json_data.c_str(), json_size); - json_reader.close(); - - // Get ref to graph exeutor - auto f_handle = tvm::runtime::Registry::Get("tvm.graph_executor.create"); - - // Greate graph runtime - r_graph_handle = - (*f_handle)(json_data, r_mod_handle, static_cast(GetTVMDevice(r_device)), 0); - - tend = std::chrono::high_resolution_clock::now(); - r_graph_load_ms = static_cast((tend - tstart).count()) / 1e6; - - // Read params binary file - tstart = std::chrono::high_resolution_clock::now(); - std::ifstream params_reader((r_model_path + "/mod.params").c_str(), std::ios::binary); - CHECK(!params_reader.fail()) << "Failed to open json file:" - << (r_model_path + "/mod.params").c_str(); - - params_reader.seekg(0, std::ios_base::end); - std::size_t param_size = params_reader.tellg(); - params_reader.seekg(0, std::ios_base::beg); - std::vector param_data(param_size / sizeof(char)); - params_reader.read((char*)¶m_data[0], param_size); - params_reader.close(); - - TVMByteArray params_arr; - params_arr.data = (char*)¶m_data[0]; - params_arr.size = param_size; - - tend = std::chrono::high_resolution_clock::now(); - r_param_read_ms = static_cast((tend - tstart).count()) / 1e6; - - // Load parameters - tstart = std::chrono::high_resolution_clock::now(); - r_graph_handle.GetFunction("load_params")(params_arr); - tend = std::chrono::high_resolution_clock::now(); - r_param_load_ms = static_cast((tend - tstart).count()) / 1e6; - - return 0; -} - -/*! - * \brief Specify if the run programs should be dumped to binary and reused in the next runs. - * \param file_name File name where pre-compiled programs should be stored. - */ -void TVMRunner::UsePreCompiledPrograms(std::string file_name) { - auto tstart = std::chrono::high_resolution_clock::now(); - if (r_run_was_called) { - LOG(INFO) << "TVMRunner UsePreCompiledPrograms: should be called before first run"; - return; - } - auto f_get = r_mod_handle->GetFunction("opencl.GetPreCompiledPrograms", true); - auto f_set = r_mod_handle->GetFunction("opencl.SetPreCompiledPrograms", true); - if (f_get != nullptr && f_set != nullptr) { - std::ifstream ifs(file_name, std::ios::in | std::ios::binary); - if (ifs.fail()) { - std::string ss = f_get(); - auto bytes = tvm::String(ss); - std::ofstream fs(file_name, std::ofstream::binary); - fs.write(bytes.c_str(), bytes.size()); - } else { - ifs.seekg(0, std::ios_base::end); - std::size_t blob_size = ifs.tellg(); - ifs.seekg(0, std::ios_base::beg); - std::string blob_data; - blob_data.reserve(blob_size); - blob_data.resize(blob_size); - ifs.read((char*)blob_data.c_str(), blob_size); - ifs.close(); - f_set(String(blob_data)); - } - } - auto tend = std::chrono::high_resolution_clock::now(); - r_pre_compiled_load_ms = static_cast((tend - tstart).count()) / 1e6; -} - -/*! - * \brief Calculated the memory size for the NDArray. - * \param NDArray object. - * \return size of the memory. - */ -inline size_t GetMemSize(NDArray& narr) { - size_t size = 1; - for (tvm_index_t i = 0; i < narr->ndim; ++i) { - size *= static_cast(narr->shape[i]); - } - size *= (narr->dtype.bits * narr->dtype.lanes + 7) / 8; - return size; -} - -/*! - * \brief Get the input alloc mem size. - * \param input_id The input id to query the mem size. - * \return The memory size. - */ -size_t TVMRunner::GetInputMemSize(std::string input_id) { - NDArray in_arr = r_graph_handle.GetFunction("get_input")(input_id); - auto ssize = GetMemSize(in_arr); - - return ssize; -} - -/*! - * \brief Get the output alloc mem size. - * \param output_id The output id to query the mem size. - * \return The memory size. - */ -size_t TVMRunner::GetOutputMemSize(std::string output_id) { - NDArray out_arr = r_graph_handle.GetFunction("get_output")(output_id); - auto ssize = GetMemSize(out_arr); - - return ssize; -} - -/*! - * \brief Set the model inputs from npz file. - * \param inputfile the npz file from where we read input tensor data. - * \param 0 on success else error code. - */ -int TVMRunner::SetInput(std::string inputfile) { - LOG(INFO) << "TVMRunner::SetInput (Numpy):" << inputfile; - cnpy::npz_t npz_input = cnpy::npz_load(inputfile); - - for (auto& elem : mInfo.input_info) { - LOG(INFO) << "Set Numpy Input for :" << elem.first; - NDArray in_arr = r_graph_handle.GetFunction("get_input")(elem.first); - auto ssize = GetMemSize(in_arr); - - if (npz_input.find(elem.first) != npz_input.end()) { - in_arr.CopyFromBytes(npz_input[elem.first].data(), ssize); - } else { - LOG(WARNING) << "Couldn't find input " << elem.first << " in npy input file"; - } - } - - return 0; -} - -/*! - * \brief Set the model input from the given binary buffer. - * \param input_id input node name. - * \param raw_input binary input buffer to copy over input NDArray. - * \param 0 on success else error code. - */ -int TVMRunner::SetInput(std::string input_id, char* raw_input) { - NDArray in_arr = r_graph_handle.GetFunction("get_input")(input_id); - auto ssize = GetMemSize(in_arr); - in_arr.CopyFromBytes(raw_input, ssize); - return 0; -} - -/*! - * \brief Set the model input from given NDArray with zero copy. - * \param input_id input node name. - * \param ndarr NDArray. - * \param 0 on success else error code. - */ -int TVMRunner::SetInput(std::string input_id, NDArray& ndarr) { - r_graph_handle.GetFunction("set_input_zero_copy")(input_id, ndarr); - return 0; -} - -/*! - * \brief Get the model outputs and dump them to npz file. - * \param outputfile the npz file to where we dump the output data. - * \param 0 on success else error code. - */ -int TVMRunner::GetOutput(std::string outputfile) { - LOG(INFO) << "TVMRunner::GetOutput (Numpy):" << outputfile; - - for (auto& elem : mInfo.output_info) { - LOG(INFO) << "Get Output for :" << elem.first; - NDArray out_arr = r_graph_handle.GetFunction("get_output")(elem.first); - auto ssize = GetMemSize(out_arr); - LOG(INFO) << "Output Size:" << ssize << " bytes"; - - void* data = (void*)malloc(ssize * (out_arr->dtype.bits * out_arr->dtype.lanes + 7) / 8); - out_arr.CopyToBytes(data, ssize); - std::vector shape; - - for (int j = 0; j < out_arr->ndim; ++j) shape.push_back(out_arr->shape[j]); - if (!elem.second.second.compare("float32")) { - cnpy::npz_save(outputfile, elem.first, (float*)data, shape, "a"); - } else if (!elem.second.second.compare("int8")) { - cnpy::npz_save(outputfile, elem.first, (int8_t*)data, shape, "a"); - } else { - LOG(WARNING) << "DType:" << elem.second.second << " is not supported for npy_save"; - } - free(data); - } - - return 0; -} - -/*! - * \brief Get output of the model as a binary buffer. - * \param output_id output node name to read the data. - * \param raw_output the buffer to copy the data to. - * \param 0 on success else error code. - */ -int TVMRunner::GetOutput(std::string output_id, char* raw_output) { - NDArray out_arr = r_graph_handle.GetFunction("get_output")(output_id); - auto ssize = GetMemSize(out_arr); - out_arr.CopyToBytes(raw_output, ssize); - return 0; -} - -/*! - * \brief Set the model output from given NDArray with zero copy. - * \param output_id output node name. - * \param ndarr NDArray. - * \param 0 on success else error code. - */ -int TVMRunner::SetOutput(std::string output_id, NDArray& ndarr) { - r_graph_handle.GetFunction("set_output_zero_copy")(output_id, ndarr); - return 0; -} - -/*! - * \brief Call one cycle of execution for the model. - * \param 0 on success else error code. - */ -int TVMRunner::Run(void) { - r_run_was_called = true; - r_graph_handle.GetFunction("run")(); - return 0; -} - -/*! - * \brief Query various metadata from the grsph runtime. - * \param 0 on success else error code. - */ -TVMMetaInfo TVMRunner::GetMetaInfo(void) { - LOG(INFO) << "TVMRunner::GetMetaInfo"; - - mInfo.n_inputs = r_graph_handle.GetFunction("get_num_inputs")(); - mInfo.n_outputs = r_graph_handle.GetFunction("get_num_outputs")(); - - Map tvm_input_info = r_graph_handle.GetFunction("get_input_info")(); - auto shape_info = GetRef>(tvm_input_info["shape"].as()); - auto dtype_info = GetRef>(tvm_input_info["dtype"].as()); - for (const auto& kv : shape_info) { - auto stuple = GetRef(kv.second.as()); - std::vector vshape; - vshape.assign(stuple.begin(), stuple.end()); - auto dtype = GetRef(dtype_info[kv.first].as()); - std::pair, std::string> value = std::make_pair(vshape, dtype); - mInfo.input_info.insert({kv.first, value}); - } - - tvm_input_info = r_graph_handle.GetFunction("get_output_info")(); - shape_info = GetRef>(tvm_input_info["shape"].as()); - dtype_info = GetRef>(tvm_input_info["dtype"].as()); - for (const auto& kv : shape_info) { - auto stuple = GetRef(kv.second.as()); - std::vector vshape; - vshape.assign(stuple.begin(), stuple.end()); - auto dtype = GetRef(dtype_info[kv.first].as()); - std::pair, std::string> value = std::make_pair(vshape, dtype); - mInfo.output_info.insert({kv.first, value}); - } - - return mInfo; -} - -/*! - * \brief Print the meta information. - * \param 0 on success else error code. - */ -void TVMRunner::PrintMetaInfo(void) { - LOG(INFO) << "Meta Information:" << r_model_path; - LOG(INFO) << " Number of Inputs:" << mInfo.n_inputs; - LOG(INFO) << " Number of Outputs:" << mInfo.n_outputs; - LOG(INFO) << " Input MetaInfo:"; - for (auto& elem : mInfo.input_info) { - std::ostringstream stream; - stream << "["; - copy(elem.second.first.begin(), elem.second.first.end() - 1, - std::ostream_iterator(stream, ", ")); - stream << elem.second.first.back() << "]"; - LOG(INFO) << " Input:" << elem.first; - LOG(INFO) << " DType:" << elem.second.second; - LOG(INFO) << " Shape:" << stream.str(); - } - LOG(INFO) << " Output MetaInfo:"; - for (auto& elem : mInfo.output_info) { - std::ostringstream stream; - stream << "["; - copy(elem.second.first.begin(), elem.second.first.end() - 1, - std::ostream_iterator(stream, ", ")); - stream << elem.second.first.back() << "]"; - LOG(INFO) << " Output:" << elem.first; - LOG(INFO) << " DType:" << elem.second.second; - LOG(INFO) << " Shape:" << stream.str(); - } -} - -/*! - * \brief Print stats information. - */ -void TVMRunner::PrintStats(void) { - LOG(INFO) << "Performance Stats:" << r_model_path; - LOG(INFO) << " Module Load :" << r_module_load_ms << " ms"; - LOG(INFO) << " Graph Runtime Create :" << r_graph_load_ms << " ms"; - LOG(INFO) << " Params Read :" << r_param_read_ms << " ms"; - LOG(INFO) << " Params Set :" << r_param_load_ms << " ms"; - LOG(INFO) << " Pre Compiled Progs Load :" << r_pre_compiled_load_ms << " ms"; - LOG(INFO) << "Total Load Time :" - << r_module_load_ms + r_graph_load_ms + r_param_read_ms + r_param_load_ms + - r_pre_compiled_load_ms - << " ms"; -} - -} // namespace runtime -} // namespace tvm diff --git a/apps/cpp_rtvm/tvm_runner.h b/apps/cpp_rtvm/tvm_runner.h deleted file mode 100644 index e93b63ae85a8..000000000000 --- a/apps/cpp_rtvm/tvm_runner.h +++ /dev/null @@ -1,116 +0,0 @@ -/* - * 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 tvm_runner.h - * \brief TVM model runner. - */ -#ifndef TVM_APPS_CPP_RTVM_RUNNER_H_ -#define TVM_APPS_CPP_RTVM_RUNNER_H_ - -#include -#include -#include - -#include - -#include "tvm/runtime/c_runtime_api.h" - -namespace tvm { -namespace runtime { - -/*! - * \brief various meta information related to the compiled TVM model. - */ -typedef struct _TVMMetaInfo { - int n_inputs; - int n_outputs; - std::map, std::string>> input_info; - std::map, std::string>> output_info; -} TVMMetaInfo; - -/*! - * \brief encapsulates TVM graph runtime functionality with simplified API interface. - */ -class TVMRunner { - public: - /*! \brief Constructor */ - TVMRunner(std::string path, std::string device); - - /*! \brief Initiates graph runtime and with the compiled model */ - int Load(void); - /*! \brief Specify if the run programs should be dumped to binary and reused in the next runs */ - void UsePreCompiledPrograms(std::string); - /*! \brief Executes one inference cycle */ - int Run(void); - /*! \brief To set the inputs from given npz file */ - int SetInput(std::string); - /*! \brief To set the input from binary data */ - int SetInput(std::string, char*); - /*! \brief To set the input from NDArray */ - int SetInput(std::string, NDArray& ndarr); - /*! \brief Save the model output into given npz file */ - int GetOutput(std::string); - /*! \brief Get the model output in binary format */ - int GetOutput(std::string, char*); - /*! \brief Swap output NDArray with given one */ - int SetOutput(std::string, NDArray& ndarr); - /*! \brief To get the input mem size */ - size_t GetInputMemSize(std::string); - /*! \brief To get the output mem size */ - size_t GetOutputMemSize(std::string); - /*! \brief Populates various meta information from graph runtime */ - TVMMetaInfo GetMetaInfo(void); - /*! \brief Print function to show all meta information */ - void PrintMetaInfo(void); - - /*! \brief Print function to show all stats information */ - void PrintStats(void); - - // Public profiling information - /*! Module load time */ - int r_module_load_ms{0}; - /*! Graph runtime creatint time */ - int r_graph_load_ms{0}; - /*! Params read time */ - int r_param_read_ms{0}; - /*! Params load time */ - int r_param_load_ms{0}; - /*! Pre compiled programs load time */ - int r_pre_compiled_load_ms{0}; - - private: - /*! \brief Module handle for the shared object */ - Module r_mod_handle; - /*! \brief Graph runtime module handle */ - Module r_graph_handle; - /*! \brief The local model path from where we load the model */ - std::string r_model_path; - /*! \brief The target device */ - std::string r_device; - /*! \brief Holds meta information queried from graph runtime */ - TVMMetaInfo mInfo; - /*! \brief Mark if the run method was called */ - bool r_run_was_called; -}; - -DLDeviceType GetTVMDevice(std::string device); -} // namespace runtime -} // namespace tvm -#endif // TVM_APPS_CPP_RTVM_RUNNER_H_ diff --git a/apps/hexagon_launcher/cmake/android/CMakeLists.txt b/apps/hexagon_launcher/cmake/android/CMakeLists.txt index 84ff4add284b..78d3cb396cfd 100644 --- a/apps/hexagon_launcher/cmake/android/CMakeLists.txt +++ b/apps/hexagon_launcher/cmake/android/CMakeLists.txt @@ -76,12 +76,10 @@ ExternalProject_Add(android_tvm_runtime "-DCMAKE_CXX_STANDARD=17" "-DCMAKE_TOOLCHAIN_FILE=${CMAKE_TOOLCHAIN_FILE}" "-DUSE_HEXAGON=ON" - "-DUSE_GRAPH_EXECUTOR=OFF" "-DUSE_HEXAGON_ARCH=${USE_HEXAGON_ARCH}" "-DUSE_HEXAGON_SDK=${USE_HEXAGON_SDK}" "-DUSE_LIBBACKTRACE=OFF" "-DUSE_LLVM=OFF" - "-DUSE_PROFILER=OFF" "-DUSE_RPC=OFF" INSTALL_COMMAND "" BUILD_ALWAYS ON diff --git a/cmake/modules/CUDA.cmake b/cmake/modules/CUDA.cmake index ad83ebe26b8c..e9e552b92b70 100644 --- a/cmake/modules/CUDA.cmake +++ b/cmake/modules/CUDA.cmake @@ -110,6 +110,7 @@ if(USE_CUDA) tvm_file_glob(GLOB CONTRIB_THRUST_SRC src/runtime/contrib/thrust/*.cu) add_library(tvm_thrust_objs OBJECT ${CONTRIB_THRUST_SRC}) target_compile_options(tvm_thrust_objs PRIVATE $<$:--expt-extended-lambda>) + target_compile_definitions(tvm_thrust_objs PUBLIC DMLC_USE_LOGGING_LIBRARY=) if (NOT USE_THRUST MATCHES ${IS_TRUE_PATTERN}) find_package(CCCL REQUIRED COMPONENTS Thrust) target_link_libraries(tvm_thrust_objs PRIVATE CCCL::Thrust) @@ -135,18 +136,6 @@ if(USE_CUDA) list(APPEND TVM_RUNTIME_LINKER_LIBS ${CUDA_NVTX_LIBRARY}) endif(USE_NVTX) - if(USE_GRAPH_EXECUTOR_CUDA_GRAPH) - if(NOT USE_GRAPH_EXECUTOR) - message(FATAL_ERROR "CUDA Graph is only supported by graph executor, please set USE_GRAPH_EXECUTOR=ON") - endif() - if(CUDAToolkit_VERSION_MAJOR LESS "10") - message(FATAL_ERROR "CUDA Graph requires CUDA 10 or above, got=" ${CUDAToolkit_VERSION}) - endif() - message(STATUS "Build with Graph executor with CUDA Graph support...") - tvm_file_glob(GLOB RUNTIME_CUDA_GRAPH_SRCS src/runtime/graph_executor/cuda_graph/*.cc) - list(APPEND RUNTIME_SRCS ${RUNTIME_CUDA_GRAPH_SRCS}) - endif() - # Add CUDA builtins to RelaxVM tvm_file_glob(GLOB RELAX_VM_CUDA_BUILTIN_SRC_CC src/runtime/relax_vm/cuda/*.cc) list(APPEND RUNTIME_SRCS ${RELAX_VM_CUDA_BUILTIN_SRC_CC}) diff --git a/cmake/modules/LibInfo.cmake b/cmake/modules/LibInfo.cmake index feef618dc2fe..cfea9b20d84a 100644 --- a/cmake/modules/LibInfo.cmake +++ b/cmake/modules/LibInfo.cmake @@ -58,7 +58,6 @@ function(add_lib_info src_file) TVM_INFO_ROCM_PATH="${ROCM_PATH}" TVM_INFO_SUMMARIZE="${SUMMARIZE}" TVM_INFO_USE_ALTERNATIVE_LINKER="${USE_ALTERNATIVE_LINKER}" - TVM_INFO_USE_AOT_EXECUTOR="${USE_AOT_EXECUTOR}" TVM_INFO_USE_ARM_COMPUTE_LIB_GRAPH_EXECUTOR="${USE_ARM_COMPUTE_LIB_GRAPH_EXECUTOR}" TVM_INFO_USE_ARM_COMPUTE_LIB="${USE_ARM_COMPUTE_LIB}" TVM_INFO_USE_BLAS="${USE_BLAS}" @@ -79,8 +78,6 @@ function(add_lib_info src_file) TVM_INFO_USE_AMX="${USE_AMX}" TVM_INFO_USE_DNNL="${USE_DNNL}" TVM_INFO_USE_FALLBACK_STL_MAP="${USE_FALLBACK_STL_MAP}" - TVM_INFO_USE_GRAPH_EXECUTOR_CUDA_GRAPH="${USE_GRAPH_EXECUTOR_CUDA_GRAPH}" - TVM_INFO_USE_GRAPH_EXECUTOR="${USE_GRAPH_EXECUTOR}" TVM_INFO_USE_GTEST="${USE_GTEST}" TVM_INFO_USE_HEXAGON="${USE_HEXAGON}" TVM_INFO_USE_HEXAGON_RPC="${USE_HEXAGON_RPC}" @@ -104,8 +101,6 @@ function(add_lib_info src_file) TVM_INFO_USE_OPENCL_GTEST="${USE_OPENCL_GTEST}" TVM_INFO_USE_OPENMP="${USE_OPENMP}" TVM_INFO_USE_PAPI="${USE_PAPI}" - TVM_INFO_USE_PROFILER="${USE_PROFILER}" - TVM_INFO_USE_PT_TVMDSOOP="${USE_PT_TVMDSOOP}" TVM_INFO_USE_RANDOM="${USE_RANDOM}" TVM_INFO_USE_RELAY_DEBUG="${USE_RELAY_DEBUG}" TVM_INFO_TVM_DEBUG_WITH_ABI_CHANGE="${TVM_DEBUG_WITH_ABI_CHANGE}" @@ -124,7 +119,6 @@ function(add_lib_info src_file) TVM_INFO_USE_TENSORFLOW_PATH="${USE_TENSORFLOW_PATH}" TVM_INFO_USE_TENSORRT_CODEGEN="${USE_TENSORRT_CODEGEN}" TVM_INFO_USE_TENSORRT_RUNTIME="${USE_TENSORRT_RUNTIME}" - TVM_INFO_USE_TF_TVMDSOOP="${USE_TF_TVMDSOOP}" TVM_INFO_USE_TFLITE="${USE_TFLITE}" TVM_INFO_USE_THREADS="${USE_THREADS}" TVM_INFO_USE_THRUST="${USE_THRUST}" diff --git a/cmake/modules/RustExt.cmake b/cmake/modules/RustExt.cmake deleted file mode 100644 index e30caf0a0b04..000000000000 --- a/cmake/modules/RustExt.cmake +++ /dev/null @@ -1,43 +0,0 @@ -# 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(USE_RUST_EXT) - set(RUST_SRC_DIR "${CMAKE_CURRENT_SOURCE_DIR}/rust") - set(CARGO_OUT_DIR "${CMAKE_CURRENT_SOURCE_DIR}/rust/target") - - if(USE_RUST_EXT STREQUAL "STATIC") - set(COMPILER_EXT_PATH "${CARGO_OUT_DIR}/release/libcompiler_ext.a") - elseif(USE_RUST_EXT STREQUAL "DYNAMIC") - set(COMPILER_EXT_PATH "${CARGO_OUT_DIR}/release/libcompiler_ext.so") - else() - message(FATAL_ERROR "invalid setting for USE_RUST_EXT, STATIC, DYNAMIC or OFF") - endif() - - add_custom_command( - OUTPUT "${COMPILER_EXT_PATH}" - COMMAND cargo build --release - MAIN_DEPENDENCY "${RUST_SRC_DIR}" - WORKING_DIRECTORY "${RUST_SRC_DIR}/compiler-ext") - - add_custom_target(rust_ext ALL DEPENDS "${COMPILER_EXT_PATH}") - - # TODO(@jroesch, @tkonolige): move this to CMake target - # target_link_libraries(tvm "${COMPILER_EXT_PATH}" PRIVATE) - list(APPEND TVM_LINKER_LIBS ${COMPILER_EXT_PATH}) - - add_definitions(-DRUST_COMPILER_EXT=1) -endif() diff --git a/cmake/modules/contrib/CODEGENC.cmake b/cmake/modules/contrib/CODEGENC.cmake deleted file mode 100644 index b461176e6a84..000000000000 --- a/cmake/modules/contrib/CODEGENC.cmake +++ /dev/null @@ -1,19 +0,0 @@ -# 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_file_glob(GLOB CSOURCE_RELAY_CONTRIB_SRC src/relay/backend/contrib/codegen_c/*.cc) -list(APPEND COMPILER_SRCS ${CSOURCE_RELAY_CONTRIB_SRC}) diff --git a/cmake/modules/contrib/CUTLASS.cmake b/cmake/modules/contrib/CUTLASS.cmake index 11224a8d1f90..b302622cbce8 100644 --- a/cmake/modules/contrib/CUTLASS.cmake +++ b/cmake/modules/contrib/CUTLASS.cmake @@ -20,7 +20,6 @@ if(USE_CUDA AND USE_CUTLASS) set(CUTLASS_RUNTIME_OBJS "") tvm_file_glob(GLOB CUTLASS_CONTRIB_SRC - src/relay/backend/contrib/cutlass/*.cc src/relax/backend/contrib/cutlass/*.cc ) list(APPEND COMPILER_SRCS ${CUTLASS_CONTRIB_SRC}) diff --git a/cmake/modules/contrib/HybridDump.cmake b/cmake/modules/contrib/HybridDump.cmake deleted file mode 100644 index 156325d5d8ec..000000000000 --- a/cmake/modules/contrib/HybridDump.cmake +++ /dev/null @@ -1,20 +0,0 @@ -# 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. - -message(STATUS "Build with contrib.hybriddump") -tvm_file_glob(GLOB HYBRID_CONTRIB_SRC src/contrib/hybrid/*.cc) -list(APPEND COMPILER_SRCS ${HYBRID_CONTRIB_SRC}) diff --git a/cmake/modules/contrib/PT_TVMDSOOP.cmake b/cmake/modules/contrib/PT_TVMDSOOP.cmake deleted file mode 100644 index a73d3f38e939..000000000000 --- a/cmake/modules/contrib/PT_TVMDSOOP.cmake +++ /dev/null @@ -1,96 +0,0 @@ -# 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(PythonInterp REQUIRED) - execute_process(COMMAND ${PYTHON_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}") - message(STATUS "PyTorch path: ${PT_PATH}") - - execute_process(COMMAND ${PYTHON_EXECUTABLE} -c "import torch;print(torch.compiled_with_cxx11_abi())" - OUTPUT_VARIABLE PT_CXX_FLAG - RESULT_VARIABLE PT_STATUS) - - string(REGEX REPLACE "\n" "" PT_CXX_FLAG "${PT_CXX_FLAG}") - message(STATUS "Found TORCH_BUILT_WITH_CXX_ABI=${PT_CXX_FLAG} ") - - if(${PT_CXX_FLAG} STREQUAL "False") - set(CXX_ABI_ENABLED 0) - else() - set(CXX_ABI_ENABLED 1) - endif() - - set_property( - SOURCE - ${CMAKE_CURRENT_SOURCE_DIR}/src/contrib/torch/tvm_module_wrapper/RuntimeModuleWrapperTorch.cc - APPEND PROPERTY - COMPILE_OPTIONS - "-D_GLIBCXX_USE_CXX11_ABI=${CXX_ABI_ENABLED}" - "-I${PT_PATH}/include" - ) - - set_property( - SOURCE - ${CMAKE_CURRENT_SOURCE_DIR}/src/contrib/torch/pt_call_tvm/tvm_class.cc - APPEND PROPERTY - COMPILE_OPTIONS - "-I${PT_PATH}/include" - ) - - 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) - separate_arguments(PT_LINK_FLAGS UNIX_COMMAND ${PT_LINK_FLAGS_STR}) - - # This old version is depereated and will be removed after tvm 0.11 - set(LIBRARY_OLD_NAME pt_tvmdsoop) - - # This new library is set for pytorch integration, which solves the c++ abi imcompability issue - set(LIBRARY_NEW_NAME pt_tvmdsoop_new) - tvm_file_glob(GLOB_RECURSE PTTVM_TORCH ${CMAKE_CURRENT_SOURCE_DIR}/src/contrib/torch/tvm_module_wrapper/*.cc) - - tvm_file_glob(GLOB_RECURSE PTTVM_SRCS ${CMAKE_CURRENT_SOURCE_DIR}/src/contrib/torch/pt_call_tvm/*.cc) - - add_library(${LIBRARY_OLD_NAME} SHARED ${PTTVM_SRCS}) - add_library(${LIBRARY_NEW_NAME} SHARED ${PTTVM_TORCH}) - set(PTTVM_LINK_FLAGS -ltvm -L${CMAKE_CURRENT_BINARY_DIR}) - - if(NOT BUILD_PT_TVMDSOOP_ONLY STREQUAL "ON") - add_dependencies(${LIBRARY_OLD_NAME} tvm) - add_dependencies(${LIBRARY_NEW_NAME} tvm) - endif() - - target_compile_options(${LIBRARY_OLD_NAME} PUBLIC ${PTTVM_COMPILE_FLAGS} ${PT_COMPILE_FLAGS}) - target_link_libraries(${LIBRARY_OLD_NAME} PUBLIC ${PTTVM_LINK_FLAGS} ${PT_LINK_FLAGS}) - target_compile_definitions(${LIBRARY_OLD_NAME} PUBLIC DMLC_USE_LOGGING_LIBRARY=) - - target_compile_options(${LIBRARY_NEW_NAME} PUBLIC ${PTTVM_COMPILE_FLAGS} ${PT_COMPILE_FLAGS}) - target_link_libraries(${LIBRARY_NEW_NAME} PUBLIC ${PTTVM_LINK_FLAGS} ${PT_LINK_FLAGS}) - target_compile_definitions(${LIBRARY_NEW_NAME} PUBLIC DMLC_USE_LOGGING_LIBRARY=) -endif() diff --git a/cmake/modules/contrib/TF_TVMDSOOP.cmake b/cmake/modules/contrib/TF_TVMDSOOP.cmake deleted file mode 100644 index f5f3f036690f..000000000000 --- a/cmake/modules/contrib/TF_TVMDSOOP.cmake +++ /dev/null @@ -1,56 +0,0 @@ -# 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_TF_TVMDSOOP STREQUAL "OFF") - find_package(Python3 COMPONENTS Interpreter) - - execute_process(COMMAND ${Python3_EXECUTABLE} -c "import tensorflow as tf; print(' '.join(tf.sysconfig.get_compile_flags()))" - OUTPUT_VARIABLE TF_COMPILE_FLAGS_STR - RESULT_VARIABLE TF_STATUS) - if (NOT ${TF_STATUS} EQUAL 0) - message(FATAL_ERROR "Fail to get TensorFlow compile flags") - endif() - - if(NOT USE_CUDA STREQUAL "OFF") - add_definitions(-DTF_TVMDSOOP_ENABLE_GPU) - endif() - - execute_process(COMMAND ${Python3_EXECUTABLE} -c "import tensorflow as tf; print(' '.join(tf.sysconfig.get_link_flags()))" - OUTPUT_VARIABLE TF_LINK_FLAGS_STR - RESULT_VARIABLE TF_STATUS) - if (NOT ${TF_STATUS} EQUAL 0) - message(FATAL_ERROR "Fail to get TensorFlow link flags") - endif() - - string(REGEX REPLACE "\n" " " TF_FLAGS "${TF_COMPILE_FLAGS} ${TF_LINK_FLAGS}") - separate_arguments(TF_COMPILE_FLAGS UNIX_COMMAND ${TF_COMPILE_FLAGS_STR}) - separate_arguments(TF_LINK_FLAGS UNIX_COMMAND ${TF_LINK_FLAGS_STR}) - - - set(OP_LIBRARY_NAME tvm_dso_op) - tvm_file_glob(GLOB_RECURSE TFTVM_SRCS ${CMAKE_CURRENT_SOURCE_DIR}/src/contrib/tf_op/*.cc) - add_library(${OP_LIBRARY_NAME} SHARED ${TFTVM_SRCS}) - set(TFTVM_LINK_FLAGS -ltvm -L${CMAKE_CURRENT_BINARY_DIR}) - - if (NOT BUILD_TVMDSOOP_ONLY STREQUAL "ON") - add_dependencies(${OP_LIBRARY_NAME} tvm) - endif() - - target_compile_options(${OP_LIBRARY_NAME} PUBLIC ${TFTVM_COMPILE_FLAGS} ${TF_COMPILE_FLAGS}) - target_link_libraries(${OP_LIBRARY_NAME} PUBLIC ${TFTVM_LINK_FLAGS} ${TF_LINK_FLAGS}) - -endif() diff --git a/conda/recipe/bld.bat b/conda/recipe/bld.bat index 57ce3666eaee..084dd85560da 100644 --- a/conda/recipe/bld.bat +++ b/conda/recipe/bld.bat @@ -29,7 +29,6 @@ cmake ^ -DUSE_CPP_RPC=ON ^ -DUSE_SORT=ON ^ -DUSE_RANDOM=ON ^ - -DUSE_PROFILER=ON ^ -DINSTALL_DEV=ON ^ %SRC_DIR% || exit /b diff --git a/docker/Dockerfile.demo_android b/docker/Dockerfile.demo_android index 36aadbf1ee42..bbe8f7d82b01 100644 --- a/docker/Dockerfile.demo_android +++ b/docker/Dockerfile.demo_android @@ -73,7 +73,6 @@ RUN cd /usr && \ -DUSE_LLVM=llvm-config-15 \ -DUSE_RPC=ON \ -DUSE_SORT=ON \ - -DUSE_GRAPH_EXECUTOR=ON \ -DUSE_VULKAN=ON \ .. && \ make -j10 diff --git a/docker/install/install_tvm_cpu.sh b/docker/install/install_tvm_cpu.sh index 48e6df3597db..ce9218992419 100755 --- a/docker/install/install_tvm_cpu.sh +++ b/docker/install/install_tvm_cpu.sh @@ -27,7 +27,6 @@ cd /usr/tvm git checkout 4b13bf668edc7099b38d463e5db94ebc96c80470 echo set\(USE_LLVM llvm-config-8\) >> config.cmake -echo set\(USE_GRAPH_EXECUTOR ON\) >> config.cmake echo set\(USE_BLAS openblas\) >> config.cmake mkdir -p build cd build diff --git a/docs/arch/debugger.rst b/docs/arch/debugger.rst deleted file mode 100644 index 3a9f198f0837..000000000000 --- a/docs/arch/debugger.rst +++ /dev/null @@ -1,193 +0,0 @@ -.. 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. - -================= -Debugger -================= - -TVM Debugger is an interface for debugging TVM's computation graph execution. It helps to provide access to graph structures and tensor values at the TVM runtime. - -******************************************* -Debug Exchange Format -******************************************* - -1. Computational Graph -====================== -The optimized graph build by relay in json -serialized format is dumped as it is. This contains the whole -information about the graph. The UX can either use this graph directly -or transform this graph to the format UX can understand. - -The Graph JSON format is explained below - -1. ``nodes`` -Nodes are either placeholders or computational nodes in json. The nodes are stored -as a list. A node contains the below information - -- ``op`` - operation type, ``null`` means it is a placeholder/variable/input node and``tvm_op`` means this node can be executed -- ``name`` - Name of the node -- ``inputs`` - Position of the inputs for this operation, Inputs is a list of tuples with (nodeid, index, version). (Optional) -- ``attrs`` - Attributes of the node which contains the following information - - - ``flatten_data`` - Whether this data need to be flattened before execution - - ``func_name`` - Fused function name, corresponds to the symbol in the lib generated by relay compilation process. - - ``num_inputs`` - Number of inputs for this node - - ``num_outputs`` - Number of outputs this node produces - -2. ``arg_nodes`` -arg_nodes is a list of indices of nodes which is placeholder/variable/input or constant/param to the graph. - -3. ``heads`` -heads is a list of entries as the output of the graph. - -4. ``node_row_ptr`` -node\_row\_ptr stores the history of forward path, so you can skip constructing the entire graph in inference tasks. - -5. ``attrs`` -attrs can contain version numbers or similar helpful information. - -- ``storage_id`` - Memory slot id for each node in the storage layout. -- ``dtype`` - Datatype of each node (enum value). -- ``dltype`` - Datatype of each node in order. -- ``shape`` - Shape of each node k order. -- ``device_index`` - Device assignment for each entry in the graph. - -Example of dumped graph: - -:: - - { - "nodes": [ # List of nodes - { - "op": "null", # operation type = null, this is a placeholder/variable/input or constant/param node - "name": "x", # Name of the argument node - "inputs": [] # inputs for this node, its none since this is an argument node - }, - { - "op": "tvm_op", # operation type = tvm_op, this node can be executed - "name": "relu0", # Name of the node - "attrs": { # Attributes of the node - "flatten_data": "0", # Whether this data need to be flattened - "func_name": "fuse_l2_normalize_relu", # Fused function name, corresponds to the symbol in the lib generated by compilation process - "num_inputs": "1", # Number of inputs for this node - "num_outputs": "1" # Number of outputs this node produces - }, - "inputs": [[0, 0, 0]] # Position of the inputs for this operation - } - ], - "arg_nodes": [0], # Which all nodes in this are argument nodes - "node_row_ptr": [0, 1, 2], # Row indices for faster depth first search - "heads": [[1, 0, 0]], # Position of the output nodes for this operation - "attrs": { # Attributes for the graph - "storage_id": ["list_int", [1, 0]], # memory slot id for each node in the storage layout - "dtype": ["list_int", [0, 0]], # Datatype of each node (enum value) - "dltype": ["list_str", [ # Datatype of each node in order - "float32", - "float32"]], - "shape": ["list_shape", [ # Shape of each node k order - [1, 3, 20, 20], - [1, 3, 20, 20]]], - "device_index": ["list_int", [1, 1]], # Device assignment for each node in order - } - } - -2. Tensor dumping -================= - -The tensor received after execution is in ``tvm.ndarray`` type. All the tensors will -be saved as binary bytes in serialized format. The result binary bytes can be loaded by the -API "load_params". - -Example of loading the parameters - :: - with open(path_params, "rb") as fi: - loaded_params = bytearray(fi.read()) - - module.load_params(loaded_params) - -*************************************** -How to use Debugger? -*************************************** - -1. In ``config.cmake`` set the ``USE_PROFILER`` flag to ``ON`` - - :: - - # Whether enable additional graph debug functions - set(USE_PROFILER ON) - -2. Do 'make' tvm, so that it will make the ``libtvm_runtime.so`` - -3. In frontend script file instead of - ``from tvm.contrib import graph_executor`` import the - ``GraphModuleDebug`` - ``from tvm.contrib.debugger.debug_executor import GraphModuleDebug`` - -:: - - from tvm.contrib.debugger.debug_executor import GraphModuleDebug - m = GraphModuleDebug( - lib["debug_create"]("default", dev), - [dev], - lib.graph_json, - dump_root="/tmp/tvmdbg", - ) - # set inputs - m.set_input('data', tvm.nd.array(data.astype(dtype))) - m.set_input(**params) - # execute - m.run() - tvm_out = m.get_output(0, tvm.nd.empty(out_shape, dtype)).numpy() - -4. If network previously was exported to external library using ``lib.export_library("network.so")`` - like shared object file/dynamic linked library, the initialization - of debug runtime will be slightly different - -:: - - lib = tvm.runtime.load_module("network.so") - m = graph_executor.create(lib["get_graph_json"](), lib, dev, dump_root="/tmp/tvmdbg") - # set inputs - m.set_input('data', tvm.nd.array(data.astype(dtype))) - m.set_input(**params) - # execute - m.run() - tvm_out = m.get_output(0, tvm.nd.empty(out_shape, dtype)).numpy() - - -The outputs are dumped to a temporary folder in ``/tmp`` folder or the -folder specified while creating the runtime. - -*************************************** -Sample Output -*************************************** - -The below is the an example output of the debugger. - -:: - - Node Name Ops Time(us) Time(%) Start Time End Time Shape Inputs Outputs - --------- --- -------- ------- ---------- -------- ----- ------ ------- - 1_NCHW1c fuse___layout_transform___4 56.52 0.02 15:24:44.177475 15:24:44.177534 (1, 1, 224, 224) 1 1 - _contrib_conv2d_nchwc0 fuse__contrib_conv2d_NCHWc 12436.11 3.4 15:24:44.177549 15:24:44.189993 (1, 1, 224, 224, 1) 2 1 - relu0_NCHW8c fuse___layout_transform___broadcast_add_relu___layout_transform__ 4375.43 1.2 15:24:44.190027 15:24:44.194410 (8, 1, 5, 5, 1, 8) 2 1 - _contrib_conv2d_nchwc1 fuse__contrib_conv2d_NCHWc_1 213108.6 58.28 15:24:44.194440 15:24:44.407558 (1, 8, 224, 224, 8) 2 1 - relu1_NCHW8c fuse___layout_transform___broadcast_add_relu___layout_transform__ 2265.57 0.62 15:24:44.407600 15:24:44.409874 (64, 1, 1) 2 1 - _contrib_conv2d_nchwc2 fuse__contrib_conv2d_NCHWc_2 104623.15 28.61 15:24:44.409905 15:24:44.514535 (1, 8, 224, 224, 8) 2 1 - relu2_NCHW2c fuse___layout_transform___broadcast_add_relu___layout_transform___1 2004.77 0.55 15:24:44.514567 15:24:44.516582 (8, 8, 3, 3, 8, 8) 2 1 - _contrib_conv2d_nchwc3 fuse__contrib_conv2d_NCHWc_3 25218.4 6.9 15:24:44.516628 15:24:44.541856 (1, 8, 224, 224, 8) 2 1 - reshape1 fuse___layout_transform___broadcast_add_reshape_transpose_reshape 1554.25 0.43 15:24:44.541893 15:24:44.543452 (64, 1, 1) 2 1 diff --git a/docs/arch/index.rst b/docs/arch/index.rst index cf4829268ee2..717876d2db12 100644 --- a/docs/arch/index.rst +++ b/docs/arch/index.rst @@ -221,7 +221,6 @@ for learning-based optimizations. .. toctree:: :maxdepth: 1 - debugger introduction_to_module_serialization device_target_interactions diff --git a/docs/deep_dive/tensor_ir/learning.rst b/docs/deep_dive/tensor_ir/learning.rst index f85c64c57b4c..b76f87f58c38 100644 --- a/docs/deep_dive/tensor_ir/learning.rst +++ b/docs/deep_dive/tensor_ir/learning.rst @@ -90,9 +90,9 @@ Function Parameters and Buffers .. code:: python # TensorIR - def mm_relu(A: T.Buffer[(128, 128), "float32"], - B: T.Buffer[(128, 128), "float32"], - C: T.Buffer[(128, 128), "float32"]): + def mm_relu(A: T.Buffer((128, 128), "float32"), + B: T.Buffer((128, 128), "float32"), + C: T.Buffer((128, 128), "float32")): ... # NumPy def lnumpy_mm_relu(A: np.ndarray, B: np.ndarray, C: np.ndarray): diff --git a/golang/src/tvm_runtime_pack.cc b/golang/src/tvm_runtime_pack.cc index e4056742eef4..475abf5a3e36 100644 --- a/golang/src/tvm_runtime_pack.cc +++ b/golang/src/tvm_runtime_pack.cc @@ -45,7 +45,6 @@ #include "src/runtime/system_library.cc" // Graph executor -#include "src/runtime/graph_executor/graph_executor.cc" #include "src/runtime/memory/memory_manager.cc" // Uncomment the following lines to enable RPC diff --git a/include/tvm/auto_scheduler/auto_schedule.h b/include/tvm/auto_scheduler/auto_schedule.h deleted file mode 100755 index 2d7e5949aea4..000000000000 --- a/include/tvm/auto_scheduler/auto_schedule.h +++ /dev/null @@ -1,104 +0,0 @@ -/* - * 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 tvm/auto_scheduler/auto_schedule.h - * \brief The user interface of the auto scheduler. - */ - -#ifndef TVM_AUTO_SCHEDULER_AUTO_SCHEDULE_H_ -#define TVM_AUTO_SCHEDULER_AUTO_SCHEDULE_H_ - -#include -#include - -#include - -namespace tvm { -namespace auto_scheduler { - -/*! \brief Tuning and measurement options. */ -class TuningOptionsNode : public Object { - public: - /*! \brief The number of total measurement trials. */ - int num_measure_trials; - /*! \brief Stops the tuning early if no improvement after n measurements. */ - int early_stopping; - /*! \brief The number of programs to be measured at each search round. */ - int num_measures_per_round; - /*! \brief Verbosity level. 0 for silent, 1 to output information during schedule searching. */ - int verbose; - /*! \brief ProgramBuilder which builds the program */ - ProgramBuilder builder; - /*! \brief ProgramRunner which runs the program and measures time costs */ - ProgramRunner runner; - /*! \brief MeasureCallback functions to be called after each measure batch */ - Optional> measure_callbacks; - - void VisitAttrs(tvm::AttrVisitor* v) { - v->Visit("num_measure_trials", &num_measure_trials); - v->Visit("early_stopping", &early_stopping); - v->Visit("num_measures_per_round", &num_measures_per_round); - v->Visit("verbose", &verbose); - v->Visit("builder", &builder); - v->Visit("runner", &runner); - v->Visit("measure_callbacks", &measure_callbacks); - } - - static constexpr const char* _type_key = "auto_scheduler.TuningOptions"; - TVM_DECLARE_FINAL_OBJECT_INFO(TuningOptionsNode, Object); -}; - -/*! - * \brief Managed reference to TuningOptionsNode. - * \sa TuningOptionsNode - */ -class TuningOptions : public ObjectRef { - public: - /*! - * \brief The constructor - * \param num_measure_trials The number of total measurement trials. - * \param early_stopping Stops the tuning early if no improvement after n measurements. - * \param num_measures_per_round The number of programs to be measured at each search round. - * \param verbose Verbosity level. 0 for silent, 1 to output information during schedule - * search. - * \param builder ProgramBuilder which builds the program. - * \param runner ProgramRunner which runs the program and measure time costs. - * \param measure_callbacks MeasureCallback functions to be called after each measure batch. - */ - TuningOptions(int num_measure_trials, int early_stopping, int num_measures_per_round, int verbose, - ProgramBuilder builder, ProgramRunner runner, - Optional> measure_callbacks); - - TVM_DEFINE_OBJECT_REF_METHODS(TuningOptions, ObjectRef, TuningOptionsNode); -}; - -/*! - * \brief Run schedule search for a given compute declaration. - * \param search_policy The search policy. - * \param tuning_options Tuning and measurement options. - * \return A `te::schedule` and an Array of `te::Tensor` to be used in `tvm.lower` or - * `tvm.build`. - */ -TVM_DLL std::pair> AutoSchedule(SearchPolicy search_policy, - TuningOptions tuning_options); -} // namespace auto_scheduler -} // namespace tvm - -#endif // TVM_AUTO_SCHEDULER_AUTO_SCHEDULE_H_ diff --git a/include/tvm/auto_scheduler/compute_dag.h b/include/tvm/auto_scheduler/compute_dag.h deleted file mode 100755 index a87563e348f7..000000000000 --- a/include/tvm/auto_scheduler/compute_dag.h +++ /dev/null @@ -1,324 +0,0 @@ -/*r - * 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 tvm/auto_scheduler/compute_dag.h - * \brief The auto-scheduler's computational graph and related program analyses. - * - * We convert a compute declaration described by `tvm.compute` (could be a single operator or a - * subgraph) to a ComputeDAG. It keeps the input/output tensors, all operations in the DAG, and - * some static analysis results for the DAG (e.g. the total float operation count, consumer/producer - * relations of operations, whether an operation stage should be tiled/compute inlined ...). - * These analyses can help the search policy to make decisions during the search. - * ComputeDAG is also responsible for the interaction between auto-scheduler's `LoopState` and - * TVM schedule (e.g. applying the `LoopState` transform steps to a TVM schedule, providing - * `LoopState` with extra information got from TVM schedule ...). - */ - -#ifndef TVM_AUTO_SCHEDULER_COMPUTE_DAG_H_ -#define TVM_AUTO_SCHEDULER_COMPUTE_DAG_H_ - -#include -#include -#include - -#include -#include -#include -#include - -namespace tvm { -namespace auto_scheduler { - -/*! \brief Static analyzer for a ComputeDAG */ -class AccessAnalyzerNode : public Object { - public: - template - using OperationMap = std::unordered_map; - - /*! \brief Map an operation to all operations it reads from. - * For each operation pair, use a two-dimensional array for multiple multi-dimensional accesses - * The inner vector represents the indices of multi-dimensional access.*/ - OperationMap>>> read_from; - /*! \brief Map an operation to all operations it is read by. - * For each operation pair, use a two-dimensional array for multiple multi-dimensional accesses - * The inner vector represents the indices of multi-dimensional access.*/ - OperationMap>>> read_by; - /*! \brief Store the number of common outer iterators for operation pairs that have - * read-write relations. */ - OperationMap> num_common_outer_iterators; - /*! \brief Store whether the operation is an op with only simple access. - * (e.g., injective, broadcast and elementwise ops without reduction) */ - OperationMap is_simple_access; - /*! \brief Store whether the operation is strictly inlineable - * (e.g., injective, broadcast and elementwise without reduction, branch or expensive operations) - */ - OperationMap is_strictly_inlineable; - /*! \brief Store whether the operation needs multi-level tiling - * (e.g., computation-intensive ops with data reuse opportunity like matmul, conv2d) */ - OperationMap needs_multi_level_tiling; - /*! \brief Store whether the operation is an output operation */ - OperationMap is_output; - /*! \brief Store the topological order of operations */ - Array ops_topo_order; - - static constexpr const char* _type_key = "auto_scheduler.AccessAnalyzer"; - TVM_DECLARE_FINAL_OBJECT_INFO(AccessAnalyzerNode, Object); -}; - -/*! - * \brief Managed reference to AccessAnalyzerNode. - * \sa AccessAnalyzerNode - */ -class AccessAnalyzer : public ObjectRef { - public: - explicit AccessAnalyzer(const Array& tensors); - - /*! - * \brief Return whether this operation is an op with simple access - * (e.g., injective, broadcast and elementwise ops without reduction) - * \param op The operation - */ - TVM_DLL bool IsSimpleAccess(const te::Operation& op) const; - - /*! - * \brief Return whether this operation is strictly inlineable - * (e.g., injective, broadcast and elementwise without reduction, branch or expensive operations) - * \param op The operation - */ - TVM_DLL bool IsStrictlyInlineable(const te::Operation& op) const; - - /*! - * \brief Return whether this operation needs multi-level tiling - * (e.g., computation-intensive ops with data reuse opportunity like matmul, conv2d) - * \param op The operation - */ - TVM_DLL bool NeedsMultiLevelTiling(const te::Operation& op) const; - - /*! - * \brief Return whether this operation is an output operation - * \param op The operation - */ - TVM_DLL bool IsOutput(const te::Operation& op) const; - - /*! - * \brief Get all consumers of an operation - * \param state The current loop state - * \param op The operation - * \return The set of consumers - * \note This function propagates the relation for inlined ops - */ - TVM_DLL std::unordered_set GetConsumers( - const State& state, const te::Operation& op) const; - - /*! - * \brief Get all producers of an operation - * \param state The current loop state - * \param op The operation - * \return The set of producers - * \note This function propagates the relation for inlined ops - */ - TVM_DLL std::unordered_set GetProducers( - const State& state, const te::Operation& op) const; - - /*! - * \brief Get all direct producers of an operation - * \param op The operation - * \return The set of direct producers - * \note This function DOES NOT propagate the relation for inlined ops - */ - TVM_DLL std::unordered_set GetDirectProducers( - const te::Operation& op) const; - - /*! - * \brief Get the number of common outer iterators. - * \param op The operation - * \param target_op The target operation - * \note This function propagates the relation for chains with multiple ops. - */ - TVM_DLL int GetNumCommonOuterIterator(const te::Operation& op, - const te::Operation& target_op) const; - - /*! - * \brief Return whether two operations are elementwise-matched - * (e.g. conv2d and relu are elementwise-matched) - * \note This function propagates the relation for chains with multiple ops. - */ - TVM_DLL bool ElementWiseMatch(const te::Operation& op, const te::Operation& target_op) const; - - TVM_DEFINE_OBJECT_REF_METHODS(AccessAnalyzer, ObjectRef, AccessAnalyzerNode); -}; - -/*! \brief The auto-scheduler's computational graph and related program analyses. */ -class ComputeDAGNode : public Object { - public: - /*! - * \brief Input and output tensors. - * This is used as the input of `tvm.lower` or `tvm.build`. - */ - Array tensors; - /*! \brief All used operations in topo order. */ - Array ops; - /*! \brief The number of float operations in this ComputeDAG. */ - double flop_ct; - /*! \brief The initial state without any transform steps. */ - State init_state; - /*! \brief The static read-write access analyzer. */ - AccessAnalyzer access_analyzer; - - void VisitAttrs(tvm::AttrVisitor* v) { - v->Visit("tensors", &tensors); - v->Visit("ops", &ops); - v->Visit("flop_ct", &flop_ct); - v->Visit("init_state", &init_state); - v->Visit("access_analyzer", &access_analyzer); - } - - static constexpr const char* _type_key = "auto_scheduler.ComputeDAG"; - TVM_DECLARE_FINAL_OBJECT_INFO(ComputeDAGNode, Object); -}; - -/*! - * \brief Options for applying layout rewrite. - * This is an optimization to rewrite the layout of input tensors according to the schedule we get. - */ -enum class LayoutRewriteOption : int { - /*! \brief Do not perform layout rewrite. */ - NoRewrite = 0, - /*! \brief Insert layout transformation stages for input placeholders in the compute DAG */ - InsertTransformStage = 1, - /*! - * \brief Do not insert layout transformation stages and assume the input placeholders - * are pre-transformed. - * \note The lowered function with this option does not accept the origial input shapes, - * so this option must be used along with `AutoSchedulerLayoutRewrite` pass in Relay. - */ - RewriteForPreTransformed = 2, -}; - -/*! - * \brief Managed reference to ComputeDAGNode. - * \sa ComputeDAGNode - */ -class ComputeDAG : public ObjectRef { - public: - /*! \brief Construct a DAG from a list of output tensors. - * \param tensors `te::Tensor`s for a compute declaration. - */ - TVM_DLL explicit ComputeDAG(Array tensors); - - /*! \brief Construct a DAG based on a schedule. - * \param sch `te::Schedule`s for a compute declaration. - */ - TVM_DLL explicit ComputeDAG(const te::Schedule& sch); - - /*! - * \brief Rewrite the layout of placeholder specified by attr `layout_free_placeholders` - * according to the loop nest derived with `transform_steps`. - * \param transform_steps Transform steps of a state. - * \param layout_rewrite Different options in layout rewrite. - * \return The updated ComputeDAG after layout rewrite. - */ - ComputeDAG RewriteLayout(Array* transform_steps, LayoutRewriteOption layout_rewrite) const; - - /*! - * \brief Apply the history transform steps to get a TVM schedule. - * \param transform_steps Transform steps of a state. - * \param stages The list of stages after applying the steps. - * Pass a valid pointer if this information needs to be used outside this function. - * \param stage_to_axes The map that stores all axes for one stage. - * Pass a valid pointer if this information needs to be used outside this function. - * \param layout_rewrite Rewrite the layout of placeholders specified by - * attr `layout_free_placeholders`. - * \return A `te.schedule` and the an Array of `te.Tensor` to be used in `tvm.lower` - * or `tvm.build`. - */ - std::pair> ApplySteps( - const Array& transform_steps, Array* stages = nullptr, - StageToAxesMap* stage_to_axes = nullptr, - LayoutRewriteOption layout_rewrite = LayoutRewriteOption::NoRewrite) const; - - /*! - * \brief Print transform steps as equivalent python schedule API. - * This can be used for debugging. - * \param transform_steps Transform steps of a state. - * \return The Python schedule code. - */ - String PrintStepsAsPython(const Array& transform_steps) const; - - /*! - * \brief Print the compute DAG to a string. This is also used to generate the ComputeDAG hash. - * \param simple_mode Simple mode will only include the op names and brief compute. - * \return The ComputeDAG in a string. - */ - String PrintDAG(bool simple_mode = false) const; - - /*! - * \brief Fill the correct bound information for a given state by calling ir_pass::InferBound. - * The states can lose complete bound information after some transform steps (e.g., compute_at). - * We can call this function to infer and fill all the bound information. - * This function calls TVM InferBound pass internally to get the bound. - * The returned state of this function is guaranteed to have complete bound information. - * \param state The input state. - * \return The State with complete bound information - */ - State InferBound(const State& state) const; - - /*! - * \brief Fill the correct bound information for the given states by calling ir_pass::InferBound. - * The states can lose complete bound information after some transform steps (e.g., compute_at). - * We can call this function to infer and fill all the bound information. - * This function calls TVM InferBound pass internally to get the bound. - * The returned state of this function is guaranteed to have complete bound information. - * \param states The input states. - * \return The States with complete bound information. - * \note The returned array will contains empty State, if there're infer bound failure on some - * states. - */ - Array InferBound(const Array& states) const; - - /*! - * \brief Since some steps may change the ComputeDAG (e.g. CacheRead/CacheWrite), the initial - * ComputeDAG may not be up-to-date. This function replays the given transform steps from the - * initial state and returns an up-to-date ComputeDAG. - * \param steps The steps to be replayed. Usually we'll filter out the unused steps to speed up - * the replay process, since we only intend to get a ComputeDAG with the up-to-date op stage - * structure. - * \return The up-to-date ComputeDAG. - */ - ComputeDAG ReplayAndGetDAG(const Array& steps) const; - - static constexpr const char* layout_free_placeholders_key = "layout_free_placeholders"; - - TVM_DEFINE_OBJECT_REF_METHODS(ComputeDAG, ObjectRef, ComputeDAGNode); - TVM_DEFINE_OBJECT_REF_COW_METHOD(ComputeDAGNode); -}; - -/*! - * \brief Get the orginal shape from a rewritten layout string. - * \param rewritten_layout The layout after auto-scheduler's layout rewrite. - * \param axis_names Specifiy the names of axes. - * \return shape The original shape. - */ -Array GetShapeFromRewrittenLayout(String rewritten_layout, Array axis_names); - -} // namespace auto_scheduler -} // namespace tvm - -#endif // TVM_AUTO_SCHEDULER_COMPUTE_DAG_H_ diff --git a/include/tvm/auto_scheduler/cost_model.h b/include/tvm/auto_scheduler/cost_model.h deleted file mode 100755 index a52c6797b6d5..000000000000 --- a/include/tvm/auto_scheduler/cost_model.h +++ /dev/null @@ -1,165 +0,0 @@ -/* - * 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 auto_scheduler/cost_model.h - * \brief Cost models that estimate the performance of programs - */ - -#ifndef TVM_AUTO_SCHEDULER_COST_MODEL_H_ -#define TVM_AUTO_SCHEDULER_COST_MODEL_H_ - -#include -#include -#include -#include - -#include - -namespace tvm { -namespace auto_scheduler { - -using runtime::PackedFunc; -using runtime::TypedPackedFunc; - -/*! \brief The base class for cost model */ -class CostModelNode : public Object { - public: - /*! - * \brief Update the cost model according to new measurement results (training data). - * \param inputs The measure inputs - * \param results The measure results - */ - virtual void Update(const Array& inputs, const Array& results) = 0; - - /*! - * \brief Predict the scores of states - * \param task The search task of states - * \param states The input states - * \param scores The predicted scores for all states - */ - virtual void Predict(const SearchTask& task, const Array& states, - std::vector* scores) = 0; - - /*! - * \brief Predict the scores of all stages in states. This is the breakdown version of `Predict` - * \param task The search task - * \param states The input states - * \param state_scores The predicted scores for all states - * \param stage_scores The predicted scores for all stages in all stages - */ - virtual void PredictStages(const SearchTask& task, const Array& states, - std::vector* state_scores, - std::vector>* stage_scores) { - LOG(FATAL) << "Not implemented"; - } - - /*! - * \brief Default virtual destructor - */ - virtual ~CostModelNode() {} - - static constexpr const char* _type_key = "auto_scheduler.CostModel"; - TVM_DECLARE_BASE_OBJECT_INFO(CostModelNode, Object); -}; - -/*! - * \brief Managed reference to CostModelNode. - * \sa CostModelNode - */ -class CostModel : public ObjectRef { - public: - TVM_DEFINE_MUTABLE_OBJECT_REF_METHODS(CostModel, ObjectRef, CostModelNode); -}; - -/*! \brief The cost model returning random value for all predictions */ -class RandomModelNode : public CostModelNode { - public: - /*! \brief Pointer to a random number generator function */ - const TypedPackedFunc* random_number_func; - - void Update(const Array& inputs, const Array& results) final; - - void Predict(const SearchTask& task, const Array& states, - std::vector* scores) final; - - static constexpr const char* _type_key = "auto_scheduler.RandomModel"; - TVM_DECLARE_FINAL_OBJECT_INFO(RandomModelNode, CostModelNode); -}; - -/*! - * \brief Managed reference to RandomModelNode. - * \sa RandomModelNode - */ -class RandomModel : public CostModel { - public: - RandomModel(); - explicit RandomModel(::tvm::runtime::ObjectPtr<::tvm::runtime::Object> n) : CostModel(n) {} - - RandomModelNode* operator->() const { return static_cast(data_.get()); } - - TVM_DEFINE_DEFAULT_COPY_MOVE_AND_ASSIGN(RandomModel); - using ContainerType = RandomModelNode; -}; - -/*! \brief A wrapper for cost model defined by python code - * This class will call functions defined in the python */ -class PythonBasedModelNode : public CostModelNode { - public: - /*! \brief Pointer to the update function in python */ - PackedFunc update_func; - /*! \brief Pointer to the predict function in python */ - PackedFunc predict_func; - /*! \brief Pointer to the predict function in python */ - PackedFunc predict_stage_func; - - void Update(const Array& inputs, const Array& results) final; - - void Predict(const SearchTask& task, const Array& states, - std::vector* scores) final; - - void PredictStages(const SearchTask& task, const Array& states, - std::vector* state_scores, - std::vector>* stage_scores) final; - - static constexpr const char* _type_key = "auto_scheduler.PythonBasedModel"; - TVM_DECLARE_FINAL_OBJECT_INFO(PythonBasedModelNode, CostModelNode); -}; - -/*! - * \brief Managed reference to PythonBasedModelNode. - * \sa PythonBasedModelNode - */ -class PythonBasedModel : public CostModel { - public: - /*! - * \brief The constructor. - * \param update_func The pointer to the update function defined in python - * \param predict_func The pointer to the prediction function defined in python - * \param predict_stage_func The pointer to the prediction function defined in python - */ - PythonBasedModel(PackedFunc update_func, PackedFunc predict_func, PackedFunc predict_stage_func); - - TVM_DEFINE_MUTABLE_OBJECT_REF_METHODS(PythonBasedModel, CostModel, PythonBasedModelNode); -}; - -} // namespace auto_scheduler -} // namespace tvm - -#endif // TVM_AUTO_SCHEDULER_COST_MODEL_H_ diff --git a/include/tvm/auto_scheduler/feature.h b/include/tvm/auto_scheduler/feature.h deleted file mode 100644 index a8b88b7f11f9..000000000000 --- a/include/tvm/auto_scheduler/feature.h +++ /dev/null @@ -1,124 +0,0 @@ -/* - * 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 auto_scheduler/feature.h - * \brief Feature extraction for the cost model. - * We extract one feature vector per BufferStoreNode statement in a TIR Stmt, - * so we call this feature as "per-store" feature. - * The cost model also does prediction for each BufferStoreNode statement and aggregates - * the predictions as the whole score for a TVM IR (Stmt). - * - * The feature specification is defined by `src/auto_scheduler/feature.cc:: FeatureSet` - */ - -#ifndef TVM_AUTO_SCHEDULER_FEATURE_H_ -#define TVM_AUTO_SCHEDULER_FEATURE_H_ - -#include -#include -#include - -#include -#include - -namespace tvm { -namespace auto_scheduler { - -/*! - * \brief Get per-store features from a TIR PrimFunc - * \param func The input lowered TIR PrimFunc - * \param cache_line_size The size of cache line in bytes - * \param max_n_bufs The maximum number of extracted buffers for one statement - * \param ret The returned feature vector - * \param log_scale Should the outputs be scaled by log2(1+x). - */ -void GetPerStoreFeature(const PrimFunc& func, int cache_line_size, int max_n_bufs, - std::vector* ret, bool log_scale = true); - -/* - * \brief Get the names of elements in the feature vector. Use this for debug and inspection. - * \param max_n_bufs The maximum number of extracted buffers for one statement - * \param ret The returned names. - */ -void GetPerStoreFeatureName(int max_n_bufs, std::vector* ret); - -/*! - * \brief Get per-store feature from states of the same task - * \param states The input states - * \param task The same search task for all states - * \param skip_first_n_feature_extraction Skip feature extraction for the first n states - * \param max_n_bufs The maximum number of extracted buffers for one statement - * \param features The returned feature vector. The innermost vector contains the - * feature vectors for all BufferStoreNode statements - */ -void GetPerStoreFeaturesFromStates(const Array& states, const SearchTask& task, - int skip_first_n_feature_extraction, int max_n_bufs, - std::vector>* features); - -/*! - * \brief Get per-store feature from states of different tasks - * \param states The input states - * \param tasks The search tasks corresponding to the input states - * \param skip_first_n_feature_extraction Skip feature extraction for the first n states - * \param max_n_bufs The maximum number of extracted buffers for one statement - * \param features The returned feature vector. The innermost vector contains the - * feature vectors for all BufferStoreNode statements - */ -void GetPerStoreFeaturesFromStates(const Array& states, const std::vector& tasks, - int skip_first_n_feature_extraction, int max_n_bufs, - std::vector>* features); - -/*! - * \brief Get per-store features from a log file - * \param filename The name of log file - * \param max_lines Only read the first n lines of the file - * \param max_n_bufs The maximum number of extracted buffers for one statement - * \param features The returned feature vector. The innermost vector contains the - * feature vectors for all BufferStoreNode statements - * \param normalized_throughputs The normalized throughputs for all states - * \param task_ids The task ids for all states - */ -void GetPerStoreFeaturesFromFile(const std::string& filename, int max_lines, int max_n_bufs, - std::vector>* features, - std::vector* normalized_throughputs, - std::vector* task_ids); - -/*! - * \brief Get per-store features from measurement input/result pairs - * \param inputs The measurement inputs - * \param results The measurement results - * \param skip_first_n_feature_extraction Skip feature extraction for the first n measurement pairs - * \param max_n_bufs The maximum number of extracted buffers for one statement - * \param features The returned feature vector. The innermost vector contains the - * feature vectors for all BufferStoreNode statements - * \param normalized_throughputs The normalized throughputs for all states - * \param task_ids The task ids for all states - */ -void GetPerStoreFeaturesFromMeasurePairs(const Array& inputs, - const Array& results, - int skip_first_n_feature_extraction, int max_n_bufs, - std::vector>* features, - std::vector* normalized_throughputs, - std::vector* task_ids); - -} // namespace auto_scheduler -} // namespace tvm - -#endif // TVM_AUTO_SCHEDULER_FEATURE_H_ diff --git a/include/tvm/auto_scheduler/loop_state.h b/include/tvm/auto_scheduler/loop_state.h deleted file mode 100755 index 0ca14c43eb47..000000000000 --- a/include/tvm/auto_scheduler/loop_state.h +++ /dev/null @@ -1,482 +0,0 @@ -/* - * 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 auto_scheduler/loop_state.h - * \brief The definition of the "state" in the search. - * - * Each LoopState corresponds to a schedule for its ComputeDAG. - * A LoopState consists of: 1. a current loop structure; 2. a list of transformation steps used to - * construct the loop structure. - * The loop structure keeps a preview of how the schedule will finally look like after lowering the - * current state (e.g. number of iterators, the extent of each iterator, the compute_at locations - * ...). - * During the schedule search process, the loop structure can provide search policy with necessary - * information on how to manipulate the current state. - * The transform history is a sequence of `TransformStep` which will finally be mapped to TVM - * schedule primitives. The steps are also used for the serialization of a state. - * - * The LoopState can be seen as a lightweight loop structure IR specifically for schedule search. - * We don't use the existing TVM IR but to extend a new structure on it is because: - * 1. We want fast incremental change to the loop structures. The search policy needs to get the - * immediate loop structures update rather than after TVM lowering; - * 2. We want serializable transform history for replay, backtracking, and mutation; - * 3. We may create some macro schedule primitives that represent the combination of several - * TVM schedule primitives. - * - * When the search is finished, we will lower the state to TVM IR with TVM's schedule primitives. - * Since we share a lot of common objects during search, the transformation is implemented in - * copy on write style. All objects are immutable, which is similar to TVM IR. - */ - -#ifndef TVM_AUTO_SCHEDULER_LOOP_STATE_H_ -#define TVM_AUTO_SCHEDULER_LOOP_STATE_H_ - -#include -#include - -#include -#include -#include -#include - -namespace tvm { -namespace auto_scheduler { - -using namespace tvm::tir; - -class ComputeDAG; - -/*! \brief The type of a stage. */ -enum class StageKind : int { - /*! \brief A placeholder stage. */ - kPlaceholder = 0, - /*! \brief A compute stage. */ - kCompute = 1 -}; - -/*! \brief The type of compute location. */ -enum class ComputeAtKind : int { - /*! \brief Compute at root. */ - kRoot = 0, - /*! \brief Compute inlined. */ - kInlined = 1, - /*! \brief Compute at some iterator. */ - kIter = 2, -}; - -/*! \brief Stage-level attributes. */ -struct StageAttributes { - /*! \brief The maximum steps for the pragma `auto_unroll_max_step`. */ - int auto_unroll_max_step; - /*! \brief The storage offset for the schedule primitive `storage_align`. */ - int storage_offset; -}; - -/*! - * \brief A op stage in the compute declaration. - * Similar to te::Stage in `include/tvm/te/schedule.h`. - */ -class StageNode : public Object { - public: - /*! \brief The operator of this stage */ - te::Operation op; - /*! \brief The iterators in this stage. */ - Array iters; - /*! \brief The type of this stage. */ - StageKind op_type; - /*! \brief The compute location of this stage. */ - ComputeAtKind compute_at; - /*! \brief Other stage-level attributes. */ - StageAttributes attrs; - - void VisitAttrs(tvm::AttrVisitor* v) { - v->Visit("op", &op); - v->Visit("iters", &iters); - v->Visit("op_type", &op_type); - v->Visit("compute_at", &compute_at); - } - - static constexpr const char* _type_key = "auto_scheduler.Stage"; - TVM_DECLARE_FINAL_OBJECT_INFO(StageNode, Object); -}; - -/*! - * \brief Managed reference to StageNode. - * \sa StageNode - */ -class Stage : public ObjectRef { - public: - /*! - * \brief The constructor. - * \param op A `te::Operation`. - */ - explicit Stage(te::Operation op); - /*! - * \brief The constructor. - * \param op The source operation - * \param op_type The stage type of this op. - * \param iters The iterators of this op. - * \param compute_at The compute at type of this op. - * \param attrs Other stage-level attributes. - */ - Stage(te::Operation op, StageKind op_type, const Array& iters, ComputeAtKind compute_at, - StageAttributes attrs); - - TVM_DEFINE_OBJECT_REF_METHODS(Stage, ObjectRef, StageNode); - TVM_DEFINE_OBJECT_REF_COW_METHOD(StageNode); -}; - -/*! \brief Use stage_id to represent a stage. */ -using StageKey = int; -/*! \brief Use stage_id and iter_id to represent a iterator. */ -using IterKey = std::pair; - -/*! - * \brief stores the compute_at relation between stages - * This stores a bi-directional mapping from stages and iter: - * 1. Stage to its attached iterator - * 2. Iterator to the stage attached to it - * You can use AttachMapNode::stage_to_attach_iter and AttachMapNode::iter_to_attached_stages - * to query the relations - */ -class AttachMapNode : public Object { - public: - struct IterKeyHash { - std::size_t operator()(const IterKey& k) const { - return ::dmlc::HashCombine(std::hash()(k.first), std::hash()(k.second)); - } - }; - - /*! \brief A Map to store the mapping of stage to its attached iterator. */ - std::unordered_map stage_to_attach_iter; - /*! \brief A Map to store the mapping of iterator to the stages attached to it. */ - std::unordered_map, IterKeyHash> iter_to_attached_stages; - - static constexpr const char* _type_key = "auto_scheduler.AttachMap"; - TVM_DECLARE_FINAL_OBJECT_INFO(AttachMapNode, Object); -}; - -/*! - * \brief Managed reference to AttachMapNode. - * \sa AttachMapNode - */ -class AttachMap : public ObjectRef { - public: - /*! - * \brief Process the stage/iterator mapping after compute at. - * \param stage_id The index of the source stage of computed at. - * \param target_stage_id The index of stage that this step will compute at to. - * \param target_iter_id The index of target iterator in the target stage. - */ - void SetComputeAtIter(int stage_id, int target_stage_id, int target_iter_id); - - /*! - * \brief Delete the entry of a specific stage. This is a public wrapper of `DeleteStageEntry`. - * \param stage_id The index of the stage to be deleted. - */ - void DeleteStage(int stage_id); - - /*! - * \brief Find the relations of original iterators in AttachMap, and update them with the new - * iterators. Both `stage_to_attach_iter` and `iter_to_attached_stages` will be updated. - * \param original_iters The original IterKey. - * \param new_iters The new IterKey for replacing the old ones. - */ - void UpdateIters(const std::vector& original_iters, - const std::vector& new_iters); - - /*! - * \brief Traverse through `stage_to_attach_iter` and `iter_to_attached_stages` map, add offset - * to stage indexes that are larger than the start_id. Used for steps that insert new stages to - * ComputeDAG (e.g., CacheRead/CacheWrite step). - * \param start_id The index threshold. This function only adds offset for stages - * with indices larger then this threshold. - * \param offset The index offset to be added to the stage index. - * \return The updated AttachMap after applying stage index offset. - */ - AttachMap ApplyStageIdOffset(int start_id, int offset = 1) const; - - TVM_DEFINE_OBJECT_REF_METHODS(AttachMap, ObjectRef, AttachMapNode); - TVM_DEFINE_OBJECT_REF_COW_METHOD(AttachMapNode); - - private: - /*! - * \brief Delete the entry of a specific stage. This will remove the items related to this - * stage in both `stage_to_attach_iter` and `iter_to_attached_stages` map. - * \param pnode A mutable pointer to AttachMapNode. - * \param stage_id The index of stage that will be removed from the map. - */ - static void DeleteStageEntry(AttachMapNode* pnode, int stage_id); -}; - -/*! - * \brief A state in the search process. - * It consists of the current loop structure and a list of transformation steps used to construct - * it. - * Each State corresponds to a specific schedule for its ComputeDAG. - */ -class StateNode : public Object { - public: - /*! \brief Current stages and loop structures. */ - Array stages; - /*! \brief History transformation steps. */ - Array transform_steps; - /*! - * \brief The attach relations of stages and iterators. This is used to track the compute at - * operation. - */ - AttachMap attach_map; - /*! \brief The up-to-date ComputeDAG of this state. The default value is an empty NullOpt, - * meaning the dag of this state is the same as the original ComputeDAG in the SearchTask. - * Otherwise, the stored value is the up-to-date ComputeDAG for this state, meaning some steps - * (e.g., CacheReadStep/CacheWriteStep) have modified the ComputeDAG. - */ - Optional current_compute_dag; - /*! - * \brief Indicate whether this state has unfilled tile sizes. A concrete state means that all - * tile sizes of the state is filled. Only concrete state can be apply to TVM schedule. - */ - bool concrete; - - void VisitAttrs(tvm::AttrVisitor* v) { - v->Visit("stages", &stages); - v->Visit("transform_steps", &transform_steps); - v->Visit("concrete", &concrete); - } - - static constexpr const char* _type_key = "auto_scheduler.State"; - TVM_DECLARE_FINAL_OBJECT_INFO(StateNode, Object); -}; - -/*! - * \brief Managed reference to StateNode. - * \sa StateNode - */ -class State : public ObjectRef { - public: - /*! - * \brief The constructor. - * \param ops `te::Operation`s for a compute declaration. - */ - explicit State(const Array& ops); - - /*! - * \brief Pretty-print the state to a human readable string. - * \param delete_trivial_loop True for skipping the trivial loops. - * (undefined or extent == 1, default set to True) - * \return The human readable string. - */ - String ToStr(bool delete_trivial_loop = true) const; - - /********** Step APIs working on a single stage **********/ - /*! - * \brief The schedule primitive corresponding to `te::Stage::bind`. - * \param stage_id The index of the stage to be binded. - * \param it The iterator to be binded. - * \param thread_type The thread type. - * \return The new iterator after binding. - */ - TVM_DLL Iterator bind(int stage_id, const Iterator& it, IteratorAnnotation thread_type); - /*! - * \brief The schedule primitive corresponding to `te::Stage::parallel`. - * \param stage_id The index of the stage to be paralleled. - * \param it The iterator to be paralleled. - * \return The new iterator after parallel. - */ - TVM_DLL Iterator parallel(int stage_id, const Iterator& it); - /*! - * \brief The schedule primitive corresponding to `te::Stage::unroll`. - * \param stage_id The index of the stage to be unrolled. - * \param it The iterator to be unrolled. - * \param max_unroll The max unroll limit. Iterator with extent larger than this limit will be - * skipped. - * \return The new iterator after unroll. - */ - TVM_DLL Iterator unroll(int stage_id, const Iterator& it, int max_unroll = -1); - /*! - * \brief The schedule primitive corresponding to `te::Stage::vectorize`. - * \param stage_id The index of the stage to be vectorized. - * \param it The iterator to be vectorized. - * \return The new iterator after vectorization. - */ - TVM_DLL Iterator vectorize(int stage_id, const Iterator& it); - /*! - * \brief The schedule primitive corresponding to `te::Stage::fuse`. - * \param stage_id The index of the stage to be fused. - * \param iters The iterators to be fused. - * \return The iterator result after fuse. - * \note If the iterators to be fused have stages attached at them(by compute_at), the fused - * result will become the new attach point. - */ - TVM_DLL Iterator fuse(int stage_id, const Array& iters); - /*! - * \brief The schedule primitive corresponding to `te.Stage.pragma`. - * \param stage_id The index of the stage to add pragma. - * \param it The iterator to add pragma. - * \param pragma_type The pragma string. - */ - TVM_DLL void pragma(int stage_id, const Iterator& it, const String& pragma_type); - /*! - * \brief The schedule primitive corresponding to `te::Stage::reorder`. - * \param stage_id The index of the stage to be reordered. - * \param order The expected iterator order. - */ - TVM_DLL void reorder(int stage_id, const Array& order); - /*! - * \brief The schedule primitive corresponding to `te::Stage::split`. - * \param stage_id The index of the stage to be split. - * \param it The iterator to be split. - * \param lengths The multiple split factors. Can be None to be filled by search policy. - * \param inner_to_outer Whether the factors go from inner to outer, or from outer to inner. - * \return The new iterator after splitting. - * \note If we do split on an iterator which has stages attached at it(by compute_at), the inner - * most iterator of split results will become the new attach point. - */ - TVM_DLL Array split(int stage_id, const Iterator& it, - const Array>& lengths, - bool inner_to_outer = true); - /*! - * \brief The schedule primitive similar to split, but uses split factors from previous steps. - * \param stage_id The index of the stage to be split. - * \param it The iterator to be split. - * \param src_step_id The index of the split step to be followed in the history. - * \param n_split The number of split level. - * \return The split new Iterators. - */ - TVM_DLL Array follow_split(int stage_id, const Iterator& it, int src_step_id, - int n_split); - /*! - * \brief The schedule primitive similar to split, but uses split factors from - * fused previous steps. - * \param stage_id The index of the stage to be split. - * \param it The iterator to be split. - * \param src_step_ids The indices of the split steps to be followed in the history. - * \param level Use the length in this split level. - * \param factor_or_nparts True to use `factor` for split from inner to outer, - False to use `nparts` for split from outer to inner. - * \return The split new Iterators. - */ - TVM_DLL Array follow_fused_split(int stage_id, const Iterator& it, - const Array& src_step_ids, int level, - bool factor_or_nparts); - /*! - * \brief The schedule primitive corresponding to `te.Stage.storage_align`. - * \param stage_id The index of the stage to be aligned. - * \param it The iterator to be aligned. - * \param factor The factor in alignment specification. - * \param offset The offset in the alignment specification. - */ - TVM_DLL void storage_align(int stage_id, const Iterator& it, int factor, int offset); - - /********** Step APIs working on multiple stages **********/ - /*! - * \brief The schedule primitive corresponding to `te::Stage::compute_at`. - * \param stage_id The index of the source stage of computed at. - * \param target_stage_id The index of stage that this step will compute at to. - * \param target_iter The indiex of the target iterator in the target stage. - * \note After compute_at, we need careful dependency analysis to compute the accurate bound - * information. However, it is relatively expensive and complicated, so we just fill "None" as - * bound for the newly created iterators. - * Call ComputeDAG::InferBound on the updated state if you need the complete bound information. - */ - TVM_DLL void compute_at(int stage_id, int target_stage_id, const Iterator& target_iter); - /*! - * \brief The schedule primitive corresponding to `te::Stage::compute_inline`. - * \param stage_id The index of the stage to be marked compute inlined. - */ - TVM_DLL void compute_inline(int stage_id); - /*! - * \brief The schedule primitive corresponding to `te::Stage::compute_root`. - * \param stage_id The index of the stage to be marked compute at root. - * \note After compute_root, we need careful dependency analysis to compute the accurate bound - * information. However, it is relatively expensive and complicated, so we just fill "None" as - * bound for the newly created iterators. - * Call ComputeDAG::InferBound on the updated state if you need the complete bound information. - */ - TVM_DLL void compute_root(int stage_id); - - /********** Step APIs adding new stages **********/ - /*! - * \brief The schedule primitive corresponding to `te::Schedule::cache_read`. - * \param stage_id The index of the stage to be cache_read. - * \param scope_name The scope name of the newly added stage. - * \param reader_stage_ids The indices of reader stages. - * \param dag The original ComputeDAG of this state. - * \note Cache read step will add an extra stage to the original ComputeDAG (at the back of the - * target stage), an up-to-date ComputeDAG is stored in State's `current_compute_dag`. - */ - TVM_DLL int cache_read(int stage_id, const String& scope_name, - const Array& reader_stage_ids, const ComputeDAG& dag); - /*! - * \brief The schedule primitive corresponding to `te::Schedule::cache_write`. - * \param stage_id The index of the stage to be cache_write. - * \param scope_name The scope name of the newly added stage. - * \param dag The original ComputeDAG of this state. - * \note Cache write step will add an extra stage to the original ComputeDAG (in the front of the - * target stage), an up-to-date ComputeDAG is stored in State's `current_compute_dag`. - * This step will cache write all output tensors of the target stage. - */ - TVM_DLL int cache_write(int stage_id, const String& scope_name, const ComputeDAG& dag); - /*! - * \brief The schedule primitive corresponding to `te::Schedule::rfactor`. - * \param stage_id The index of the iterator to be factored. - * \param it The iterator to be factored. - * \param factor_iter_id The position where the new iterator is placed. - * \param dag The original ComputeDAG of this state. - * \note Rfactor step will add an extra stage to the original ComputeDAG (in the front of the - * target stage), an up-to-date ComputeDAG is stored in State's `current_compute_dag`. - */ - TVM_DLL int rfactor(int stage_id, const Iterator& it, int factor_iter_id, const ComputeDAG& dag); - - TVM_DEFINE_OBJECT_REF_METHODS(State, ObjectRef, StateNode); - TVM_DEFINE_OBJECT_REF_COW_METHOD(StateNode); -}; - -} // namespace auto_scheduler -} // namespace tvm - -// Hash and equal function for State -namespace std { - -/*! - * \brief The equal_to function for auto_scheduler::State. - * This function checks the equality by looking at the lowered string format of states. - * If two states with different transform history have the same lowered string format, - * they will be considered being equal. - */ -template <> -struct equal_to<::tvm::auto_scheduler::State> { - bool operator()(const ::tvm::auto_scheduler::State& lhs, - const ::tvm::auto_scheduler::State& rhs) const { - return lhs.ToStr() == rhs.ToStr(); - } -}; - -/*! \brief The hash function for auto_scheduler::State. */ -template <> -struct hash<::tvm::auto_scheduler::State> { - std::size_t operator()(const ::tvm::auto_scheduler::State& state) const { - return tvm::runtime::ObjectHash()(state.ToStr()); - } -}; - -} // namespace std - -#endif // TVM_AUTO_SCHEDULER_LOOP_STATE_H_ diff --git a/include/tvm/auto_scheduler/measure.h b/include/tvm/auto_scheduler/measure.h deleted file mode 100755 index 8576468816cb..000000000000 --- a/include/tvm/auto_scheduler/measure.h +++ /dev/null @@ -1,542 +0,0 @@ -/* - * 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 auto_scheduler/measure.h - * \brief Distributed measurement infrastructure to measure the runtime costs of tensor programs. - * These functions are responsible for building the tvm module, uploading it to remote devices, - * recording the running time costs, and checking the correctness of the output. - * - * The measurement is separated into two steps: build and run. - * A builder builds the executable binary files and a runner runs the binary files to get the - * measurement results. The flow of data structures is - * - * `ProgramBuilder` `ProgramRunner` - * `MeasureInput` -----------------> `BuildResult` ----------------> `MeasureResult` - * - * The core functions is implemented in python to utilize python's multiprocessing - * and error handling (see also `python/tvm/auto_scheduler/measure.py`). - * This c++ file is just a wrapper for the python functions. - */ - -#ifndef TVM_AUTO_SCHEDULER_MEASURE_H_ -#define TVM_AUTO_SCHEDULER_MEASURE_H_ - -#include -#include - -#include -#include -#include -#include - -namespace tvm { -namespace auto_scheduler { - -class SearchPolicy; -class MeasureInput; -class MeasureResult; - -/*! \brief The error code of one measurement */ -enum class MeasureErrorNO : int { - /*! \brief No error. */ - kNoError = 0, - /*! \brief Errors happen when apply transform steps from init state. */ - kInstantiationError = 1, - /*! \brief Errors happen when compiling code on host. (when build module) */ - kCompileHostError = 2, - /*! \brief Errors happen when compiling code on device. (when load module) */ - kCompileDeviceError = 3, - /*! \brief Errors happen when run program on device. */ - kRuntimeDeviceError = 4, - /*! \brief Answer is wrong when compared to a reference output. */ - kWrongAnswerError = 5, - /*! \brief Timeout during compilation. */ - kBuildTimeoutError = 6, - /*! \brief Timeout during run. */ - kRunTimeoutError = 7, - /*! \brief Unknown error. */ - kUnknownError = 8, -}; - -// Inputs and results of one measurement - -/*! \brief Store the input of a measurement */ -class MeasureInputNode : public Object { - public: - /*! \brief The search task. */ - SearchTask task; - /*! \brief The program state to be measured. */ - State state; - - void VisitAttrs(tvm::AttrVisitor* v) { - v->Visit("task", &task); - v->Visit("state", &state); - } - - /*! \brief Do shallow copy. */ - MeasureInput copy() const; - - static constexpr const char* _type_key = "auto_scheduler.MeasureInput"; - TVM_DECLARE_FINAL_OBJECT_INFO(MeasureInputNode, Object); -}; - -/*! - * \brief Managed reference to MeasureInputNode. - * \sa MeasureInputNode - */ -class MeasureInput : public ObjectRef { - public: - /*! - * \brief The constructor. - * \param task The SearchTask of this measure. - * \param state The State to be measured. - */ - MeasureInput(SearchTask task, State state); - - TVM_DEFINE_OBJECT_REF_METHODS(MeasureInput, ObjectRef, MeasureInputNode); -}; - -/*! \brief Store the result of a build. */ -class BuildResultNode : public Object { - public: - /*! \brief The filename of built binary file. */ - String filename; - /*! \brief The arguments. */ - Array args; - /*! \brief The error code. (0 means no error, see MeasureErrorNO) */ - int error_no; - /*! \brief The error message if there is any error. */ - String error_msg; - /*! \brief The time cost of build. */ - double time_cost; - - void VisitAttrs(tvm::AttrVisitor* v) { - v->Visit("filename", &filename); - v->Visit("args", &args); - v->Visit("error_no", &error_no); - v->Visit("error_msg", &error_msg); - v->Visit("time_cost", &time_cost); - } - - static constexpr const char* _type_key = "auto_scheduler.BuildResult"; - TVM_DECLARE_FINAL_OBJECT_INFO(BuildResultNode, Object); -}; - -/*! - * \brief Managed reference to BuildResultNode. - * \sa BuildResultNode - */ -class BuildResult : public ObjectRef { - public: - /*! - * \brief The constructor. - * \param filename The filename of built binary file. - * \param args The arguments. - * \param error_no The error code. - * \param error_msg The error message if there is any error. - * \param time_cost The time cost of build. - */ - BuildResult(String filename, Array args, int error_no, String error_msg, - double time_cost); - TVM_DEFINE_OBJECT_REF_METHODS(BuildResult, ObjectRef, BuildResultNode); -}; - -/*! \brief Store the results of a measurement. */ -class MeasureResultNode : public Object { - public: - /*! \brief The time costs of execution. */ - Array costs; - /*! \brief The error code. (0 means no error, see MeasureErrorNO) */ - int error_no; - /*! \brief The error message if there is any error. */ - String error_msg; - /*! \brief The time cost of build and run. */ - double all_cost; - /*! \brief The time stamps of this measurement. */ - double timestamp; - - void VisitAttrs(tvm::AttrVisitor* v) { - v->Visit("costs", &costs); - v->Visit("error_no", &error_no); - v->Visit("error_msg", &error_msg); - v->Visit("all_cost", &all_cost); - v->Visit("timestamp", ×tamp); - } - - /*! \brief Do shallow copy. */ - MeasureResult copy() const; - - static constexpr const char* _type_key = "auto_scheduler.MeasureResult"; - TVM_DECLARE_FINAL_OBJECT_INFO(MeasureResultNode, Object); -}; - -/*! - * \brief Managed reference to MeasureResultNode. - * \sa MeasureResultNode - */ -class MeasureResult : public ObjectRef { - public: - /*! - * \brief The constructor. - * \param costs The time costs of execution. - * \param error_no The error code. - * \param error_msg The error message if there is any error. - * \param all_cost The time cost of build and run. - * \param timestamp The time stamps of this measurement. - */ - MeasureResult(Array costs, int error_no, String error_msg, double all_cost, - double timestamp); - - TVM_DEFINE_OBJECT_REF_METHODS(MeasureResult, ObjectRef, MeasureResultNode); -}; - -/*! \brief Bass class of measurement callbacks */ -class MeasureCallbackNode : public Object { - public: - /*! - * \brief Callback function that will be called on measurement input/result pairs - * after each measurement batch. - * \param policy The current search policy. - * \param inputs An Array of MeasureInput. - * \param results An Array of MeasureResult. - */ - virtual void Callback(const SearchPolicy& policy, const Array& inputs, - const Array& results) = 0; - static constexpr const char* _type_key = "auto_scheduler.MeasureCallback"; - TVM_DECLARE_BASE_OBJECT_INFO(MeasureCallbackNode, Object); -}; - -/*! - * \brief Managed reference to MeasureCallbackNode. - * \sa MeasureCallbackNode - */ -class MeasureCallback : public ObjectRef { - public: - TVM_DEFINE_MUTABLE_OBJECT_REF_METHODS(MeasureCallback, ObjectRef, MeasureCallbackNode); -}; - -/*! \brief A wrapper for measure callback defined by python code - * This class will call functions defined in the python */ -class PythonBasedMeasureCallbackNode : public MeasureCallbackNode { - public: - /*! \brief Pointer to the callback function in python */ - PackedFunc callback_func; - - void Callback(const SearchPolicy& policy, const Array& inputs, - const Array& results) final; - static constexpr const char* _type_key = "auto_scheduler.PythonBasedMeasureCallback"; - TVM_DECLARE_FINAL_OBJECT_INFO(PythonBasedMeasureCallbackNode, MeasureCallbackNode); -}; - -/*! - * \brief Managed reference to PythonBasedMeasureCallbackNode. - * \sa PythonBasedMeasureCallbackNode - */ -class PythonBasedMeasureCallback : public MeasureCallback { - public: - /*! - * \brief The constructor. - * \param callback_func The pointer to the callback function defined in python - */ - explicit PythonBasedMeasureCallback(PackedFunc callback_func); - - TVM_DEFINE_MUTABLE_OBJECT_REF_METHODS(PythonBasedMeasureCallback, MeasureCallback, - PythonBasedMeasureCallbackNode); -}; - -// The base class of ProgramBuilders and ProgramRunners. - -/*! \brief ProgramBuilder that builds the programs */ -class ProgramBuilderNode : public Object { - public: - /*! \brief The number of build processes to run in parallel */ - int n_parallel; - /*! \brief Timeout of a build */ - int timeout; - - /*! - * \brief Build programs and return results. - * \param inputs An Array of MeasureInput. - * \param verbose Verbosity level. 0 for silent, 1 to output information during program - * building. - * \return An Array of MeasureResult. - */ - virtual Array Build(const Array& inputs, int verbose) = 0; - - static constexpr const char* _type_key = "auto_scheduler.ProgramBuilder"; - TVM_DECLARE_BASE_OBJECT_INFO(ProgramBuilderNode, Object); -}; - -/*! - * \brief Managed reference to ProgramBuilderNode. - * \sa ProgramBuilderNode - */ -class ProgramBuilder : public ObjectRef { - public: - TVM_DEFINE_MUTABLE_OBJECT_REF_METHODS(ProgramBuilder, ObjectRef, ProgramBuilderNode); -}; - -/*! \brief ProgramRunner that runs the built programs and measure the time cost. */ -class ProgramRunnerNode : public Object { - public: - /*! \brief Timeout of a run. */ - int timeout; - /*! \brief The number of times to run the generated code for taking average. */ - int number; - /*! \brief The number of times to repeat the measurement. */ - int repeat; - /*! \brief The minimum duration of one repeat in milliseconds. */ - int min_repeat_ms; - /*! \brief The cool down interval between two measurements. */ - double cooldown_interval; - /*! \brief Whether to flush cache on CPU between repeated measurements. */ - bool enable_cpu_cache_flush; - /*! \brief Which device to run on if multiple are avaialble. */ - int device; - - /*! - * \brief Run measurement and return results. - * \param inputs An Array of MeasureInput. - * \param build_results An Array of BuildResult. - * \param verbose Verbosity level. 0 for silent, 1 to output information during program - * running. - * \return An Array of MeasureResult. - */ - virtual Array Run(const Array& inputs, - const Array& build_results, int verbose) = 0; - - static constexpr const char* _type_key = "auto_scheduler.ProgramRunner"; - TVM_DECLARE_BASE_OBJECT_INFO(ProgramRunnerNode, Object); -}; - -/*! - * \brief Managed reference to ProgramRunnerNode. - * \sa ProgramRunnerNode - */ -class ProgramRunner : public ObjectRef { - public: - TVM_DEFINE_MUTABLE_OBJECT_REF_METHODS(ProgramRunner, ObjectRef, ProgramRunnerNode); -}; - -// Implementation of various builders and runners - -/*! \brief LocalBuilder use local CPU cores to build programs in parallel */ -class LocalBuilderNode : public ProgramBuilderNode { - public: - /*! \brief Build function. */ - String build_func; - - Array Build(const Array& inputs, int verbose) final; - - static constexpr const char* _type_key = "auto_scheduler.LocalBuilder"; - TVM_DECLARE_FINAL_OBJECT_INFO(LocalBuilderNode, ProgramBuilderNode); -}; - -/*! - * \brief Managed reference to LocalBuilderNode. - * \sa LocalBuilderNode - */ -class LocalBuilder : public ProgramBuilder { - public: - /*! - * \brief The constructor. - * \param timeout The timeout limit (in second) for each build thread. - * This will be used in a wrapper of the multiprocessing.Process.join(). - * \param n_parallel The number of threads used to build in parallel. - * \param build_func The name of the registered build function. - */ - LocalBuilder(int timeout, int n_parallel, const String& build_func); - - TVM_DEFINE_OBJECT_REF_METHODS(LocalBuilder, ProgramBuilder, LocalBuilderNode); -}; - -/*! \brief LocalRunner that uses local CPU/GPU to measure the time cost of programs */ -class LocalRunnerNode : public ProgramRunnerNode { - public: - Array Run(const Array& inputs, - const Array& build_results, int verbose) final; - - static constexpr const char* _type_key = "auto_scheduler.LocalRunner"; - TVM_DECLARE_FINAL_OBJECT_INFO(LocalRunnerNode, ProgramRunnerNode); -}; - -/*! - * \brief Managed reference to LocalRunnerNode. - * \sa LocalRunnerNode - */ -class LocalRunner : public ProgramRunner { - public: - /*! - * \brief The constructor. See the corresponding class in python/tvm/auto_scheduler/measure.py - * for more detailed parameter explanation. - * \param timeout The timeout limit (in second) for each run. - * This is used in a wrapper of the multiprocessing.Process.join(). - * \param number The number of times to run the generated code for taking average. - * \param repeat The number of times to repeat the measurement. - * \param min_repeat_ms The minimum duration of one repeat in milliseconds. - * \param cooldown_interval The cool down interval between two measurements. - * \param enable_cpu_cache_flush Whether to flush cache on CPU between repeated measurements. - * \param device Which device to run on if multiple are available. - */ - LocalRunner(int timeout, int number, int repeat, int min_repeat_ms, double cooldown_interval, - bool enable_cpu_cache_flush, int device); - - TVM_DEFINE_MUTABLE_OBJECT_REF_METHODS(LocalRunner, ProgramRunner, LocalRunnerNode); -}; - -/*! - * \brief RPCRunner that uses RPC call to measures the time cost of programs on remote devices. - * Or sometime we may need to use RPC even in local running to insulate the thread environment. - * (e.g. running CUDA programs) - */ -class RPCRunnerNode : public ProgramRunnerNode { - public: - /*! \brief The key of the device registered in the RPC tracker. */ - String key; - /*! \brief The host address of the RPC Tracker. */ - String host; - /*! \brief The port of the RPC Tracker. */ - int port; - /*! \brief The priority of this run request, larger is more prior. */ - int priority; - /*! \brief The number of tasks run in parallel. */ - int n_parallel; - - Array Run(const Array& inputs, - const Array& build_results, int verbose) final; - - static constexpr const char* _type_key = "auto_scheduler.RPCRunner"; - TVM_DECLARE_FINAL_OBJECT_INFO(RPCRunnerNode, ProgramRunnerNode); -}; - -/*! - * \brief Managed reference to RPCRunnerNode. - * \sa RPCRunnerNode - */ -class RPCRunner : public ProgramRunner { - public: - /*! - * \brief The constructor. See the corresponding class in python/tvm/auto_scheduler/measure.py - * for more detailed parameter explanation. - * \param key The key of the device registered in the RPC tracker. - * \param host The host address of the RPC Tracker. - * \param port The port of RPC Tracker. - * \param priority The priority of this run request, larger is more prior. - * \param n_parallel The number of tasks run in parallel. - * \param timeout Timeout of a run. - * \param number The number of times to run the generated code for taking average. - * \param repeat The number of times to repeat the measurement. - * \param min_repeat_ms The minimum duration of one repeat in milliseconds. - * \param cooldown_interval The cool down interval between two measurements. - * \param enable_cpu_cache_flush Whether to flush cache on CPU between repeated measurements. - * \param device Which device to run on if multiple are available. - */ - RPCRunner(const String& key, const String& host, int port, int priority, int n_parallel, - int timeout, int number, int repeat, int min_repeat_ms, double cooldown_interval, - bool enable_cpu_cache_flush, int device); - - TVM_DEFINE_MUTABLE_OBJECT_REF_METHODS(RPCRunner, ProgramRunner, RPCRunnerNode); -}; - -/*! - * \brief Measurer that measures the time costs of tvm programs - * This class combines ProgramBuilder and ProgramRunner, and provides a simpler API */ -class ProgramMeasurerNode : public Object { - public: - /*! \brief Measured programs counter. */ - int ct; - /*! \brief Continuous error counter. */ - int error_ct; - /*! \brief Workload key to best flops map. */ - std::unordered_map best_flops; - /*! \brief Workload key to best state map. */ - std::unordered_map best_state; - /*! \brief Workload key to best state's count index map. */ - std::unordered_map best_ct; - /*! \brief The set of workloads that have at least one valid schedule */ - std::unordered_set has_valid; - /*! \brief The ProgramBuilder to build each program. */ - ProgramBuilder builder; - /*! \brief The ProgramRunner to measure each program. */ - ProgramRunner runner; - /*! \brief MeasureCallback to be called after each measure batch. */ - Optional> callbacks; - /*! \brief Verbosity level. 0 for silent, 1 to output information during program measuring. */ - int verbose; - /*! \brief The number of allowed maximum continuous error before forcely stopping the tuning */ - int max_continuous_error; - - /*! \brief Reset book keeping variables */ - void Reset(); - - /*! - * \brief Do measurement. - * \param task The current SearchTask. - * \param policy The current SearchPolicy. - * \param inputs The inputs of measurement. - * \param batch_size Number of programs to be measured in one batch. - * \return results The results of measurement. - */ - Array Measure(const SearchTask& task, const SearchPolicy& policy, - const Array& inputs, int batch_size = -1); - /*! - * \brief Do measurement silently. - * This API will not print the measure results to screen. - * \param task The current SearchTask. - * \param inputs The MeasureInputs. - * \param results A pointer to a MeasureResult Array, this is used as output. - */ - void SilentMeasure(const SearchTask& task, const Array& inputs, - Array* results); - - /*! \brief The default max continuous error setting. */ - static const int DEFAULT_MAX_CONTINUOUS_ERROR = 150; - - static constexpr const char* _type_key = "auto_scheduler.ProgramMeasurer"; - TVM_DECLARE_FINAL_OBJECT_INFO(ProgramMeasurerNode, Object); -}; - -/*! - * \brief Managed reference to ProgramMeasurerNode. - * \sa ProgramMeasurerNode - */ -class ProgramMeasurer : public ObjectRef { - public: - /*! - * \brief The constructor. - * \param builder The ProgramBuilder to build programs. - * \param runner The ProgramRunner to measure programs. - * \param callbacks MeasureCallback to be called after each measurement batch. - * \param verbose Verbosity level. 0 for silent, 1 to output information during program - * measuring. - * \param max_continuous_error The number of allowed maximum continuous error before - * forcely stopping the tuning. - */ - ProgramMeasurer(ProgramBuilder builder, ProgramRunner runner, - Optional> callbacks, int verbose, - int max_continuous_error = -1); - - TVM_DEFINE_MUTABLE_OBJECT_REF_METHODS(ProgramMeasurer, ObjectRef, ProgramMeasurerNode); -}; - -} // namespace auto_scheduler -} // namespace tvm - -#endif // TVM_AUTO_SCHEDULER_MEASURE_H_ diff --git a/include/tvm/auto_scheduler/measure_record.h b/include/tvm/auto_scheduler/measure_record.h deleted file mode 100755 index c82ed076eca7..000000000000 --- a/include/tvm/auto_scheduler/measure_record.h +++ /dev/null @@ -1,140 +0,0 @@ -/* - * 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 tvm/auto_scheduler/measure_record.h - * \brief Json serialization format for dumping and loading measurement records. - */ - -#ifndef TVM_AUTO_SCHEDULER_MEASURE_RECORD_H_ -#define TVM_AUTO_SCHEDULER_MEASURE_RECORD_H_ - -#include - -#include -#include -#include - -namespace tvm { -namespace auto_scheduler { - -const std::string AUTO_SCHEDULER_LOG_VERSION = "v0.6"; // NOLINT(*) - -/*! \brief Callback for logging the input and results of measurements to file */ -class RecordToFileNode : public MeasureCallbackNode { - public: - /*! \brief The name of output file. */ - String filename; - - void Callback(const SearchPolicy& policy, const Array& inputs, - const Array& results) final; - - static constexpr const char* _type_key = "auto_scheduler.RecordToFile"; - TVM_DECLARE_FINAL_OBJECT_INFO(RecordToFileNode, MeasureCallbackNode); -}; - -/*! - * \brief Managed reference to RecordToFileNode. - * \sa RecordToFileNode - */ -class RecordToFile : public MeasureCallback { - public: - /*! - * \brief The constructor. - * \param filename The name of output file - */ - explicit RecordToFile(String filename); - - TVM_DEFINE_MUTABLE_OBJECT_REF_METHODS(RecordToFile, MeasureCallback, RecordToFileNode); -}; - -/*! \brief Log reader to load step logs from a file.*/ -class RecordReaderNode : public Object { - public: - /*! \brief The name of input file. */ - String filename; - /*! \brief The reading file stream. */ - std::ifstream infile; - - ~RecordReaderNode(); - - /*! - * \brief Read next line in the log file. - * \param inp A pointer to a MeasureInputNode, this is used as output. - * \param res A pointer to a MeasureResultNode, this is used as output. - * \return Whether the read is successful. */ - bool ReadNext(MeasureInputNode* inp, MeasureResultNode* res); - - /*! - * \brief Read multiple lines from the log file. - * \param max_size The maximum number of lines. -1 means read all lines. - * \param skip_size Skip the first n lines. - * \return The MeasureInputs and MeasureResults loaded from the log file. - */ - std::pair, Array> ReadLines(int max_size = -1, - int skip_size = 0); - - static constexpr const char* _type_key = "auto_scheduler.RecordReader"; - TVM_DECLARE_FINAL_OBJECT_INFO(RecordReaderNode, Object); - - private: - /*! \brief A string storing the current line. */ - std::string cur_line_; -}; - -/*! - * \brief Managed reference to RecordReaderNode. - * \sa RecordReaderNode - */ -class RecordReader : public ObjectRef { - public: - /*! - * \brief The constructor. - * \param filename The name of input file - */ - explicit RecordReader(String filename); - - TVM_DEFINE_MUTABLE_OBJECT_REF_METHODS(RecordReader, ObjectRef, RecordReaderNode); -}; - -/*! - * \brief Append measure records to an output stream. - * \param os A pointer to a output stream. - * \param inputs The MeasureInputs to be written. - * \param results The MeasureResults to be written. - * \param log_version The log version for the given record. - */ -void WriteMeasureRecords(std::ostream* os, const Array& inputs, - const Array& results, - const std::string log_version = AUTO_SCHEDULER_LOG_VERSION); - -/*! - * \brief Read one measure record from a string. - * \param str The record string to be parsed. - * \param inp A pointer to a MeasureInputNode used to store the return value. - * \param res A pointer to a MeasureResultNode used to store the return value. - * \param log_version A pointer to a string used to store the log version. - */ -void ReadMeasureRecord(const std::string& str, MeasureInputNode* inp, MeasureResultNode* res, - std::string* log_version); - -} // namespace auto_scheduler -} // namespace tvm - -#endif // TVM_AUTO_SCHEDULER_MEASURE_RECORD_H_ diff --git a/include/tvm/auto_scheduler/search_policy.h b/include/tvm/auto_scheduler/search_policy.h deleted file mode 100755 index e433799b7fa5..000000000000 --- a/include/tvm/auto_scheduler/search_policy.h +++ /dev/null @@ -1,206 +0,0 @@ -/* - * 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 tvm/auto_scheduler/search_policy.h - * \brief The base class of search policies, including the abstract definition of search policy and - * other supporting data structures. - * - * \note How to add a new search policy. - * In design, there's no need for users to implement their own search policy, our formal search - * policy(will be brought later) should be enough to cover most use cases. Meanwhile, a custom rule - * mechanism will be provided to enable user-defined template search to serve the same functionality - * as the current AutoTVM template. - * - * This guide is for advanced uses who have special requirements. - * 1. The only function that must be implemented is Search(), which takes a task as input and - * returns the best states found. - * 2. Information about the compute declaration of ops/subgraphs can be acquired from SearchTask. - * This structure also contains some information about the target device. (e.g. knowing the width - * of the device vector unit, we can limit the max vectorize size during schedule search) - * 3. SearchCallback provides more flexibility to do extra affairs before/after the search process. - * 4. ProgramMeasurer provides a simple but useful api to help check the performance of states got - * during the search process. - */ - -#ifndef TVM_AUTO_SCHEDULER_SEARCH_POLICY_H_ -#define TVM_AUTO_SCHEDULER_SEARCH_POLICY_H_ - -#include -#include -#include - -#include -#include -#include -#include - -namespace tvm { -namespace auto_scheduler { - -class ProgramMeasurer; -class SearchPolicyNode; - -/*! - * \brief Callback function to be called by the search process. - * This interface allows to do extra initializations before schedule search or extra - * check during/after the schedule search. - */ -class SearchCallbackNode : public Object { - public: - /*! - * \brief Run the registered callback function. - * \param policy A pointer to a SearchPolicyNode. - */ - virtual void Callback(SearchPolicyNode* policy) = 0; - - static constexpr const char* _type_key = "auto_scheduler.SearchCallback"; - TVM_DECLARE_BASE_OBJECT_INFO(SearchCallbackNode, Object); -}; - -/*! - * \brief Managed reference to SearchCallbackNode. - * \sa SearchCallbackNode - */ -class SearchCallback : public ObjectRef { - public: - TVM_DEFINE_MUTABLE_OBJECT_REF_METHODS(SearchCallback, ObjectRef, SearchCallbackNode); -}; - -/*! \brief Preload measured states from a log file. - * This can resume the state of the search policy */ -class PreloadMeasuredStatesNode : public SearchCallbackNode { - public: - /*! \brief The name of the record log file. */ - String filename; - - void Callback(SearchPolicyNode* policy) final; - - static constexpr const char* _type_key = "auto_scheduler.PreloadMeasuredStates"; - TVM_DECLARE_FINAL_OBJECT_INFO(PreloadMeasuredStatesNode, SearchCallbackNode); -}; - -/*! - * \brief Managed reference to PreloadMeasuredStatesNode. - * \sa PreloadMeasuredStatesNode - */ -class PreloadMeasuredStates : public SearchCallback { - public: - /*! - * \brief The constructor. - * \param filename The name of the record log file. - */ - explicit PreloadMeasuredStates(String filename); - - TVM_DEFINE_MUTABLE_OBJECT_REF_METHODS(PreloadMeasuredStates, SearchCallback, - PreloadMeasuredStatesNode); -}; - -/*! \brief Attribute keys of ops used for SearchPolicy. */ -struct SearchPolicyKey { - /*! \brief Always apply unroll to the inner most iterator of the specificed iterators. */ - static constexpr const char* always_unroll_inner = "auto_scheduler_always_unroll_inner"; - /*! \brief The specified iterators will be placed in the inner most tile without split. */ - static constexpr const char* no_split_at_inner = "auto_scheduler_no_split_at_inner"; - /*! \brief The specified iterators are indices of const tensors in "fake reduction". */ - static constexpr const char* simplify_const_tensor_indices = - "auto_scheduler_simplify_const_tensor_indices"; -}; - -/*! - * \brief The base class of search policies. - */ -class SearchPolicyNode : public Object { - public: - /*! \brief The current search task. */ - SearchTask search_task; - /*! - * \brief Verbose level to control the screen output during schedule search. - * 0 for silent, 1 to output state & measure information during search process. - */ - int verbose; - - void VisitAttrs(AttrVisitor* v) { - v->Visit("search_task", &search_task); - v->Visit("verbose", &verbose); - } - - /*! - * \brief Do schedule search for a task. Takes the SearchTask as input and returns the best state - * found during the search. - * \param num_measure_trials The number of total measurement trials. - * \param early_stopping Stops the tuning early if no improvement after n measurements. - * \param num_measures_per_round The number of programs to be measured at each search round. - * \param measurer A ProgramMeasurer to build and measure programs - * \return The best state found. - */ - virtual State Search(int num_measure_trials, int early_stopping, int num_measures_per_round, - ProgramMeasurer measurer) = 0; - - /*! - * \brief Continue the search by doing an additional search round. - * \param num_measure The number of measurements - * \param measurer The measurer to measure programs - * \return The measurement records for measurements in this search round - */ - virtual std::pair, Array> ContinueSearchOneRound( - int num_measure, ProgramMeasurer measurer) = 0; - - /*! - * \brief Preload measured states from a log file to resume the state of the search policy. - * \param log_file The name of the record log file. - */ - void PreloadMeasuredStates(const String& log_file); - - /*! - * \brief Call SearchCallback with the current SearchPolicyNode - * \param callbacks SearchCallback to be called. - */ - void RunCallbacks(const Array& callbacks); - - static constexpr const char* _type_key = "auto_scheduler.SearchPolicy"; - TVM_DECLARE_BASE_OBJECT_INFO(SearchPolicyNode, Object); - - protected: - /*! - * \brief The set of already measured states. - * We store the string format of a state for redundancy check. This is used to make sure a - * measured state will never be measured again. - */ - std::unordered_set measured_states_set_; - /*! \brief The array of already measured states. - * The good states can be used as the initial population in evolutionary search. */ - std::vector measured_states_vector_; - /*! \brief The throughputs of already measured states */ - std::vector measured_states_throughputs_; -}; - -/*! - * \brief Managed reference to SearchPolicyNode. - * \sa SearchPolicyNode - */ -class SearchPolicy : public ObjectRef { - public: - TVM_DEFINE_MUTABLE_OBJECT_REF_METHODS(SearchPolicy, ObjectRef, SearchPolicyNode); -}; - -} // namespace auto_scheduler -} // namespace tvm - -#endif // TVM_AUTO_SCHEDULER_SEARCH_POLICY_H_ diff --git a/include/tvm/auto_scheduler/search_task.h b/include/tvm/auto_scheduler/search_task.h deleted file mode 100755 index efbc2529592b..000000000000 --- a/include/tvm/auto_scheduler/search_task.h +++ /dev/null @@ -1,171 +0,0 @@ -/* - * 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 auto_scheduler/search_task.h - * \brief Meta information and hardware parameters for a search task. - */ - -#ifndef TVM_AUTO_SCHEDULER_SEARCH_TASK_H_ -#define TVM_AUTO_SCHEDULER_SEARCH_TASK_H_ - -#include -#include -#include - -namespace tvm { -namespace auto_scheduler { - -class HardwareParams; - -/*! \brief The parameters of target hardware used to guide the SearchPolicy. */ -class HardwareParamsNode : public Object { - public: - /*! \brief The number of cores. */ - int num_cores; - /*! \brief The width of vector units in bytes. */ - int vector_unit_bytes; - /*! \brief The size of cache line in bytes. */ - int cache_line_bytes; - - // GPU related parameters got from device query API - /*! \brief The max shared memory per block in bytes. */ - int max_shared_memory_per_block; - /*! \brief The max local memory per block in bytes. */ - int max_local_memory_per_block; - /*! \brief The max number of threads per block. */ - int max_threads_per_block; - /*! \brief The max vthread extent. */ - int max_vthread_extent; - /*! \brief The thread numbers of a warp. */ - int warp_size; - - void VisitAttrs(tvm::AttrVisitor* v) { - v->Visit("num_cores", &num_cores); - v->Visit("vector_unit_bytes", &vector_unit_bytes); - v->Visit("cache_line_bytes", &cache_line_bytes); - v->Visit("max_shared_memory_per_block", &max_shared_memory_per_block); - v->Visit("max_local_memory_per_block", &max_local_memory_per_block); - v->Visit("max_threads_per_block", &max_threads_per_block); - v->Visit("max_vthread_extent", &max_vthread_extent); - v->Visit("warp_size", &warp_size); - } - - /*! - * \brief Get the default hardware params. - * \param target A `tvm.target`. - * \param target_host A `tvm.target` for host device. - * \return A HardwareParams object. - */ - static HardwareParams GetDefaultHardwareParams(const Target& target, const Target& target_host); - - static constexpr const char* _type_key = "auto_scheduler.HardwareParams"; - TVM_DECLARE_FINAL_OBJECT_INFO(HardwareParamsNode, Object); -}; - -/*! - * \brief Managed reference to HardwareParamsNode. - * \sa HardwareParamsNode - */ -class HardwareParams : public ObjectRef { - public: - /*! - * \brief The constructor. - * \param num_cores The number of cores. - * \param vector_unit_bytes The width of vector units in bytes. - * \param cache_line_bytes The size of cache line in bytes. - * \param max_shared_memory_per_block The max amount of shared memory per block for GPU. - * \param max_local_memory_per_block The max amount of local memory per block for GPU. - * \param max_threads_per_block The max number of threads per block for GPU. - * \param max_vthread_extent The max extent of vthread for GPU. - * \param warp_size The warp size for GPU - */ - HardwareParams(int num_cores, int vector_unit_bytes, int cache_line_bytes, - int max_shared_memory_per_block, int max_local_memory_per_block, - int max_threads_per_block, int max_vthread_extent, int warp_size); - - TVM_DEFINE_OBJECT_REF_METHODS(HardwareParams, ObjectRef, HardwareParamsNode); - TVM_DEFINE_OBJECT_REF_COW_METHOD(HardwareParamsNode); -}; - -/*! - * \brief The computation information and hardware parameters for a specific schedule search task. - */ -class SearchTaskNode : public Object { - public: - /*! \brief The ComputeDAG for the compute declaration. */ - ComputeDAG compute_dag; - /*! \brief The workload key for the compute declaration. */ - String workload_key; - /*! \brief The description string of this task. */ - String desc; - /*! \brief The target device of this search task. */ - Target target; - /*! \brief The target host device of this search task. */ - Target target_host; - /*! \brief Hardware parameters used in this search task. */ - HardwareParams hardware_params; - /*! \brief The layout rewrite option used for measuring programs. */ - LayoutRewriteOption layout_rewrite_option; - /*! \brief Names of some user defined input data used in program measuring. */ - Array task_input_names; - - void VisitAttrs(tvm::AttrVisitor* v) { - v->Visit("compute_dag", &compute_dag); - v->Visit("workload_key", &workload_key); - v->Visit("desc", &desc); - v->Visit("target", &target); - v->Visit("target_host", &target_host); - v->Visit("hardware_params", &hardware_params); - v->Visit("layout_rewrite_option", &layout_rewrite_option); - v->Visit("task_input_names", &task_input_names); - } - - static constexpr const char* _type_key = "auto_scheduler.SearchTask"; - TVM_DECLARE_FINAL_OBJECT_INFO(SearchTaskNode, Object); -}; - -/*! - * \brief Managed reference to SearchTaskNode. - * \sa SearchTaskNode - */ -class SearchTask : public ObjectRef { - public: - /*! - * \brief The constructor. - * \param compute_dag The ComputeDAG for the compute declaration. - * \param workload_key The workload key for the compute declaration. - * \param target The target device of this search task. - * \param target_host The target host device of this search task. - * \param hardware_params Hardware parameters used in this search task. - * \param layout_rewrite_option The layout rewrite option used for measuring programs. - * \param task_input_names Names of some user defined input data used in program measuring. - * \param desc The description string of this task. - */ - SearchTask(ComputeDAG compute_dag, String workload_key, Target target, Target target_host, - Optional hardware_params, LayoutRewriteOption layout_rewrite_option, - Array task_input_names, String desc = ""); - - TVM_DEFINE_OBJECT_REF_METHODS(SearchTask, ObjectRef, SearchTaskNode); -}; - -} // namespace auto_scheduler -} // namespace tvm - -#endif // TVM_AUTO_SCHEDULER_SEARCH_TASK_H_ diff --git a/include/tvm/auto_scheduler/transform_step.h b/include/tvm/auto_scheduler/transform_step.h deleted file mode 100755 index fa770a3ed36e..000000000000 --- a/include/tvm/auto_scheduler/transform_step.h +++ /dev/null @@ -1,1195 +0,0 @@ -/* - * 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 auto_scheduler/transform_step.h - * \brief Transformation steps. These steps are used to manipulate `LoopState`. - * They are similar to the schedule primitives in te::Stage. - * - * \note How to add a new transform step: - * Take fuse step for example: - * 1. Define class `FuseStepNode`, `FuseStep` in `transform_steps.h`, and implement its first - * construction function `FuseStep::FuseStep()` in `transform_steps.cc`. - * 2. Implement `FuseStepNode::ApplyToSchedule()` and `FuseStepNode::PrintAsPythonAPI()`. - * - In these two functions you need to lower this step with tvm's te schedule API - * 3. Implement `FuseStepNode::ApplyToState` and the state API `State::fuse`. - * - In these two functions you need to incrementally update all data structures in State with - * CopyOnWrite style. - * 4. Add your step to `StepApplyToState`, `StepApplyToSchedule`, and `StepPrintAsPythonAPI`. - * 5. Log record serialization support: - * - Add `FuseStepNode::WriteToRecord` which takes a mutable JSONWriter pointer as input and - * output the record to it. - * - Add another construction function that takes a mutable JSONReader as input, this will get a - * step record from the reader and create the step. - * - Add the step implementation to `StepReadFromRecord`. - * 6. Add its corresponding Python API to `loop_state.py` with necessary unit tests. The test should - * at lease cover two parts: the functional test and the record serialization test. - */ - -#ifndef TVM_AUTO_SCHEDULER_TRANSFORM_STEP_H_ -#define TVM_AUTO_SCHEDULER_TRANSFORM_STEP_H_ - -#include -#include -#include -#include - -#include - -namespace tvm { -namespace auto_scheduler { - -typedef Map, ObjectHash, ObjectEqual> StageToAxesMap; - -/*! - * \brief Update the current stage IterVar information to StageToAxesMap. - * \param stage The stage to be updated. - * \param stage_to_axes The map to be updated. - */ -void UpdateStageToAxesMap(const te::Stage& stage, StageToAxesMap* stage_to_axes); - -/*! \brief The type of an iterator. */ -enum class IteratorKind : int { - /*! \brief Spatial iterator. */ - kSpatial = 0, - /*! \brief Reduction iterator. */ - kReduction = 1, - /*! \brief Fused spatial and reduction iterator. */ - kMixed = 2, - /*! \brief Special iterator. (e.g. virtual root iterator) */ - kSpecial = 3 -}; - -/*! \brief The type of an iterator's annotation. */ -enum class IteratorAnnotation : int { - /*! \brief This iterator has no annotation. */ - kNone = 0, - /*! \brief This iterator has been unrolled. */ - kUnroll = 1, - /*! \brief This iterator has been vectorized. */ - kVectorize = 2, - /*! \brief This iterator has been paralleld. */ - kParallel = 3, - /*! \brief This iterator has been bind to vthread. */ - kVThread = 4, - /*! \brief This iterator has been bind to blockIdx.x. */ - kBlockX = 5, - /*! \brief This iterator has been bind to threadIdx.x. */ - kThreadX = 6, - /*! \brief This iterator has been bind to blockIdx.y. */ - kBlockY = 7, - /*! \brief This iterator has been bind to threadIdx.y. */ - kThreadY = 8, - /*! \brief This iterator has been bind to blockIdx.y. */ - kBlockZ = 9, - /*! \brief This iterator has been bind to threadIdx.y. */ - kThreadZ = 10, - /*! \brief This iterator has been mapped with a tensorize intrinsic. */ - kTensorize = 11 -}; - -extern const char* IteratorAnnotationString[]; - -// forward declaration -class Iterator; - -/*! - * \brief An iterator of a for-loop - * Similar to tvm::IterVar in `include/tvm/tir/expr.h` - */ -class IteratorNode : public Object { - public: - /*! \brief The name of this iterator. */ - String name; - /*! \brief The range of this iterator. */ - Range range; - /*! \brief The iterator type of this iterator. */ - IteratorKind iter_kind; - /*! \brief The annotation type of this iterator. */ - IteratorAnnotation annotation; - /*! The original iterators before fusion. */ - std::vector orig_iters; - - void VisitAttrs(tvm::AttrVisitor* v) { - v->Visit("name", &name); - v->Visit("range", &range); - v->Visit("iter_kind", &iter_kind); - v->Visit("annotation", &annotation); - } - - static constexpr const char* _type_key = "auto_scheduler.Iterator"; - TVM_DECLARE_FINAL_OBJECT_INFO(IteratorNode, Object); -}; - -/*! - * \brief Managed reference to IteratorNode. - * \sa IteratorNode - */ -class Iterator : public ObjectRef { - public: - /*! - * \brief The constructor. - * \param name The name of this iterator. - * \param range The range of this iterator. - * \param iter_kind The iterator type of this iterator. - * \param annotation The annotation type of this iterator. - * \param orig_iters The original iterators before fusion - */ - Iterator(String name, Range range, IteratorKind iter_kind, IteratorAnnotation annotation, - const std::vector* orig_iters = nullptr); - - TVM_DEFINE_OBJECT_REF_METHODS(Iterator, ObjectRef, IteratorNode); -}; - -/*! - * \brief The base class of transformation steps. Each step has its corresponding tvm.te - * schedule primitives. - */ -class StepNode : public Object { - public: - /*! \brief The index of the stage. */ - int stage_id; - - /*! - * \brief Serialize the current step record to JSONWriter. - * \param writer The output JSONWriter. - */ - virtual void WriteToRecord(dmlc::JSONWriter* writer) const = 0; - - static constexpr const char* _type_key = "auto_scheduler.Step"; - TVM_DECLARE_BASE_OBJECT_INFO(StepNode, Object); -}; - -/*! - * \brief Managed reference to StepNode. - * \sa StepNode - */ -class Step : public ObjectRef { - public: - /*! - * \brief CopyOnWrite function for Step. - * This works almost the same as a normal ObjectRef.CopyOnWrite(), but can dispatch to different - * steps. - * \return A base StepNode pointer, need to cast to its real StepNode type before doing any - * modifications. - * \code - * - * SplitStep ref; - * StepNode* mutable_ref = ref.CopyOnWrite(); - * dynamic_cast(mutable_ref)->... = ...; - * - * \endcode - */ - StepNode* CopyOnWrite(); - - TVM_DEFINE_OBJECT_REF_METHODS(Step, ObjectRef, StepNode); -}; - -// Forward declaration -class State; -class ComputeDAG; - -/*! - * \brief Read a step record from JSONReader and create the corresponding step. - * \param reader The input JSONReader. - */ -Step StepReadFromRecord(dmlc::JSONReader* reader); - -/*! - * \brief Apply a general step to a State with runtime dynamic dispatching. - * \param step The step to be applied to State. - * \param state A mutable pointer to state, which will be updated. - * \param dag The original ComputeDAG of this state. - */ -void StepApplyToState(const Step& step, State* state, const ComputeDAG& dag); - -/*! - * \brief Apply a general step to tvm.schedule with runtime dynamic dispatching. - * \param step The step to be applied to tvm.schedule. - * \param stages The list of current stages - * \param stage_to_axes A map that maps stage ot all its iterators. - * \param schedule A mutable point to the current schedule - * \param transform_steps An array of all history transform steps. - */ -void StepApplyToSchedule(const Step& step, Array* stages, StageToAxesMap* stage_to_axes, - te::Schedule* schedule, const Array& transform_steps); - -/*! - * \brief Print a general step as equivalent python schedule API with runtime dynamic dispatching. - * \param step The step to be printed as python API. - * \param stages The list of current stages - * \param stage_to_axes A map that maps stage ot all its iterators. - * \param schedule A mutable point to the current schedule - * \param transform_steps An array of all history transform steps. - * \return Python schedule code. - */ -String StepPrintAsPythonAPI(const Step& step, Array* stages, - StageToAxesMap* stage_to_axes, te::Schedule* schedule, - const Array& transform_steps); - -/********** Steps working on single stage **********/ - -/*! - * \brief Annotation step that corresponds to vectorize, parallel, unroll and thread binding. - * (i.e. te::Stage::vectorize, te::Stage::parallel, te::Stage::vectorize, te::Stage::bind) - */ -class AnnotationStepNode : public StepNode { - public: - /*! \brief The index of the iterator to add annotation. */ - int iter_id; - /*! \brief The annotation type of this step. */ - IteratorAnnotation annotation; - - void WriteToRecord(dmlc::JSONWriter* writer) const final; - - /*! - * \brief Apply the current step to State. - * \param state A mutable pointer to state, which will be updated. - * \return The iterator result after annotate. - */ - Iterator ApplyToState(State* state) const; - - /*! - * \brief Apply the current step to tvm.schedule. - * \param stages The list of current stages - * \param stage_to_axes A map that maps stage ot all its iterators. - */ - void ApplyToSchedule(Array* stages, StageToAxesMap* stage_to_axes) const; - - /*! - * \brief Print the current step as equivalent python schedule API. - * \param stages The list of current stages - * \param stage_to_axes A map that maps stage ot all its iterators. - * \return Python schedule code. - */ - String PrintAsPythonAPI(Array* stages, StageToAxesMap* stage_to_axes) const; - - static constexpr const char* record_prefix_str = "AN"; - - static constexpr const char* _type_key = "auto_scheduler.AnnotationStep"; - TVM_DECLARE_FINAL_OBJECT_INFO(AnnotationStepNode, StepNode); -}; - -/*! - * \brief Managed reference to AnnotationStepNode. - * \sa AnnotationStepNode - */ -class AnnotationStep : public Step { - public: - /*! - * \brief The constructor. - * \param stage_id The index of the stage to add annotation. - * \param iter_id The index of the iterator to add annotation. - * \param ann The annotation type of this step. - */ - AnnotationStep(int stage_id, int iter_id, IteratorAnnotation ann); - - /*! - * \brief The constructor used to read a step record from JSONReader and create the - * corresponding step. - * \param reader The input JSONReader. - */ - explicit AnnotationStep(dmlc::JSONReader* reader); - - TVM_DEFINE_OBJECT_REF_METHODS(AnnotationStep, Step, AnnotationStepNode); -}; - -/*! \brief Fuse step that corresponds to te::Stage::fuse */ -class FuseStepNode : public StepNode { - public: - /*! \brief The ids of iterators to fuse. */ - Array fused_ids; - - void WriteToRecord(dmlc::JSONWriter* writer) const final; - - /*! - * \brief Apply the current step to State. - * \param state A mutable pointer to state, which will be updated. - * \return The iterator result after fuse. - * \note If the iterators to be fused have stages attached at them(by compute_at), the fused - * result will become the new attach point. - */ - Iterator ApplyToState(State* state) const; - - /*! - * \brief Apply the current step to tvm.schedule. - * \param stages The list of current stages - * \param stage_to_axes A map that maps stage ot all its iterators. - * \return The iterator result after fuse. - */ - tir::IterVar ApplyToSchedule(Array* stages, StageToAxesMap* stage_to_axes) const; - - /*! - * \brief Print the current step as equivalent python schedule API. - * \param stages The list of current stages - * \param stage_to_axes A map that maps stage ot all its iterators. - * \return Python schedule code. - */ - String PrintAsPythonAPI(Array* stages, StageToAxesMap* stage_to_axes) const; - - static constexpr const char* record_prefix_str = "FU"; - - static constexpr const char* _type_key = "auto_scheduler.FuseStep"; - TVM_DECLARE_FINAL_OBJECT_INFO(FuseStepNode, StepNode); -}; - -/*! - * \brief Managed reference to FuseStepNode. - * \sa FuseStepNode - */ -class FuseStep : public Step { - public: - /*! - * \brief The constructor. - * \param stage_id The index of the stage to be fused. - * \param fused_ids The index of the iterators to be fused. - */ - FuseStep(int stage_id, const Array& fused_ids); - - /*! - * \brief The constructor used to read a step record from JSONReader and create the - * corresponding step. - * \param reader The input JSONReader. - */ - explicit FuseStep(dmlc::JSONReader* reader); - - TVM_DEFINE_OBJECT_REF_METHODS(FuseStep, Step, FuseStepNode); -}; - -/*! \brief Pragma step that corresponds to te::Stage::pragma */ -class PragmaStepNode : public StepNode { - public: - /*! \brief The index of the iterator to add pragma. */ - int iter_id; - /*! \brief The pragma string. */ - String pragma_type; - - void WriteToRecord(dmlc::JSONWriter* writer) const final; - - /*! - * \brief Apply the current step to State. - * \param state A mutable pointer to state, which will be updated. - */ - void ApplyToState(State* state) const; - - /*! - * \brief Apply the current step to tvm.schedule. - * \param stages The list of current stages - * \param stage_to_axes A map that maps stage ot all its iterators. - */ - void ApplyToSchedule(Array* stages, StageToAxesMap* stage_to_axes) const; - - /*! - * \brief Print the current step as equivalent python schedule API. - * \param stages The list of current stages - * \param stage_to_axes A map that maps stage ot all its iterators. - * \return Python schedule code. - */ - String PrintAsPythonAPI(Array* stages, StageToAxesMap* stage_to_axes) const; - - static constexpr const char* record_prefix_str = "PR"; - - static constexpr const char* _type_key = "auto_scheduler.PragmaStep"; - TVM_DECLARE_FINAL_OBJECT_INFO(PragmaStepNode, StepNode); -}; - -/*! - * \brief Managed reference to PragmaStepNode. - * \sa PragmaStepNode - */ -class PragmaStep : public Step { - public: - /*! - * \brief The constructor. - * \param stage_id The index of the stage to be fused. - * \param iter_id The index of the iterator to add pragma. - * \param pragma_type The pragma string. - */ - PragmaStep(int stage_id, int iter_id, String pragma_type); - - /*! - * \brief The constructor used to read a step record from JSONReader and create the - * corresponding step. - * \param reader The input JSONReader. - */ - explicit PragmaStep(dmlc::JSONReader* reader); - - TVM_DEFINE_OBJECT_REF_METHODS(PragmaStep, Step, PragmaStepNode); -}; - -/*! \brief Reorder step that corresponds to te::Stage::reorder */ -class ReorderStepNode : public StepNode { - public: - /*! - * \brief The iterator ids after reorder. - * This array should specify the order of all iterators. - */ - Array after_ids; - - void WriteToRecord(dmlc::JSONWriter* writer) const final; - - /*! - * \brief Apply the current step to State. - * \param state A mutable pointer to state, which will be updated. - */ - void ApplyToState(State* state) const; - - /*! - * \brief Apply the current step to tvm.schedule. - * \param stages The list of current stages - * \param stage_to_axes A map that maps stage ot all its iterators. - */ - void ApplyToSchedule(Array* stages, StageToAxesMap* stage_to_axes) const; - - /*! - * \brief Print the current step as equivalent python schedule API. - * \param stages The list of current stages - * \param stage_to_axes A map that maps stage ot all its iterators. - * \return Python schedule code. - */ - String PrintAsPythonAPI(Array* stages, StageToAxesMap* stage_to_axes) const; - - static constexpr const char* record_prefix_str = "RE"; - - static constexpr const char* _type_key = "auto_scheduler.ReorderStep"; - TVM_DECLARE_FINAL_OBJECT_INFO(ReorderStepNode, StepNode); -}; - -/*! - * \brief Managed reference to ReorderStepNode. - * \sa ReorderStepNode - */ -class ReorderStep : public Step { - public: - /*! - * \brief The constructor. - * \param stage_id The index of the stage to be reordered. - * \param after_ids The expected indexes of the iterators after reorder. - */ - ReorderStep(int stage_id, const Array& after_ids); - - /*! - * \brief The constructor used to read a step record from JSONReader and create the - * corresponding step. - * \param reader The input JSONReader. - */ - explicit ReorderStep(dmlc::JSONReader* reader); - - TVM_DEFINE_OBJECT_REF_METHODS(ReorderStep, Step, ReorderStepNode); -}; - -/*! - * \brief Split step that corresponds to te::Stage::split with additional - * support of multiple-level of factors - */ -class SplitStepNode : public StepNode { - public: - /*! \brief The id of the iter to split. */ - int iter_id; - /*! \brief The extent length of the axis to split. */ - Optional extent; - /*! \brief The split factors. */ - Array> lengths; - /*! - * \brief If true, the `lengths` denote the lengths of iterators - * from inner level to outer level - */ - bool inner_to_outer; - - void WriteToRecord(dmlc::JSONWriter* writer) const final; - - /*! - * \brief Apply the current step to State. - * \param state A mutable pointer to state, which will be updated. - * \return The iterator results after split. - * \note If we do split on an iterator which has stages attached at it(by compute_at), the inner - * most iterator of split results will become the new attach point. - */ - Array ApplyToState(State* state) const; - - /*! - * \brief Apply the current step to tvm.schedule. - * \param stages The list of current stages - * \param stage_to_axes A map that maps stage ot all its iterators. - * \return The iterator results after split. - */ - Array ApplyToSchedule(Array* stages, - StageToAxesMap* stage_to_axes) const; - - /*! - * \brief Print the current step as equivalent python schedule API. - * \param stages The list of current stages - * \param stage_to_axes A map that maps stage ot all its iterators. - * \return Python schedule code. - */ - String PrintAsPythonAPI(Array* stages, StageToAxesMap* stage_to_axes) const; - - static constexpr const char* record_prefix_str = "SP"; - - static constexpr const char* _type_key = "auto_scheduler.SplitStep"; - TVM_DECLARE_FINAL_OBJECT_INFO(SplitStepNode, StepNode); -}; - -/*! - * \brief Managed reference to SplitStepNode. - * \sa SplitStepNode - */ -class SplitStep : public Step { - public: - /*! - * \brief The constructor. - * \param stage_id The index of the stage to be split. - * \param iter_id The index of the iterator to be split. - * \param extent The extent length of the axis to split. - * \param lengths The multiple split factors. Can be None to be filled by search policy. - * \param inner_to_outer The split direction. - */ - SplitStep(int stage_id, int iter_id, Optional extent, - const Array>& lengths, bool inner_to_outer); - - /*! - * \brief The constructor used to read a step record from JSONReader and create the - * corresponding step. - * \param reader The input JSONReader. - */ - explicit SplitStep(dmlc::JSONReader* reader); - - TVM_DEFINE_OBJECT_REF_METHODS(SplitStep, Step, SplitStepNode); -}; - -/*! \brief Similar to SplitStepNode, but uses split factors from another step - * (i.e. Follow another split step) */ -class FollowSplitStepNode : public StepNode { - public: - /*! \brief The id of the iter to be split. */ - int iter_id; - /*! \brief The index of the split step to be followed in the history. */ - int src_step_id; - /*! \brief The number of split level. */ - int n_split; - - void WriteToRecord(dmlc::JSONWriter* writer) const final; - - /*! - * \brief Extract split lengths. - * \param transform_steps An array of history transform steps. - * \return The multiple split factors. - */ - Array> ExtractSplitLengths(const Array& transform_steps) const; - - /*! - * \brief Apply the current step to State. - * \param state A mutable pointer to state, which will be updated. - * \return The iterator results after split. - */ - Array ApplyToState(State* state) const; - - /*! - * \brief Apply the current step to tvm.schedule. - * \param stages The list of current stages - * \param stage_to_axes A map that maps stage ot all its iterators. - * \param transform_steps An array of history transform steps. - * \return The iterator results after split. - */ - Array ApplyToSchedule(Array* stages, StageToAxesMap* stage_to_axes, - const Array& transform_steps) const; - - /*! - * \brief Print the current step as equivalent python schedule API. - * \param stages The list of current stages - * \param stage_to_axes A map that maps stage ot all its iterators. - * \param transform_steps An array of history transform steps. - * \return Python schedule code. - */ - String PrintAsPythonAPI(Array* stages, StageToAxesMap* stage_to_axes, - const Array& transform_steps) const; - - static constexpr const char* record_prefix_str = "FSP"; - - static constexpr const char* _type_key = "auto_scheduler.FollowSplitStep"; - TVM_DECLARE_FINAL_OBJECT_INFO(FollowSplitStepNode, StepNode); -}; - -/*! - * \brief Managed reference to FollowSplitStepNode. - * \sa FollowSplitStepNode - */ -class FollowSplitStep : public Step { - public: - /*! - * \brief The constructor. - * \param stage_id The index of the stage to be split. - * \param iter_id The index of the iterator to be split. - * \param src_step_id The index of the split step to be followed in the history. - * \param n_split The number of split level. - */ - FollowSplitStep(int stage_id, int iter_id, int src_step_id, int n_split); - - /*! - * \brief The constructor used to read a step record from JSONReader and create the - * corresponding step. - * \param reader The input JSONReader. - */ - explicit FollowSplitStep(dmlc::JSONReader* reader); - - TVM_DEFINE_OBJECT_REF_METHODS(FollowSplitStep, Step, FollowSplitStepNode); -}; - -/*! \brief Similar to FollowSplitStep, but uses split factors from multiple steps. - * \note This can be used for the split in cooperative fetching. - */ -class FollowFusedSplitStepNode : public StepNode { - public: - /*! \brief The id of the iter to split. */ - int iter_id; - /*! \brief The indices of the split steps to be followed in the history. */ - Array src_step_ids; - /*! \brief Use the length in this split level. */ - int level; - /*! \brief If this is true, use factor. Otherwise, use nparts. */ - bool factor_or_nparts; - - void WriteToRecord(dmlc::JSONWriter* writer) const final; - - /*! - * \brief Extract split length. - * \param transform_steps An array of history transform steps. - * \return Split factor. - */ - Optional ExtractSplitLength(const Array& transform_steps) const; - - /*! - * \brief Apply the current step to State. - * \param state A mutable pointer to state, which will be updated. - * \return The iterator results after split. - */ - Array ApplyToState(State* state) const; - - /*! - * \brief Apply the current step to tvm.schedule. - * \param stages The list of current stages - * \param stage_to_axes A map that maps stage ot all its iterators. - * \param transform_steps An array of history transform steps. - * \return The iterator results after split. - */ - Array ApplyToSchedule(Array* stages, StageToAxesMap* stage_to_axes, - const Array& transform_steps) const; - - /*! - * \brief Print the current step as equivalent python schedule API. - * \param stages The list of current stages - * \param stage_to_axes A map that maps stage ot all its iterators. - * \param transform_steps An array of history transform steps. - * \return Python schedule code. - */ - String PrintAsPythonAPI(Array* stages, StageToAxesMap* stage_to_axes, - const Array& transform_steps) const; - - static constexpr const char* record_prefix_str = "FFSP"; - - static constexpr const char* _type_key = "auto_scheduler.FollowFusedSplitStep"; - TVM_DECLARE_FINAL_OBJECT_INFO(FollowFusedSplitStepNode, StepNode); -}; - -/*! - * \brief Managed reference to FollowFusedSplitStepNode. - * \sa FollowFusedSplitStepNode - */ -class FollowFusedSplitStep : public Step { - public: - /*! - * \brief The constructor. - * \param stage_id The index of the stage to be split. - * \param iter_id The index of the iterator to be split. - * \param src_step_ids An array of index for split step to be followed in the history. - * \param level Use the length in this split level. - * \param factor_or_nparts If this is true, use factor. Otherwise, use nparts. - */ - FollowFusedSplitStep(int stage_id, int iter_id, const Array& src_step_ids, int level, - bool factor_or_nparts); - - /*! - * \brief The constructor used to read a step record from JSONReader and create the - * corresponding step. - * \param reader The input JSONReader. - */ - explicit FollowFusedSplitStep(dmlc::JSONReader* reader); - - TVM_DEFINE_OBJECT_REF_METHODS(FollowFusedSplitStep, Step, FollowFusedSplitStepNode); -}; - -/*! \brief Storage align step that corresponds to te::Stage::storage_align */ -class StorageAlignStepNode : public StepNode { - public: - /*! \brief The iterator to be aligned. */ - int iter_id; - /*! \brief The factor in alignment specification. */ - int factor; - /*! \brief The offset in the alignment specification. */ - int offset; - - void WriteToRecord(dmlc::JSONWriter* writer) const final; - - /*! - * \brief Apply the current step to State. - * \param state A mutable pointer to State, which will be updated. - */ - void ApplyToState(State* state) const; - - /*! - * \brief Apply the current step to tvm.schedule. - * \param stages The list of current stages - * \param stage_to_axes A map that maps stage ot all its iterators. - */ - void ApplyToSchedule(Array* stages, StageToAxesMap* stage_to_axes) const; - - /*! - * \brief Print the current step as equivalent python schedule API. - * \param stages The list of current stages - * \param stage_to_axes A map that maps stage ot all its iterators. - * \return Python schedule code. - */ - String PrintAsPythonAPI(Array* stages, StageToAxesMap* stage_to_axes) const; - - static constexpr const char* record_prefix_str = "SA"; - - static constexpr const char* _type_key = "auto_scheduler.StorageAlignStep"; - TVM_DECLARE_FINAL_OBJECT_INFO(StorageAlignStepNode, StepNode); -}; - -/*! - * \brief Managed reference to StorageAlignStepNode. - * \sa StorageAlignStepNode - */ -class StorageAlignStep : public Step { - public: - /*! - * \brief The constructor. - * \param stage_id The index of the stage to be aligned. - * \param iter_id The index of the iterator to be aligned. - * \param factor The factor in alignment specification. - * \param offset The offset in the alignment specification. - */ - StorageAlignStep(int stage_id, int iter_id, int factor, int offset); - - /*! - * \brief The constructor used to read a step record from JSONReader and create the - * corresponding step. - * \param reader The input JSONReader. - */ - explicit StorageAlignStep(dmlc::JSONReader* reader); - - TVM_DEFINE_OBJECT_REF_METHODS(StorageAlignStep, Step, StorageAlignStepNode); -}; - -/********** Steps working on multiple stages **********/ - -/*! \brief Compute at step that corresponds to te::Stage::compute_at */ -class ComputeAtStepNode : public StepNode { - public: - /*! \brief The index of stage that this step will compute at to. */ - int target_stage_id; - /*! \brief The index of iterator in target stage that this step will compute at to. */ - int target_iter_id; - - void WriteToRecord(dmlc::JSONWriter* writer) const final; - - /*! - * \brief Apply the current step to State. - * \param state A mutable pointer to state, which will be updated. - * \note After compute_at, we need careful dependency analysis to compute the accurate bound - * information. However, it is relatively expensive and complicated, so we just fill "None" as - * bound for the newly created iterators. - * Call ComputeDAG::InferBound on the updated state if you need the complete bound information. - */ - void ApplyToState(State* state) const; - - /*! - * \brief Apply the current step to tvm.schedule. - * \param stages The list of current stages - * \param stage_to_axes A map that maps stage ot all its iterators. - */ - void ApplyToSchedule(Array* stages, StageToAxesMap* stage_to_axes) const; - - /*! - * \brief Print the current step as equivalent python schedule API. - * \param stages The list of current stages - * \param stage_to_axes A map that maps stage ot all its iterators. - * \return Python schedule code. - */ - String PrintAsPythonAPI(Array* stages, StageToAxesMap* stage_to_axes) const; - - static constexpr const char* record_prefix_str = "CA"; - - static constexpr const char* _type_key = "auto_scheduler.ComputeAtStep"; - TVM_DECLARE_FINAL_OBJECT_INFO(ComputeAtStepNode, StepNode); -}; - -/*! - * \brief Managed reference to ComputeAtStepNode. - * \sa ComputeAtStepNode - */ -class ComputeAtStep : public Step { - public: - /*! - * \brief The constructor. - * \param stage_id The index of the source stage. - * \param target_stage_id The index of stage that this step will compute at to. - * \param target_iter_id The index of iterator in target stage that this step will compute at to. - */ - ComputeAtStep(int stage_id, int target_stage_id, int target_iter_id); - - /*! - * \brief The constructor used to read a step record from JSONReader and create the - * corresponding step. - * \param reader The input JSONReader. - */ - explicit ComputeAtStep(dmlc::JSONReader* reader); - - TVM_DEFINE_OBJECT_REF_METHODS(ComputeAtStep, Step, ComputeAtStepNode); -}; - -/*! \brief Compute inline step that corresponds to te::Stage::compute_inline */ -class ComputeInlineStepNode : public StepNode { - public: - void WriteToRecord(dmlc::JSONWriter* writer) const final; - - /*! - * \brief Apply the current step to State. - * \param state A mutable pointer to state, which will be updated. - */ - void ApplyToState(State* state) const; - - /*! - * \brief Apply the current step to tvm.schedule. - * \param stages The list of current stages - * \param stage_to_axes A map that maps stage ot all its iterators. - */ - void ApplyToSchedule(Array* stages, StageToAxesMap* stage_to_axes) const; - - /*! - * \brief Print the current step as equivalent python schedule API. - * \param stages The list of current stages - * \param stage_to_axes A map that maps stage ot all its iterators. - * \return Python schedule code. - */ - String PrintAsPythonAPI(Array* stages, StageToAxesMap* stage_to_axes) const; - - static constexpr const char* record_prefix_str = "CI"; - - static constexpr const char* _type_key = "auto_scheduler.ComputeInlineStep"; - TVM_DECLARE_FINAL_OBJECT_INFO(ComputeInlineStepNode, StepNode); -}; - -/*! - * \brief Managed reference to ComputeInlineStepNode. - * \sa ComputeInlineStepNode - */ -class ComputeInlineStep : public Step { - public: - /*! - * \brief The constructor. - * \param stage_id The index of the stage to be marked compute inlined. - */ - explicit ComputeInlineStep(int stage_id); - - /*! - * \brief The constructor used to read a step record from JSONReader and create the - * corresponding step. - * \param reader The input JSONReader. - */ - explicit ComputeInlineStep(dmlc::JSONReader* reader); - - TVM_DEFINE_OBJECT_REF_METHODS(ComputeInlineStep, Step, ComputeInlineStepNode); -}; - -/*! \brief Compute root step that corresponds to te::Stage::compute_root */ -class ComputeRootStepNode : public StepNode { - public: - void WriteToRecord(dmlc::JSONWriter* writer) const final; - - /*! - * \brief Apply the current step to State. - * \param state A mutable pointer to state, which will be updated. - * \note After compute_root, we need careful dependency analysis to compute the accurate bound - * information. However, it is relatively expensive and complicated, so we just fill "None" as - * bound for the newly created iterators. - * Call ComputeDAG::InferBound on the updated state if you need the complete bound information. - */ - void ApplyToState(State* state) const; - - /*! - * \brief Apply the current step to tvm.schedule. - * \param stages The list of current stages - * \param stage_to_axes A map that maps stage ot all its iterators. - */ - void ApplyToSchedule(Array* stages, StageToAxesMap* stage_to_axes) const; - - /*! - * \brief Print the current step as equivalent python schedule API. - * \param stages The list of current stages - * \param stage_to_axes A map that maps stage ot all its iterators. - * \return Python schedule code. - */ - String PrintAsPythonAPI(Array* stages, StageToAxesMap* stage_to_axes) const; - - static constexpr const char* record_prefix_str = "CR"; - - static constexpr const char* _type_key = "auto_scheduler.ComputeRootStep"; - TVM_DECLARE_FINAL_OBJECT_INFO(ComputeRootStepNode, StepNode); -}; - -/*! - * \brief Managed reference to ComputeRootStepNode. - * \sa ComputeRootStepNode - */ -class ComputeRootStep : public Step { - public: - /*! - * \brief The constructor. - * \param stage_id The index of the stage to be marked compute at root. - */ - explicit ComputeRootStep(int stage_id); - - /*! - * \brief The constructor used to read a step record from JSONReader and create the - * corresponding step. - * \param reader The input JSONReader. - */ - explicit ComputeRootStep(dmlc::JSONReader* reader); - - TVM_DEFINE_OBJECT_REF_METHODS(ComputeRootStep, Step, ComputeRootStepNode); -}; - -/********** Steps adding new stages **********/ - -/*! - * \brief Cache read step that corresponds to te::Schedule::cache_read. - * \note Cache read step adds an extra stage to the original ComputeDAG, - * an up-to-date ComputeDAG will be stored in State's `current_compute_dag`. - */ -class CacheReadStepNode : public StepNode { - public: - /*! \brief The scope name of the newly added read stage. (e.g., local, shared, global) */ - String scope_name; - /*! \brief The indices of read stages. */ - Array reader_stage_ids; - - void WriteToRecord(dmlc::JSONWriter* writer) const final; - - /*! - * \brief Apply the current step to State. - * \param state A mutable pointer to state, which will be updated. - * \param dag The original ComputeDAG of this state. - * \return The index of the new added stage. - */ - int ApplyToState(State* state, const ComputeDAG& dag) const; - - /*! - * \brief Apply the current step to tvm.schedule. - * \param stages The list of current stages - * \param stage_to_axes A map that maps stage ot all its iterators. - * \param schedule A mutable pointer to a te::Schedule. - * \return The output Tensor of the new added stage. - */ - te::Tensor ApplyToSchedule(Array* stages, StageToAxesMap* stage_to_axes, - te::Schedule* schedule) const; - - /*! - * \brief Print the current step as equivalent python schedule API. - * \param stages The list of current stages - * \param stage_to_axes A map that maps stage ot all its iterators. - * \param schedule A mutable pointer to a te::Schedule. - * \return Python schedule code. - */ - String PrintAsPythonAPI(Array* stages, StageToAxesMap* stage_to_axes, - te::Schedule* schedule) const; - - static constexpr const char* record_prefix_str = "CHR"; - - static constexpr const char* _type_key = "auto_scheduler.CacheReadStep"; - TVM_DECLARE_FINAL_OBJECT_INFO(CacheReadStepNode, StepNode); -}; - -/*! - * \brief Managed reference to CacheReadStepNode. - * \sa CacheReadStepNode - */ -class CacheReadStep : public Step { - public: - /*! - * \brief The constructor. - * \param stage_id The index of the stage to be cache_read. - * \param scope_name The scope name of the newly added stage. - * \param reader_stage_ids The indices of reader stages. - */ - CacheReadStep(int stage_id, String scope_name, const Array& reader_stage_ids); - - /*! - * \brief The constructor used to read a step record from JSONReader and create the - * corresponding step. - * \param reader The input JSONReader. - */ - explicit CacheReadStep(dmlc::JSONReader* reader); - - TVM_DEFINE_OBJECT_REF_METHODS(CacheReadStep, Step, CacheReadStepNode); -}; - -/*! - * \brief Cache write step that corresponds to te::Schedule::cache_write. - * \note Cache write step will add an extra stage to the original ComputeDAG, a up-to-date - * ComputeDAG is stored in State's `current_compute_dag`. - * This step will cache write all output tensors of the target stage. - */ -class CacheWriteStepNode : public StepNode { - public: - /*! \brief The scope name of the newly added compute stage. (e.g. local, shared, global) */ - String scope_name; - - void WriteToRecord(dmlc::JSONWriter* writer) const final; - - /*! - * \brief Apply the current step to State. - * \param state A mutable pointer to state, which will be updated. - * \param dag The original ComputeDAG of this state. - * \return The index of the new added stage. - */ - int ApplyToState(State* state, const ComputeDAG& dag) const; - - /*! - * \brief Apply the current step to tvm.schedule. - * \param stages The list of current stages - * \param stage_to_axes A map that maps stage ot all its iterators. - * \param schedule A mutable pointer to a te::Schedule. - * \return The output Tensors of the new added stage. - */ - Array ApplyToSchedule(Array* stages, StageToAxesMap* stage_to_axes, - te::Schedule* schedule) const; - - /*! - * \brief Print the current step as equivalent python schedule API. - * \param stages The list of current stages - * \param stage_to_axes A map that maps stage ot all its iterators. - * \param schedule A mutable pointer to a te::Schedule. - * \return Python schedule code. - */ - String PrintAsPythonAPI(Array* stages, StageToAxesMap* stage_to_axes, - te::Schedule* schedule) const; - - static constexpr const char* record_prefix_str = "CHW"; - - static constexpr const char* _type_key = "auto_scheduler.CacheWriteStep"; - TVM_DECLARE_FINAL_OBJECT_INFO(CacheWriteStepNode, StepNode); -}; - -/*! - * \brief Managed reference to CacheWriteStepNode. - * \sa CacheWriteStepNode - */ -class CacheWriteStep : public Step { - public: - /*! - * \brief The constructor. - * \param stage_id The index of the stage to be cache_write. - * \param scope_name The scope name of the newly added stage. - */ - CacheWriteStep(int stage_id, String scope_name); - - /*! - * \brief The constructor used to read a step record from JSONReader and create the - * corresponding step. - * \param reader The input JSONReader. - */ - explicit CacheWriteStep(dmlc::JSONReader* reader); - - TVM_DEFINE_OBJECT_REF_METHODS(CacheWriteStep, Step, CacheWriteStepNode); -}; - -/*! \brief Reduction factor step that corresponds to te::Schedule::rfactor */ -class RfactorStepNode : public StepNode { - public: - /*! \brief The index of the iterator to be factored. */ - int iter_id; - /*! \brief The position where the new iterator is placed. */ - int factor_iter_id; - - void WriteToRecord(dmlc::JSONWriter* writer) const final; - - /*! - * \brief Apply the current step to State. - * \param state A mutable pointer to State, which will be updated. - * \param dag The original ComputeDAG of this state. - * \return The index of the new added stage. - */ - int ApplyToState(State* state, const ComputeDAG& dag) const; - - /*! - * \brief Apply the current step to tvm.schedule. - * \param stages The list of current stages - * \param stage_to_axes A map that maps stage ot all its iterators. - * \param schedule A mutable pointer to a te::Schedule. - * \return The output Tensors of the new added stage. - */ - Array ApplyToSchedule(Array* stages, StageToAxesMap* stage_to_axes, - te::Schedule* schedule) const; - - /*! - * \brief Print the current step as equivalent python schedule API. - * \param stages The list of current stages - * \param stage_to_axes A map that maps stage ot all its iterators. - * \param schedule A mutable pointer to a te::Schedule. - * \return Python schedule code. - */ - String PrintAsPythonAPI(Array* stages, StageToAxesMap* stage_to_axes, - te::Schedule* schedule) const; - - static constexpr const char* record_prefix_str = "RF"; - - static constexpr const char* _type_key = "auto_scheduler.RfactorStep"; - TVM_DECLARE_FINAL_OBJECT_INFO(RfactorStepNode, StepNode); -}; - -/*! - * \brief Managed reference to RfactorStepNode. - * \sa RfactorStepNode - */ -class RfactorStep : public Step { - public: - /*! - * \brief The constructor. - * \param stage_id The index of the stage to be factored. - * \param iter_id The index of the iterator to be factored. - * \param factor_iter_id The position where the new iterator is placed. - */ - RfactorStep(int stage_id, int iter_id, int factor_iter_id); - - /*! - * \brief The constructor used to read a step record from JSONReader and create the - * corresponding step. - * \param reader The input JSONReader. - */ - explicit RfactorStep(dmlc::JSONReader* reader); - - TVM_DEFINE_OBJECT_REF_METHODS(RfactorStep, Step, RfactorStepNode); -}; - -} // namespace auto_scheduler -} // namespace tvm - -#endif // TVM_AUTO_SCHEDULER_TRANSFORM_STEP_H_ diff --git a/include/tvm/ir/adt.h b/include/tvm/ir/adt.h deleted file mode 100644 index 50e9bcbab273..000000000000 --- a/include/tvm/ir/adt.h +++ /dev/null @@ -1,163 +0,0 @@ -/* - * 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 tvm/ir/adt.h - * \brief Algebraic data type definitions. - * - * We adopt relay's ADT definition as a unified class - * for decripting structured data. - */ -#ifndef TVM_IR_ADT_H_ -#define TVM_IR_ADT_H_ - -#include -#include -#include -#include -#include -#include -#include - -#include - -namespace tvm { - -/*! - * \brief ADT constructor. - * Constructors compare by pointer equality. - * \sa Constructor - */ -class ConstructorNode : public RelayExprNode { - public: - /*! \brief The name (only a hint) */ - String name_hint; - /*! \brief Input to the constructor. */ - Array inputs; - /*! \brief The datatype the constructor will construct. */ - GlobalTypeVar belong_to; - /*! \brief Index in the table of constructors (set when the type is registered). */ - mutable int32_t tag = -1; - - ConstructorNode() {} - - void VisitAttrs(AttrVisitor* v) { - v->Visit("name_hint", &name_hint); - v->Visit("inputs", &inputs); - v->Visit("belong_to", &belong_to); - v->Visit("tag", &tag); - v->Visit("span", &span); - v->Visit("_checked_type_", &checked_type_); - } - - bool SEqualReduce(const ConstructorNode* other, SEqualReducer equal) const { - // Use namehint for now to be consistent with the legacy relay impl - // TODO(tvm-team) revisit, need to check the type var. - return equal(name_hint, other->name_hint) && equal(inputs, other->inputs); - } - - void SHashReduce(SHashReducer hash_reduce) const { - hash_reduce(name_hint); - hash_reduce(inputs); - } - - static constexpr const char* _type_key = "relay.Constructor"; - TVM_DECLARE_FINAL_OBJECT_INFO(ConstructorNode, RelayExprNode); -}; - -/*! - * \brief Managed reference to ConstructorNode - * \sa ConstructorNode - */ -class Constructor : public RelayExpr { - public: - /*! - * \brief Constructor - * \param name_hint the name of the constructor. - * \param inputs The input types. - * \param belong_to The data type var the constructor will construct. - */ - TVM_DLL Constructor(String name_hint, Array inputs, GlobalTypeVar belong_to); - - TVM_DEFINE_OBJECT_REF_METHODS(Constructor, RelayExpr, ConstructorNode); -}; - -/*! \brief TypeData container node */ -class TypeDataNode : public TypeNode { - public: - /*! - * \brief The header is simply the name of the ADT. - * We adopt nominal typing for ADT definitions; - * that is, differently-named ADT definitions with same constructors - * have different types. - */ - GlobalTypeVar header; - /*! \brief The type variables (to allow for polymorphism). */ - Array type_vars; - /*! \brief The constructors. */ - Array constructors; - - void VisitAttrs(AttrVisitor* v) { - v->Visit("header", &header); - v->Visit("type_vars", &type_vars); - v->Visit("constructors", &constructors); - v->Visit("span", &span); - } - - bool SEqualReduce(const TypeDataNode* other, SEqualReducer equal) const { - return equal.DefEqual(header, other->header) && equal.DefEqual(type_vars, other->type_vars) && - equal(constructors, other->constructors); - } - - void SHashReduce(SHashReducer hash_reduce) const { - hash_reduce.DefHash(header); - hash_reduce.DefHash(type_vars); - hash_reduce(constructors); - } - - static constexpr const char* _type_key = "relay.TypeData"; - TVM_DECLARE_FINAL_OBJECT_INFO(TypeDataNode, TypeNode); -}; - -/*! - * \brief Stores all data for an Algebraic Data Type (ADT). - * - * In particular, it stores the handle (global type var) for an ADT - * and the constructors used to build it and is kept in the module. Note - * that type parameters are also indicated in the type data: this means that - * for any instance of an ADT, the type parameters must be indicated. That is, - * an ADT definition is treated as a type-level function, so an ADT handle - * must be wrapped in a TypeCall node that instantiates the type-level arguments. - * The kind checker enforces this. - */ -class TypeData : public Type { - public: - /*! - * \brief Constructor - * \param header the name of ADT. - * \param type_vars type variables. - * \param constructors constructors field. - */ - TVM_DLL TypeData(GlobalTypeVar header, Array type_vars, Array constructors); - - TVM_DEFINE_OBJECT_REF_METHODS(TypeData, Type, TypeDataNode); -}; - -} // namespace tvm -#endif // TVM_IR_ADT_H_ diff --git a/include/tvm/ir/env_func.h b/include/tvm/ir/env_func.h index 386666a2c50c..dd93af6852fe 100644 --- a/include/tvm/ir/env_func.h +++ b/include/tvm/ir/env_func.h @@ -25,6 +25,7 @@ #define TVM_IR_ENV_FUNC_H_ #include +#include #include #include diff --git a/include/tvm/ir/function.h b/include/tvm/ir/function.h index 381ea6b8d6d3..3845409968e7 100644 --- a/include/tvm/ir/function.h +++ b/include/tvm/ir/function.h @@ -83,7 +83,7 @@ enum class LinkageType : int { /*! * \brief Generic attribute names that can be attached to any function. * - * \sa tvm::tir::attr, tvm::relay::attr + * \sa tvm::tir::attr, tvm::relax::attr */ namespace attr { /*! diff --git a/include/tvm/ir/module.h b/include/tvm/ir/module.h index 8fd87a6304dd..b3895ee725d1 100644 --- a/include/tvm/ir/module.h +++ b/include/tvm/ir/module.h @@ -24,7 +24,6 @@ #ifndef TVM_IR_MODULE_H_ #define TVM_IR_MODULE_H_ -#include #include #include #include @@ -58,8 +57,6 @@ class IRModuleNode : public Object { public: /*! \brief A map from ids to all global functions. */ Map functions; - /*! \brief A map from global type vars to ADT type data. */ - Map type_definitions; /*! \brief The source map for the module. */ SourceMap source_map; /* \brief Additional attributes storing meta-data about the module. */ @@ -72,21 +69,6 @@ class IRModuleNode : public Object { */ Map global_var_map_; - /*! \brief A map from string names to global type variables (ADT names) - * that ensures global uniqueness. - */ - Map global_type_var_map_; - - /*! \brief A map from constructor tags to constructor objects - * for convenient access - */ - std::unordered_map constructor_tag_map_; - - /*! \brief The files previously imported, required to ensure - importing is idempotent for each module. - */ - std::unordered_set import_set_; - /*! * \brief Get a module attribute. * @@ -149,9 +131,7 @@ class IRModuleNode : public Object { void VisitAttrs(AttrVisitor* v) { v->Visit("functions", &functions); - v->Visit("type_definitions", &type_definitions); v->Visit("global_var_map_", &global_var_map_); - v->Visit("global_type_var_map_", &global_type_var_map_); v->Visit("source_map", &source_map); v->Visit("attrs", &attrs); v->Visit("global_infos", &global_infos); @@ -179,27 +159,6 @@ class IRModuleNode : public Object { */ TVM_DLL void AddUnchecked(const GlobalVar& var, const BaseFunc& func); - /*! - * \brief Add a type-level definition to the global environment. - * \param var The var of the global type definition. - * \param type The ADT. - * \param update Controls whether you can replace a definition in the - * environment. - */ - TVM_DLL void AddTypeDef(const GlobalTypeVar& var, const TypeData& type, bool update = false); - - /*! - * \brief Add a type-level definition to the global environment. - * \param var The var of the global type definition. - * \param type The ADT. - * \param update Controls whether you can replace a definition in the - * environment. - * - * It does not do type checking as AddTypeDef does. - */ - TVM_DLL void AddTypeDefUnchecked(const GlobalTypeVar& var, const TypeData& type, - bool update = false); - /*! * \brief Update a function in the global environment. * \param var The name of the global function to update. @@ -207,13 +166,6 @@ class IRModuleNode : public Object { */ TVM_DLL void Update(const GlobalVar& var, const BaseFunc& func); - /*! - * \brief Update a type definition in the global environment. - * \param var The name of the global type definition to update. - * \param type The new ADT. - */ - TVM_DLL void UpdateTypeDef(const GlobalTypeVar& var, const TypeData& type); - /*! * \brief Update an array of global infos in the global environment. * \param name The name of the global info. @@ -234,13 +186,6 @@ class IRModuleNode : public Object { */ TVM_DLL bool ContainGlobalVar(const String& name) const; - /*! - * \brief Check if the global_type_var_map_ contains a global type variable. - * \param name The variable name. - * \returns true if contains, otherise false. - */ - TVM_DLL bool ContainGlobalTypeVar(const String& name) const; - /*! * \brief Lookup a global function by its variable. * \param str The unique string specifying the global variable. @@ -255,27 +200,6 @@ class IRModuleNode : public Object { */ TVM_DLL Array GetGlobalVars() const; - /*! - * \brief Look up a global function by its name. - * \param str The unique string specifying the global variable. - * \returns The global variable. - */ - TVM_DLL GlobalTypeVar GetGlobalTypeVar(const String& str) const; - - /*! - * \brief Collect all global type vars defined in this module. - * \returns An array of global type vars - */ - TVM_DLL Array GetGlobalTypeVars() const; - - /*! - * \brief Find constructor of ADT using name - * \param adt name of the ADT the constructor belongs to - * \param cons name of the constructor - * \returns Constructor of ADT, error if not found - */ - TVM_DLL Constructor GetConstructor(const String& adt, const String& cons) const; - /*! * \brief Look up a global function by its variable. * \param var The global var to lookup. @@ -290,27 +214,6 @@ class IRModuleNode : public Object { */ TVM_DLL BaseFunc Lookup(const String& name) const; - /*! - * \brief Look up a global type definition by its variable. - * \param var The var of the global type definition. - * \return The type definition. - */ - TVM_DLL TypeData LookupTypeDef(const GlobalTypeVar& var) const; - - /*! - * \brief Look up a global type definition by its name. - * \param var The name of the global type definition. - * \return The type definition. - */ - TVM_DLL TypeData LookupTypeDef(const String& var) const; - - /*! - * \brief Look up a constructor by its tag. - * \param tag The tag for the constructor. - * \return The constructor object. - */ - TVM_DLL Constructor LookupTag(const int32_t tag); - /*! * \brief Update the functions inside this environment by * functions in another environment. @@ -323,24 +226,6 @@ class IRModuleNode : public Object { * \returns The shallow copy of the IRModule. */ TVM_DLL IRModule ShallowCopy(); - - /*! - * \brief Import Relay code from the file at path. - * \param path The path of the Relay code to import. - * - * \note The path resolution behavior is standard, - * if abosolute will be the absolute file, if - * relative it will be resovled against the current - * working directory. - */ - TVM_DLL void Import(const String& path); - - /*! - * \brief Import Relay code from the file at path, relative to the standard library. - * \param path The path of the Relay code to import. - */ - TVM_DLL void ImportFromStd(const String& path); - /*! * \brief The set of imported files. */ @@ -354,8 +239,6 @@ class IRModuleNode : public Object { TVM_DECLARE_FINAL_OBJECT_INFO(IRModuleNode, Object); private: - /*! \brief Helper function for registering a typedef's constructors */ - void RegisterConstructors(const GlobalTypeVar& var, const TypeData& type); friend class IRModule; }; @@ -368,15 +251,11 @@ class IRModule : public ObjectRef { /*! * \brief constructor * \param functions Functions in the module. - * \param type_definitions Type definitions in the module. - * \param import_set Set of imported files in the module. * \param map The module source map. * \param attrs The module meta-data attributes. * \param global_infos Global infos in the module. */ - TVM_DLL explicit IRModule(Map functions, - Map type_definitions = {}, - std::unordered_set import_set = {}, SourceMap map = {}, + TVM_DLL explicit IRModule(Map functions, SourceMap map = {}, DictAttrs attrs = DictAttrs(), Map> global_infos = {}); @@ -394,51 +273,12 @@ class IRModule : public ObjectRef { return static_cast(ptr); } - /*! - * \brief Constructs a module from a standalone expression \p expr. - * - * If \p expr is a function it will be bound directly. Otherwise a function over the free - * variables of \p expr (possibly none) with \p expr as body is created and bound. - * - * The function is bound to, in preference order: - * - The "global_symbol" attribute of \p expr, if it is a function with that attribute. - * - 'main' - * - A unique name derived from 'main' if 'main' is already bound in \p global_funcs. - * - * Additional global functions and type definitions may be included in the result module. - * - * See also \p FromExpr. - * - * \param expr The expression to set as the main function to the module. - * \param global_funcs The global function map. Default empty. - * \param type_definitions The global type definition map. Default empty. - * \param import_set Set of external modules already imported. Default empty. - * - * \returns A module with \p expr set as the main function, and the global var to which - * \p expr was bound (typcially 'main'). - * - * TODO(mbs): Does import_set and the bound global var need to be exposed via ffi? - */ - static std::pair FromExprInContext( - const RelayExpr& expr, const Map& global_funcs = {}, - const Map& type_definitions = {}, - std::unordered_set import_set = {}); - /*! * \brief As for \p FromExprInContext, but assuming \p expr is bound to 'main' and no * imports. */ TVM_DLL static IRModule FromExpr(const RelayExpr& expr, - const Map& global_funcs = {}, - const Map& type_definitions = {}); - - /*! - * \brief Parse text format source file into an IRModule. - * \param text A string of Relay source code. - * \param source_path The path to the source file. - * \return A Relay module. - */ - TVM_DLL static IRModule FromText(const String& text, const String& source_path); + const Map& global_funcs = {}); /*! * \brief Create a shallow copy of an IRModule. diff --git a/include/tvm/ir/op.h b/include/tvm/ir/op.h index 6e6b8bee5fc3..a703f16f5a3b 100644 --- a/include/tvm/ir/op.h +++ b/include/tvm/ir/op.h @@ -25,12 +25,12 @@ #ifndef TVM_IR_OP_H_ #define TVM_IR_OP_H_ -#include #include +#include #include #include -#include #include +#include #include #include @@ -110,17 +110,6 @@ class OpNode : public RelayExprNode { hash_reduce(name); } - /*! - * \brief Check that if current op is a "primtive operator". - * That is the arguments are all type variables, and there is a single - * type relation applied to the input and output types. - */ - bool IsPrimitiveOp() const { - if (is_primitive_ != -1) return is_primitive_ != 0; - is_primitive_ = this->IsPrimitiveOp_() ? 1 : 0; - return is_primitive_ != 0; - } - static constexpr const char* _type_key = "Op"; TVM_DECLARE_FINAL_OBJECT_INFO(OpNode, RelayExprNode); @@ -137,25 +126,9 @@ class OpNode : public RelayExprNode { friend class AttrRegistry; friend class OpRegEntry; - friend bool IsPrimitiveOp(const RelayExpr&); // Program internal unique index of operator. // Used to help index the program. uint32_t index_{0}; - // whether this is a primitive op. -1 means unknown. - mutable int is_primitive_{-1}; - // Internal function to compute if it is primitive op - bool IsPrimitiveOp_() const { - const auto& fn_ty = this->op_type; - ICHECK(fn_ty.get() != nullptr) << "op_type of " << this->name << " is not registered"; - if (fn_ty->type_constraints.size() != 1) return false; - const TypeRelationNode* rel = fn_ty->type_constraints[0].as(); - if (rel == nullptr) return false; - // validate if the type parameter matches up - for (size_t i = 0; i < fn_ty->type_params.size(); ++i) { - if (!fn_ty->type_params[i].same_as(rel->args[i])) return false; - } - return true; - } }; /*! @@ -222,17 +195,6 @@ class OpRegEntry { */ inline OpRegEntry& add_argument(const std::string& name, const std::string& type, const std::string& description); - /*! - * \brief Attach the type function corresponding to the return type. - * \param rel_name The type relation name to register. - * \param type_rel_func The backing relation function which can solve an arbitrary - * relation on variables. - * \return reference to self. - */ - inline OpRegEntry& add_type_rel( - const std::string& rel_name, - runtime::TypedPackedFunc&, int, const Attrs&, const TypeReporter&)> - type_rel_func); /*! * \brief Set the attrs type key and index to be AttrsType. * \tparam AttrsType the attribute type to b set. @@ -383,60 +345,6 @@ inline OpRegEntry& OpRegEntry::add_argument(const std::string& name, const std:: return *this; } -inline OpRegEntry& OpRegEntry::add_type_rel( - const std::string& rel_name, - runtime::TypedPackedFunc&, int, const Attrs&, const TypeReporter&)> - type_rel_func) { - auto func_name = std::string("tvm.relay.type_relation.") + rel_name; - TypeRelationFn env_type_rel_func; - - if (runtime::Registry::Get(func_name)) { - auto env_func = EnvFunc::Get(func_name); - env_type_rel_func = env_func; - } else { - runtime::Registry::Register(func_name).set_body(type_rel_func.packed()); - auto env_func = EnvFunc::Get(func_name); - env_type_rel_func = env_func; - } - - Array type_params; - Array arg_types; - - // Add inputs. - std::string input_name_prefix = "in"; - for (int i = 0; i < get()->num_inputs; i++) { - auto name = input_name_prefix + std::to_string(i); - auto param = TypeVar(name, TypeKind::kType); - type_params.push_back(param); - arg_types.push_back(param); - } - - Array ty_call_args = arg_types; - - // Add output type. - auto out_param = TypeVar("out", TypeKind::kType); - type_params.push_back(out_param); - // this will trigger copy on write. - ty_call_args.push_back(out_param); - - // The attributes of primitive op is nullptr - // - // The attributes of primitive operator can vary at the call site. - // The type of sum is also dependent on Attrs being passed. - // So puting nullptr in the Attrs means that the operator is polymorphic on Attrs. - // - // A common example is sum(x, axis), where the choice of axis - // can affect the type of the function. - TypeConstraint type_rel = - TypeRelation(env_type_rel_func, ty_call_args, arg_types.size(), Attrs()); - - auto func_type = FuncType(arg_types, out_param, type_params, {type_rel}); - - get()->op_type = func_type; - - return *this; -} - inline OpRegEntry& OpRegEntry::set_num_inputs(int32_t n) { // NOLINT(*) get()->num_inputs = n; return *this; @@ -482,23 +390,5 @@ inline ValueType OpAttrMap::get(const RelayExpr& expr, ValueType def_ } } -/*! - * \brief Check that an expression is a "primitive operator". - * - * Will return true if the expression is an operator which - * matches the form of primitive operators registered directly - * by the Relay codebase. - * - * That is the arguments are all type variables, and there is a single - * type relation applied to the input and output types. - * - * \param expr An expression. - * \return Whether the expression is primitive op. - */ -inline bool IsPrimitiveOp(const RelayExpr& expr) { - const auto* op = expr.as(); - return op != nullptr && op->IsPrimitiveOp(); -} - } // namespace tvm #endif // TVM_IR_OP_H_ diff --git a/include/tvm/ir/si_builder.h b/include/tvm/ir/si_builder.h deleted file mode 100644 index ab5f2d450fe4..000000000000 --- a/include/tvm/ir/si_builder.h +++ /dev/null @@ -1,103 +0,0 @@ -/* - * 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 tvm/ir/si_builder.h - * \brief build a source info during rewriting expressions. - */ -#ifndef TVM_IR_SI_BUILDER_H_ -#define TVM_IR_SI_BUILDER_H_ - -#include -#include -#include - -#include -#include - -namespace tvm { - -/*! - * \brief Source Information Builder, SIBuilder provides helper APIs for filling spans, - * particularly useful for one-to-many, many-to-one and many-to-many IR transformations. - */ -class SIBuilder { - public: - /*! - * \brief Create SIBuilder from a given span - */ - explicit SIBuilder(const Span& span = Span()); - - /*! - * \brief Create SIBuilder from a given span sequence - */ - explicit SIBuilder(const Array& spans = Array()); - explicit SIBuilder(const std::initializer_list& init); - - /*! - * \brief Create SIBuilder via a subgraph, - * Will construct span based on the exprs in the subgraph. Including the inputs exprs. - * - * \param entry Entry expr for subgraph - * \param inputs End exprs for subgraph - */ - template ::value>> - explicit SIBuilder(const T& entry, const tvm::Array& inputs = {}); - explicit SIBuilder(const tir::Stmt& entry, const tvm::Array& inputs = {}); - explicit SIBuilder(const tir::Stmt& entry, const tvm::Array& inputs = {}); - - ~SIBuilder(); - - SIBuilder(const SIBuilder&) = delete; - SIBuilder& operator=(const SIBuilder&) = delete; - - /*! - * \brief build a span of source information, which is based on the given span or subgraph. - * - * \return the built span - */ - Span Build() const; - - /*! - * \brief Recursively fill all span of exprs in subgraph from entry until inputs. - * - * \param entry Entry expr for subgraph. - * \param inputs End exprs for subgraph, will not be filled with new span. - */ - template ::value>> - void RecursivelyFillSpan( - const T& entry, const std::unordered_set& inputs) const; - - void RecursivelyFillSpan( - const tir::Stmt& entry, - const std::unordered_set& inputs) const; - void RecursivelyFillSpan( - const tir::Stmt& entry, - const std::unordered_set& inputs) const; - - private: - struct Impl; - std::unique_ptr impl_; - - std::unique_ptr CreateImpl(const Span& span); -}; - -} // namespace tvm - -#endif // TVM_IR_SI_BUILDER_H_ diff --git a/include/tvm/ir/type.h b/include/tvm/ir/type.h index ec13635a2643..04a59c7c422f 100644 --- a/include/tvm/ir/type.h +++ b/include/tvm/ir/type.h @@ -197,152 +197,6 @@ class PointerType : public Type { TVM_DEFINE_OBJECT_REF_METHODS(PointerType, Type, PointerTypeNode); }; -/*! \brief Possible kinds of TypeVars. */ -enum TypeKind : int { - kType = 0, - /*! \brief Template variable in shape expression. */ - kShapeVar = 1, - kBaseType = 2, - kConstraint = 4, - kAdtHandle = 5, - kTypeData = 6 -}; - -/*! \brief Converts a TypeKind to a string. */ -inline String TypeKind2String(TypeKind kind) { - switch (kind) { - case TypeKind::kType: - return "Type"; - case TypeKind::kShapeVar: - return "ShapeVar"; - case TypeKind::kBaseType: - return "BaseType"; - case TypeKind::kConstraint: - return "Constraint"; - case TypeKind::kAdtHandle: - return "AdtHandle"; - case TypeKind::kTypeData: - return "TypeData"; - } - LOG(FATAL) << "ValueError: Unknown TypeKind: " << static_cast(kind); -} - -/*! - * \brief Type parameter in functions. - * - * A type variable can be viewed as template parameter in c++ template function. - * - * For example, in the following pesudo code, - * the TypeVar of f is TypeVar("n", kind=kShapeVar). - * This function can take in a Tensor with shape=(3, 3) and - * returns a Tensor with shape=(9,) - * - * \code - * - * template - * f(x : Tensor[i32, (n, n)]) -> Tensor[i32, (n * n)] - * - * \endcode - * \sa TypeVar, TypeKind - */ -class TypeVarNode : public TypeNode { - public: - /*! - * \brief The name of the variable, - * this only acts as a hint to the user, - * and is not used for equality. - */ - String name_hint; - /*! \brief The kind of type parameter */ - TypeKind kind; - - void VisitAttrs(AttrVisitor* v) { - v->Visit("name_hint", &name_hint); - v->Visit("kind", &kind); - v->Visit("span", &span); - } - - bool SEqualReduce(const TypeVarNode* other, SEqualReducer equal) const { - return equal(kind, other->kind) && equal.FreeVarEqualImpl(this, other); - } - - void SHashReduce(SHashReducer hash_reduce) const { - hash_reduce(kind); - hash_reduce.FreeVarHashImpl(this); - } - - static constexpr const char* _type_key = "TypeVar"; - TVM_DECLARE_FINAL_OBJECT_INFO(TypeVarNode, TypeNode); -}; - -/*! - * \brief Managed reference to TypeVarNode - * \sa TypeVarNode - */ -class TypeVar : public Type { - public: - /*! - * \brief Constructor - * \param name_hint The name of the type var. - * \param kind The kind of the type var. - * \param span The span information. - */ - TVM_DLL TypeVar(String name_hint, TypeKind kind, Span span = Span()); - - TVM_DEFINE_OBJECT_REF_METHODS(TypeVar, Type, TypeVarNode); -}; - -/*! - * \brief A global type variable that is used for defining new types or type aliases. - * \sa GlobalTypeVar - */ -class GlobalTypeVarNode : public TypeNode { - public: - /*! - * \brief The name of the variable, - * this only acts as a hint to the user, - * and is not used for equality. - */ - String name_hint; - /*! \brief The kind of type parameter */ - TypeKind kind; - - void VisitAttrs(AttrVisitor* v) { - v->Visit("name_hint", &name_hint); - v->Visit("kind", &kind); - } - - bool SEqualReduce(const GlobalTypeVarNode* other, SEqualReducer equal) const { - // name matters for now in global type var. - return equal(name_hint, other->name_hint) && equal.FreeVarEqualImpl(this, other); - } - - void SHashReduce(SHashReducer hash_reduce) const { - hash_reduce(name_hint); - hash_reduce.FreeVarHashImpl(this); - } - - static constexpr const char* _type_key = "GlobalTypeVar"; - TVM_DECLARE_FINAL_OBJECT_INFO(GlobalTypeVarNode, TypeNode); -}; - -/*! - * \brief Managed reference to GlobalTypeVarNode - * \sa GlobalTypeVarNode - */ -class GlobalTypeVar : public Type { - public: - /*! - * \brief Constructor - * \param name_hint The name of the type var. - * \param kind The kind of the type var. - * \param span The span of the type. - */ - TVM_DLL GlobalTypeVar(String name_hint, TypeKind kind, Span span = Span()); - - TVM_DEFINE_OBJECT_REF_METHODS(GlobalTypeVar, Type, GlobalTypeVarNode); -}; - /*! * \brief The type of tuple values. * \sa TupleType @@ -405,33 +259,13 @@ inline bool IsVoidType(const Type& type) { return n && n->fields.size() == 0; } -/*! - * \brief Potential Constraints in a function. - * \sa TypeConstraint - */ -class TypeConstraintNode : public TypeNode { - public: - static constexpr const char* _type_key = "TypeConstraint"; - static constexpr const uint32_t _type_child_slots = 1; - TVM_DECLARE_BASE_OBJECT_INFO(TypeConstraintNode, TypeNode); -}; - -/*! - * \brief Managed reference to TypeConstraintNode. - * \sa TypeConstraintNode, TypeRelation - */ -class TypeConstraint : public Type { - public: - TVM_DEFINE_OBJECT_REF_METHODS(TypeConstraint, Type, TypeConstraintNode); -}; - /*! * \brief Function type. * * We support polymorphic function type. * This can be roughly viewed as template function in C++. * - * \sa FuncType, TypeVar, TypeConstraint + * \sa FuncType, TypeConstraint */ class FuncTypeNode : public TypeNode { public: @@ -439,35 +273,21 @@ class FuncTypeNode : public TypeNode { Array arg_types; /*! \brief The type of return value. */ Type ret_type; - // The following fields are used in polymorphic(template) functions - // For normal functions, the following two fields will be empty. - /*! \brief The type parameters of the function */ - Array type_params; - /*! - * \brief potential constraint the type need to obey - * \note this field is reserved for further purposes. - */ - Array type_constraints; void VisitAttrs(AttrVisitor* v) { v->Visit("arg_types", &arg_types); v->Visit("ret_type", &ret_type); - v->Visit("type_params", &type_params); - v->Visit("type_constraints", &type_constraints); v->Visit("span", &span); } bool SEqualReduce(const FuncTypeNode* other, SEqualReducer equal) const { // type params first as they defines type vars. - return equal.DefEqual(type_params, other->type_params) && equal(arg_types, other->arg_types) && - equal(ret_type, other->ret_type) && equal(type_constraints, other->type_constraints); + return equal(arg_types, other->arg_types) && equal(ret_type, other->ret_type); } void SHashReduce(SHashReducer hash_reduce) const { - hash_reduce.DefHash(type_params); hash_reduce(arg_types); hash_reduce(ret_type); - hash_reduce(type_constraints); } static constexpr const char* _type_key = "FuncType"; @@ -484,100 +304,13 @@ class FuncType : public Type { * \brief Constructor * \param arg_types The types of the arguments. * \param ret_type The type of the return value. - * \param type_params The type parameters. - * \param type_constraints The type constraints. * \param span The span information. * \sa FuncTypeNode for more docs about these fields. */ - TVM_DLL FuncType(Array arg_types, Type ret_type, Array type_params, - Array type_constraints, Span span = Span()); + TVM_DLL FuncType(Array arg_types, Type ret_type, Span span = Span()); TVM_DEFINE_OBJECT_REF_METHODS(FuncType, Type, FuncTypeNode); }; -/*! - * \brief Intermediate values that is used to indicate incomplete type - * during type inference. - * - * If we view the type relations as "computational graph of types", - * then IncompleteType represents intermediate values of the graph, - * TypeVar represents the input to the graph. - * - * \sa IncompleteType - */ -class IncompleteTypeNode : public TypeNode { - public: - /*! \brief kind of the type. */ - TypeKind kind; - - void VisitAttrs(tvm::AttrVisitor* v) { - v->Visit("kind", &kind); - v->Visit("span", &span); - } - - bool SEqualReduce(const IncompleteTypeNode* other, SEqualReducer equal) const { - return equal(kind, other->kind) && equal.FreeVarEqualImpl(this, other); - } - - void SHashReduce(SHashReducer hash_reduce) const { hash_reduce(kind); } - - static constexpr const char* _type_key = "IncompleteType"; - TVM_DECLARE_FINAL_OBJECT_INFO(IncompleteTypeNode, TypeNode); -}; - -/*! - * \brief Managed reference to IncompleteTypeNode. - * \sa IncompleteTypeNode - */ -class IncompleteType : public Type { - public: - /*! - * \brief Constructor. - * \param kind kind of the type. - * \param span The span information. - */ - TVM_DLL explicit IncompleteType(TypeKind kind, Span span = Span()); - - TVM_DEFINE_OBJECT_REF_METHODS(IncompleteType, Type, IncompleteTypeNode); -}; - -/*! - * \brief Reference Type High-level Relay IR. - * - * \sa RelayRefType. - */ -class RelayRefTypeNode : public TypeNode { - public: - /*! \brief The type of value in the Reference. */ - Type value; - - RelayRefTypeNode() {} - - void VisitAttrs(tvm::AttrVisitor* v) { - v->Visit("value", &value); - v->Visit("span", &span); - } - - bool SEqualReduce(const RelayRefTypeNode* other, SEqualReducer equal) const { - return equal(value, other->value); - } - - void SHashReduce(SHashReducer hash_reduce) const { hash_reduce(value); } - - // Keep the relay prefix in the type as this type is specific - // to the relay itself. - static constexpr const char* _type_key = "relay.RefType"; - TVM_DECLARE_FINAL_OBJECT_INFO(RelayRefTypeNode, TypeNode); -}; - -/*! - * \brief Managed reference to RelayRefTypeNode. - * \sa RelayRefTypeNode. - */ -class RelayRefType : public Type { - public: - TVM_DLL explicit RelayRefType(Type value, Span span = Span()); - TVM_DEFINE_OBJECT_REF_METHODS(RelayRefType, Type, RelayRefTypeNode); -}; } // namespace tvm #endif // TVM_IR_TYPE_H_ diff --git a/include/tvm/ir/type_functor.h b/include/tvm/ir/type_functor.h index 334a35d052e1..eb213b17dfbc 100644 --- a/include/tvm/ir/type_functor.h +++ b/include/tvm/ir/type_functor.h @@ -25,7 +25,6 @@ #define TVM_IR_TYPE_FUNCTOR_H_ #include -#include #include #include @@ -77,16 +76,8 @@ class TypeFunctor { } // Functions that can be overriden by subclass virtual R VisitType_(const TensorTypeNode* op, Args... args) TYPE_FUNCTOR_DEFAULT; - virtual R VisitType_(const TypeVarNode* op, Args... args) TYPE_FUNCTOR_DEFAULT; - virtual R VisitType_(const TypeConstraintNode* op, Args... args) TYPE_FUNCTOR_DEFAULT; virtual R VisitType_(const FuncTypeNode* op, Args... args) TYPE_FUNCTOR_DEFAULT; - virtual R VisitType_(const TypeRelationNode* op, Args... args) TYPE_FUNCTOR_DEFAULT; virtual R VisitType_(const TupleTypeNode* op, Args... args) TYPE_FUNCTOR_DEFAULT; - virtual R VisitType_(const IncompleteTypeNode* op, Args... args) TYPE_FUNCTOR_DEFAULT; - virtual R VisitType_(const RelayRefTypeNode* op, Args... args) TYPE_FUNCTOR_DEFAULT; - virtual R VisitType_(const GlobalTypeVarNode* op, Args... args) TYPE_FUNCTOR_DEFAULT; - virtual R VisitType_(const TypeCallNode* op, Args... args) TYPE_FUNCTOR_DEFAULT; - virtual R VisitType_(const TypeDataNode* op, Args... args) TYPE_FUNCTOR_DEFAULT; virtual R VisitType_(const PrimTypeNode* op, Args... args) TYPE_FUNCTOR_DEFAULT; virtual R VisitType_(const PointerTypeNode* op, Args... args) TYPE_FUNCTOR_DEFAULT; virtual R VisitTypeDefault_(const Object* op, Args...) { @@ -100,16 +91,8 @@ class TypeFunctor { FType vtable; // Set dispatch TVM_TYPE_FUNCTOR_DISPATCH(TensorTypeNode); - TVM_TYPE_FUNCTOR_DISPATCH(TypeVarNode); - TVM_TYPE_FUNCTOR_DISPATCH(TypeConstraintNode); TVM_TYPE_FUNCTOR_DISPATCH(FuncTypeNode); - TVM_TYPE_FUNCTOR_DISPATCH(TypeRelationNode); TVM_TYPE_FUNCTOR_DISPATCH(TupleTypeNode); - TVM_TYPE_FUNCTOR_DISPATCH(IncompleteTypeNode); - TVM_TYPE_FUNCTOR_DISPATCH(RelayRefTypeNode); - TVM_TYPE_FUNCTOR_DISPATCH(GlobalTypeVarNode); - TVM_TYPE_FUNCTOR_DISPATCH(TypeCallNode); - TVM_TYPE_FUNCTOR_DISPATCH(TypeDataNode); TVM_TYPE_FUNCTOR_DISPATCH(PrimTypeNode); TVM_TYPE_FUNCTOR_DISPATCH(PointerTypeNode); return vtable; @@ -123,16 +106,9 @@ class TypeFunctor { */ class TVM_DLL TypeVisitor : public TypeFunctor { public: - void VisitType_(const TypeVarNode* op) override; - void VisitType_(const IncompleteTypeNode* op) override; void VisitType_(const TensorTypeNode* op) override; void VisitType_(const FuncTypeNode* op) override; void VisitType_(const TupleTypeNode* op) override; - void VisitType_(const TypeRelationNode* op) override; - void VisitType_(const RelayRefTypeNode* op) override; - void VisitType_(const GlobalTypeVarNode* op) override; - void VisitType_(const TypeCallNode* op) override; - void VisitType_(const TypeDataNode* op) override; void VisitType_(const PrimTypeNode* op) override; void VisitType_(const PointerTypeNode* op) override; }; @@ -143,16 +119,9 @@ class TVM_DLL TypeVisitor : public TypeFunctor { class TVM_DLL TypeMutator : public TypeFunctor { public: Type VisitType(const Type& t) override; - Type VisitType_(const TypeVarNode* op) override; Type VisitType_(const TensorTypeNode* op) override; - Type VisitType_(const IncompleteTypeNode* op) override; Type VisitType_(const FuncTypeNode* op) override; Type VisitType_(const TupleTypeNode* op) override; - Type VisitType_(const TypeRelationNode* type_rel) override; - Type VisitType_(const RelayRefTypeNode* op) override; - Type VisitType_(const GlobalTypeVarNode* op) override; - Type VisitType_(const TypeCallNode* op) override; - Type VisitType_(const TypeDataNode* op) override; Type VisitType_(const PrimTypeNode* op) override; Type VisitType_(const PointerTypeNode* op) override; @@ -160,12 +129,5 @@ class TVM_DLL TypeMutator : public TypeFunctor { Array MutateArray(Array arr); }; -/*! - * \brief Bind free type variables in the type. - * \param type The type to be updated. - * \param args_map The binding map. - */ -Type Bind(const Type& type, const Map& args_map); - } // namespace tvm #endif // TVM_IR_TYPE_FUNCTOR_H_ diff --git a/include/tvm/ir/type_relation.h b/include/tvm/ir/type_relation.h deleted file mode 100644 index dd6861750a10..000000000000 --- a/include/tvm/ir/type_relation.h +++ /dev/null @@ -1,243 +0,0 @@ -/* - * 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 tvm/ir/type_relation.h - * \brief Type relation and function for type inference(checking). - */ -#ifndef TVM_IR_TYPE_RELATION_H_ -#define TVM_IR_TYPE_RELATION_H_ - -#include -#include -#include -#include -#include -#include - -namespace tvm { - -/*! - * \brief Type function application. - * \sa TypeCall - */ -class TypeCallNode : public TypeNode { - public: - /*! - * \brief The type-level function (ADT that takes type params). - */ - Type func; - /*! \brief The arguments. */ - Array args; - - void VisitAttrs(AttrVisitor* v) { - v->Visit("func", &func); - v->Visit("args", &args); - v->Visit("span", &span); - } - - bool SEqualReduce(const TypeCallNode* other, SEqualReducer equal) const { - return equal(func, other->func) && equal(args, other->args); - } - - void SHashReduce(SHashReducer hash_reduce) const { - hash_reduce(func); - hash_reduce(args); - } - - static constexpr const char* _type_key = "TypeCall"; - TVM_DECLARE_FINAL_OBJECT_INFO(TypeCallNode, TypeNode); -}; - -/*! - * \brief Managed reference to TypeCallNode. - * \sa TypeCallNode - */ -class TypeCall : public Type { - public: - /*! - * \brief Constructor - * \param func The type function to apply. - * \param args The arguments to the type function. - */ - TVM_DLL TypeCall(Type func, Array args); - - TVM_DEFINE_OBJECT_REF_METHODS(TypeCall, Type, TypeCallNode); -}; - -/*! - * \brief reporter that reports back to the - * type resolution information. - */ -class TypeReporterNode : public Object { - public: - /*! \brief virtual destructor */ - virtual ~TypeReporterNode() {} - /*! - * \brief Create a type equality constraint. - * - * The "assign direction" acts as a hint to the solver - * showing that it is more likely to resolve dst by src. - * But it is possible for the solver to resolve src by dst as well. - */ - TVM_DLL virtual void Assign(const Type& dst, const Type& src) = 0; - - /*! - * \brief assert shape expression comparison. - * \note Use assert only if any of the condition input is symbolic. - * \param cond The condition of operation. - * \return false if assertion can be proven to have failed - * true if solver can still proceed. - */ - TVM_DLL virtual bool Assert(const PrimExpr& cond) = 0; - /*! - * \brief assert shape expression equals each other. - * \param lhs The left operand. - * \param rhs The right operand. - * \return false if assertion can be proven to have failed - * true if solver can still proceed. - */ - TVM_DLL virtual bool AssertEQ(const PrimExpr& lhs, const PrimExpr& rhs) = 0; - - /*! - * \brief Set the location at which to report unification errors. - * \param span The span at which to report the error. - */ - TVM_DLL virtual void SetSpan(const Span& span) = 0; - - TVM_DLL virtual Span GetSpan() = 0; - - TVM_DLL virtual DiagnosticContext GetDiagCtx() = 0; - - /*! - * \brief Retrieve the current global module. - * \return The global module. - */ - TVM_DLL virtual IRModule GetModule() = 0; - - // solver is not serializable. - void VisitAttrs(AttrVisitor* v) {} - - static constexpr const char* _type_key = "TypeReporter"; - TVM_DECLARE_FINAL_OBJECT_INFO(TypeReporterNode, Object); -}; - -/*! - * \brief Container class of TypeReporter. - * \sa TypeReporterNode - */ -class TypeReporter : public ObjectRef { - public: - TypeReporter() {} - explicit TypeReporter(ObjectPtr n) : ObjectRef(n) {} - TypeReporterNode* operator->() const { - return const_cast(static_cast(get())); - } - using ContainerType = TypeReporterNode; -}; - -/*! - * \brief User defined type constraint function. - * - * If the input type information can be used to fully decide - * the IncompleteTypes, then the function should call - * reporter.Assign to report the new types, and return true. - * Otherwise, the function should return false. - * - * \param args The arguments to the relation. - * The types are stored in the form of - * [input_type_0, input_type_1, ... input_type_n, - * output_type_0, output_type_1, ... output_type_m] - * - * \param num_inputs Number of input types in the args. - * \param attrs The additional attributes of the operator. - * \param reporter The reporter to report solution to. - * \return false if This relation cannot be resolved. - * true if this relation has been resolved. - */ -using TypeRelationFn = TypedEnvFunc& args, int num_inputs, - const Attrs& attrs, const TypeReporter& reporter)>; - -/*! - * \brief User defined type relation, it is an input-output relation on types. - * - * TypeRelation is more generalized than type call as it allows inference - * of both inputs and outputs. - * - * \sa TypeRelation - */ -class TypeRelationNode : public TypeConstraintNode { - public: - /*! - * \brief The function on input and output variables which - * this is not directly serializable, - * need to be looked-up in the module. - */ - TypeRelationFn func; - /*! \brief The type arguments to the type function. */ - Array args; - /*! \brief Number of inputs arguments */ - int num_inputs; - /*! \brief Attributes to the relation function */ - Attrs attrs; - - void VisitAttrs(AttrVisitor* v) { - v->Visit("func", &func); - v->Visit("args", &args); - v->Visit("num_inputs", &num_inputs); - v->Visit("attrs", &attrs); - v->Visit("span", &span); - } - - bool SEqualReduce(const TypeRelationNode* other, SEqualReducer equal) const { - return equal(func, other->func) && equal(args, other->args) && - equal(num_inputs, other->num_inputs) && equal(attrs, other->attrs); - } - - void SHashReduce(SHashReducer hash_reduce) const { - hash_reduce(func); - hash_reduce(args); - hash_reduce(num_inputs); - hash_reduce(attrs); - } - - static constexpr const char* _type_key = "TypeRelation"; - TVM_DECLARE_FINAL_OBJECT_INFO(TypeRelationNode, TypeConstraintNode); -}; - -/*! - * \brief Managed reference to TypeRelationNode. - * \sa TypeRelationNode - */ -class TypeRelation : public TypeConstraint { - public: - /*! - * \brief Constructor - * \param func The relation function. - * \param args The arguments to the type relation. - * \param num_inputs Number of inputs. - * \param attrs Attributes to the relation function. - * \sa TypeRelationNode for more docs about these fields. - */ - TVM_DLL TypeRelation(TypeRelationFn func, Array args, int num_inputs, Attrs attrs); - - TVM_DEFINE_OBJECT_REF_METHODS(TypeRelation, TypeConstraint, TypeRelationNode); -}; -} // namespace tvm -#endif // TVM_IR_TYPE_RELATION_H_ diff --git a/include/tvm/node/functor.h b/include/tvm/node/functor.h index fa0ab552e6bb..58d59c81cb16 100644 --- a/include/tvm/node/functor.h +++ b/include/tvm/node/functor.h @@ -23,7 +23,7 @@ #ifndef TVM_NODE_FUNCTOR_H_ #define TVM_NODE_FUNCTOR_H_ -#include +#include #include #include diff --git a/include/tvm/relax/analysis.h b/include/tvm/relax/analysis.h index 527327d56a42..2de2f4fd36d5 100644 --- a/include/tvm/relax/analysis.h +++ b/include/tvm/relax/analysis.h @@ -28,8 +28,8 @@ #include #include #include +#include #include -#include #include #include @@ -511,7 +511,7 @@ TVM_DLL Expr RemoveAllUnused(Expr expr); * \note This analysis applies on TIR function but is primarily used by relax passes. * As a result we place it under the relax namespace. */ -TVM_DLL relay::OpPatternKind AnalyzeOpPatternKind(const tir::PrimFunc& func); +TVM_DLL OpPatternKind AnalyzeOpPatternKind(const tir::PrimFunc& func); /*! * \brief Check if the given PrimFunc is essentially doing a reshape operation. diff --git a/include/tvm/relax/attrs/op.h b/include/tvm/relax/attrs/op.h index 8e3e9d92f554..10c267b21faa 100644 --- a/include/tvm/relax/attrs/op.h +++ b/include/tvm/relax/attrs/op.h @@ -24,6 +24,7 @@ #ifndef TVM_RELAX_ATTRS_OP_H_ #define TVM_RELAX_ATTRS_OP_H_ +#include #include namespace tvm { diff --git a/include/tvm/relax/block_builder.h b/include/tvm/relax/block_builder.h index ad2b9820707a..070aef2fcb6d 100644 --- a/include/tvm/relax/block_builder.h +++ b/include/tvm/relax/block_builder.h @@ -25,6 +25,7 @@ #define TVM_RELAX_BLOCK_BUILDER_H_ #include +#include #include #include #include diff --git a/include/tvm/relax/dataflow_matcher.h b/include/tvm/relax/dataflow_matcher.h index 8f2024f26403..15d1c9c5fbda 100644 --- a/include/tvm/relax/dataflow_matcher.h +++ b/include/tvm/relax/dataflow_matcher.h @@ -26,6 +26,7 @@ #include #include +#include #include @@ -69,7 +70,8 @@ TVM_DLL Optional> MatchGraph(const PatternContext& ctx, */ TVM_DLL Function RewriteBindings( const PatternContext& ctx, - TypedPackedFunc(Map, Map)> rewriter, Function f); + runtime::TypedPackedFunc(Map, Map)> rewriter, + Function f); /** * \brief Rewrite a function with the given pattern and the rewriter function. @@ -95,7 +97,7 @@ TVM_DLL Function RewriteBindings( * \return The updated function, if any updates were applied. */ TVM_DLL Function RewriteCall(const DFPattern& pattern, - TypedPackedFunc)> rewriter, + runtime::TypedPackedFunc)> rewriter, Function func); } // namespace relax diff --git a/include/tvm/relax/expr.h b/include/tvm/relax/expr.h index 60032c34622f..9afddeb807ab 100644 --- a/include/tvm/relax/expr.h +++ b/include/tvm/relax/expr.h @@ -20,6 +20,7 @@ #define TVM_RELAX_EXPR_H_ #include +#include #include #include #include diff --git a/include/tvm/relax/expr_functor.h b/include/tvm/relax/expr_functor.h index c3aea24dcb50..9c867129fdd2 100644 --- a/include/tvm/relax/expr_functor.h +++ b/include/tvm/relax/expr_functor.h @@ -30,14 +30,10 @@ #include #include #include -#include #include -#include -#include #include #include -#include namespace tvm { namespace relax { diff --git a/include/tvm/relax/op_attr_types.h b/include/tvm/relax/op_attr_types.h index 0ddc2baefbef..434a89a28871 100644 --- a/include/tvm/relax/op_attr_types.h +++ b/include/tvm/relax/op_attr_types.h @@ -29,11 +29,31 @@ #include #include -#include - namespace tvm { namespace relax { +enum OpPatternKind { + // Elementwise operation + kElemWise = 0, + // Broadcasting operator, can always map output axis to the input in order. + // for example :code:`out[i, ax1, j, ax2] = input[i, j]`. + // Note that the axis need to be in order so transpose is not a bcast operator. + kBroadcast = 1, + // Injective operator, can always injectively map output axis to a single input axis. + // All injective operator can still be safely fused to injective and reduction. + kInjective = 2, + // Communicative reduction operator. + kCommReduce = 3, + // Complex operation, can still fuse elemwise operations into its output. + // but cannot chain another complex op + kOutEWiseFusable = 4, + // The pattern for tuple nodes. Can fuse into subsequent injective ops, + // but treated specially + kTuple = 7, + // Opaque operation, cannot fuse anything. + kOpaque = 8 +}; + /*! * \brief Infer output struct info given the call * diff --git a/include/tvm/relax/type.h b/include/tvm/relax/type.h index 9c20a524353a..210730bec644 100644 --- a/include/tvm/relax/type.h +++ b/include/tvm/relax/type.h @@ -28,7 +28,6 @@ #include #include #include -#include #include #include diff --git a/include/tvm/relay/adt.h b/include/tvm/relay/adt.h deleted file mode 100644 index cdb8e52d2359..000000000000 --- a/include/tvm/relay/adt.h +++ /dev/null @@ -1,344 +0,0 @@ -/* - * 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 tvm/relay/adt.h - * \brief Algebraic data types for Relay - */ -#ifndef TVM_RELAY_ADT_H_ -#define TVM_RELAY_ADT_H_ - -#include -#include -#include -#include -#include - -#include -#include -#include - -namespace tvm { -namespace relay { - -using Constructor = tvm::Constructor; -using ConstructorNode = tvm::ConstructorNode; - -using TypeData = tvm::TypeData; -using TypeDataNode = tvm::TypeDataNode; - -/*! \brief Base type for declaring relay pattern. */ -class PatternNode : public RelayNode { - public: - static constexpr const char* _type_key = "relay.Pattern"; - static constexpr const bool _type_has_method_sequal_reduce = true; - static constexpr const bool _type_has_method_shash_reduce = true; - TVM_DECLARE_BASE_OBJECT_INFO(PatternNode, Object); -}; - -/*! - * \brief Pattern is the base type for an ADT match pattern in Relay. - * - * Given an ADT value, a pattern might accept it and bind the pattern variable to some value - * (typically a subnode of the input or the input). Otherwise, the pattern rejects the value. - * - * ADT pattern matching thus takes a list of values and binds to the first that accepts the value. - */ -class Pattern : public ObjectRef { - public: - Pattern() {} - explicit Pattern(ObjectPtr p) : ObjectRef(p) {} - - using ContainerType = PatternNode; -}; - -/*! \brief A wildcard pattern: Accepts all input and binds nothing. */ -class PatternWildcard; -/*! \brief PatternWildcard container node */ -class PatternWildcardNode : public PatternNode { - public: - void VisitAttrs(tvm::AttrVisitor* v) { v->Visit("span", &span); } - - bool SEqualReduce(const PatternNode* other, SEqualReducer equal) const { return true; } - - void SHashReduce(SHashReducer hash_reduce) const {} - - static constexpr const char* _type_key = "relay.PatternWildcard"; - TVM_DECLARE_FINAL_OBJECT_INFO(PatternWildcardNode, PatternNode); -}; - -class PatternWildcard : public Pattern { - public: - /* \brief Overload the default constructors. */ - TVM_DLL PatternWildcard(); - explicit PatternWildcard(ObjectPtr n) : Pattern(n) {} - /* \brief Copy constructor. */ - PatternWildcard(const PatternWildcard& pat) : PatternWildcard(pat.data_) {} - /* \brief Move constructor. */ - PatternWildcard(PatternWildcard&& pat) : PatternWildcard(std::move(pat.data_)) {} - /* \brief Copy assignment. */ - PatternWildcard& operator=(const PatternWildcard& other) { - (*this).data_ = other.data_; - return *this; - } - /* \brief Move assignment. */ - PatternWildcard& operator=(PatternWildcard&& other) { - (*this).data_ = std::move(other.data_); - return *this; - } - - const PatternWildcardNode* operator->() const { - return static_cast(get()); - } - - using ContainerType = PatternWildcardNode; -}; - -/*! \brief A var pattern. Accept all input and bind to a var. */ -class PatternVar; -/*! \brief PatternVar container node */ -class PatternVarNode : public PatternNode { - public: - /*! \brief Variable that stores the matched value. */ - tvm::relay::Var var; - - void VisitAttrs(tvm::AttrVisitor* v) { - v->Visit("var", &var); - v->Visit("span", &span); - } - - bool SEqualReduce(const PatternVarNode* other, SEqualReducer equal) const { - return equal.DefEqual(var, other->var); - } - - void SHashReduce(SHashReducer hash_reduce) const { hash_reduce.DefHash(var); } - - static constexpr const char* _type_key = "relay.PatternVar"; - TVM_DECLARE_FINAL_OBJECT_INFO(PatternVarNode, PatternNode); -}; - -class PatternVar : public Pattern { - public: - /*! - * \brief Constructor - * \param var The var to construct a pattern - */ - TVM_DLL explicit PatternVar(tvm::relay::Var var); - - TVM_DEFINE_OBJECT_REF_METHODS(PatternVar, Pattern, PatternVarNode); -}; - -/*! \brief A constructor pattern. Matches a value with the given constructor, binds recursively. */ -class PatternConstructor; -/*! \brief PatternVar container node */ -class PatternConstructorNode : public PatternNode { - public: - /*! Constructor matched by the pattern. */ - Constructor constructor; - /*! Sub-patterns to match against each input to the constructor. */ - tvm::Array patterns; - - void VisitAttrs(tvm::AttrVisitor* v) { - v->Visit("constructor", &constructor); - v->Visit("patterns", &patterns); - v->Visit("span", &span); - } - - bool SEqualReduce(const PatternConstructorNode* other, SEqualReducer equal) const { - return equal(constructor, other->constructor) && equal(patterns, other->patterns); - } - - void SHashReduce(SHashReducer hash_reduce) const { - hash_reduce(constructor); - hash_reduce(patterns); - } - - static constexpr const char* _type_key = "relay.PatternConstructor"; - TVM_DECLARE_FINAL_OBJECT_INFO(PatternConstructorNode, PatternNode); -}; - -class PatternConstructor : public Pattern { - public: - /*! - * \brief Constructor - * \param constructor The constructor of a pattern - * \param patterns The sub-patterns for matching - */ - TVM_DLL PatternConstructor(Constructor constructor, tvm::Array patterns); - - TVM_DEFINE_OBJECT_REF_METHODS(PatternConstructor, Pattern, PatternConstructorNode); -}; - -/*! \brief A tuple pattern. Matches a tuple, binds recursively. */ -class PatternTuple; -/*! \brief PatternVar container node */ -class PatternTupleNode : public PatternNode { - public: - /* TODO(@jroesch): rename to field_pats */ - /*! Sub-patterns to match against each value of the tuple. */ - tvm::Array patterns; - - void VisitAttrs(tvm::AttrVisitor* v) { - v->Visit("patterns", &patterns); - v->Visit("span", &span); - } - - bool SEqualReduce(const PatternTupleNode* other, SEqualReducer equal) const { - return equal(patterns, other->patterns); - } - - void SHashReduce(SHashReducer hash_reduce) const { hash_reduce(patterns); } - - static constexpr const char* _type_key = "relay.PatternTuple"; - TVM_DECLARE_FINAL_OBJECT_INFO(PatternTupleNode, PatternNode); -}; - -class PatternTuple : public Pattern { - public: - /*! - * \brief Constructor - * \param patterns The sub-patterns to match against each value of the tuple - */ - TVM_DLL explicit PatternTuple(tvm::Array patterns); - - TVM_DEFINE_OBJECT_REF_METHODS(PatternTuple, Pattern, PatternTupleNode); -}; - -/*! \brief A clause in a match expression. */ -class Clause; -/*! \brief Clause container node. */ -class ClauseNode : public Object { - public: - /*! \brief The pattern the clause matches. */ - Pattern lhs; - /*! \brief The resulting value. */ - Expr rhs; - - void VisitAttrs(tvm::AttrVisitor* v) { - v->Visit("lhs", &lhs); - v->Visit("rhs", &rhs); - } - - bool SEqualReduce(const ClauseNode* other, SEqualReducer equal) const { - return equal(lhs, other->lhs) && equal(rhs, other->rhs); - } - - void SHashReduce(SHashReducer hash_reduce) const { - hash_reduce(lhs); - hash_reduce(rhs); - } - - static constexpr const char* _type_key = "relay.Clause"; - static constexpr const bool _type_has_method_sequal_reduce = true; - static constexpr const bool _type_has_method_shash_reduce = true; - TVM_DECLARE_FINAL_OBJECT_INFO(ClauseNode, Object); -}; - -class Clause : public ObjectRef { - public: - /*! - * \brief Constructor - * \param lhs The pattern matched by the clause. - * \param rhs The resulting value - */ - TVM_DLL explicit Clause(Pattern lhs, Expr rhs); - - TVM_DEFINE_OBJECT_REF_METHODS(Clause, ObjectRef, ClauseNode); - TVM_DEFINE_OBJECT_REF_COW_METHOD(ClauseNode); -}; - -/*! - * \brief Returns \p clause with the given properties. A null property denotes 'no change'. - * Returns \p clause if all properties are unchanged. Otherwise, returns a copy with the new - * fields. - */ -Clause WithFields(Clause clause, Optional opt_lhs = Optional(), - Optional opt_rhs = Optional()); - -/*! \brief ADT pattern matching exression. */ -class Match; -/*! \brief Match container node. */ -class MatchNode : public ExprNode { - public: - /*! \brief The input being deconstructed. */ - Expr data; - - /*! \brief The match node clauses. */ - tvm::Array clauses; - - /*! \brief Should this match be complete (cover all cases)? - * If yes, the type checker will generate an error if there are any missing cases. - */ - bool complete; - - void VisitAttrs(tvm::AttrVisitor* v) { - v->Visit("data", &data); - v->Visit("clauses", &clauses); - v->Visit("complete", &complete); - v->Visit("virtual_device_", &virtual_device_); - v->Visit("span", &span); - v->Visit("_checked_type_", &checked_type_); - } - - bool SEqualReduce(const MatchNode* other, SEqualReducer equal) const { - equal->MarkGraphNode(); - return equal(data, other->data) && equal(clauses, other->clauses) && - equal(complete, other->complete); - } - - void SHashReduce(SHashReducer hash_reduce) const { - hash_reduce->MarkGraphNode(); - hash_reduce(data); - hash_reduce(clauses); - hash_reduce(complete); - } - - static constexpr const char* _type_key = "relay.Match"; - TVM_DECLARE_FINAL_OBJECT_INFO(MatchNode, ExprNode); -}; - -class Match : public Expr { - public: - /*! - * \brief Constructor - * \param data the input being deconstructed. - * \param clauses The clauses for matching. - * \param complete Indicate if this match is complete. - * \param span The span of the expression. - */ - TVM_DLL Match(Expr data, tvm::Array clauses, bool complete = true, Span span = Span()); - - TVM_DEFINE_OBJECT_REF_METHODS(Match, RelayExpr, MatchNode); - TVM_DEFINE_OBJECT_REF_COW_METHOD(MatchNode); -}; - -/*! - * \brief Returns \p match with the given properties. A null property denotes 'no change'. - * Returns \p match if all properties are unchanged. Otherwise, returns a copy with the new - * fields. - */ -Match WithFields(Match match, Optional opt_data = Optional(), - Optional> opt_clauses = Optional>(), - Optional opt_complete = Optional(), - Optional opt_span = Optional()); - -} // namespace relay -} // namespace tvm - -#endif // TVM_RELAY_ADT_H_ diff --git a/include/tvm/relay/analysis.h b/include/tvm/relay/analysis.h deleted file mode 100644 index 0f85587262ac..000000000000 --- a/include/tvm/relay/analysis.h +++ /dev/null @@ -1,256 +0,0 @@ -/* - * 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 tvm/relay/analysis.h - * \brief The set of Relay analysis passes written in C++. - */ -#ifndef TVM_RELAY_ANALYSIS_H_ -#define TVM_RELAY_ANALYSIS_H_ - -#include -#include -#include -#include -#include -#include - -#include -#include - -namespace tvm { -namespace relay { - -/*! - * \brief Check that types are well kinded by applying "kinding rules". - * - * This pass ensures we do not do things that violate the design of the - * type system when writing down types. - * - * For example tensors are not allowed to contain functions in Relay. - * - * We check this by ensuring the `dtype` field of a Tensor always contains - * a data type such as `int`, `float`, `uint`. - * - * \param t The type to check. - * \param mod The global module. - * \param diag_ctx The Diagnostic context. - * - * \return The kind of the passed type. - */ -TVM_DLL Kind KindCheck(const Type& t, const IRModule& mod, - Optional diag_ctx = Optional()); - -/*! - * \brief Check whether an expression is constant. - * - * If the inputs of an expression are all constant, it means the expression - * itself is constant also. - * - * \param e the expression. - * - * \return whether the expression is constant. - */ -TVM_DLL bool ConstantCheck(const Expr& e); - -/*! - * \brief Check whether an expression is in the basic block normal form. - * - * \param e the expression. - * - * \return whether the expression is in the basic block normal form. - */ -TVM_DLL bool BasicBlockNormalFormCheck(const Expr& e); - -/*! - * \brief Check that each Var is only bound once. - * - * For example, the expression `let x = 1 in let x = 2 in 3` bound x twice. - * - * `let f = (x -> x) in let g = (x -> x + 1) in f(g(2))` also bound x twice, - * although x is not shadowed. - * - * \param expr the expression to check. - * \param diag_ctx the diagnostic context - * - * \return true iff all Var in expr is bound at most once. - */ -TVM_DLL bool WellFormed(const Expr& expr, - Optional diag_ctx = Optional()); - -/*! - * \brief Get all bound variables from expression expr. - * - * Bound variables are all variables that are declared in the expr. - * They only have meaning inside that expr, and can only be used in it. - * - * \param expr the expression. - * - * \return List of bound vars, in the PostDFS order in the expression. - */ -TVM_DLL tvm::Array BoundVars(const Expr& expr); - -/*! - * \brief Get all bound variables from pattern pat. - * - * Bound variables are all variables that got bound by the pat. - * They only have meaning inside that expr, and can only be used in it. - * - * \param pat the Pattern. - * - * \return List of bound vars, in the PostDFS order in the expression. - */ -TVM_DLL tvm::Array BoundVars(const Pattern& pat); - -/*! - * \brief Get free type parameters from expression expr. - * - * Free variables are variables that are not bound by a - * let or a function parameter in the context. - * - * \param expr the expression. - * - * \return List of free vars, in the PostDFS order in the expression. - */ -TVM_DLL tvm::Array FreeVars(const Expr& expr); - -/*! - * \brief Get all variables from expression expr. - * - * \param expr the expression. - * - * \return List of all vars, in the PostDFS order in the expression. - */ -TVM_DLL tvm::Array AllVars(const Expr& expr); - -/*! - * \brief Get free TypeVars from expression expr. - * - * Free type parameters are type parameters that are not bound by a function - * type in the context. - * - * \param expr the expression. - * \param mod the module. - * - * \return List of free vars, in the PostDFS order visited by expr. - */ -TVM_DLL tvm::Array FreeTypeVars(const Expr& expr, const IRModule& mod); - -/*! - * \brief Get free TypeVars from type t. - * - * Free type parameters are type parameters that are not bound by a function - * type in the context. - * - * \param t the type. - * \param mod the module. - * - * \return List of free type vars, in the PostDFS order visited by type. - */ -TVM_DLL tvm::Array FreeTypeVars(const Type& t, const IRModule& mod); - -/*! - * \brief Get all bound type variables from expression expr. - * - * Bound variables are all type variables that are declared in the expr. - * They only have meaning inside that expr, and can only be used in it. - * - * \param expr the expression. - * \param mod the module. - * - * \return List of bound type vars, in the PostDFS order in the expression. - */ -TVM_DLL tvm::Array BoundTypeVars(const Expr& expr, const IRModule& mod); - -/*! - * \brief Get all bound type variables from type t. - * - * Bound variables are all type variables that are declared in the type. - * They only have meaning inside that type, and can only be used in it. - * - * \param t the type - * \param mod the module. - * - * \return List of bound type vars, in the PostDFS order visited by type. - */ -TVM_DLL tvm::Array BoundTypeVars(const Type& t, const IRModule& mod); - -/*! - * \brief Get all type variables in expression expr. - * - * \param expr the expression. - * \param mod the module. - * - * \return List of type vars, in the PostDFS order in the expression. - */ -TVM_DLL tvm::Array AllTypeVars(const Expr& expr, const IRModule& mod); - -/*! - * \brief Get all type variables in type t. - * - * \param t the type. - * \param mod the module. - * - * \return List of type vars, in the PostDFS order visited by type. - */ -TVM_DLL tvm::Array AllTypeVars(const Type& t, const IRModule& mod); - -/*! - * \brief Finds cases that the given match expression does not catch, if any. - * - * \param match the match expression to test - * - * \param mod The module used for accessing global type var definitions, can be None. - * - * \return Returns a list of cases (as patterns) that are not handled by the match - * expression. - */ -TVM_DLL Array UnmatchedCases(const Match& match, const IRModule& mod); - -/*! - * \brief Get reference counter of each internal ExprNode in body. - * - * \param body The body expression. - * - * \return The reference count mapping. - */ -TVM_DLL std::unordered_map GetExprRefCount(const Expr& body); - -/*! - * \brief Get the updated module for collecting calibration data. - * - * \param mod The module to be updated. - * - * \return The updated module. - */ -TVM_DLL IRModule GetCalibrateModule(IRModule mod); - -/*! - * \brief Get the output map between subgrpahs and its inputs/output. - * - * \param mod The module for running calibration. - * - * \return The mapping between a subgraph name and its postition in the output tuple. - */ -TVM_DLL Map> GetCalibrateOutputMap(const IRModule& mod); - -} // namespace relay -} // namespace tvm - -#endif // TVM_RELAY_ANALYSIS_H_ diff --git a/include/tvm/relay/attrs/algorithm.h b/include/tvm/relay/attrs/algorithm.h deleted file mode 100644 index 3652a09e9168..000000000000 --- a/include/tvm/relay/attrs/algorithm.h +++ /dev/null @@ -1,97 +0,0 @@ -/* - * 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 tvm/relay/attrs/vision.h - * \brief Auxiliary attributes for vision operators. - */ -#ifndef TVM_RELAY_ATTRS_ALGORITHM_H_ -#define TVM_RELAY_ATTRS_ALGORITHM_H_ - -#include -#include -#include - -#include - -namespace tvm { -namespace relay { - -/*! \brief Attributes used in argsort operators */ -struct ArgsortAttrs : public tvm::AttrsNode { - int axis; - bool is_ascend; - DataType dtype; - - TVM_DECLARE_ATTRS(ArgsortAttrs, "relay.attrs.ArgsortAttrs") { - TVM_ATTR_FIELD(axis).set_default(-1).describe( - "Axis along which to sort the input tensor." - "If not given, the flattened array is used."); - TVM_ATTR_FIELD(is_ascend).set_default(true).describe( - "Whether to sort in ascending or descending order." - "By default, sort in ascending order"); - TVM_ATTR_FIELD(dtype) - .set_default(NullValue()) - .describe("DType of the output indices."); - } -}; - -struct TopKAttrs : public tvm::AttrsNode { - Optional k; - int axis; - bool is_ascend; - std::string ret_type; - DataType dtype; - - TVM_DECLARE_ATTRS(TopKAttrs, "relay.attrs.TopkAttrs") { - TVM_ATTR_FIELD(k).describe("Number of top elements to select"); - TVM_ATTR_FIELD(axis).set_default(-1).describe("Axis along which to sort the input tensor."); - TVM_ATTR_FIELD(ret_type).set_default("both").describe( - "The return type [both, values, indices]." - "both - return both top k data and indices." - "values - return top k data only." - "indices - return top k indices only."); - TVM_ATTR_FIELD(is_ascend).set_default(false).describe( - "Whether to sort in ascending or descending order." - "By default, sort in descending order"); - TVM_ATTR_FIELD(dtype) - .set_default(NullValue()) - .describe("Data type of the output indices."); - } -}; - -struct SearchSortedAttrs : public tvm::AttrsNode { - bool right; - DataType dtype; - - TVM_DECLARE_ATTRS(SearchSortedAttrs, "relay.attrs.SearchSortedAttrs") { - TVM_ATTR_FIELD(right).set_default(false).describe( - "Controls which index is returned if a value lands exactly on one of sorted values. If " - " false, the index of the first suitable location found is given. If true, return the " - "last such index. If there is no suitable index, return either 0 or N (where N is the " - "size of the innermost dimension)."); - TVM_ATTR_FIELD(dtype) - .set_default(DataType::Int(32)) - .describe("Data type of the output indices."); - } -}; - -} // namespace relay -} // namespace tvm -#endif // TVM_RELAY_ATTRS_ALGORITHM_H_ diff --git a/include/tvm/relay/attrs/annotation.h b/include/tvm/relay/attrs/annotation.h deleted file mode 100644 index 1066416838b5..000000000000 --- a/include/tvm/relay/attrs/annotation.h +++ /dev/null @@ -1,59 +0,0 @@ -/* - * 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 tvm/relay/attrs/annotation.h - * \brief Attribute for annotation operators. - */ -#ifndef TVM_RELAY_ATTRS_ANNOTATION_H_ -#define TVM_RELAY_ATTRS_ANNOTATION_H_ - -#include - -#include - -namespace tvm { -namespace relay { - -/*! - * \brief Annotate an expression to be cast into specific data type. - */ -struct CastHintAttrs : public tvm::AttrsNode { - DataType dtype; - - TVM_DECLARE_ATTRS(CastHintAttrs, "relay.attrs.CastHintAttrs") { - TVM_ATTR_FIELD(dtype).describe("The data type denoted to be cast."); - } -}; - -/*! - * \brief Options for the operators used to annotate a compiler. - */ -struct CompilerAttrs : public tvm::AttrsNode { - /*! \brief A 3rd party compiler for code generation. */ - std::string compiler; - - TVM_DECLARE_ATTRS(CompilerAttrs, "relay.attrs.CompilerAttrs") { - TVM_ATTR_FIELD(compiler).describe("A 3rd party compiler used for code generation."); - } -}; - -} // namespace relay -} // namespace tvm -#endif // TVM_RELAY_ATTRS_ANNOTATION_H_ diff --git a/include/tvm/relay/attrs/bitserial.h b/include/tvm/relay/attrs/bitserial.h deleted file mode 100644 index ed04c59ec865..000000000000 --- a/include/tvm/relay/attrs/bitserial.h +++ /dev/null @@ -1,133 +0,0 @@ -/* - * 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 tvm/relay/attrs/bitserial.h - * \brief Auxiliary attributes for bitserial operators. - */ - -#ifndef TVM_RELAY_ATTRS_BITSERIAL_H_ -#define TVM_RELAY_ATTRS_BITSERIAL_H_ - -#include -#include - -#include - -namespace tvm { -namespace relay { - -/*! \brief Attributes used in bitpack operators */ -struct BitPackAttrs : public tvm::AttrsNode { - int bits; - int pack_axis; - int bit_axis; - DataType pack_type; - std::string name; - - TVM_DECLARE_ATTRS(BitPackAttrs, "relay.attrs.BitPackAttrs") { - TVM_ATTR_FIELD(bits).set_default(1).describe("Number of bits to quantize with."); - TVM_ATTR_FIELD(pack_axis).set_default(1).describe( - "Axis that should be compressed, typically channels."); - TVM_ATTR_FIELD(bit_axis).set_default(-1).describe("New axis for packed bits."); - TVM_ATTR_FIELD(pack_type) - .set_default(NullValue()) - .describe("Type of int to pack bits into."); - TVM_ATTR_FIELD(name).set_default("BitPack").describe("Name of operation."); - } -}; - -/*! \brief Attribues used in bitserial convolution operators */ -struct BinaryConv2DAttrs : public tvm::AttrsNode { - Array strides; - Array padding; - IndexExpr channels; - Array kernel_size; - int activation_bits; - int weight_bits; - std::string data_layout; - std::string kernel_layout; - DataType pack_dtype; - DataType out_dtype; - bool unipolar; - - TVM_DECLARE_ATTRS(BinaryConv2DAttrs, "relay.attrs.BinaryConv2DAttrs") { - TVM_ATTR_FIELD(strides) - .set_default(Array({1, 1})) - .describe("Specifies the strides of the convolution."); - TVM_ATTR_FIELD(padding) - .set_default(Array({0, 0})) - .describe( - "If padding is non-zero the input is implicitly zero-padded" - "on both sides for padding number of points."); - TVM_ATTR_FIELD(kernel_size) - .set_default(Array({3, 3})) - .describe("Specifies the dimensions of the convolution window."); - TVM_ATTR_FIELD(channels) - .set_default(NullValue()) - .describe("Number of output channels, needed for shape inference."); - TVM_ATTR_FIELD(activation_bits) - .set_default(1) - .describe("Number of bits activation should be packed with."); - TVM_ATTR_FIELD(weight_bits) - .set_default(1) - .describe("Number of bits kernel should be packed with."); - TVM_ATTR_FIELD(data_layout) - .set_default("NCHW") - .describe("Dimension ordering of input data, can be 'NCHW' or NHWC'."); - TVM_ATTR_FIELD(kernel_layout) - .set_default("OIHW") - .describe("Dimension ordering of kernel data, can be 'OIHW' or HWIO'."); - TVM_ATTR_FIELD(pack_dtype) - .set_default(NullValue()) - .describe("Datatype to pack bits into."); - TVM_ATTR_FIELD(out_dtype).set_default(NullValue()).describe("Output datatype."); - TVM_ATTR_FIELD(unipolar).set_default(true).describe( - "Whether to use unipolar or bipolar quantization."); - } -}; - -/*~ \brief Attributes for bitserial dense operator */ -struct BinaryDenseAttrs : public tvm::AttrsNode { - IndexExpr units; - int data_bits; - int weight_bits; - DataType pack_dtype; - DataType out_dtype; - bool unipolar; - - TVM_DECLARE_ATTRS(BinaryDenseAttrs, "relay.attrs.BinaryDenseAttrs") { - TVM_ATTR_FIELD(units).describe("Number of hidden units of the dense transformation."); - TVM_ATTR_FIELD(data_bits).set_default(1).describe( - "Number of bits to pack for incoming tensor."); - TVM_ATTR_FIELD(weight_bits) - .set_default(1) - .describe("Number of bits to pack for weight tensor."); - TVM_ATTR_FIELD(pack_dtype) - .set_default(NullValue()) - .describe("Datatype to pack bits into before computation."); - TVM_ATTR_FIELD(out_dtype).set_default(NullValue()).describe("Output data type."); - TVM_ATTR_FIELD(unipolar).set_default(true).describe( - "Whether to use unipolar or bipolar quantization for inputs."); - } -}; - -} // namespace relay -} // namespace tvm -#endif // TVM_RELAY_ATTRS_BITSERIAL_H_ diff --git a/include/tvm/relay/attrs/call.h b/include/tvm/relay/attrs/call.h deleted file mode 100644 index e0b347de1783..000000000000 --- a/include/tvm/relay/attrs/call.h +++ /dev/null @@ -1,50 +0,0 @@ -/* - * 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 tvm/relay/attrs/call.h - * \brief Attribute for call_lowered operator. - */ -#ifndef TVM_RELAY_ATTRS_CALL_H_ -#define TVM_RELAY_ATTRS_CALL_H_ - -#include - -#include - -namespace tvm { -namespace relay { - -/*! - * \brief Metadata for calls to TIR functions, useful for program analysis crossing Relay and TIR. - */ -struct CallLoweredAttrs : public tvm::AttrsNode { - /*! \brief Additional metadata attached to the call node. Should be replaced by explict fields. */ - Map metadata; - - TVM_DECLARE_ATTRS(CallLoweredAttrs, "relay.attrs.CallLoweredAttrs") { - TVM_ATTR_FIELD(metadata) - .describe("Metadata attached to the lowered function call.") - .set_default(Map()); - } -}; - -} // namespace relay -} // namespace tvm -#endif // TVM_RELAY_ATTRS_CALL_H_ diff --git a/include/tvm/relay/attrs/debug.h b/include/tvm/relay/attrs/debug.h deleted file mode 100644 index 112228bb41ee..000000000000 --- a/include/tvm/relay/attrs/debug.h +++ /dev/null @@ -1,48 +0,0 @@ -/* - * 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 tvm/relay/attrs/debug.h - * \brief Auxiliary attributes for debug operators. - */ -#ifndef TVM_RELAY_ATTRS_DEBUG_H_ -#define TVM_RELAY_ATTRS_DEBUG_H_ - -#include -#include - -#include - -namespace tvm { -namespace relay { - -/*! - * \brief Options for the debug operators. - */ -struct DebugAttrs : public tvm::AttrsNode { - EnvFunc debug_func; - - TVM_DECLARE_ATTRS(DebugAttrs, "relay.attrs.DebugAttrs") { - TVM_ATTR_FIELD(debug_func).describe("The function to use when debugging."); - } -}; - -} // namespace relay -} // namespace tvm -#endif // TVM_RELAY_ATTRS_DEBUG_H_ diff --git a/include/tvm/relay/attrs/device_copy.h b/include/tvm/relay/attrs/device_copy.h deleted file mode 100644 index fe0534a8a2b4..000000000000 --- a/include/tvm/relay/attrs/device_copy.h +++ /dev/null @@ -1,52 +0,0 @@ -/* - * 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 tvm/relay/attrs/device_copy.h - * \brief Attribute for the device copy operator. - */ -#ifndef TVM_RELAY_ATTRS_DEVICE_COPY_H_ -#define TVM_RELAY_ATTRS_DEVICE_COPY_H_ - -#include -#include - -#include - -namespace tvm { -namespace relay { - -/*! - * \brief Options for the device copy operators. - */ -struct DeviceCopyAttrs : public tvm::AttrsNode { - VirtualDevice src_virtual_device = VirtualDevice::FullyUnconstrained(); - VirtualDevice dst_virtual_device = VirtualDevice::FullyUnconstrained(); - - TVM_DECLARE_ATTRS(DeviceCopyAttrs, "relay.attrs.DeviceCopyAttrs") { - TVM_ATTR_FIELD(src_virtual_device) - .describe("The (virtual) device and scope where the op copies data from."); - TVM_ATTR_FIELD(dst_virtual_device) - .describe("The (virtual) device and scope where the op copies data to."); - } -}; - -} // namespace relay -} // namespace tvm -#endif // TVM_RELAY_ATTRS_DEVICE_COPY_H_ diff --git a/include/tvm/relay/attrs/image.h b/include/tvm/relay/attrs/image.h deleted file mode 100644 index 43510ea68501..000000000000 --- a/include/tvm/relay/attrs/image.h +++ /dev/null @@ -1,322 +0,0 @@ -/* - * 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 tvm/relay/attrs/image.h - * \brief Auxiliary attributes for image operators. - */ -#ifndef TVM_RELAY_ATTRS_IMAGE_H_ -#define TVM_RELAY_ATTRS_IMAGE_H_ - -#include -#include - -#include - -namespace tvm { -namespace relay { - -/*! \brief Attributes used in image resize1d operator */ -struct Resize1DAttrs : public tvm::AttrsNode { - Array size; - Array roi; - std::string layout; - std::string method; - std::string coordinate_transformation_mode; - std::string rounding_method; - double cubic_alpha; - int cubic_exclude; - double extrapolation_value; - DataType out_dtype; - - TVM_DECLARE_ATTRS(Resize1DAttrs, "relay.attrs.Resize1DAttrs") { - TVM_ATTR_FIELD(size).set_default(NullValue>()).describe("Output Size."); - TVM_ATTR_FIELD(roi) - .set_default(NullValue>()) - .describe("Region of Interest for coordinate transformation mode 'tf_crop_and_resize'"); - TVM_ATTR_FIELD(layout).set_default("NCW").describe( - "Dimension ordering of input data. Can be 'NCW', 'NWC', etc." - "'N', 'C', 'W' stands for batch, channel and width" - "dimensions respectively. Resize is applied on the" - "'W' dimension."); - TVM_ATTR_FIELD(method).set_default("linear").describe( - "Specify the mode to use for scaling." - "nearest_neighbor - Nearest Neighbor" - "linear - Linear Interpolation" - "cubic - Cubic Interpolation"); - TVM_ATTR_FIELD(coordinate_transformation_mode) - .set_default("half_pixel") - .describe( - "Describes how to transform the coordinate in the resized tensor" - "to the coordinate in the original tensor." - "Refer to the ONNX Resize operator specification for details" - "Available options are half_pixel, align_corners and asymmetric"); - TVM_ATTR_FIELD(rounding_method) - .set_default("round") - .describe( - "indicates how to find the \"nearest\" pixel in nearest_neighbor method" - "Available options are round, floor, and ceil."); - TVM_ATTR_FIELD(cubic_alpha) - .set_default(-0.5) - .describe("Spline Coefficient for cubic interpolation"); - TVM_ATTR_FIELD(cubic_exclude) - .set_default(0) - .describe("Flag to exclude exterior of the image during cubic interpolation"); - TVM_ATTR_FIELD(extrapolation_value) - .set_default(0.0) - .describe("Value to return when roi is outside of the image"); - TVM_ATTR_FIELD(out_dtype).set_default(NullValue()).describe("Output data type."); - } -}; - -/*! \brief Attributes used in image resize2d operator */ -struct Resize2DAttrs : public tvm::AttrsNode { - Array size; - Array roi; - std::string layout; - std::string method; - std::string coordinate_transformation_mode; - std::string rounding_method; - double cubic_alpha; - int cubic_exclude; - double extrapolation_value; - DataType out_dtype; - - TVM_DECLARE_ATTRS(Resize2DAttrs, "relay.attrs.Resize2DAttrs") { - TVM_ATTR_FIELD(size).set_default(NullValue>()).describe("Output Size."); - TVM_ATTR_FIELD(roi) - .set_default(NullValue>()) - .describe("Region of Interest for coordinate transformation mode 'tf_crop_and_resize'"); - TVM_ATTR_FIELD(layout).set_default("NCHW").describe( - "Dimension ordering of input data. Can be 'NCHW', 'NHWC', etc." - "'N', 'C', 'H', 'W' stands for batch, channel, height, and width" - "dimensions respectively. Resize is applied on the 'H' and" - "'W' dimensions."); - TVM_ATTR_FIELD(method).set_default("linear").describe( - "Specify the mode to use for scaling." - "nearest_neighbor - Nearest Neighbor" - "linear - Bilinear Interpolation" - "cubic - Bicubic Interpolation"); - TVM_ATTR_FIELD(coordinate_transformation_mode) - .set_default("half_pixel") - .describe( - "Describes how to transform the coordinate in the resized tensor" - "to the coordinate in the original tensor." - "Refer to the ONNX Resize operator specification for details" - "Available options are half_pixel, align_corners and asymmetric"); - TVM_ATTR_FIELD(rounding_method) - .set_default("round") - .describe( - "indicates how to find the \"nearest\" pixel in nearest_neighbor method" - "Available options are round, floor, and ceil."); - TVM_ATTR_FIELD(cubic_alpha) - .set_default(-0.5) - .describe("Spline Coefficient for Bicubic Interpolation"); - TVM_ATTR_FIELD(cubic_exclude) - .set_default(0) - .describe("Flag to exclude exterior of the image during bicubic interpolation"); - TVM_ATTR_FIELD(extrapolation_value) - .set_default(0.0) - .describe("Value to return when roi is outside of the image"); - TVM_ATTR_FIELD(out_dtype).set_default(NullValue()).describe("Output data type."); - } -}; - -/*! \brief Attributes used in image resize3d operator */ -struct Resize3DAttrs : public tvm::AttrsNode { - Array size; - Array roi; - std::string layout; - std::string method; - std::string coordinate_transformation_mode; - std::string rounding_method; - double cubic_alpha; - int cubic_exclude; - double extrapolation_value; - DataType out_dtype; - - TVM_DECLARE_ATTRS(Resize3DAttrs, "relay.attrs.Resize3DAttrs") { - TVM_ATTR_FIELD(size).set_default(NullValue>()).describe("Output Size."); - TVM_ATTR_FIELD(roi) - .set_default(NullValue>()) - .describe("Region of Interest for coordinate transformation mode 'tf_crop_and_resize'"); - TVM_ATTR_FIELD(layout).set_default("NCDHW").describe( - "Dimension ordering of input data. Can be 'NCDHW', 'NDHWC', etc." - "'N', 'C', 'D', 'H', 'W' stands for batch, channel, depth, height, and width" - "dimensions respectively. Resize3d is applied on the 'D', 'H' and" - "'W' dimensions."); - TVM_ATTR_FIELD(method).set_default("linear").describe( - "Specify the mode to use for scaling." - "nearest_neighbor - Nearest Neighbor" - "linear - Trilinear Interpolation" - "cubic - Tricubic Interpolation"); - TVM_ATTR_FIELD(coordinate_transformation_mode) - .set_default("half_pixel") - .describe( - "Describes how to transform the coordinate in the resized tensor" - "to the coordinate in the original tensor." - "Refer to the ONNX Resize operator specification for details" - "Available options are half_pixel, align_corners and asymmetric"); - TVM_ATTR_FIELD(rounding_method) - .set_default("round") - .describe( - "indicates how to find the \"nearest\" pixel in nearest_neighbor method" - "Available options are round, floor, and ceil."); - TVM_ATTR_FIELD(cubic_alpha) - .set_default(-0.5) - .describe("Spline Coefficient for Tricubic Interpolation"); - TVM_ATTR_FIELD(cubic_exclude) - .set_default(0) - .describe("Flag to exclude exterior of the image during tricubic interpolation"); - TVM_ATTR_FIELD(extrapolation_value) - .set_default(0.0) - .describe("Value to return when roi is outside of the image"); - TVM_ATTR_FIELD(out_dtype).set_default(NullValue()).describe("Output data type."); - } -}; - -/*! \brief Attributes used in image crop_and_resize operator */ -struct CropAndResizeAttrs : public tvm::AttrsNode { - Array crop_size; - std::string layout; - std::string method; - double extrapolation_value; - DataType out_dtype; - - TVM_DECLARE_ATTRS(CropAndResizeAttrs, "relay.attrs.CropAndResizeAttrs") { - TVM_ATTR_FIELD(crop_size).set_default(NullValue>()).describe("Target Size."); - TVM_ATTR_FIELD(layout).set_default("NCHW").describe( - "Dimension ordering of input data. Can be 'NCHW', 'NHWC', etc." - "'N', 'C', 'H', 'W' stands for batch, channel, height, and width" - "dimensions respectively. Resize is applied on the 'H' and" - "'W' dimensions."); - TVM_ATTR_FIELD(method) - .set_default("bilinear") - .describe( - "Specify the mode to use for scaling." - "nearest_neighbor - Nearest Neighbor" - "bilinear - Bilinear Interpolation"); - TVM_ATTR_FIELD(extrapolation_value) - .set_default(0.0) - .describe("Specify value for extrapolation."); - TVM_ATTR_FIELD(out_dtype).set_default(NullValue()).describe("Output data type."); - } -}; - -/*! \brief Attributes used in dilation operators */ -struct Dilation2DAttrs : public tvm::AttrsNode { - Array strides; - Array padding; - Array dilations; - std::string data_layout; - std::string kernel_layout; - DataType out_dtype; - - TVM_DECLARE_ATTRS(Dilation2DAttrs, "relay.attrs.Dilation2DAttrs") { - TVM_ATTR_FIELD(strides) - .set_default(Array({1, 1})) - .describe("Specifies the strides of the sliding window. [stride_height, stride_width]."); - TVM_ATTR_FIELD(padding) - .set_default(Array({0, 0})) - .describe( - "If padding is non-zero, then the input is implicitly zero-padded" - "Padding support both symmetric and asymmetric as" - "one int : same padding used on all sides" - "two int : bottom, right will use same padding as top, left" - "four int : padding width in the order of (top, left, bottom, right)"); - TVM_ATTR_FIELD(dilations) - .set_default(Array({1, 1})) - .describe("Specifies the dilation rate to use. [dilation_height, dilation_width]"); - TVM_ATTR_FIELD(data_layout) - .set_default("NCHW") - .describe( - "Dimension ordering of input data. Can be 'NCHW', 'NHWC', etc." - "'N', 'C', 'H', 'W' stands for batch, channel, height, and width" - "dimensions respectively. Convolution is applied on the 'H' and" - "'W' dimensions."); - TVM_ATTR_FIELD(kernel_layout) - .set_default("IHW") - .describe( - "Dimension ordering of weight. Can be 'IHW', 'HWI', etc." - "'I', 'H', 'W' stands for input_channel, height, and width" - "dimensions respectively."); - TVM_ATTR_FIELD(out_dtype) - .set_default(NullValue()) - .describe("Output data type, set to explicit type under mixed precision setting"); - } -}; - -/*! \brief Attributes used in image affine_grid operator */ -struct AffineGridAttrs : public tvm::AttrsNode { - Array target_shape; - - TVM_DECLARE_ATTRS(AffineGridAttrs, "relay.attrs.AffineGridAttrs") { - TVM_ATTR_FIELD(target_shape).describe("Specifies the output shape (H, W)."); - } -}; - -/*! \brief Attributes used in image grid_sample operator */ -struct GridSampleAttrs : public tvm::AttrsNode { - String method; - String layout; - String padding_mode; - bool align_corners; - - TVM_DECLARE_ATTRS(GridSampleAttrs, "relay.attrs.GridSampleAttrs") { - TVM_ATTR_FIELD(method) - .set_default("bilinear") - .describe( - "Specify the mode to use for scaling." - "nearest - 2D or 3D Nearest Interpolation." - "bilinear - '2D Bilinear' or '3D Trilinear' Interpolation." - "bicubic - 2D Bicubic Interpolation."); - TVM_ATTR_FIELD(layout).set_default("NCHW").describe( - "Dimension ordering of input data. Can be 'NCHW', 'NCDHW', etc." - "'N', 'C', 'D', 'H', 'W' stands for batch, channel, depth, height, and width" - "dimensions respectively." - "2D Resize is applied on the 'H' and 'W' dimensions." - "3D Resize is applied on the 'D' and 'H' and 'W' dimensions."); - TVM_ATTR_FIELD(padding_mode) - .set_default("zeros") - .describe( - "If :attr:'grid' has values outside the range of '[-1, 1]', the corresponding" - "outputs are handled as defined by padding_mode. Options are" - "padding_mode='zeros': use '0' for out-of-bound grid locations," - "padding_mode='border': use border values for out-of-bound grid locations" - "padding_mode='reflection': use values at locations reflected by" - "the border for out-of-bound grid locations. For location far away" - "from the border, it will keep being reflected until becoming in bound," - "e.g., (normalized) pixel location 'x = -3.5' reflects by border '-1'" - "and becomes 'x' = 1.5, then reflects by border '1' and becomes" - "'x' = -0.5"); - TVM_ATTR_FIELD(align_corners) - .set_default(true) - .describe( - "Geometrically, we consider the pixels of the" - "input as squares rather than points." - "If set to True, the extrema (-1 and 1) are considered as referring" - "to the center points of the input's corner pixels. If set to False, they" - "are instead considered as referring to the corner points of the input's corner" - "pixels, making the sampling more resolution agnostic."); - } -}; - -} // namespace relay -} // namespace tvm -#endif // TVM_RELAY_ATTRS_IMAGE_H_ diff --git a/include/tvm/relay/attrs/memory.h b/include/tvm/relay/attrs/memory.h deleted file mode 100644 index 07d6cc7e271e..000000000000 --- a/include/tvm/relay/attrs/memory.h +++ /dev/null @@ -1,78 +0,0 @@ -/* - * 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 tvm/relay/attrs/memory.h - * \brief Attributes for memory operators. - */ -#ifndef TVM_RELAY_ATTRS_MEMORY_H_ -#define TVM_RELAY_ATTRS_MEMORY_H_ - -#include -#include -#include - -#include -#include - -namespace tvm { -namespace relay { - -std::vector FlattenTupleType(const Type& type); -std::vector FromTupleType(const Type& type, const Expr& expr); -Expr ToTupleType(const Type& t, const std::vector& exprs); - -/*! - * \brief Options for allocating storage. - */ -struct AllocStorageAttrs : public tvm::AttrsNode { - DataType dtype; - VirtualDevice virtual_device = VirtualDevice::FullyUnconstrained(); - - TVM_DECLARE_ATTRS(AllocStorageAttrs, "relay.attrs.AllocStorageAttrs") { - TVM_ATTR_FIELD(dtype) - .describe("The dtype of the tensor to allocate.") - .set_default(DataType::Float(32, 1)); - TVM_ATTR_FIELD(virtual_device).describe("The virtual device on which to allocate memory."); - } -}; - -/*! - * \brief Options for allocating tensors. - */ -struct AllocTensorAttrs : public tvm::AttrsNode { - Constant const_shape; - Array assert_shape; - DataType dtype; - - TVM_DECLARE_ATTRS(AllocTensorAttrs, "relay.attrs.AllocTensorAttrs") { - TVM_ATTR_FIELD(dtype) - .describe("The dtype of the tensor to allocate.") - .set_default(DataType::Float(32, 1)); - TVM_ATTR_FIELD(const_shape).describe("The shape of constant used to aid in type inference."); - TVM_ATTR_FIELD(assert_shape) - .describe( - "The shape to cast the return type of the allocation to, " - "used to specify the shape obtained via further analysis."); - } -}; - -} // namespace relay -} // namespace tvm -#endif // TVM_RELAY_ATTRS_MEMORY_H_ diff --git a/include/tvm/relay/attrs/nn.h b/include/tvm/relay/attrs/nn.h deleted file mode 100644 index 58edb9df8b97..000000000000 --- a/include/tvm/relay/attrs/nn.h +++ /dev/null @@ -1,1593 +0,0 @@ -/* - * 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 tvm/relay/attrs/nn.h - * \brief Auxiliary attributes for nn operators. - */ -#ifndef TVM_RELAY_ATTRS_NN_H_ -#define TVM_RELAY_ATTRS_NN_H_ - -#include -#include - -#include - -namespace tvm { -namespace relay { - -/*! - * \brief Add a 1D Tensor to an axis of a data. - * - * \note bias_add is a special add operator that is in nn - * and enables automatic derivation of bias's shape. - * You can directly use add for more generalized case. - */ -struct BiasAddAttrs : public tvm::AttrsNode { - int axis; - - TVM_DECLARE_ATTRS(BiasAddAttrs, "relay.attrs.BiasAddAttrs") { - TVM_ATTR_FIELD(axis).describe("The axis to add the bias").set_default(1); - } -}; - -/*! \brief Attributes used in 1D convolution operators */ -struct Conv1DAttrs : public tvm::AttrsNode { - Array strides; - Array padding; - Array dilation; - int groups; - IndexExpr channels; - Array kernel_size; - tvm::String data_layout; - tvm::String kernel_layout; - tvm::String out_layout; - DataType out_dtype; - - TVM_DECLARE_ATTRS(Conv1DAttrs, "relay.attrs.Conv1DAttrs") { - TVM_ATTR_FIELD(strides) - .set_default(Array({ - 1, - })) - .describe("Specifies the stride of the convolution."); - TVM_ATTR_FIELD(padding) - .set_default(Array({0, 0})) - .describe( - "If padding is non-zero, then the input is implicitly zero-padded" - "on both sides for padding number of points"); - TVM_ATTR_FIELD(dilation) - .set_default(Array({ - 1, - })) - .describe("Specifies the dilation rate to use for dilated convolution."); - TVM_ATTR_FIELD(groups).set_default(1).describe( - "Currently unused but may be added in the future."); - TVM_ATTR_FIELD(channels) - .describe( - "The number of output channels in the convolution." - " If it is not set, inferred by shape of the weight.") - .set_default(NullValue()); - TVM_ATTR_FIELD(kernel_size) - .describe("Specifies the dimensions of the convolution window.") - .set_default(NullValue>()); - TVM_ATTR_FIELD(data_layout) - .set_default("NCW") - .describe( - "Dimension ordering of input data. Can be 'NCW', 'NWC', etc." - "'N', 'C', 'W' stands for batch, channel, and width" - "dimensions respectively. Convolution is applied on the 'W'" - "dimension."); - TVM_ATTR_FIELD(kernel_layout) - .set_default("OIW") - .describe( - "Dimension ordering of weight. Can be 'OIW', or 'WIO', etc." - "'O', 'I', 'W' stands for num_filter, input_channel, and width" - "dimensions respectively."); - TVM_ATTR_FIELD(out_layout) - .set_default("") - .describe( - "Dimension ordering of output. Can be 'NCW', 'NWC', etc." - "'N', 'C', 'W' stands for batch, channel, and width" - "dimensions respectively. Default to be same as input layout."); - - // use 0 bits to indicate none. - TVM_ATTR_FIELD(out_dtype) - .set_default(NullValue()) - .describe("Output data type, set to explicit type under mixed precision setting"); - } -}; - -/*! \brief Attributes used in convolution operators */ -struct Conv2DAttrs : public tvm::AttrsNode { - Array strides; - Array padding; - Array dilation; - int groups; - IndexExpr channels; - Array kernel_size; - tvm::String data_layout; - tvm::String kernel_layout; - tvm::String out_layout; - tvm::String auto_scheduler_rewritten_layout; // The layout after auto-scheduler's layout rewrite - Array meta_schedule_original_shape; // The original shape of the weights - DataType out_dtype; - - TVM_DECLARE_ATTRS(Conv2DAttrs, "relay.attrs.Conv2DAttrs") { - TVM_ATTR_FIELD(strides) - .set_default(Array({1, 1})) - .describe("Specifies the strides of the convolution."); - TVM_ATTR_FIELD(padding) - .set_default(Array({0, 0})) - .describe( - "If padding is non-zero, then the input is implicitly zero-padded" - "Padding support both symmetric and asymmetric as" - "one int : same padding used on all sides" - "two int : bottom, right will use same padding as top, left" - "four int : padding width in the order of (top, left, bottom, right)"); - TVM_ATTR_FIELD(dilation) - .set_default(Array({1, 1})) - .describe("Specifies the dilation rate to use for dilated convolution."); - TVM_ATTR_FIELD(groups).set_default(1).describe( - "Controls the connections between inputs and outputs." - "At groups=1, all inputs are convolved to all outputs." - "At groups=2, the operation becomes equivalent to having two convolution" - "layers side by side, each seeing half the input channels, and producing" - "half the output channels, and both subsequently concatenated."); - TVM_ATTR_FIELD(channels) - .describe( - "The number of output channels in the convolution." - " If it is not set, inferred by shape of the weight.") - .set_default(NullValue()); - TVM_ATTR_FIELD(kernel_size) - .describe("Specifies the dimensions of the convolution window.") - .set_default(NullValue>()); - TVM_ATTR_FIELD(data_layout) - .set_default("NCHW") - .describe( - "Dimension ordering of input data. Can be 'NCHW', 'NHWC', etc." - "'N', 'C', 'H', 'W' stands for batch, channel, height, and width" - "dimensions respectively. Convolution is applied on the 'H' and" - "'W' dimensions."); - TVM_ATTR_FIELD(kernel_layout) - .set_default("OIHW") - .describe( - "Dimension ordering of weight. Can be 'OIHW', 'OIHW16o16i', etc." - "'O', 'I', 'H', 'W' stands for num_filter, input_channel, height, and width" - "dimensions respectively."); - TVM_ATTR_FIELD(out_layout) - .set_default("") - .describe( - "Dimension ordering of output. Can be 'NCHW', 'NHWC', etc." - "'N', 'C', 'H', 'W' stands for batch, channel, height, and width" - "dimensions respectively. Default to be same as input layout."); - - // use 0 bits to indicate none. - TVM_ATTR_FIELD(out_dtype) - .set_default(NullValue()) - .describe("Output data type, set to explicit type under mixed precision setting"); - } -}; - -/*! \brief Attributes used in winograd weight transformation operators */ -struct ConvWinogradWeightTransformAttrs : public tvm::AttrsNode { - int tile_size; - - TVM_DECLARE_ATTRS(ConvWinogradWeightTransformAttrs, - "relay.attrs.ConvWinogradWeightTransformAttrs") { - TVM_ATTR_FIELD(tile_size).describe( - "Tile size of winograd. E.g. 2 for F(2x2, 3x3) and 4 for F(4x4, 3x3)"); - } -}; - -/*! \brief Attributes used in gemm weight transformation operators */ -struct ConvGemmWeightTransformAttrs : public tvm::AttrsNode { - int tile_N; - int tile_K; - - TVM_DECLARE_ATTRS(ConvGemmWeightTransformAttrs, "relay.attrs.ConvGemmWeightTransformAttrs") { - TVM_ATTR_FIELD(tile_N).describe( - "Tile size across N axis of the weight transformation for ConvGemm. (N = OC)"); - TVM_ATTR_FIELD(tile_K).describe( - "Tile size across K axis of the weight transformation for ConvGemm. (K = KW * KH * IC)"); - } -}; - -/*! \brief Attributes used in convolution operators with winograd algorithm */ -struct Conv2DWinogradAttrs : public tvm::AttrsNode { - int tile_size; - Array strides; - Array padding; - Array dilation; - int groups; - IndexExpr channels; - Array kernel_size; - tvm::String data_layout; - tvm::String kernel_layout; - tvm::String out_layout; - tvm::String auto_scheduler_rewritten_layout; // The layout after auto-scheduler's layout rewrite - Array meta_schedule_original_shape; // The original shape of the weights - DataType out_dtype; - - TVM_DECLARE_ATTRS(Conv2DWinogradAttrs, "relay.attrs.Conv2DWinogradAttrs") { - TVM_ATTR_FIELD(tile_size).describe( - "The tile size of winograd. E.g. 2 for F(2x2, 3x3) and 4 for F(4x4, 3x3)"); - TVM_ATTR_FIELD(strides) - .set_default(Array({1, 1})) - .describe("Specifies the strides of the convolution."); - TVM_ATTR_FIELD(padding) - .set_default(Array({0, 0})) - .describe( - "If padding is non-zero, then the input is implicitly zero-padded" - "Padding support both symmetric and asymmetric as" - "one int : same padding used on all sides" - "two int : bottom, right will use same padding as top, left" - "four int : padding width in the order of (top, left, bottom, right)"); - TVM_ATTR_FIELD(dilation) - .set_default(Array({1, 1})) - .describe("Specifies the dilation rate to use for dilated convolution."); - TVM_ATTR_FIELD(groups).set_default(1).describe( - "Controls the connections between inputs and outputs." - "At groups=1, all inputs are convolved to all outputs." - "At groups=2, the operation becomes equivalent to having two convolution" - "layers side by side, each seeing half the input channels, and producing" - "half the output channels, and both subsequently concatenated."); - TVM_ATTR_FIELD(channels) - .describe( - "The number of output channels in the convolution." - " If it is not set, inferred by shape of the weight.") - .set_default(NullValue()); - TVM_ATTR_FIELD(kernel_size) - .describe("Specifies the dimensions of the convolution window.") - .set_default(NullValue>()); - TVM_ATTR_FIELD(data_layout) - .set_default("NCHW") - .describe( - "Dimension ordering of input data. Can be 'NCHW', 'NHWC', etc." - "'N', 'C', 'H', 'W' stands for batch, channel, height, and width" - "dimensions respectively. Convolution is applied on the 'H' and" - "'W' dimensions."); - TVM_ATTR_FIELD(kernel_layout) - .set_default("OIHW") - .describe( - "Dimension ordering of weight. Can be 'OIHW', 'OIHW16o16i', etc." - "'O', 'I', 'H', 'W' stands for num_filter, input_channel, height, and width" - "dimensions respectively."); - TVM_ATTR_FIELD(out_layout) - .set_default("") - .describe( - "Dimension ordering of output. Can be 'NCHW', 'NHWC', etc." - "'N', 'C', 'H', 'W' stands for batch, channel, height, and width" - "dimensions respectively. Default to be same as input layout."); - - // use 0 bits to indicate none. - TVM_ATTR_FIELD(out_dtype) - .set_default(NullValue()) - .describe("Output data type, set to explicit type under mixed precision setting"); - } -}; - -/*! \brief Attributes used in winograd weight transformation operators */ -struct Conv2DWinogradNNPACKWeightTransformAttrs - : public tvm::AttrsNode { - int convolution_algorithm; - DataType out_dtype; - - TVM_DECLARE_ATTRS(Conv2DWinogradNNPACKWeightTransformAttrs, - "relay.attrs.Conv2DWinogradNNPACKWeightTransformAttrs") { - TVM_ATTR_FIELD(convolution_algorithm) - .describe( - "The convolution algorithm for Winograd NNPACK. " - "E.g. tvm.contrib.nnpack.ConvolutionAlgorithm.WT_8x8 for WT_8x8, " - "tvm.contrib.nnpack.ConvolutionAlgorithm.WT_8x8_FP16 for WT_8x8_FP16"); - TVM_ATTR_FIELD(out_dtype) - .set_default(NullValue()) - .describe("Output data type, set to explicit type under mixed precision setting"); - } -}; - -/*! \brief Attributes used in convolution operators */ -struct Conv3DAttrs : public tvm::AttrsNode { - Array strides; - Array padding; - Array dilation; - int groups; - IndexExpr channels; - Array kernel_size; - tvm::String data_layout; - tvm::String kernel_layout; - tvm::String out_layout; - tvm::String auto_scheduler_rewritten_layout; // The layout after auto-scheduler's layout rewrite - Array meta_schedule_original_shape; // The original shape of the weights - DataType out_dtype; - - TVM_DECLARE_ATTRS(Conv3DAttrs, "relay.attrs.Conv3DAttrs") { - TVM_ATTR_FIELD(strides) - .set_default(Array({1, 1, 1})) - .describe("Specifies the strides of the convolution."); - TVM_ATTR_FIELD(padding) - .set_default(Array({0, 0, 0})) - .describe( - "If padding is non-zero, then the input is implicitly zero-padded" - "Padding support both symmetric and asymmetric as" - "one int : same padding used on all sides" - "three int : back, bottom, right will use same padding as front, top, left" - "six int : padding width in the order of (front, top, left, back, bottom," - "right)"); - TVM_ATTR_FIELD(dilation) - .set_default(Array({1, 1, 1})) - .describe("Specifies the dilation rate to use for dilated convolution."); - TVM_ATTR_FIELD(groups).set_default(1).describe( - "Controls the connections between inputs and outputs." - "At groups=1, all inputs are convolved to all outputs." - "At groups=2, the operation becomes equivalent to having two convolution" - "layers side by side, each seeing half the input channels, and producing" - "half the output channels, and both subsequently concatenated."); - TVM_ATTR_FIELD(channels) - .describe( - "The number of output channels in the convolution." - " If it is not set, inferred by shape of the weight.") - .set_default(NullValue()); - TVM_ATTR_FIELD(kernel_size) - .describe("Specifies the dimensions of the convolution window.") - .set_default(NullValue>()); - TVM_ATTR_FIELD(data_layout) - .set_default("NCDHW") - .describe( - "Dimension ordering of input data. Can be 'NCDHW', 'NDHWC', etc." - "'N', 'C', 'D', 'H', 'W' stands for batch, channel, depth, height, and width" - "dimensions respectively. Convolution is applied on the 'D', 'H' and" - "'W' dimensions."); - TVM_ATTR_FIELD(kernel_layout) - .set_default("OIDHW") - .describe( - "Dimension ordering of weight. Can be 'OIDHW', 'OIDHW16o16i', etc." - "'O', 'I', 'D', 'H', 'W' stands for num_filter, input_channel, depth, height," - "and width dimensions respectively."); - TVM_ATTR_FIELD(out_layout) - .set_default("") - .describe( - "Dimension ordering of output. Can be 'NCDHW', 'NDHWC', etc." - "'N', 'C', 'D', 'H', 'W' stands for batch, channel, depth, height, and width" - "dimensions respectively. Default to be same as input layout."); - - // use 0 bits to indicate none. - TVM_ATTR_FIELD(out_dtype) - .set_default(NullValue()) - .describe("Output data type, set to explicit type under mixed precision setting"); - } -}; - -/*! \brief Attributes used in transposed convolution operator */ -struct Conv3DTransposeAttrs : public tvm::AttrsNode { - IndexExpr channels; - Array kernel_size; - Array strides; - Array padding; - Array output_padding; - Array dilation; - int groups; - tvm::String data_layout; - tvm::String kernel_layout; - tvm::String out_layout; - DataType out_dtype; - - TVM_DECLARE_ATTRS(Conv3DTransposeAttrs, "relay.attrs.Conv3DTransposeAttrs") { - TVM_ATTR_FIELD(channels) - .set_default(NullValue()) - .describe( - "The dimensionality of the output space" - "i.e. the number of output channels in the convolution."); - TVM_ATTR_FIELD(kernel_size) - .describe("The dimensions of the convolution window.") - .set_default(NullValue>()); - TVM_ATTR_FIELD(strides) - .set_default(Array({1, 1, 1})) - .describe("The strides of the convolution."); - TVM_ATTR_FIELD(output_padding) - .set_default(Array({0, 0, 0})) - .describe( - "Zero-padding added to one side of the output." - "Padding support both symmetric and asymmetric as" - "one int : same padding used on all sides" - "three int : front, bottom, right will use same padding as back, top, left" - "six int : padding width in the order of (front, top, left, back, bottom, right)"); - TVM_ATTR_FIELD(padding) - .set_default(Array({0, 0, 0})) - .describe( - "If padding is non-zero, then the input is implicitly zero-padded" - "Padding support both symmetric and asymmetric as" - "one int : same padding used on all sides" - "three int : front, bottom, right will use same padding as back, top, left" - "six int : padding width in the order of (front, top, left, back, bottom, right)"); - TVM_ATTR_FIELD(dilation) - .set_default(Array({1, 1, 1})) - .describe("Specifies the dilation rate to use for dilated convolution."); - TVM_ATTR_FIELD(groups).set_default(1).describe( - "Controls the connections between inputs and outputs." - "At groups=1, all inputs are convolved to all outputs." - "At groups=2, the operation becomes equivalent to having two convolution" - "layers side by side, each seeing half the input channels, and producing" - "half the output channels, and both subsequently concatenated."); - TVM_ATTR_FIELD(data_layout) - .set_default("NCDHW") - .describe( - "Dimension ordering of data. Can be 'NCDHW', 'NDHWC', etc." - "'N', 'C', 'D', 'H', 'W' stands for batch, channel, depth, height, and width" - "dimensions respectively. Convolution is applied on the 'D', 'H' and" - "'W' dimensions."); - TVM_ATTR_FIELD(kernel_layout) - .set_default("IODHW") - .describe( - "Dimension ordering of data and weight. Can be 'IODHW', 'IODHW16i16o', etc." - "'I', 'O', 'D', 'H', 'W' stands for input_channel, num_filter, depth, height, and width" - "dimensions respectively."); - TVM_ATTR_FIELD(out_layout) - .set_default("") - .describe( - "Dimension ordering of output. Can be 'NCDHW', 'NDHWC', etc." - "'N', 'C', 'D', 'H', 'W' stands for batch, channel, depth, height, and width" - "dimensions respectively. Default to be same as input layout."); - TVM_ATTR_FIELD(out_dtype) - .set_default(NullValue()) - .describe("Output data type, set to explicit type under mixed precision setting"); - } -}; - -/*! \brief Attributes used in 3d winograd convolution operators */ -struct Conv3DWinogradAttrs : public tvm::AttrsNode { - int tile_size; - Array strides; - Array padding; - Array dilation; - int groups; - IndexExpr channels; - Array kernel_size; - std::string data_layout; - std::string kernel_layout; - std::string out_layout; - DataType out_dtype; - - TVM_DECLARE_ATTRS(Conv3DWinogradAttrs, "relay.attrs.Conv3DWinogradAttrs") { - TVM_ATTR_FIELD(tile_size).describe( - "The tile size of winograd. E.g. 2 for F(2x2x2, 3x3x3) and 4 for F(4x4x4, 3x3x3)"); - TVM_ATTR_FIELD(strides) - .set_default(Array({1, 1, 1})) - .describe("Specifies the strides of the convolution."); - TVM_ATTR_FIELD(padding) - .set_default(Array({0, 0, 0})) - .describe( - "If padding is non-zero, then the input is implicitly zero-padded" - "Padding support both symmetric and asymmetric as" - "one int : same padding used on all sides" - "three int : back, bottom, right will use same padding as front, top, left" - "six int : padding width in the order of (front, top, left, back, bottom," - "right)"); - TVM_ATTR_FIELD(dilation) - .set_default(Array({1, 1, 1})) - .describe("Specifies the dilation rate to use for dilated convolution."); - TVM_ATTR_FIELD(groups).set_default(1).describe( - "Controls the connections between inputs and outputs." - "At groups=1, all inputs are convolved to all outputs." - "At groups=2, the operation becomes equivalent to having two convolution" - "layers side by side, each seeing half the input channels, and producing" - "half the output channels, and both subsequently concatenated."); - TVM_ATTR_FIELD(channels) - .describe( - "The number of output channels in the convolution." - " If it is not set, inferred by shape of the weight.") - .set_default(NullValue()); - TVM_ATTR_FIELD(kernel_size) - .describe("Specifies the dimensions of the convolution window.") - .set_default(NullValue>()); - TVM_ATTR_FIELD(data_layout) - .set_default("NCDHW") - .describe( - "Dimension ordering of input data. Can be 'NCDHW', 'NDHWC', etc." - "'N', 'C', 'D', 'H', 'W' stands for batch, channel, depth, height, and width" - "dimensions respectively. Convolution is applied on the 'D', 'H' and" - "'W' dimensions."); - TVM_ATTR_FIELD(kernel_layout) - .set_default("OIDHW") - .describe( - "Dimension ordering of weight. Can be 'OIDHW', 'OIDHW16o16i', etc." - "'O', 'I', 'D', 'H', 'W' stands for num_filter, input_channel, depth, height," - "and width dimensions respectively."); - TVM_ATTR_FIELD(out_layout) - .set_default("") - .describe( - "Dimension ordering of output. Can be 'NCDHW', 'NDHWC', etc." - "'N', 'C', 'D', 'H', 'W' stands for batch, channel, depth, height, and width" - "dimensions respectively. Default to be same as input layout."); - - // use 0 bits to indicate none. - TVM_ATTR_FIELD(out_dtype) - .set_default(NullValue()) - .describe("Output data type, set to explicit type under mixed precision setting"); - } -}; - -/*! \brief Attributes used in softmax operators */ -struct SoftmaxAttrs : public tvm::AttrsNode { - int axis; - - TVM_DECLARE_ATTRS(SoftmaxAttrs, "relay.attrs.SoftmaxAttrs") { - TVM_ATTR_FIELD(axis).set_default(-1).describe("The axis to sum over when computing softmax."); - } -}; - -/*! \brief Attributes used in transposed convolution operator */ -struct Conv2DTransposeAttrs : public tvm::AttrsNode { - IndexExpr channels; - Array kernel_size; - Array strides; - Array padding; - Array output_padding; - Array dilation; - int groups; - std::string data_layout; - std::string kernel_layout; - std::string out_layout; - DataType out_dtype; - - TVM_DECLARE_ATTRS(Conv2DTransposeAttrs, "relay.attrs.Conv2DTransposeAttrs") { - TVM_ATTR_FIELD(channels) - .set_default(NullValue()) - .describe( - "The dimensionality of the output space" - "i.e. the number of output channels in the convolution."); - TVM_ATTR_FIELD(kernel_size) - .describe("The dimensions of the convolution window.") - .set_default(NullValue>()); - TVM_ATTR_FIELD(strides) - .set_default(Array({1, 1})) - .describe("The strides of the convolution."); - TVM_ATTR_FIELD(output_padding) - .set_default(Array({0, 0})) - .describe( - "Zero-padding added to one side of the output." - "Padding support both symmetric and asymmetric as" - "one int : same padding used on all sides" - "two int : bottom, right will use same padding as top, left" - "four int : padding width in the order of (top, left, bottom, right)"); - TVM_ATTR_FIELD(padding) - .set_default(Array({0, 0})) - .describe( - "If padding is non-zero, then the input is implicitly zero-padded" - "Padding support both symmetric and asymmetric as" - "one int : same padding used on all sides" - "two int : bottom, right will use same padding as top, left" - "four int : padding width in the order of (top, left, bottom, right)"); - TVM_ATTR_FIELD(dilation) - .set_default(Array({1, 1})) - .describe("Specifies the dilation rate to use for dilated convolution."); - TVM_ATTR_FIELD(groups).set_default(1).describe( - "Controls the connections between inputs and outputs." - "At groups=1, all inputs are convolved to all outputs." - "At groups=2, the operation becomes equivalent to having two convolution" - "layers side by side, each seeing half the input channels, and producing" - "half the output channels, and both subsequently concatenated."); - TVM_ATTR_FIELD(data_layout) - .set_default("NCHW") - .describe( - "Dimension ordering of data. Can be 'NCHW', 'NHWC', etc." - "'N', 'C', 'H', 'W' stands for batch, channel, height, and width" - "dimensions respectively. Convolution is applied on the 'H' and" - "'W' dimensions."); - TVM_ATTR_FIELD(kernel_layout) - .set_default("IOHW") - .describe( - "Dimension ordering of data and weight. Can be 'IOHW', 'OIHW16o16i', etc." - "'I', 'O', 'H', 'W' stands for input_channel, num_filter, height, and width" - "dimensions respectively."); - TVM_ATTR_FIELD(out_layout) - .set_default("") - .describe( - "Dimension ordering of output. Can be 'NCHW', 'NHWC', etc." - "'N', 'C', 'H', 'W' stands for batch, channel, height, and width" - "dimensions respectively. Default to be same as input layout."); - TVM_ATTR_FIELD(out_dtype) - .set_default(NullValue()) - .describe("Output data type, set to explicit type under mixed precision setting"); - } -}; - -/*! \brief Attributes used in dilate operator */ -struct DilateAttrs : public tvm::AttrsNode { - Array strides; - double dilation_value; - - TVM_DECLARE_ATTRS(DilateAttrs, "relay.attrs.DilateAttrs") { - TVM_ATTR_FIELD(strides) - .set_default(Array({1, 1})) - .describe("Dilation stride on each dimension, 1 means no dilation."); - TVM_ATTR_FIELD(dilation_value).set_default(0.0).describe("Value used to dilate the input."); - } -}; - -/*! \brief Attributes used in 1D transposed convolution operator */ -struct Conv1DTransposeAttrs : public tvm::AttrsNode { - IndexExpr channels; - Array kernel_size; - Array strides; - Array padding; - Array output_padding; - Array dilation; - int groups; - std::string data_layout; - std::string kernel_layout; - std::string out_layout; - DataType out_dtype; - - TVM_DECLARE_ATTRS(Conv1DTransposeAttrs, "relay.attrs.Conv1DTransposeAttrs") { - TVM_ATTR_FIELD(channels) - .set_default(NullValue()) - .describe( - "The dimensionality of the output space" - "i.e. the number of output channels in the convolution."); - TVM_ATTR_FIELD(kernel_size) - .describe("The dimensions of the convolution window.") - .set_default(NullValue>()); - TVM_ATTR_FIELD(strides) - .set_default(Array({1})) - .describe("The strides of the convolution."); - TVM_ATTR_FIELD(output_padding) - .set_default(Array({0})) - .describe("Zero-padding added to one side of the output."); - TVM_ATTR_FIELD(padding) - .set_default(Array({0})) - .describe( - "Symmetric or asymmetric padding." - "Single value: the input is implicitly zero-padded on both sides." - "Two values: padding[0] is used for left input padding, " - "padding[1] is used for right input padding,"); - TVM_ATTR_FIELD(dilation) - .set_default(Array({1})) - .describe("Specifies the dilation rate to use for dilated convolution."); - TVM_ATTR_FIELD(groups).set_default(1).describe( - "Controls the connections between inputs and outputs." - "At groups=1, all inputs are convolved to all outputs." - "At groups=2, the operation becomes equivalent to having two convolution" - "layers side by side, each seeing half the input channels, and producing" - "half the output channels, and both subsequently concatenated."); - TVM_ATTR_FIELD(data_layout) - .set_default("NCW") - .describe( - "Dimension ordering of data. Can be 'NCW', 'NWC', etc." - "'N', 'C', 'W' stands for batch, channel, and width" - "dimensions respectively. Convolution is applied on the" - "'W' dimension."); - TVM_ATTR_FIELD(kernel_layout) - .set_default("IOW") - .describe( - "Dimension ordering of data and weight. Can be 'IOW', 'IOW16o16i', etc." - "'I', 'O', 'W' stands for input_channel, num_filter and width" - "dimensions respectively."); - TVM_ATTR_FIELD(out_layout) - .set_default("") - .describe( - "Dimension ordering of output. Can be 'NCW', 'NWC', etc." - "'N', 'C', 'W' stands for batch, channel, and width" - "dimensions respectively. Default to be same as input layout."); - TVM_ATTR_FIELD(out_dtype) - .set_default(NullValue()) - .describe("Output data type, set to explicit type under mixed precision setting"); - } -}; - -/*! \brief Attributes for max pool operator */ -struct MaxPool2DAttrs : public tvm::AttrsNode { - Array pool_size; - Array strides; - Array padding; - Array dilation; - tvm::String layout; - tvm::String out_layout; - bool ceil_mode; - - TVM_DECLARE_ATTRS(MaxPool2DAttrs, "relay.attrs.MaxPool2DAttrs") { - TVM_ATTR_FIELD(pool_size).describe("Size of the pooling windows."); - TVM_ATTR_FIELD(strides) - .set_default(Array({1, 1})) - .describe("Specifies the strides of the convolution."); - TVM_ATTR_FIELD(dilation) - .set_default(Array({1, 1})) - .describe("Specifies the dilation of the convolution."); - TVM_ATTR_FIELD(padding) - .set_default(Array({0, 0})) - .describe( - "If padding is non-zero, then the input is implicitly zero-padded" - "Padding support both symmetric and asymmetric as" - "one int : same padding used on all sides" - "two int : bottom, right will use same padding as top, left" - "four int : padding width in the order of (top, left, bottom, right)"); - TVM_ATTR_FIELD(layout).set_default("NCHW").describe( - "Dimension ordering of input data. Can be 'NCHW', 'NHWC', etc." - "'N', 'C', 'H', 'W' stands for batch, channel, height, and width" - "dimensions respectively. Pooling is applied on the 'H' and" - "'W' dimensions."); - TVM_ATTR_FIELD(out_layout) - .set_default("") - .describe( - "Dimension ordering of output data. Can be 'NCHW', 'NHWC', etc." - "'N', 'C', 'H', 'W' stands for batch, channel, height, and width" - "dimensions respectively. Pooling is applied on the 'H' and" - "'W' dimensions."); - TVM_ATTR_FIELD(ceil_mode).set_default(false).describe( - "When true, will use ceil instead of floor to compute the output shape."); - } -}; - -/*! \brief Attributes for avg pool operator */ -struct AvgPool2DAttrs : public tvm::AttrsNode { - Array pool_size; - Array strides; - Array padding; - Array dilation; - tvm::String layout; - tvm::String out_layout; - bool ceil_mode; - bool count_include_pad; - - TVM_DECLARE_ATTRS(AvgPool2DAttrs, "relay.attrs.AvgPool2DAttrs") { - TVM_ATTR_FIELD(pool_size).describe("Size of the pooling windows."); - TVM_ATTR_FIELD(strides) - .set_default(Array({1, 1})) - .describe("Specifies the strides of the convolution."); - TVM_ATTR_FIELD(dilation) - .set_default(Array({1, 1})) - .describe("Specifies the dilation of the convolution."); - TVM_ATTR_FIELD(padding) - .set_default(Array({0, 0})) - .describe( - "If padding is non-zero, then the input is implicitly zero-padded" - "Padding support both symmetric and asymmetric as" - "one int : same padding used on all sides" - "two int : bottom, right will use same padding as top, left" - "four int : padding width in the order of (top, left, bottom, right)"); - TVM_ATTR_FIELD(layout).set_default("NCHW").describe( - "Dimension ordering of input data. Can be 'NCHW', 'NHWC', etc." - "'N', 'C', 'H', 'W' stands for batch, channel, height, and width" - "dimensions respectively. Pooling is applied on the 'H' and" - "'W' dimensions."); - TVM_ATTR_FIELD(out_layout) - .set_default("") - .describe( - "Dimension ordering of output data. Can be 'NCHW', 'NHWC', etc." - "'N', 'C', 'H', 'W' stands for batch, channel, height, and width" - "dimensions respectively. Pooling is applied on the 'H' and" - "'W' dimensions."); - TVM_ATTR_FIELD(ceil_mode).set_default(false).describe( - "When true, will use ceil instead of floor to compute the output shape."); - TVM_ATTR_FIELD(count_include_pad) - .set_default(false) - .describe("When true, will include padding to compute the average"); - } -}; - -/*! \brief Attributes for global pool operator */ -struct GlobalPool2DAttrs : public tvm::AttrsNode { - tvm::String layout; - tvm::String out_layout; - - TVM_DECLARE_ATTRS(GlobalPool2DAttrs, "relay.attrs.GlobalPool2DAttrs") { - TVM_ATTR_FIELD(layout).set_default("NCHW").describe( - "Dimension ordering of input data. Can be 'NCHW', 'NHWC', etc." - "'N', 'C', 'H', 'W' stands for batch, channel, height, and width" - "dimensions respectively. Pooling is applied on the 'H' and" - "'W' dimensions."); - TVM_ATTR_FIELD(out_layout) - .set_default("") - .describe( - "Dimension ordering of output data. Can be 'NCHW', 'NHWC', etc." - "'N', 'C', 'H', 'W' stands for batch, channel, height, and width" - "dimensions respectively. Pooling is applied on the 'H' and" - "'W' dimensions."); - } -}; - -/*! \brief Attributes for 1d adaptive pool operator */ -struct AdaptivePool1DAttrs : public tvm::AttrsNode { - Array output_size; - std::string layout; - tvm::String out_layout; - - TVM_DECLARE_ATTRS(AdaptivePool1DAttrs, "relay.attrs.AdaptivePool1DAttrs") { - TVM_ATTR_FIELD(output_size).set_default(Array({})).describe("Output width."); - TVM_ATTR_FIELD(layout).set_default("NCW").describe( - "Dimension ordering of input data. Can be 'NCW', 'NWC', etc." - "'N', 'C', 'W' stands for batch, channel, and width" - "dimensions respectively. Pooling is applied on the" - "'W' dimension."); - TVM_ATTR_FIELD(out_layout) - .set_default("") - .describe( - "Dimension ordering of output data. Can be 'NCW', 'NWC', etc." - "'N', 'C', 'W' stands for batch, channel, and width" - "dimensions respectively. Pooling is applied on the" - "'W' dimension."); - } -}; - -/*! \brief Attributes for 2d adaptive pool operator */ -struct AdaptivePool2DAttrs : public tvm::AttrsNode { - Array output_size; - std::string layout; - tvm::String out_layout; - - TVM_DECLARE_ATTRS(AdaptivePool2DAttrs, "relay.attrs.AdaptivePool2DAttrs") { - TVM_ATTR_FIELD(output_size) - .set_default(Array({})) - .describe("Output height and width."); - TVM_ATTR_FIELD(layout).set_default("NCHW").describe( - "Dimension ordering of input data. Can be 'NCHW', 'NHWC', etc." - "'N', 'C', 'H', 'W' stands for batch, channel, height, and width" - "dimensions respectively. Pooling is applied on the 'H' and" - "'W' dimensions."); - TVM_ATTR_FIELD(out_layout) - .set_default("") - .describe( - "Dimension ordering of output data. Can be 'NCHW', 'NHWC', etc." - "'N', 'C', 'H', 'W' stands for batch, channel, height, and width" - "dimensions respectively. Pooling is applied on the 'H' and" - "'W' dimensions."); - } -}; - -/*! \brief Attributes for 3d adaptive pool operator */ -struct AdaptivePool3DAttrs : public tvm::AttrsNode { - Array output_size; - std::string layout; - tvm::String out_layout; - - TVM_DECLARE_ATTRS(AdaptivePool3DAttrs, "relay.attrs.AdaptivePool3DAttrs") { - TVM_ATTR_FIELD(output_size) - .set_default(Array({})) - .describe("Output depth, height and width."); - TVM_ATTR_FIELD(layout).set_default("NCDHW").describe( - "Dimension ordering of input data. Can be 'NCDHW', 'NDHWC', etc." - "'N', 'C', 'D', 'H', 'W' stands for batch, channel, depth, height, and width" - "dimensions respectively. Pooling is applied on 'D', 'H' and" - "'W' dimensions."); - TVM_ATTR_FIELD(out_layout) - .set_default("") - .describe( - "Dimension ordering of output data. Can be 'NCDHW', 'NDHWC', etc." - "'N', 'C', 'D', 'H', 'W' stands for batch, channel, depth, height, and width" - "dimensions respectively. Pooling is applied on 'D', 'H' and" - "'W' dimensions."); - } -}; - -/*! \brief Attributes for 1D max pool operator */ -struct MaxPool1DAttrs : public tvm::AttrsNode { - Array pool_size; - Array strides; - Array dilation; - Array padding; - std::string layout; - tvm::String out_layout; - bool ceil_mode; - - TVM_DECLARE_ATTRS(MaxPool1DAttrs, "relay.attrs.MaxPool1DAttrs") { - TVM_ATTR_FIELD(pool_size).describe("Size of the pooling windows."); - TVM_ATTR_FIELD(strides) - .set_default(Array({1})) - .describe("Specifies the strides of the convolution."); - TVM_ATTR_FIELD(dilation) - .set_default(Array({1})) - .describe("Specifies the dilation of the convolution."); - TVM_ATTR_FIELD(padding) - .set_default(Array({0})) - .describe( - "If padding is non-zero, then the input is implicitly zero-padded" - "Padding supports both symmetric and asymmetric as" - "one int : same padding used on each side" - "two int : indicates left padding, right padding"); - TVM_ATTR_FIELD(layout).set_default("NCW").describe( - "Dimension ordering of input data. Can be 'NCW', 'NWC', etc." - "'N', 'C', 'W' stands for batch, channel, and width" - "dimensions respectively. Pooling is applied on the 'W' dimensions."); - TVM_ATTR_FIELD(out_layout) - .set_default("") - .describe( - "Dimension ordering of output data. Can be 'NCW', 'NWC', etc." - "'N', 'C', 'W' stands for batch, channel, and width" - "dimensions respectively. Pooling is applied on the 'W' dimensions."); - TVM_ATTR_FIELD(ceil_mode).set_default(false).describe( - "When true, will use ceil instead of floor to compute the output shape."); - } -}; - -/*! \brief Attributes for 1D avg pool operator */ -struct AvgPool1DAttrs : public tvm::AttrsNode { - Array pool_size; - Array strides; - Array dilation; - Array padding; - std::string layout; - tvm::String out_layout; - bool ceil_mode; - bool count_include_pad; - - TVM_DECLARE_ATTRS(AvgPool1DAttrs, "relay.attrs.AvgPool1DAttrs") { - TVM_ATTR_FIELD(pool_size).describe("Size of the pooling windows."); - TVM_ATTR_FIELD(strides) - .set_default(Array({1})) - .describe("Specifies the strides of the convolution."); - TVM_ATTR_FIELD(dilation) - .set_default(Array({1})) - .describe("Specifies the dilation of the convolution."); - TVM_ATTR_FIELD(padding) - .set_default(Array({0})) - .describe( - "If padding is non-zero, then the input is implicitly zero-padded" - "Padding supports both symmetric and asymmetric as" - "one int : same padding used on each side" - "two int : indicates left padding, right padding"); - TVM_ATTR_FIELD(layout).set_default("NCW").describe( - "Dimension ordering of input data. Can be 'NCW', 'NHC', etc." - "'N', 'C', 'W' stands for batch, channel, and width" - "dimensions respectively. Pooling is applied on the 'W' dimension."); - TVM_ATTR_FIELD(out_layout) - .set_default("") - .describe( - "Dimension ordering of output data. Can be 'NCW', 'NHC', etc." - "'N', 'C', 'W' stands for batch, channel, and width" - "dimensions respectively. Pooling is applied on the 'W' dimension."); - TVM_ATTR_FIELD(ceil_mode).set_default(false).describe( - "When true, will use ceil instead of floor to compute the output shape."); - TVM_ATTR_FIELD(count_include_pad) - .set_default(false) - .describe("When true, will include padding to compute the average"); - } -}; - -/*! \brief Attributes for 3D max pool operator */ -struct MaxPool3DAttrs : public tvm::AttrsNode { - Array pool_size; - Array strides; - Array dilation; - Array padding; - std::string layout; - tvm::String out_layout; - bool ceil_mode; - - TVM_DECLARE_ATTRS(MaxPool3DAttrs, "relay.attrs.MaxPool3DAttrs") { - TVM_ATTR_FIELD(pool_size).describe("Size of the pooling windows."); - TVM_ATTR_FIELD(strides) - .set_default(Array({1, 1, 1})) - .describe("Specifies the strides of the convolution."); - TVM_ATTR_FIELD(dilation) - .set_default(Array({1, 1, 1})) - .describe("Specifies the dilation of the convolution."); - TVM_ATTR_FIELD(padding) - .set_default(Array({0, 0, 0})) - .describe( - "If padding is non-zero, then the input is implicitly zero-padded" - "Padding support both symmetric and asymmetric as" - "one int : same padding used on all sides" - "three int : back, bottom, right will use same padding as front, top, left" - "six int : padding width in the order of (front, top, left, back, bottom, right)"); - TVM_ATTR_FIELD(layout).set_default("NCDHW").describe( - "Dimension ordering of input data. Can be 'NCDHW', 'NDHWC', etc." - "'N', 'C', 'D', 'H', 'W' stands for batch, channel, depth, height, and width" - "dimensions respectively. Pooling is applied on the 'D', 'H' and" - "'W' dimensions."); - TVM_ATTR_FIELD(out_layout) - .set_default("") - .describe( - "Dimension ordering of output data. Can be 'NCDHW', 'NDHWC', etc." - "'N', 'C', 'D', 'H', 'W' stands for batch, channel, depth, height, and width" - "dimensions respectively. Pooling is applied on the 'D', 'H' and" - "'W' dimensions."); - TVM_ATTR_FIELD(ceil_mode).set_default(false).describe( - "When true, will use ceil instead of floor to compute the output shape."); - } -}; - -/*! \brief Attributes for 3D avg pool operator */ -struct AvgPool3DAttrs : public tvm::AttrsNode { - Array pool_size; - Array strides; - Array dilation; - Array padding; - std::string layout; - tvm::String out_layout; - bool ceil_mode; - bool count_include_pad; - - TVM_DECLARE_ATTRS(AvgPool3DAttrs, "relay.attrs.AvgPool3DAttrs") { - TVM_ATTR_FIELD(pool_size).describe("Size of the pooling windows."); - TVM_ATTR_FIELD(strides) - .set_default(Array({1, 1, 1})) - .describe("Specifies the strides of the convolution."); - TVM_ATTR_FIELD(dilation) - .set_default(Array({1, 1, 1})) - .describe("Specifies the dilation of the convolution."); - TVM_ATTR_FIELD(padding) - .set_default(Array({0, 0, 0})) - .describe( - "If padding is non-zero, then the input is implicitly zero-padded" - "Padding support both symmetric and asymmetric as" - "one int : same padding used on all sides" - "three int : back, bottom, right will use same padding as front, top, left" - "six int : padding width in the order of (front, top, left, back, bottom, right)"); - TVM_ATTR_FIELD(layout).set_default("NCDHW").describe( - "Dimension ordering of input data. Can be 'NCDHW', 'NDHWC', etc." - "'N', 'C', 'D', 'H', 'W' stands for batch, channel, depth, height, and width" - "dimensions respectively. Pooling is applied on the 'D', 'H' and" - "'W' dimensions."); - TVM_ATTR_FIELD(out_layout) - .set_default("") - .describe( - "Dimension ordering of output data. Can be 'NCDHW', 'NDHWC', etc." - "'N', 'C', 'D', 'H', 'W' stands for batch, channel, depth, height, and width" - "dimensions respectively. Pooling is applied on the 'D', 'H' and" - "'W' dimensions."); - TVM_ATTR_FIELD(ceil_mode).set_default(false).describe( - "When true, will use ceil instead of floor to compute the output shape."); - TVM_ATTR_FIELD(count_include_pad) - .set_default(false) - .describe("When true, will include padding to compute the average"); - } -}; - -/*! \brief Attributes for matmul operator */ -struct MatmulAttrs : public tvm::AttrsNode { - IndexExpr units; - DataType out_dtype; - bool transpose_a; - bool transpose_b; - // layout of B after auto-scheduler's layout rewrite - tvm::String auto_scheduler_rewritten_layout; - Array meta_schedule_original_shape; // The original shape of the weights - - TVM_DECLARE_ATTRS(MatmulAttrs, "relay.attrs.MatmulAttrs") { - TVM_ATTR_FIELD(units).describe("Number of hidden units of the dense transformation."); - - // use 0 bits to indicate none. - TVM_ATTR_FIELD(out_dtype) - .set_default(NullValue()) - .describe("Output data type, set to explicit type under mixed precision setting"); - - TVM_ATTR_FIELD(transpose_a) - .set_default(false) - .describe("Whether the first input tensor is in transposed format."); - - TVM_ATTR_FIELD(transpose_b) - .set_default(false) - .describe("Whether the second input tensor is in transposed format."); - } -}; - -/*! \brief Attributes for dense operator */ -struct DenseAttrs : public tvm::AttrsNode { - IndexExpr units; - // layout of B after auto-scheduler's layout rewrite - tvm::String auto_scheduler_rewritten_layout; - Array meta_schedule_original_shape; // The original shape of the weights - DataType out_dtype; - - TVM_DECLARE_ATTRS(DenseAttrs, "relay.attrs.DenseAttrs") { - TVM_ATTR_FIELD(units).describe("Number of hidden units of the dense transformation."); - - // use 0 bits to indicate none. - TVM_ATTR_FIELD(out_dtype) - .set_default(NullValue()) - .describe("Output data type, set to explicit type under mixed precision setting"); - } -}; - -/*! \brief Attributes for dense_pack operator */ -struct DensePackAttrs : public tvm::AttrsNode { - IndexExpr units; - DataType out_dtype; - tvm::String weight_layout; - - TVM_DECLARE_ATTRS(DensePackAttrs, "relay.attrs.DensePackAttrs") { - TVM_ATTR_FIELD(units).describe("Number of hidden units of the dense transformation."); - - // use 0 bits to indicate none. - TVM_ATTR_FIELD(out_dtype) - .set_default(NullValue()) - .describe("Output data type, set to explicit type under mixed precision setting"); - TVM_ATTR_FIELD(weight_layout) - .set_default("NC") - .describe("Dimension ordering of weight. Packed layouts, such as NC8n, are possible."); - } -}; - -/*! \brief Attributes for batch matmul operator. */ -struct BatchMatmulAttrs : public tvm::AttrsNode { - DataType out_dtype; - bool transpose_a; - bool transpose_b; - tvm::String auto_scheduler_rewritten_layout; // The layout after auto-scheduler's layout rewrite - Array meta_schedule_original_shape; // The original shape of the weights - - TVM_DECLARE_ATTRS(BatchMatmulAttrs, "relay.attrs.BatchMatmulAttrs") { - // use 0 bits to indicate none. - TVM_ATTR_FIELD(out_dtype) - .set_default(NullValue()) - .describe("Output data type, set to explicit type under mixed precision setting"); - - TVM_ATTR_FIELD(transpose_a) - .set_default(false) - .describe("Whether the first input tensor is in transposed format."); - - TVM_ATTR_FIELD(transpose_b) - .set_default(false) - .describe("Whether the second input tensor is in transposed format."); - } -}; - -/*! \brief Attributes for sparse_dense operator */ -struct SparseDenseAttrs : public tvm::AttrsNode { - bool sparse_lhs; - - TVM_DECLARE_ATTRS(SparseDenseAttrs, "relay.attrs.SparseDenseAttrs") { - TVM_ATTR_FIELD(sparse_lhs) - .set_default(false) - .describe( - "Indicate whether sparse matrix is multiplied on the right or the left. If true, then " - "the operation is S * D^T (D dense, S sparse). If false, the operation is D * S^T"); - } -}; - -/*! \brief Attributes for sparse_transpose operator */ -struct SparseTransposeAttrs : public tvm::AttrsNode { - TVM_DECLARE_ATTRS(SparseTransposeAttrs, "relay.attrs.SparseTransposeAttrs") {} -}; - -/*! \brief Attributes for sparse_dense operator */ -struct SparseConv2DAttrs : public tvm::AttrsNode { - std::string layout; - Array kernel_size; - - TVM_DECLARE_ATTRS(SparseConv2DAttrs, "relay.attrs.SparseConv2DAttrs") { - TVM_ATTR_FIELD(layout).set_default("NHWC").describe( - "Dimension ordering of input data. Can be 'NCHW', 'NHWC'" - "'N', 'C', 'H', 'W' stands for batch, channel, height, and width" - "dimensions respectively."); - TVM_ATTR_FIELD(kernel_size) - .set_default(Array{1, 1}) - .describe("Kernel size for SparseConv2D, 1x1 or 3x3. "); - } -}; - -/*! \brief Attributes for FIFO buffer operator */ -struct FIFOBufferAttrs : public tvm::AttrsNode { - int axis; - - TVM_DECLARE_ATTRS(FIFOBufferAttrs, "relay.attrs.FIFOBufferAttrs") { - TVM_ATTR_FIELD(axis).set_default(0); - } -}; - -/*! \brief Attributes for upsampling operator */ -struct UpSamplingAttrs : public tvm::AttrsNode { - double scale_h; - double scale_w; - tvm::String layout; - tvm::String method; - bool align_corners; - - TVM_DECLARE_ATTRS(UpSamplingAttrs, "relay.attrs.UpSamplingAttrs") { - TVM_ATTR_FIELD(scale_h).describe("The upsampling factor for height"); - TVM_ATTR_FIELD(scale_w).describe("The upsampling factor for width"); - TVM_ATTR_FIELD(layout).set_default("NCHW").describe( - "Dimension ordering of input data. Can be 'NCHW', 'NHWC', etc." - "'N', 'C', 'H', 'W' stands for batch, channel, height, and width" - "dimensions respectively. Upsampling is applied on the 'H' and" - "'W' dimensions."); - TVM_ATTR_FIELD(method) - .set_default("nearest_neighbor") - .describe( - "Specify the mode to use for scaling." - "nearest_neighbor - Nearest Neighbor" - "bilinear - Bilinear Interpolation" - "bicubic - Bicubic Interpolation"); - TVM_ATTR_FIELD(align_corners) - .set_default(false) - .describe("Should be true to preserve the values at the corner pixels"); - } -}; - -/*! \brief Attributes for upsampling3d operator */ -struct UpSampling3DAttrs : public tvm::AttrsNode { - double scale_d; - double scale_h; - double scale_w; - std::string layout; - std::string method; - std::string coordinate_transformation_mode; - - TVM_DECLARE_ATTRS(UpSampling3DAttrs, "relay.attrs.UpSampling3DAttrs") { - TVM_ATTR_FIELD(scale_d).describe("The upsampling factor for depth"); - TVM_ATTR_FIELD(scale_h).describe("The upsampling factor for height"); - TVM_ATTR_FIELD(scale_w).describe("The upsampling factor for width"); - TVM_ATTR_FIELD(layout).set_default("NCDHW").describe( - "Dimension ordering of input data. Can be 'NCDHW', 'NDHWC', etc." - "'N', 'C', 'D', 'H', 'W' stands for batch, channel, depth, height, and width" - "dimensions respectively. Upsampling is applied on the 'D', 'H' and" - "'W' dimensions."); - TVM_ATTR_FIELD(method) - .set_default("nearest_neighbor") - .describe( - "Specify the mode to use for scaling." - "nearest_neighbor - Nearest Neighbor" - "trilinear - Trilinear Interpolation"); - TVM_ATTR_FIELD(coordinate_transformation_mode) - .set_default("half_pixel") - .describe( - "Describes how to transform the coordinate in the resized tensor" - "to the coordinate in the original tensor." - "Refer to the ONNX Resize operator specification for details" - "Available options are half_pixel, align_corners and asymmetric"); - } -}; - -/*! \brief Attributes used for the padding operator */ -struct PadAttrs : public tvm::AttrsNode { - Array> pad_width; - tvm::String pad_mode; - - TVM_DECLARE_ATTRS(PadAttrs, "relay.attrs.PadAttrs") { - TVM_ATTR_FIELD(pad_width).describe( - "Number of values padded to the edges of each axis, " - "in the format of ((before_1, after_1), ..., (before_N, after_N))"); - TVM_ATTR_FIELD(pad_mode) - .set_default("constant") - .describe( - "Padding type to use. \"constant\" pads with constant_value, " - "\"edge\" pads using the edge values of the input array, " - "\"reflect\" pads by reflecting values with respect to the edges."); - } -}; - -/*! \brief Attributes used for the MirrorPadding operator */ -struct MirrorPadAttrs : public tvm::AttrsNode { - std::string mode; - Array> pad_width; - - TVM_DECLARE_ATTRS(MirrorPadAttrs, "relay.attrs.MirrorPadAttrs") { - TVM_ATTR_FIELD(mode) - .set_default("SYMMETRIC") - .describe("Specifies how mirroring should be performed."); - TVM_ATTR_FIELD(pad_width).describe( - "Number of values padded to the edges of each axis, " - "in the format of ((before_1, after_1), ..., (before_N, after_N))"); - } -}; - -/*! \brief Attributes for leaky relu operator */ -struct LeakyReluAttrs : public tvm::AttrsNode { - double alpha; - - TVM_DECLARE_ATTRS(LeakyReluAttrs, "relay.attrs.LeakyReluAttrs") { - TVM_ATTR_FIELD(alpha).set_lower_bound(0.0).set_default(0.25).describe( - "Slope coefficient for the negative half axis."); - } -}; - -/*! \brief Attributes for prelu operator */ -struct PReluAttrs : public tvm::AttrsNode { - int axis; - - TVM_DECLARE_ATTRS(PReluAttrs, "relay.attrs.PReluAttrs") { - TVM_ATTR_FIELD(axis).set_default(1).describe( - "Specify which shape axis the channel is specified."); - } -}; - -/*! \brief Attributes used in dropout operator */ -struct DropoutAttrs : public tvm::AttrsNode { - double rate; - TVM_DECLARE_ATTRS(DropoutAttrs, "relay.attrs.DropoutAttrs") { - TVM_ATTR_FIELD(rate) - .describe("Fraction of the input that gets dropped out during training time") - .set_default(0.5); - } -}; // struct DropoutAttrs - -/*! \brief Attributes used in batch_norm operator */ -struct BatchNormAttrs : public tvm::AttrsNode { - int axis; - double epsilon; - bool center; - bool scale; - - TVM_DECLARE_ATTRS(BatchNormAttrs, "relay.attrs.BatchNormAttrs") { - TVM_ATTR_FIELD(axis).describe("Specify which shape axis denotes the channel.").set_default(1); - TVM_ATTR_FIELD(epsilon) - .describe("Small float added to variance to avoid dividing by zero") - .set_default(1e-5); - TVM_ATTR_FIELD(center) - .describe("If True, add offset of beta to normalized tensor. If False, beta is ignored") - .set_default(true); - TVM_ATTR_FIELD(scale) - .describe( - "If True, multiply by gamma. If False, gamma is not used. " - "When the next layer is piecewise linear (also, e.g., nn.relu), " - "this can be disabled since the scaling will be done by the next layer.") - .set_default(true); - } -}; // struct BatchNormAttrs - -/*! \brief Attributes used in instance_norm operator */ -struct InstanceNormAttrs : public tvm::AttrsNode { - int axis; - double epsilon; - bool center; - bool scale; - - TVM_DECLARE_ATTRS(InstanceNormAttrs, "relay.attrs.InstanceNormAttrs") { - TVM_ATTR_FIELD(axis).describe("Specify which shape axis denotes the channel.").set_default(1); - TVM_ATTR_FIELD(epsilon) - .describe("Small float added to variance to avoid dividing by zero") - .set_default(1e-5); - TVM_ATTR_FIELD(center).set_default(true).describe( - "If true, add offset of beta to normalized tensor; " - "otherwise, beta is ignored."); - TVM_ATTR_FIELD(scale).set_default(true).describe( - "If true, multiply by gamma; otherwise, gamma is ignored."); - } -}; // struct InstanceNormAttrs - -/*! \brief Attributes used in layer_norm operator */ -struct LayerNormAttrs : public tvm::AttrsNode { - int axis; - double epsilon; - bool center; - bool scale; - - TVM_DECLARE_ATTRS(LayerNormAttrs, "relay.attrs.LayerNormAttrs") { - TVM_ATTR_FIELD(axis).set_default(-1).describe("Specify which shape axis denotes the channel."); - TVM_ATTR_FIELD(epsilon).set_default(1e-5).describe( - "Small float added to variance to avoid dividing by zero"); - TVM_ATTR_FIELD(center).set_default(true).describe( - "If true, add offset of beta to normalized tensor; " - "otherwise, beta is ignored."); - TVM_ATTR_FIELD(scale).set_default(true).describe( - "If true, multiply by gamma; otherwise, gamma is ignored."); - } -}; // struct LayerNormAttrs - -/*! \brief Attributes used in group_norm operator */ -struct GroupNormAttrs : public tvm::AttrsNode { - int num_groups; - int axis; - double epsilon; - bool center; - bool scale; - - TVM_DECLARE_ATTRS(GroupNormAttrs, "relay.attrs.GroupNormAttrs") { - TVM_ATTR_FIELD(num_groups) - .set_default(0) - .describe("Specify number of groups to separate the channels into."); - TVM_ATTR_FIELD(axis).set_default(1).describe("Specify which shape axis denotes the channel."); - TVM_ATTR_FIELD(epsilon).set_default(1e-5).describe( - "Small float added to variance to avoid dividing by zero"); - TVM_ATTR_FIELD(center).set_default(true).describe( - "If true, add offset of beta to normalized tensor; " - "otherwise, beta is ignored."); - TVM_ATTR_FIELD(scale).set_default(true).describe( - "If true, multiply by gamma; otherwise, gamma is ignored."); - } -}; // struct GroupNormAttrs - -/*! \brief Attributes for LRN operator */ -struct LRNAttrs : public tvm::AttrsNode { - int size; - int axis; - double bias; - double alpha; - double beta; - - TVM_DECLARE_ATTRS(LRNAttrs, "relay.attrs.LRNAttrs") { - TVM_ATTR_FIELD(size).set_default(5).describe( - "The size of the local region to be considered for normalization."); - TVM_ATTR_FIELD(axis).set_default(1).describe("Axis of input data layout channel."); - TVM_ATTR_FIELD(bias).set_default(2).describe("The offset parameter to avoid division by 0."); - TVM_ATTR_FIELD(alpha).set_default(0.0001).describe("The scaling parameter."); - TVM_ATTR_FIELD(beta).set_default(0.75).describe("The exponent parameter."); - } -}; - -/*! \brief Attributes for L2Normalize operator */ -struct L2NormalizeAttrs : public tvm::AttrsNode { - double eps; - Array axis; - - TVM_DECLARE_ATTRS(L2NormalizeAttrs, "relay.attrs.L2NormalizeAttrs") { - TVM_ATTR_FIELD(eps).describe("A lower bound value for the norm, to avoid division by 0."); - TVM_ATTR_FIELD(axis).describe("Axis over the normalization applied."); - } -}; - -/*! \brief Attributes for DeformableConv2D operator */ -struct DeformableConv2DAttrs : public tvm::AttrsNode { - Array strides; - Array padding; - Array dilation; - int deformable_groups; - int groups; - IndexExpr channels; - Array kernel_size; - std::string data_layout; - std::string kernel_layout; - std::string out_layout; - DataType out_dtype; - - TVM_DECLARE_ATTRS(DeformableConv2DAttrs, "relay.attrs.DeformableConv2DAttrs") { - TVM_ATTR_FIELD(strides) - .set_default(Array({1, 1})) - .describe("Specifies the strides of the convolution."); - TVM_ATTR_FIELD(padding) - .set_default(Array({0, 0})) - .describe( - "If padding is non-zero, then the input is implicitly zero-padded" - "Padding support both symmetric and asymmetric as" - "one int : same padding used on all sides" - "two int : bottom, right will use same padding as top, left" - "four int : padding width in the order of (top, left, bottom, right)"); - TVM_ATTR_FIELD(dilation) - .set_default(Array({1, 1})) - .describe("Specifies the dilation rate to use for dilated convolution."); - TVM_ATTR_FIELD(deformable_groups) - .set_default(1) - .describe( - "Controls the connections between inputs and offsets." - "Input channels are partitioned into multiple deformable groups. Offsets" - "are shared across input channels in the same deformable group."); - TVM_ATTR_FIELD(groups).set_default(1).describe( - "Controls the connections between inputs and outputs." - "At groups=1, all inputs are convolved to all outputs." - "At groups=2, the operation becomes equivalent to having two convolution" - "layers side by side, each seeing half the input channels, and producing" - "half the output channels, and both subsequently concatenated."); - TVM_ATTR_FIELD(channels) - .describe( - "The number of output channels in the convolution." - " If it is not set, inferred by shape of the weight.") - .set_default(NullValue()); - TVM_ATTR_FIELD(kernel_size) - .describe("Specifies the dimensions of the convolution window.") - .set_default(NullValue>()); - TVM_ATTR_FIELD(data_layout) - .set_default("NCHW") - .describe( - "Dimension ordering of input data. Can be 'NCHW', 'NHWC', etc." - "'N', 'C', 'H', 'W' stands for batch, channel, height, and width" - "dimensions respectively. Convolution is applied on the 'H' and" - "'W' dimensions."); - TVM_ATTR_FIELD(kernel_layout) - .set_default("OIHW") - .describe( - "Dimension ordering of weight. Can be 'OIHW', 'OIHW16o16i', etc." - "'O', 'I', 'H', 'W' stands for num_filter, input_channel, height, and width" - "dimensions respectively."); - TVM_ATTR_FIELD(out_layout) - .set_default("") - .describe( - "Dimension ordering of output. Can be 'NCHW', 'NHWC', etc." - "'N', 'C', 'H', 'W' stands for batch, channel, height, and width" - "dimensions respectively. Default to be same as input layout."); - - // use 0 bits to indicate none. - TVM_ATTR_FIELD(out_dtype) - .set_default(NullValue()) - .describe("Output data type, set to explicit type under mixed precision setting"); - } -}; - -/*! \brief Attributes used in subpixel operators */ -struct SubPixelAttrs : public tvm::AttrsNode { - int block_size; - std::string layout; - std::string mode; - - TVM_DECLARE_ATTRS(SubPixelAttrs, "relay.attrs.SubPixelAttrs") { - TVM_ATTR_FIELD(block_size) - .describe("The size of subpixel blocks to compose or decompose.") - .set_default(1); - TVM_ATTR_FIELD(layout).set_default("NCHW").describe( - "Dimension ordering of input data. Can be 'NCHW', 'NHWC', etc." - "'N', 'C', 'H', 'W' stands for batch, channel, height, and width" - "dimensions respectively."); - TVM_ATTR_FIELD(mode).set_default("DCR").describe( - "Indicates order in which channels are accessed. Must be one of" - "DCR or CDR."); - } -}; // struct SubPixelAttrs - -/*! \brief Attributes used in correlation operators */ -struct CorrelationAttrs : public tvm::AttrsNode { - int kernel_size; - int max_displacement; - int stride1; - int stride2; - Array padding; - bool is_multiply; - String layout; - - TVM_DECLARE_ATTRS(CorrelationAttrs, "relay.attrs.CorrelationAttrs") { - TVM_ATTR_FIELD(kernel_size) - .describe("Kernel size for correlation, must be an odd number.") - .set_default(1); - TVM_ATTR_FIELD(max_displacement).describe("Max displacement of Correlation.").set_default(1); - TVM_ATTR_FIELD(stride1).describe("Stride for data1.").set_default(1); - TVM_ATTR_FIELD(stride2).describe("Stride for data2.").set_default(1); - TVM_ATTR_FIELD(padding) - .describe("Padding for data1 and data2.") - .set_default(Array{0, 0}); - TVM_ATTR_FIELD(is_multiply) - .describe("Operation type is either multiplication or substraction.") - .set_default(true); - TVM_ATTR_FIELD(layout).set_default("NCHW").describe( - "Dimension ordering of input data. Can be 'NCHW', 'NHWC', etc." - "'N', 'C', 'H', 'W' stands for batch, channel, height, and width" - "dimensions respectively."); - } -}; // struct CorrelationAttrs - -/*! \brief Attributes used in SpaceToBatchND operator */ -struct SpaceToBatchNDAttrs : public tvm::AttrsNode { - Array block_shape; - Array> paddings; - double pad_value; - - TVM_DECLARE_ATTRS(SpaceToBatchNDAttrs, "relay.attrs.SpaceToBatchNDAttrs") { - TVM_ATTR_FIELD(block_shape) - .set_default(Array({1, 1})) - .describe("1-D containing block size for each spatial dimension."); - TVM_ATTR_FIELD(paddings).describe("2-D containing paddings for each spatial dimension."); - TVM_ATTR_FIELD(pad_value).set_default(0.0).describe("The value used for padding."); - } -}; // struct SpaceToBatchNDAttrs - -/*! \brief Attributes used in BatchToSpaceND operator */ -struct BatchToSpaceNDAttrs : public tvm::AttrsNode { - Array block_shape; - Array> crops; - - TVM_DECLARE_ATTRS(BatchToSpaceNDAttrs, "relay.attrs.BatchToSpaceNDAttrs") { - TVM_ATTR_FIELD(block_shape) - .set_default(Array({1, 1})) - .describe("1-D containing block size for each spatial dimension."); - TVM_ATTR_FIELD(crops).describe("2-D containing amount to crop from spatial dimension."); - } -}; // struct BatchToSpaceNDAttrs - -/*! \brief Attributes used in NLLLoss operator */ -struct NLLLossAttrs : public tvm::AttrsNode { - std::string reduction; - int ignore_index; - - TVM_DECLARE_ATTRS(NLLLossAttrs, "relay.attrs.NLLLossAttrs") { - TVM_ATTR_FIELD(reduction).set_default("mean").describe( - "The reduction method to apply to the output. Can be" - "'none', 'mean' or 'sum'."); - TVM_ATTR_FIELD(ignore_index).describe("The target value to ignore."); - } -}; // struct NLLLossAttrs - -} // namespace relay -} // namespace tvm -#endif // TVM_RELAY_ATTRS_NN_H_ diff --git a/include/tvm/relay/attrs/on_device.h b/include/tvm/relay/attrs/on_device.h deleted file mode 100644 index 3facc3a597f1..000000000000 --- a/include/tvm/relay/attrs/on_device.h +++ /dev/null @@ -1,106 +0,0 @@ -/* - * 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 tvm/relay/attrs/on_device.h - * \brief Attribute for the "on_device" annotation (ie operator). - */ -#ifndef TVM_RELAY_ATTRS_ON_DEVICE_H_ -#define TVM_RELAY_ATTRS_ON_DEVICE_H_ - -#include -#include - -#include - -namespace tvm { -namespace relay { - -/*! - * \brief Attributes for the "on_device" annotation (ie operator). - * - * The Relay call: - * \code - * on_device(sub_expr, virtual_device=S) - * \endcode - * constrains \p sub_expr to execute and store its result on the \p VirtualDevice \p S. - * However the annotation itself may appear in an expression to be executed and stored on a - * different \p VirtualDevice. If so the compiler will automatically insert a "device_copy" call to - * mediate the transition between \p VirtualDevices. - * - * E.g.: Assuming %x and %y reside on the GPU and %z on the CPU then: - * \code - * multiply(on_device(add(%x, %y), virtual_device=GPU), %z) - * \endcode - * indicates the \p add should execute on the GPU but the \p multiply should execute on the CPU. - * The compiler will rewrite this to: - * \code - * multiply(device_copy(add(%x, %y), src_virtual_device=GPU, dst_virtual_device=CPU), %z) - * \endcode - * - * The \p constraint_body (default true) and \p constraint_result (default false) fields can be - * used by passes for finer-grained control over how the \p VirtualDevice constraint should be - * applied. - */ -struct OnDeviceAttrs : public tvm::AttrsNode { - /*! - * \brief The \p VirtualDevice to constraint to apply to the body, result, or both body and result - * of the "on_device" call. - */ - VirtualDevice virtual_device = VirtualDevice::FullyUnconstrained(); - - /*! - * \brief If false (the default), the result of the "on_device" call is not constrained to be - * \p virtual_device. - */ - bool constrain_result = false; - - /*! - * \brief If true (the default), the body of the "on_device" call is constrained to be \p - * virtual_device. - */ - bool constrain_body = true; - - /*! - * \brief Returns true if both the body and result are constrained. - */ - bool is_fixed() const { return constrain_result && constrain_body; } - - /*! - * \brief Returns true only the body is constrained (the 'normal' case). - */ - bool is_normal() const { return !constrain_result && constrain_body; } - - TVM_DECLARE_ATTRS(OnDeviceAttrs, "relay.attrs.OnDeviceAttrs") { - TVM_ATTR_FIELD(virtual_device) - .describe("The (virtual) device to constrain to.") - .set_default(VirtualDevice::FullyUnconstrained()); - TVM_ATTR_FIELD(constrain_result) - .describe("Whether the constraint applies to the overall expression") - .set_default(false); - TVM_ATTR_FIELD(constrain_body) - .describe("Whether the constraint applies to the body sub-expression.") - .set_default(true); - } -}; - -} // namespace relay -} // namespace tvm - -#endif // TVM_RELAY_ATTRS_ON_DEVICE_H_ diff --git a/include/tvm/relay/attrs/random.h b/include/tvm/relay/attrs/random.h deleted file mode 100644 index 4bdac4c1763d..000000000000 --- a/include/tvm/relay/attrs/random.h +++ /dev/null @@ -1,76 +0,0 @@ -/* - * 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 tvm/relay/attrs/vision.h - * \brief Auxiliary attributes for random operators. - */ -#ifndef TVM_RELAY_ATTRS_RANDOM_H_ -#define TVM_RELAY_ATTRS_RANDOM_H_ - -#include - -namespace tvm { -namespace relay { - -struct ThreefryGenerateAttrs : public tvm::AttrsNode { - Array out_shape; - - TVM_DECLARE_ATTRS(ThreefryGenerateAttrs, "relay.attrs.ThreefryGenerateAttrs") { - TVM_ATTR_FIELD(out_shape).describe("Shape of random numbers to generate"); - } -}; - -struct UniformAttrs : public tvm::AttrsNode { - Array out_shape; - DataType out_dtype; - - TVM_DECLARE_ATTRS(UniformAttrs, "relay.attrs.UniformAttrs") { - TVM_ATTR_FIELD(out_shape).describe("Shape of random numbers to generate"); - TVM_ATTR_FIELD(out_dtype) - .set_default(NullValue()) - .describe("Data type of the generated numbers"); - } -}; - -struct NormalAttrs : public tvm::AttrsNode { - Array out_shape; - DataType out_dtype; - - TVM_DECLARE_ATTRS(NormalAttrs, "relay.attrs.NormalAttrs") { - TVM_ATTR_FIELD(out_shape).describe("Shape of random numbers to generate"); - TVM_ATTR_FIELD(out_dtype) - .set_default(NullValue()) - .describe("Data type of the generated numbers"); - } -}; - -struct MultinomialAttrs : public tvm::AttrsNode { - Integer num_samples; - - TVM_DECLARE_ATTRS(MultinomialAttrs, "relay.attrs.MultinomialAttrs") { - TVM_ATTR_FIELD(num_samples) - .set_default(1) - .describe("Number of samples to draw from the distribution."); - } -}; - -} // namespace relay -} // namespace tvm -#endif // TVM_RELAY_ATTRS_RANDOM_H_ diff --git a/include/tvm/relay/attrs/reduce.h b/include/tvm/relay/attrs/reduce.h deleted file mode 100644 index d91b3594b5a3..000000000000 --- a/include/tvm/relay/attrs/reduce.h +++ /dev/null @@ -1,132 +0,0 @@ -/* - * 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 tvm/relay/attrs/reduce.h - * \brief Auxiliary attributes for reduce operators. - */ -#ifndef TVM_RELAY_ATTRS_REDUCE_H_ -#define TVM_RELAY_ATTRS_REDUCE_H_ - -#include - -#include - -namespace tvm { -namespace relay { - -/*! \brief Attributes for Reduce operators */ -struct ReduceAttrs : public tvm::AttrsNode { - Array axis; - bool keepdims; - bool exclude; - - TVM_DECLARE_ATTRS(ReduceAttrs, "relay.attrs.ReduceAttrs") { - TVM_ATTR_FIELD(axis) - .set_default(NullValue>()) - .describe(R"code(The axis or axes along which to perform the reduction. - - The default, `axis=()`, will compute over all elements into a - scalar array with shape `(1,)`. - - If `axis` is int, a reduction is performed on a particular axis. - - If `axis` is a tuple of ints, a reduction is performed on all the axes - specified in the tuple. - - If `exclude` is true, reduction will be performed on the axes that are - NOT in axis instead.)code"); - - TVM_ATTR_FIELD(keepdims).set_default(false).describe( - "If this is set to `True`, the reduced axes are left " - "in the result as dimension with size one."); - TVM_ATTR_FIELD(exclude).set_default(false).describe( - "Whether to perform reduction on axis that are NOT in axis instead."); - } -}; - -/*! \brief Attributes for Reduce operators which reduce by finding a single element. E.g. argmin */ -struct ArgReduceAttrs : public tvm::AttrsNode { - Array axis; - bool keepdims; - bool select_last_index; - bool exclude; - - TVM_DECLARE_ATTRS(ArgReduceAttrs, "relay.attrs.ArgReduceAttrs") { - TVM_ATTR_FIELD(axis) - .set_default(NullValue>()) - .describe(R"code(The axis or axes along which to perform the reduction. - - The default, `axis=()`, will compute over all elements into a - scalar array with shape `(1,)`. - - If `axis` is int, a reduction is performed on a particular axis. - - If `axis` is a tuple of ints, a reduction is performed on all the axes - specified in the tuple. - - If `exclude` is true, reduction will be performed on the axes that are - NOT in axis instead.)code"); - - TVM_ATTR_FIELD(keepdims).set_default(false).describe( - "If this is set to `True`, the reduced axes are left " - "in the result as dimension with size one."); - TVM_ATTR_FIELD(select_last_index) - .set_default(false) - .describe( - "Whether to select the last index if the target element appears multiple times, else " - "select the first index which the target element appears"); - TVM_ATTR_FIELD(exclude).set_default(false).describe( - "Whether to perform reduction on axis that are NOT in axis instead."); - } -}; - -struct VarianceAttrs : public tvm::AttrsNode { - Array axis; - bool keepdims; - bool exclude; - bool unbiased; - - TVM_DECLARE_ATTRS(VarianceAttrs, "relay.attrs.VarianceAttrs") { - TVM_ATTR_FIELD(axis) - .set_default(NullValue>()) - .describe(R"code(The axis or axes along which to perform the reduction. - - The default, `axis=()`, will compute over all elements into a - scalar array with shape `(1,)`. - - If `axis` is int, a reduction is performed on a particular axis. - - If `axis` is a tuple of ints, a reduction is performed on all the axes - specified in the tuple. - - If `exclude` is true, reduction will be performed on the axes that are - NOT in axis instead.)code"); - - TVM_ATTR_FIELD(keepdims).set_default(false).describe( - "If this is set to `True`, the reduced axes are left " - "in the result as dimension with size one."); - TVM_ATTR_FIELD(exclude).set_default(false).describe( - "Whether to perform reduction on axis that are NOT in axis instead."); - TVM_ATTR_FIELD(unbiased).set_default(false).describe("Whether to use the unbiased estimation."); - } -}; -} // namespace relay -} // namespace tvm -#endif // TVM_RELAY_ATTRS_REDUCE_H_ diff --git a/include/tvm/relay/attrs/transform.h b/include/tvm/relay/attrs/transform.h deleted file mode 100644 index 91020fc7443b..000000000000 --- a/include/tvm/relay/attrs/transform.h +++ /dev/null @@ -1,614 +0,0 @@ -/* - * 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 tvm/relay/attrs/transform.h - * \brief Transform operators. - */ -#ifndef TVM_RELAY_ATTRS_TRANSFORM_H_ -#define TVM_RELAY_ATTRS_TRANSFORM_H_ - -#include -#include -#include -#include - -#include - -namespace tvm { -namespace relay { - -/*! \brief Attributes used for the sliding_window operator */ -struct SlidingWindowAttrs : public tvm::AttrsNode { - int axis; - Array window_shape; - Array strides; - TVM_DECLARE_ATTRS(SlidingWindowAttrs, "relay.attrs.SlidingWindowAttrs") { - TVM_ATTR_FIELD(axis).describe( - "What axis the sliding window begin forming over." - "Window will be slid over this axis and all following axes." - "The axis value determines the window shape (and thus, the" - "number of strides):" - "window shape and strides must both be of length" - "`data.ndim-axis`."); - TVM_ATTR_FIELD(window_shape) - .describe( - "The window shape to form over the input." - "Window shape must be of length `data.ndim-axis`."); - TVM_ATTR_FIELD(strides).describe( - "How to stride the window along each dimension." - "Strides must be of length `data.ndim-axis`."); - } -}; // struct SlidingWindowAttrs - -/*! \brief data type cast */ -struct CastAttrs : public tvm::AttrsNode { - DataType dtype; - - TVM_DECLARE_ATTRS(CastAttrs, "relay.attrs.CastAttrs") { - TVM_ATTR_FIELD(dtype).describe("Target data type"); - } -}; // struct CastAttrs. - -/*! \brief Attributes used in expand_dims operators */ -struct ExpandDimsAttrs : public tvm::AttrsNode { - int axis; - int num_newaxis; - - TVM_DECLARE_ATTRS(ExpandDimsAttrs, "relay.attrs.ExpandDimsAttrs") { - TVM_ATTR_FIELD(axis).describe( - "The axis at which the input array is expanded." - "Should lie in range `[-data.ndim - 1, data.ndim]`." - "If `axis < 0`, it is the first axis inserted;" - "If `axis >= 0`, it is the last axis inserted in Python's negative indexing."); - TVM_ATTR_FIELD(num_newaxis) - .describe("Number of axes to be inserted. Should be >= 0.") - .set_lower_bound(0) - .set_default(1); - } -}; // struct ExpandDimsAttrs - -/*! \brief Attributes used in dynamic expand_dims operators */ -struct DynExpandDimsAttrs : public tvm::AttrsNode { - int num_newaxis; - - TVM_DECLARE_ATTRS(DynExpandDimsAttrs, "relay.attrs.DynExpandDimsAttrs") { - TVM_ATTR_FIELD(num_newaxis) - .describe("Number of axes to be inserted. Should be >= 0.") - .set_lower_bound(0) - .set_default(1); - } -}; // struct ExpandDimsAttrs - -/*! \brief Attributes used in concatenate operators */ -struct ConcatenateAttrs : public tvm::AttrsNode { - int axis; - TVM_DECLARE_ATTRS(ConcatenateAttrs, "relay.attrs.ConcatenateAttrs") { - TVM_ATTR_FIELD(axis) - .describe( - "The axis at which the input arrays are concatenated." - "Should lie in range `[-ndim, ndim)`.") - .set_default(0); - } -}; // struct ConcatenateAttrs - -/*! \brief Attributes used in transpose operators */ -struct TransposeAttrs : public tvm::AttrsNode { - Array axes; - TVM_DECLARE_ATTRS(TransposeAttrs, "relay.attrs.TransposeAttrs") { - TVM_ATTR_FIELD(axes).describe("The target axes order, reverse order if not specified."); - } -}; // struct TransposeAttrs - -/*! \brief Attributes used in reshape operators */ -struct ReshapeAttrs : public tvm::AttrsNode { - Array newshape; - bool allowzero; - TVM_DECLARE_ATTRS(ReshapeAttrs, "relay.attrs.ReshapeAttrs") { - TVM_ATTR_FIELD(newshape).describe( - "The new shape. Should be compatible with the original shape."); - TVM_ATTR_FIELD(allowzero).set_default(0).describe( - "Whether to honor the value of zero in newshape."); - } -}; // struct ReshapeAttrs - -/*! \brief Attributes used in MXNet-style reshape_like operators */ -struct ReshapeLikeAttrs : public tvm::AttrsNode { - int lhs_begin; - Integer lhs_end; // can be None - int rhs_begin; - Integer rhs_end; // can be None - TVM_DECLARE_ATTRS(ReshapeLikeAttrs, "relay.attrs.ReshapeLikeAttrs") { - TVM_ATTR_FIELD(lhs_begin).set_default(0).describe( - "The axis of the input where reshaping should begin."); - TVM_ATTR_FIELD(lhs_end) - .set_default(NullValue()) - .describe("The axis of the input where reshaping should end, exclusive."); - TVM_ATTR_FIELD(rhs_begin).set_default(0).describe( - "The axis of the shape_like tensor to begin taking dimensions from."); - TVM_ATTR_FIELD(rhs_end) - .set_default(NullValue()) - .describe("The axis of the shape_like tensor to end taking dimensions from, exclusive."); - } -}; // struct ReshapeLikeAttrs - -struct ScatterElementsAttrs : public tvm::AttrsNode { - Integer axis; - String reduction; - - TVM_DECLARE_ATTRS(ScatterElementsAttrs, "relay.attrs.ScatterElementsAttrs") { - TVM_ATTR_FIELD(axis).set_default(0).describe("The axis over which to select values."); - TVM_ATTR_FIELD(reduction).set_default("update").describe( - "Reduction mode of the scatter elements, " - "either \"update\", \"add\", \"mul\", \"mean\", \"min\" or \"max\"."); - } -}; - -struct ScatterNDAttrs : public tvm::AttrsNode { - String mode; - - TVM_DECLARE_ATTRS(ScatterNDAttrs, "relay.attrs.ScatterNDAttrs") { - TVM_ATTR_FIELD(mode).set_default("update").describe( - "Accumulation mode of the ScatterND, " - "either \"update\", \"add\", \"mul\", \"min\" or \"max\"."); - } -}; - -struct GatherAttrs : public tvm::AttrsNode { - Integer axis; - - TVM_DECLARE_ATTRS(GatherAttrs, "relay.attrs.GatherAttrs") { - TVM_ATTR_FIELD(axis) - .set_default(NullValue()) - .describe("The axis over which to select values."); - } -}; - -struct GatherNDAttrs : public tvm::AttrsNode { - Integer batch_dims; - Optional index_rank; - - TVM_DECLARE_ATTRS(GatherNDAttrs, "relay.attrs.GatherNDAttrs") { - TVM_ATTR_FIELD(batch_dims).set_default(Integer(0)).describe("The number of batch dimensions."); - TVM_ATTR_FIELD(index_rank) - .set_default(NullValue()) - .describe( - "The size of an indexing tuple, which is a fixed value. Only needed when the number of " - "indexting tuples is dynamic."); - } -}; - -struct TakeAttrs : public tvm::AttrsNode { - Integer batch_dims; - Integer axis; - tvm::String mode; - - TVM_DECLARE_ATTRS(TakeAttrs, "relay.attrs.TakeAttrs") { - TVM_ATTR_FIELD(batch_dims) - .set_default(0) - .describe("The batch_dims over which to select values."); - TVM_ATTR_FIELD(axis) - .set_default(NullValue()) - .describe("The axis over which to select values."); - TVM_ATTR_FIELD(mode).set_default("clip").describe( - "Specify how out-of-bound indices will behave." - "clip - clip to the range (default)" - "wrap - wrap around the indices" - "fast - no clip or wrap around (user must make sure indices are in-bound)"); - } -}; - -/*! \brief Attributes that specify a tensor */ -struct InitOpAttrs : public tvm::AttrsNode { - Optional> shape; - DataType dtype; - - TVM_DECLARE_ATTRS(InitOpAttrs, "relay.attrs.InitOpAttrs") { - TVM_ATTR_FIELD(shape).describe("Target shape."); - TVM_ATTR_FIELD(dtype).describe("Target data type.").set_default(NullValue()); - } -}; // struct InitOpAttrs - -/*! \brief Attributes used in arange operators */ -struct ArangeAttrs : public tvm::AttrsNode { - Expr start; - Expr stop; - Expr step; - DataType dtype; - - TVM_DECLARE_ATTRS(ArangeAttrs, "relay.attrs.ArangeAttrs") { - TVM_ATTR_FIELD(start).describe("Start of interval. The interval includes this value."); - TVM_ATTR_FIELD(stop).describe("Stop of interval. The interval does not include this value."); - TVM_ATTR_FIELD(step).describe("Spacing between values."); - TVM_ATTR_FIELD(dtype).describe("Target data type."); - } -}; // struct ArangeAttrs - -/*! \brief Attributes used in meshgrid operators */ -struct MeshgridAttrs : public tvm::AttrsNode { - std::string indexing; - - TVM_DECLARE_ATTRS(MeshgridAttrs, "relay.attrs.MeshgridAttrs") { - TVM_ATTR_FIELD(indexing) - .describe( - "Indexing mode, either \"ij\" for matrix or \"xy\" for cartesian in which first two" - "dimensions are swapped.") - .set_default("ij"); - } -}; // struct MeshgridAttrs - -/*! \brief Attributes used in stack operators */ -struct StackAttrs : public tvm::AttrsNode { - Integer axis; - TVM_DECLARE_ATTRS(StackAttrs, "relay.attrs.StackAttrs") { - TVM_ATTR_FIELD(axis).set_default(0).describe( - "The axis in the result array along which the input arrays are stacked."); - } -}; // struct StackAttrs - -/*! \brief Attributes used in repeat operators */ -struct RepeatAttrs : public tvm::AttrsNode { - Integer repeats; - Integer axis; - TVM_DECLARE_ATTRS(RepeatAttrs, "relay.attrs.RepeatAttrs") { - TVM_ATTR_FIELD(repeats).describe("The number of repetitions for each element."); - TVM_ATTR_FIELD(axis) - .set_default(NullValue()) - .describe(" The axis along which to repeat values."); - } -}; // struct RepeatAttrs - -/*! \brief Attributes used in tile operators */ -struct TileAttrs : public tvm::AttrsNode { - Array reps; - TVM_DECLARE_ATTRS(TileAttrs, "relay.attrs.TileAttrs") { - TVM_ATTR_FIELD(reps).describe( - "The number of times for repeating the tensor a." - "Each dim sizeof reps must be a positive integer."); - } -}; // struct TileAttrs - -/*! \brief Attributes used in reverse operators */ -struct ReverseAttrs : public tvm::AttrsNode { - Integer axis; - TVM_DECLARE_ATTRS(ReverseAttrs, "relay.attrs.ReverseAttrs") { - TVM_ATTR_FIELD(axis) - .set_default(NullValue()) - .describe("The axis along which to reverse elements."); - } -}; // struct ReverseAttrs - -/*! \brief Attributes used in reverse_sequence operators */ -struct ReverseSequenceAttrs : public tvm::AttrsNode { - Integer seq_axis; - Integer batch_axis; - - TVM_DECLARE_ATTRS(ReverseSequenceAttrs, "relay.attrs.ReverseSequenceAttrs") { - TVM_ATTR_FIELD(seq_axis).set_default(1).describe( - "The seq axis along which to reverse elements."); - TVM_ATTR_FIELD(batch_axis) - .set_default(0) - .describe("The batch axis along which to slice the tensor."); - } -}; // struct ReverseSequenceAttrs - -/*! \brief Attributes used in squeeze operators */ -struct SqueezeAttrs : public tvm::AttrsNode { - // use axis to make the name numpy compatible. - Array axis; - - TVM_DECLARE_ATTRS(SqueezeAttrs, "relay.attrs.SqueezeAttrs") { - TVM_ATTR_FIELD(axis) - .describe( - "The axis to squeeze in the input tensor." - "If `axis = None`, all axis of dimension 1 get squeezed;" - "Else, the dimension in axes get squeezed." - "It is an error if an axis does not has dimension 1.") - .set_default(NullValue>()); - } -}; // struct SqueezeAttrs - -struct SplitAttrs : public tvm::AttrsNode { - Variant> indices_or_sections; - int axis; - - TVM_DECLARE_ATTRS(SplitAttrs, "relay.attrs.SplitAttrs") { - TVM_ATTR_FIELD(indices_or_sections) - .describe( - "Indices or sections to split into. Accepts an int or a tuple" - "If indices_or_sections is an integer, the input will be divided equally" - "along given axis. If such a split is not possible, an error is raised." - "If indices_or_sections is a tuple of sorted integers," - "the entries indicate where along axis the array is split."); - TVM_ATTR_FIELD(axis).set_default(0).describe("the axis to be splitted."); - } -}; - -/*! \brief Attributes for StridedSlice operator */ -struct StridedSliceAttrs : public tvm::AttrsNode { - Optional> begin; - Optional> end; - Optional> strides; - tvm::String slice_mode; - Optional> axes; - - TVM_DECLARE_ATTRS(StridedSliceAttrs, "relay.attrs.StridedSliceAttrs") { - TVM_ATTR_FIELD(begin).describe("Indices for begin of slice, begin index is also inclusive"); - TVM_ATTR_FIELD(end).describe("Indices for end of slice, end index is exclusive"); - TVM_ATTR_FIELD(strides).describe( - "Stride values of the slice, a stride can be negative, which causes a reverse slice."); - TVM_ATTR_FIELD(slice_mode) - .set_default("end") - .describe( - "The slice mode [end, size]." - "end - The default slice mode, ending indices for the slice." - "size - The input strides will be ignored, input end in this mode indicates the size" - "of a slice starting at the location specified by begin. If end[i] is -1," - "all remaining elements in that dimension are included in the slice"); - TVM_ATTR_FIELD(axes).describe( - "Axes along which slicing is applied. When it is specified, the length of begin, end, " - "strides, and axes must be equal."); - } -}; - -struct SliceLikeAttrs : public tvm::AttrsNode { - Array axes; - - TVM_DECLARE_ATTRS(SliceLikeAttrs, "relay.attrs.SliceLikeAttrs") { - TVM_ATTR_FIELD(axes).describe( - "List of axes on which input data will be sliced according to the " - "corresponding size of the second input. By default will slice " - "on all axes. Negative axes mean counting in reverse."); - } -}; - -/*! \brief Attributes for Clip operator */ -struct ClipAttrs : public tvm::AttrsNode { - double a_min; - double a_max; - - TVM_DECLARE_ATTRS(ClipAttrs, "relay.attrs.ClipAttrs") { - TVM_ATTR_FIELD(a_min).describe("The minimum clip value."); - TVM_ATTR_FIELD(a_max).describe("The maximum clip value."); - } -}; - -/*! \brief Attributes for FixedPointMultiply operator */ -struct FixedPointMultiplyAttrs : public tvm::AttrsNode { - int32_t multiplier; - int32_t shift; - - TVM_DECLARE_ATTRS(FixedPointMultiplyAttrs, "relay.attrs.FixedPointMultiplyAttrs") { - TVM_ATTR_FIELD(multiplier) - .describe("Multiplier of a fixed floating point number described as multiplier*2^(shift)"); - TVM_ATTR_FIELD(shift).describe( - "Shift of a fixed floating point number described as multiplier*2^(shift)"); - } -}; - -/*! \brief Attributes for per channel/per axes FixedPointMultiply operator */ -struct FixedPointMultiplyPerAxisAttrs : public tvm::AttrsNode { - bool is_lshift_required; - bool is_rshift_required; - Array axes; - - TVM_DECLARE_ATTRS(FixedPointMultiplyPerAxisAttrs, "relay.attrs.FixedPointMultiplyPerAxisAttrs") { - TVM_ATTR_FIELD(is_lshift_required) - .describe("Whether left shift is required in fixed point multiplication.") - .set_default(false); - TVM_ATTR_FIELD(is_rshift_required) - .describe("Whether right shift is required in fixed point multiplication.") - .set_default(false); - TVM_ATTR_FIELD(axes).describe("List of axes on which input data was quantized."); - } -}; - -/*! \brief Attributes for LayoutTransform operator */ -struct LayoutTransformAttrs : public tvm::AttrsNode { - std::string src_layout; - std::string dst_layout; - - TVM_DECLARE_ATTRS(LayoutTransformAttrs, "relay.attrs.LayoutTransformAttrs") { - TVM_ATTR_FIELD(src_layout).describe("The source layout of the tensor. (e.g. NCHW)"); - TVM_ATTR_FIELD(dst_layout).describe("The destination layout of the tensor. (e.g. NCHW16c)"); - } -}; - -/*! \brief Attributes for AutoSchedulerLayoutTransform operator */ -struct AutoSchedulerLayoutTransformAttrs - : public tvm::AttrsNode { - std::string src_layout; - std::string dst_layout; - - TVM_DECLARE_ATTRS(AutoSchedulerLayoutTransformAttrs, - "relay.attrs.AutoSchedulerLayoutTransformAttrs") { - TVM_ATTR_FIELD(src_layout).describe("The source layout of the tensor. (e.g. 1N32C112H112W)"); - TVM_ATTR_FIELD(dst_layout) - .describe("The destination layout of the tensor. (e.g. 1N2C112H112W16c)"); - } -}; - -/*! \brief Attributes for MetaScheduleLayoutTransform operator */ -struct MetaScheduleLayoutTransformAttrs : public tvm::AttrsNode { - tir::IndexMap index_map; - - TVM_DECLARE_ATTRS(MetaScheduleLayoutTransformAttrs, - "relay.attrs.MetaScheduleLayoutTransformAttrs") { - TVM_ATTR_FIELD(index_map).describe( - "The order of the extents, for example, " - "let extents = [2, 3, 4], reorder = [0, 2, 1], and the shape of buffer A is (4, 6)" - "then A[i, j] will be first rewritten to " - "A[(6 * i + j) / 12, (6 * i + j) / 4 % 3 , (6 * i + j) % 4] according to the `extents`," - "and then reordered to A[(6 * i + j) / 12, (6 * i + j) % 4 , (6 * i + j) / 4 % 3]" - "according to `reorder`"); - } -}; - -/*! \brief Attributes for ShapeOf operator */ -struct ShapeOfAttrs : public tvm::AttrsNode { - DataType dtype; - - TVM_DECLARE_ATTRS(ShapeOfAttrs, "relay.attrs.ShapeOfAttrs") { - TVM_ATTR_FIELD(dtype).describe("Target data type").set_default(NullValue()); - } -}; - -struct SequenceMaskAttrs : public tvm::AttrsNode { - double mask_value; - int axis; - - TVM_DECLARE_ATTRS(SequenceMaskAttrs, "relay.attrs.SequenceMaskAttrs") { - TVM_ATTR_FIELD(mask_value).set_default(0).describe("The masking value."); - TVM_ATTR_FIELD(axis).set_default(0).describe( - "The axis of the length dimension. Can only be 0 or 1."); - } -}; // struct SequenceMaskAttrs. - -/*! \brief Attributes used in sparse_to_dense operator */ -struct SparseToDenseAttrs : public tvm::AttrsNode { - Array output_shape; - - TVM_DECLARE_ATTRS(SparseToDenseAttrs, "relay.attrs.SparseToDenseAttrs") { - TVM_ATTR_FIELD(output_shape).describe("Shape of the dense output tensor"); - } -}; // struct SparseToDenseAttrs - -/*! \brief Attributes for ndarray_size operator */ -struct NdarraySizeAttrs : public tvm::AttrsNode { - DataType dtype; - - TVM_DECLARE_ATTRS(NdarraySizeAttrs, "relay.attrs.NdarraySizeAttrs") { - TVM_ATTR_FIELD(dtype).describe("Target data type").set_default(NullValue()); - } -}; - -/*! \brief Attributes used in one-hot operator */ -struct OneHotAttrs : public tvm::AttrsNode { - int depth; - int axis; - DataType dtype; - - TVM_DECLARE_ATTRS(OneHotAttrs, "relay.attrs.OneHotAttrs") { - TVM_ATTR_FIELD(depth).set_default(1).describe("Depth of the one hot dimension."); - TVM_ATTR_FIELD(axis).set_default(-1).describe("Axis to fill."); - TVM_ATTR_FIELD(dtype).set_default(NullValue()).describe("Output data type."); - } -}; // struct OneHotAttrs - -/*! \brief Attributes used in matrix_set_diag operator */ -struct MatrixSetDiagAttrs : public tvm::AttrsNode { - int k1; - int k2; - bool super_diag_right_align; - bool sub_diag_right_align; - - TVM_DECLARE_ATTRS(MatrixSetDiagAttrs, "relay.attrs.MatrixSetDiagAttrs") { - TVM_ATTR_FIELD(k1).set_default(0).describe("Lower limit (included) of the range of diagonals."); - TVM_ATTR_FIELD(k2).set_default(0).describe("Upper limit (included) of the range of diagonals."); - TVM_ATTR_FIELD(super_diag_right_align) - .set_default(true) - .describe("Bool, true iff super-diagonal is right aligned (left-padded)."); - TVM_ATTR_FIELD(sub_diag_right_align) - .set_default(false) - .describe("Bool, true iff sub-diagonal is right aligned (left-padded)."); - } -}; // struct MatrixSetDiagAttrs - -/*! \brief Attributes used in cumsum and cumprod operator */ -struct ScanopAttrs : public tvm::AttrsNode { - Integer axis; - DataType dtype; - Bool exclusive = Bool(false); - TVM_DECLARE_ATTRS(ScanopAttrs, "relay.attrs.ScanopAttrs") { - TVM_ATTR_FIELD(axis).describe("The axis to operate over").set_default(NullValue()); - TVM_ATTR_FIELD(dtype).describe("Output data type").set_default(NullValue()); - - // Default is 0 which is "false" - TVM_ATTR_FIELD(exclusive) - .describe("The first element is not included") - .set_default(Bool(false)); - } -}; // struct ScanopAttrs - -/*! \brief Attributes used in unique operator */ -struct UniqueAttrs : public tvm::AttrsNode { - bool sorted; - bool return_counts; - TVM_DECLARE_ATTRS(UniqueAttrs, "relay.attrs.UniqueAttrs") { - TVM_ATTR_FIELD(sorted).describe("Whether the unique elements are sorted").set_default(true); - TVM_ATTR_FIELD(return_counts) - .describe("Whether to return an additional tensor with counts of each unique elements") - .set_default(false); - } -}; // struct UniqueAttrs - -/*! \brief Attributes used in einsum operator */ -struct EinsumAttrs : public tvm::AttrsNode { - String equation; - - TVM_DECLARE_ATTRS(EinsumAttrs, "relay.attrs.EinsumAttrs") { - TVM_ATTR_FIELD(equation).describe("The einsum expression string"); - } -}; // struct EinsumAttrs - -/*! \brief Attributes used in stft operator */ -struct StftAttrs : public tvm::AttrsNode { - int n_fft; - int hop_length; - int win_length; - bool normalized; - bool onesided; - - TVM_DECLARE_ATTRS(StftAttrs, "relay.attrs.StftAttrs") { - TVM_ATTR_FIELD(n_fft).set_default(-1).describe("The size of Fourier transform"); - TVM_ATTR_FIELD(hop_length) - .set_default(-1) - .describe("The distance between neighboring sliding window frames"); - TVM_ATTR_FIELD(win_length).set_default(-1).describe("The size of window frame and STFT filter"); - TVM_ATTR_FIELD(normalized) - .set_default(false) - .describe("Whether to return the normalized STFT results"); - TVM_ATTR_FIELD(onesided).set_default(true).describe( - "Whether to return onesided result or fill with conjugate symmetry"); - } -}; // struct StftAttrs - -/*! \brief Attributes used in DFT operator */ -struct DFTAttrs : public tvm::AttrsNode { - Bool inverse = Bool(false); - - TVM_DECLARE_ATTRS(DFTAttrs, "relay.attrs.DFTAttrs") { - TVM_ATTR_FIELD(inverse) - .describe("Whether to perform the inverse discrete Fourier transform") - .set_default(Bool(false)); - } -}; // struct DFTAttrs - -struct TriluAttrs : public tvm::AttrsNode { - bool upper; - - TVM_DECLARE_ATTRS(TriluAttrs, "relay.attrs.TriluAttrs") { - TVM_ATTR_FIELD(upper).set_default(true).describe( - "Whether to keep the upper or lower half of the diagonal."); - } -}; // struct TriluAttrs - -} // namespace relay -} // namespace tvm -#endif // TVM_RELAY_ATTRS_TRANSFORM_H_ diff --git a/include/tvm/relay/attrs/vision.h b/include/tvm/relay/attrs/vision.h deleted file mode 100644 index bdb9bfbb4903..000000000000 --- a/include/tvm/relay/attrs/vision.h +++ /dev/null @@ -1,251 +0,0 @@ -/* - * 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 tvm/relay/attrs/vision.h - * \brief Auxiliary attributes for vision operators. - */ -#ifndef TVM_RELAY_ATTRS_VISION_H_ -#define TVM_RELAY_ATTRS_VISION_H_ - -#include -#include - -#include - -namespace tvm { -namespace relay { - -/*! \brief Attributes used in multibox_prior operators */ -struct MultiBoxPriorAttrs : public tvm::AttrsNode { - Array sizes; - Array ratios; - Array steps; - Array offsets; - bool clip; - - TVM_DECLARE_ATTRS(MultiBoxPriorAttrs, "relay.attrs.MultiBoxPriorAttrs") { - TVM_ATTR_FIELD(sizes) - .set_default(Array({static_cast(1.0)})) - .describe("List of sizes of generated MultiBoxPriores."); - TVM_ATTR_FIELD(ratios) - .set_default(Array({static_cast(1.0)})) - .describe("List of aspect ratios of generated MultiBoxPriores."); - TVM_ATTR_FIELD(steps) - .set_default(Array({static_cast(-1.0), static_cast(-1.0)})) - .describe("Priorbox step across y and x, -1 for auto calculation."); - TVM_ATTR_FIELD(offsets) - .set_default(Array({static_cast(0.5), static_cast(0.5)})) - .describe("Priorbox center offsets, y and x respectively."); - TVM_ATTR_FIELD(clip).set_default(false).describe("Whether to clip out-of-boundary boxes."); - } -}; - -struct MultiBoxTransformLocAttrs : public tvm::AttrsNode { - bool clip; - double threshold; - Array variances; - bool keep_background; - - TVM_DECLARE_ATTRS(MultiBoxTransformLocAttrs, "relay.attrs.MultiBoxTransformLocAttrs") { - TVM_ATTR_FIELD(clip).set_default(true).describe("Clip out-of-boundary boxes."); - TVM_ATTR_FIELD(threshold).set_default(0.01).describe("Threshold to be a positive prediction."); - TVM_ATTR_FIELD(variances) - .set_default(Array({0.1f, 0.1f, 0.2f, 0.2f})) - .describe("Variances to be decoded from box regression output."); - TVM_ATTR_FIELD(keep_background) - .set_default(false) - .describe("Whether to keep boxes detected as background or not"); - } -}; - -/*! \brief Attributes used in get_valid_counts operator */ -struct GetValidCountsAttrs : public tvm::AttrsNode { - Optional score_threshold; - int id_index; - int score_index; - - TVM_DECLARE_ATTRS(GetValidCountsAttrs, "relay.attrs.GetValidCountsAttrs") { - TVM_ATTR_FIELD(score_threshold).describe("Lower limit of score for valid bounding boxes."); - TVM_ATTR_FIELD(id_index).set_default(0).describe("Axis index of id."); - TVM_ATTR_FIELD(score_index).set_default(1).describe("Index of the scores/confidence of boxes."); - } -}; - -/*! \brief Attributes used in non_maximum_suppression operator */ -struct NonMaximumSuppressionAttrs : public tvm::AttrsNode { - bool force_suppress; - int top_k; - int coord_start; - int score_index; - int id_index; - bool return_indices; - bool invalid_to_bottom; - - TVM_DECLARE_ATTRS(NonMaximumSuppressionAttrs, "relay.attrs.NonMaximumSuppressionAttrs") { - TVM_ATTR_FIELD(force_suppress) - .set_default(false) - .describe("Suppress all detections regardless of class_id."); - TVM_ATTR_FIELD(top_k).set_default(-1).describe( - "Keep maximum top k detections before nms, -1 for no limit."); - TVM_ATTR_FIELD(coord_start) - .set_default(2) - .describe("Start index of the consecutive 4 coordinates."); - TVM_ATTR_FIELD(score_index).set_default(1).describe("Index of the scores/confidence of boxes."); - TVM_ATTR_FIELD(id_index).set_default(0).describe("Axis index of id."); - TVM_ATTR_FIELD(return_indices) - .set_default(true) - .describe("Whether to return box indices in input data."); - TVM_ATTR_FIELD(invalid_to_bottom) - .set_default(false) - .describe("Whether to move all invalid bounding boxes to the bottom."); - } -}; - -/*! \brief Attributes used in all_class_non_maximum_suppression operator */ -struct AllClassNonMaximumSuppressionAttrs - : public tvm::AttrsNode { - std::string output_format; - - TVM_DECLARE_ATTRS(AllClassNonMaximumSuppressionAttrs, - "relay.attrs.AllClassNonMaximumSuppressionAttrs") { - TVM_ATTR_FIELD(output_format) - .set_default("onnx") - .describe( - "Output format, onnx or tensorflow. Returns outputs in a way that can be easily " - "consumed by each frontend."); - } -}; - -/*! \brief Attributes used in regular_non_maximum_suppression operator */ -struct RegularNonMaximumSuppressionAttrs - : public tvm::AttrsNode { - int32_t max_detections_per_class; - int32_t max_detections; - int32_t num_classes; - double iou_threshold; - double score_threshold; - - TVM_DECLARE_ATTRS(RegularNonMaximumSuppressionAttrs, - "relay.attrs.RegularNonMaximumSuppressionAttrs") { - TVM_ATTR_FIELD(max_detections_per_class) - .describe("The maxinum number of output selected boxes per class."); - TVM_ATTR_FIELD(max_detections).describe("The maxinum number of output selected boxes."); - TVM_ATTR_FIELD(num_classes).describe("The number of classes without background."); - TVM_ATTR_FIELD(iou_threshold).describe("The IoU threshold for box the overlap test."); - TVM_ATTR_FIELD(score_threshold) - .describe("Score threshold to filter out low score boxes early."); - } -}; - -/*! \brief Attributes used in roi_align operators */ -struct ROIAlignAttrs : public tvm::AttrsNode { - Array pooled_size; - double spatial_scale; - int sample_ratio; - std::string layout; - std::string mode; - TVM_DECLARE_ATTRS(ROIAlignAttrs, "relay.attrs.ROIAlignAttrs") { - TVM_ATTR_FIELD(pooled_size).describe("Output size of roi align."); - TVM_ATTR_FIELD(spatial_scale) - .describe( - "Ratio of input feature map height (or w) to raw image height (or w). " - "Equals the reciprocal of total stride in convolutional layers, which should be " - "in range (0.0, 1.0]"); - TVM_ATTR_FIELD(sample_ratio) - .set_default(-1) - .describe("Optional sampling ratio of ROI align, using adaptive size by default."); - TVM_ATTR_FIELD(layout).set_default("NCHW").describe( - "Dimension ordering of data and weight. Can be 'NCHW', 'NHWC', etc." - "'N', 'C', 'H', 'W' stands for batch, channel, height, and width" - "dimensions respectively. Convolution is applied on the 'H' and" - "'W' dimensions."); - TVM_ATTR_FIELD(mode).set_default("avg").describe( - "Mode for ROI Align. Can be 'avg' or 'max'. The default mode is 'avg'."); - } -}; - -/*! \brief Attributes used in roi_pool operators */ -struct ROIPoolAttrs : public tvm::AttrsNode { - Array pooled_size; - double spatial_scale; - std::string layout; - TVM_DECLARE_ATTRS(ROIPoolAttrs, "relay.attrs.ROIPoolAttrs") { - TVM_ATTR_FIELD(pooled_size).describe("Output size of roi pool."); - TVM_ATTR_FIELD(spatial_scale) - .describe( - "Ratio of input feature map height (or w) to raw image height (or w). " - "Equals the reciprocal of total stride in convolutional layers, which should be " - "in range (0.0, 1.0]"); - TVM_ATTR_FIELD(layout).set_default("NCHW").describe( - "Dimension ordering of data and weight. Can be 'NCHW', 'NHWC', etc." - "'N', 'C', 'H', 'W' stands for batch, channel, height, and width" - "dimensions respectively. Convolution is applied on the 'H' and" - "'W' dimensions."); - } -}; - -/*! \brief Attributes used in yolo reorg operators */ -struct YoloReorgAttrs : public tvm::AttrsNode { - Integer stride; - - TVM_DECLARE_ATTRS(YoloReorgAttrs, "relay.attrs.YoloReorgAttrs") { - TVM_ATTR_FIELD(stride).set_default(1).describe("Stride value for yolo reorg"); - } -}; - -/*! \brief Attributes used in proposal operators */ -struct ProposalAttrs : public tvm::AttrsNode { - Array scales; - Array ratios; - int feature_stride; - double threshold; - int rpn_pre_nms_top_n; - int rpn_post_nms_top_n; - int rpn_min_size; - bool iou_loss; - - TVM_DECLARE_ATTRS(ProposalAttrs, "relay.attrs.ProposalAttrs") { - TVM_ATTR_FIELD(scales) - .set_default(Array({4.0f, 8.0f, 16.0f, 32.0f})) - .describe("Used to generate anchor windows by enumerating scales"); - TVM_ATTR_FIELD(ratios) - .set_default(Array({0.5f, 1.0f, 2.0f})) - .describe("Used to generate anchor windows by enumerating ratios"); - TVM_ATTR_FIELD(feature_stride) - .set_default(16) - .describe( - "The size of the receptive field each unit in the convolution layer of the rpn," - "for example the product of all stride's prior to this layer."); - TVM_ATTR_FIELD(threshold).set_default(0.7).describe( - "IoU threshold of non-maximum suppresion (suppress boxes with IoU >= this threshold)"); - TVM_ATTR_FIELD(rpn_pre_nms_top_n) - .set_default(6000) - .describe("Number of top scoring boxes to apply NMS. -1 to use all boxes"); - TVM_ATTR_FIELD(rpn_post_nms_top_n) - .set_default(300) - .describe("Number of top scoring boxes to keep after applying NMS to RPN proposals"); - TVM_ATTR_FIELD(rpn_min_size).set_default(16).describe("Minimum height or width in proposal"); - TVM_ATTR_FIELD(iou_loss).set_default(false).describe("Usage of IoU Loss"); - } -}; - -} // namespace relay -} // namespace tvm -#endif // TVM_RELAY_ATTRS_VISION_H_ diff --git a/include/tvm/relay/attrs/vm.h b/include/tvm/relay/attrs/vm.h deleted file mode 100644 index 7eb1008004de..000000000000 --- a/include/tvm/relay/attrs/vm.h +++ /dev/null @@ -1,58 +0,0 @@ -/* - * 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 tvm/relay/attrs/vm.h - * \brief Attributes for Relay vm operators. - */ -#ifndef TVM_RELAY_ATTRS_VM_H_ -#define TVM_RELAY_ATTRS_VM_H_ - -#include - -namespace tvm { -namespace relay { - -/*! - * \brief Options for the shape function operator. - */ -struct ShapeFuncAttrs : public tvm::AttrsNode { - Array is_input; - - TVM_DECLARE_ATTRS(ShapeFuncAttrs, "relay.attrs.ShapeFuncAttrs") { - TVM_ATTR_FIELD(is_input).describe( - "A bool indicating whether the shape function should" - "expect shape or input in each position."); - } -}; - -/*! - * \brief Attributes for VM reshape_tensor operator. - */ -struct ReshapeTensorAttrs : public tvm::AttrsNode { - Array newshape; - - TVM_DECLARE_ATTRS(ReshapeTensorAttrs, "relay.attrs.ReshapeTensorAttrs") { - TVM_ATTR_FIELD(newshape).describe("The new shape of output tensor"); - } -}; - -} // namespace relay -} // namespace tvm -#endif // TVM_RELAY_ATTRS_VM_H_ diff --git a/include/tvm/relay/base.h b/include/tvm/relay/base.h deleted file mode 100644 index a66b8044998b..000000000000 --- a/include/tvm/relay/base.h +++ /dev/null @@ -1,154 +0,0 @@ -/* - * 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 tvm/relay/base.h - * \brief Base classes for the Relay IR. - */ -#ifndef TVM_RELAY_BASE_H_ -#define TVM_RELAY_BASE_H_ - -#include -#include -#include - -#include -#include - -namespace tvm { -/*! - * \brief Relay: a high level functional IR for TVM. - * - * This namespace contains the abstract syntax tree, and other - * essential data structures for the Relay IR. - * - * You can find more about Relay by reading the language reference. - */ -namespace relay { - -#define RELAY_DEBUG(...) \ - { \ - auto fdebug = runtime::Registry::Get("relay.debug"); \ - ICHECK(fdebug) << "Could not find Relay Python debugger function."; \ - (*fdebug)("RELAY_DEBUG", __FILE__, __LINE__, __VA_ARGS__); \ - } - -#define RELAY_DEBUG_INTERP(...) \ - { \ - auto fdebug = runtime::Registry::Get("relay.debug_interp"); \ - ICHECK(fdebug) << "Could not find Relay Python debugger function."; \ - (*fdebug)("RELAY_DEBUG", __FILE__, __LINE__, __VA_ARGS__); \ - } - -/*! - * \brief Symbolic expression for tensor shape. - */ -using IndexExpr = ::tvm::PrimExpr; - -using SourceName = tvm::SourceName; -using Span = tvm::Span; -using SpanNode = tvm::SpanNode; - -/*! - * \brief This is the base node container of all relay structures. - */ -class RelayNode : public Object { - public: - /*! \brief The location of the program in a SourceFragment can be null, - * check with span.defined() */ - mutable Span span; - - static constexpr const char* _type_key = "relay.Node"; - TVM_DECLARE_BASE_OBJECT_INFO(RelayNode, Object); -}; - -/*! - * \brief The unique identifier of variables. - * - * Id is like name to the variables, - * except that id is unique for each Var. - * - * \note Do not create Id directly, they are created in Var. - */ -class IdNode : public Object { - public: - /*! - * \brief The name of the variable, - * this only acts as a hint to the user, - * and is not used for equality. - */ - String name_hint; - - void VisitAttrs(tvm::AttrVisitor* v) { v->Visit("name_hint", &name_hint); } - - bool SEqualReduce(const IdNode* other, SEqualReducer equal) const { - return equal.FreeVarEqualImpl(this, other); - } - - void SHashReduce(SHashReducer hash_reduce) const { hash_reduce.FreeVarHashImpl(this); } - - static constexpr const char* _type_key = "relay.Id"; - static constexpr const bool _type_has_method_sequal_reduce = true; - static constexpr const bool _type_has_method_shash_reduce = true; - TVM_DECLARE_FINAL_OBJECT_INFO(IdNode, Object); -}; - -class Id : public ObjectRef { - public: - /*! - * \brief The constructor - * \param name_hint The name of the variable. - */ - TVM_DLL explicit Id(String name_hint); - - TVM_DEFINE_OBJECT_REF_METHODS(Id, ObjectRef, IdNode); -}; - -/*! - * \brief Pretty print a node for debug purposes. - * - * \param node The node to be printed. - * \return The text reperesentation. - * \note This function does not show version or meta-data. - * Use AsText if you want to store the text. - * \sa AsText. - */ -TVM_DLL String PrettyPrint(const ObjectRef& node); - -/*! - * \brief Render the node as a string in the text format. - * - * \param node The node to be rendered. - * \param show_meta_data Whether to print meta data section. - * \param annotate An optional callback function for attaching - * additional comment block to an expr. - * - * \note We support a limited set of IR nodes that are part of - * relay IR and - * - * \sa PrettyPrint. - * \return The text representation. - */ -TVM_DLL String AsText(const ObjectRef& node, bool show_meta_data = true, - runtime::TypedPackedFunc annotate = nullptr); - -} // namespace relay -} // namespace tvm - -#endif // TVM_RELAY_BASE_H_ diff --git a/include/tvm/relay/dataflow_matcher.h b/include/tvm/relay/dataflow_matcher.h deleted file mode 100644 index 8dd5fbdd5eac..000000000000 --- a/include/tvm/relay/dataflow_matcher.h +++ /dev/null @@ -1,123 +0,0 @@ -/* - * 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 tvm/relay/dataflow_matcher.h - * \brief A pattern matcher for matching dataflow properties. - */ -#ifndef TVM_RELAY_DATAFLOW_MATCHER_H_ -#define TVM_RELAY_DATAFLOW_MATCHER_H_ - -#include -#include - -#include -#include -#include - -namespace tvm { -namespace relay { - -class DFPatternCallback; -/*! - * \brief Base type of all dataflow pattern callbacks. - * \sa DFPatternCallback - */ -class DFPatternCallbackNode : public Object { - public: - /*! \brief Pattern this callback matches */ - DFPattern pattern; - /*! \brief Function to call when finding a matched expression */ - PackedFunc function; - /*! \brief Require InferType to be run before the callback */ - bool require_type; - /*! \brief Run the callback only once */ - bool rewrite_once; - - void VisitAttrs(tvm::AttrVisitor* v) { - v->Visit("pattern", &pattern); - v->Visit("require_type", &require_type); - v->Visit("rewrite_once", &rewrite_once); - } - - static constexpr const char* _type_key = "DFPatternCallbackNode"; - TVM_DECLARE_BASE_OBJECT_INFO(DFPatternCallbackNode, Object); -}; - -/*! - * \brief Managed reference to dataflow pattern callbacks. - * \sa DFPatternCallbackNode - */ -class DFPatternCallback : public ObjectRef { - public: - TVM_DLL DFPatternCallback(DFPattern pattern, PackedFunc callback, bool require_type, - bool rewrite_once = false); - TVM_DEFINE_OBJECT_REF_METHODS(DFPatternCallback, ObjectRef, DFPatternCallbackNode); -}; - -/*! - * \brief Determine if a pattern matches an expression - * - * \param pattern The pattern to match - * \param expr The expression to match - * - * \return Return true if the pattern and the expression match, return false otherwise. - */ -bool MatchPattern(DFPattern pattern, Expr expr); - -/*! - * \brief Rewrite an expression based on some number of DFPatternCallbacks - * - * \param callbacks An array of DFPatternCallback Nodes - * \param expr The expression to rewrite - * \param mod The module that associates with the expr - * - * \return Return An Expr with every match of the pattern inside the callbacks rewritten by the - * functions inside the callbacks - */ -Expr RewritePatterns(Array callbacks, Expr expr, IRModule mod = IRModule()); - -/*! - * \brief Partition all matches of a DFPattern inside an Expr into separate Function calls - * - * \param pattern The pattern to match - * \param expr The expression to patition - * \param attrs A set of parameter names and values to apply to the partitioned function - * \param check A callback function for checking more complicated properties of the matched - * expressions, returns true if the match is accepted and false otherwise - * - * \return Return the paritioned Expr. - */ -Expr PartitionPattern(DFPattern pattern, Expr expr, Map attrs, PackedFunc check); - -/*! - * \brief Infer the type of an expression. - * - * \param expr The expression to rewrite - * - * \return Return An Expr with unambiguous type information filled in, as well as it's - * checked type field populated with the result type. - * - */ -Expr InferType(const Expr& expr); - -} // namespace relay -} // namespace tvm - -#endif // TVM_RELAY_DATAFLOW_MATCHER_H_ diff --git a/include/tvm/relay/dataflow_pattern.h b/include/tvm/relay/dataflow_pattern.h deleted file mode 100644 index 040372db3533..000000000000 --- a/include/tvm/relay/dataflow_pattern.h +++ /dev/null @@ -1,589 +0,0 @@ -/* - * 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 tvm/relay/dataflow_pattern.h - * \brief A pattern language for matching dataflow properties. - */ -#ifndef TVM_RELAY_DATAFLOW_PATTERN_H_ -#define TVM_RELAY_DATAFLOW_PATTERN_H_ - -#include -#include - -#include -#include -#include -#include - -namespace tvm { -namespace relay { - -/*! - * \brief Base type of all dataflow patterns. - * \sa DFPattern - */ -class DFPatternNode : public Object { - public: - static constexpr const char* _type_key = "DFPatternNode"; - TVM_DECLARE_BASE_OBJECT_INFO(DFPatternNode, Object); -}; - -/*! - * \brief Managed reference to dataflow patterns. - * \sa DFPatternNode - */ -class DFPattern : public ObjectRef { - public: - /*! \brief Syntatic Sugar for creating a CallPattern */ - DFPattern operator()(const std::vector& args) const; - /*! \brief Syntatic Sugar for creating a CallPattern with an "add" op */ - DFPattern operator+(const DFPattern& other) const; - /*! \brief Syntatic Sugar for creating a CallPattern with a "subtract" op */ - DFPattern operator-(const DFPattern& other) const; - /*! \brief Syntatic Sugar for creating a CallPattern with a "multiply" op */ - DFPattern operator*(const DFPattern& other) const; - /*! \brief Syntatic Sugar for creating a CallPattern with a "divide" op */ - DFPattern operator/(const DFPattern& other) const; - /*! \brief Syntatic Sugar for creating an AltPattern */ - DFPattern operator||(const DFPattern& other) const; - /*! \brief Syntatic Sugar for creating an Optional Pattern */ - DFPattern Optional(const std::function& func) const; - /*! \brief Syntatic Sugar for creating an AttrPattern */ - DFPattern HasAttr(const Map& attrs) const; - /*! \brief Syntatic Sugar for creating a TypePattern */ - DFPattern HasType(const Type& type) const; - /*! \brief Syntatic Sugar for creating a DataTypePattern with a DataType */ - DFPattern HasDtype(const DataType& dtype) const; - /*! \brief Syntatic Sugar for creating a DataTypePattern with a data type's name */ - DFPattern HasDtype(const std::string& dtype) const; - /*! \brief Syntatic Sugar for creating a ShapePattern */ - DFPattern HasShape(const Array shape) const; - - TVM_DEFINE_OBJECT_REF_METHODS(DFPattern, ObjectRef, DFPatternNode); -}; - -/*! - * \brief Pattern for Relay Expression. - */ -class ExprPatternNode : public DFPatternNode { - public: - /*! \brief The expression to match. */ - Expr expr; - - void VisitAttrs(tvm::AttrVisitor* v) { v->Visit("expr", &expr); } - - static constexpr const char* _type_key = "relay.dataflow_pattern.ExprPattern"; - TVM_DECLARE_FINAL_OBJECT_INFO(ExprPatternNode, DFPatternNode); -}; - -/*! - * \brief A pattern which matches a literal expression. - * - * \note Uses structural equality on expressions to check equality. - * - */ -class ExprPattern : public DFPattern { - public: - TVM_DLL explicit ExprPattern(Expr expr); - TVM_DEFINE_OBJECT_REF_METHODS(ExprPattern, DFPattern, ExprPatternNode); -}; - -/*! - * \brief A Pattern to Match a Relay Variable - */ -class VarPattern; -/*! \brief Container for Var */ -class VarPatternNode : public DFPatternNode { - public: - /*! - * \brief The name of the Var (optional). - */ - String name; - - /*! \return The name hint of the variable */ - const String& name_hint() const { return name; } - - void VisitAttrs(tvm::AttrVisitor* v) { v->Visit("name", &name); } - - static constexpr const char* _type_key = "relay.dataflow_pattern.VarPattern"; - TVM_DECLARE_FINAL_OBJECT_INFO(VarPatternNode, DFPatternNode); -}; - -class VarPattern : public DFPattern { - public: - TVM_DLL VarPattern(String name_hint); - TVM_DEFINE_OBJECT_REF_METHODS(VarPattern, DFPattern, VarPatternNode); -}; - -/*! - * \brief A Pattern to Match a Relay Constant - */ -class ConstantPattern; -/*! \brief Container for Constant */ -class ConstantPatternNode : public DFPatternNode { - public: - void VisitAttrs(tvm::AttrVisitor* v) {} - - static constexpr const char* _type_key = "relay.dataflow_pattern.ConstantPattern"; - TVM_DECLARE_FINAL_OBJECT_INFO(ConstantPatternNode, DFPatternNode); -}; - -class ConstantPattern : public DFPattern { - public: - TVM_DEFINE_OBJECT_REF_METHODS(ConstantPattern, DFPattern, ConstantPatternNode); -}; - -/*! - * \brief Call corresponds to operator invocation. - * Corresponds to the operator in computational graph terminology. - */ -class CallPattern; -/*! \brief CallPattern container. */ -class CallPatternNode : public DFPatternNode { - public: - /*! - * \brief The operator(function) being invoked - * - * - It can be relay::Op which corresponds to the primitive operators. - * - It can also be user defined functions (Function, GlobalVar, Var). - */ - DFPattern op; - - /*! \brief The arguments(inputs) of the call */ - tvm::Array args; - - void VisitAttrs(tvm::AttrVisitor* v) { - v->Visit("op", &op); - v->Visit("args", &args); - } - - static constexpr const char* _type_key = "relay.dataflow_pattern.CallPattern"; - TVM_DECLARE_FINAL_OBJECT_INFO(CallPatternNode, DFPatternNode); -}; - -class CallPattern : public DFPattern { - public: - TVM_DLL CallPattern(DFPattern op, Array args); - TVM_DEFINE_OBJECT_REF_METHODS(CallPattern, DFPattern, CallPatternNode); -}; - -/*! - * \brief Relay Function container - * \sa Function - */ -class FunctionPatternNode : public DFPatternNode { - public: - /*! \brief Function parameters */ - tvm::Array params; - /*! - * \brief - * The expression which represents the computation of the function, - * the expression may reference the parameters, and the type of it - * or sub-expressions may reference the type variables. - */ - DFPattern body; - - void VisitAttrs(tvm::AttrVisitor* v) { - v->Visit("params", ¶ms); - v->Visit("body", &body); - } - - static constexpr const char* _type_key = "relay.dataflow_pattern.FunctionPattern"; - TVM_DECLARE_FINAL_OBJECT_INFO(FunctionPatternNode, DFPatternNode); -}; - -/*! - * \brief Managed reference to FunctionNode. - * \sa FunctionNode - */ -class FunctionPattern : public DFPattern { - public: - /*! - * \brief Constructor - * \param params The parameters of the function. - * \param body The body of the function. - */ - TVM_DLL FunctionPattern(tvm::Array params, DFPattern body); - - TVM_DEFINE_OBJECT_REF_METHODS(FunctionPattern, DFPattern, FunctionPatternNode); - TVM_DEFINE_OBJECT_REF_COW_METHOD(FunctionPatternNode); -}; - -/*! \brief A binding of a sub-network. */ -class LetPatternNode : public DFPatternNode { - public: - /*! \brief The variable we bind to */ - DFPattern var; - /*! \brief The value we bind var to */ - DFPattern value; - /*! \brief The body of the let binding */ - DFPattern body; - - void VisitAttrs(tvm::AttrVisitor* v) { - v->Visit("var", &var); - v->Visit("value", &value); - v->Visit("body", &body); - } - - static constexpr const char* _type_key = "relay.dataflow_pattern.LetPattern"; - TVM_DECLARE_FINAL_OBJECT_INFO(LetPatternNode, DFPatternNode); -}; - -/*! - * \brief Let binding that binds a local var - */ -class LetPattern : public DFPattern { - public: - /*! - * \brief The constructor - * \param var The variable that is bound to. - * \param value The value used to bind to the variable. - * \param body The body of the let binding. - */ - TVM_DLL LetPattern(DFPattern var, DFPattern value, DFPattern body); - - TVM_DEFINE_OBJECT_REF_METHODS(LetPattern, DFPattern, LetPatternNode); -}; - -/*! \brief Tuple of multiple Exprs */ -class TuplePattern; -/*! \brief Tuple container */ -class TuplePatternNode : public DFPatternNode { - public: - /*! \brief the fields of the tuple */ - tvm::Array fields; - - void VisitAttrs(tvm::AttrVisitor* v) { v->Visit("fields", &fields); } - - static constexpr const char* _type_key = "relay.dataflow_pattern.TuplePattern"; - TVM_DECLARE_FINAL_OBJECT_INFO(TuplePatternNode, DFPatternNode); -}; - -class TuplePattern : public DFPattern { - public: - TVM_DLL explicit TuplePattern(tvm::Array fields); - TVM_DEFINE_OBJECT_REF_METHODS(TuplePattern, DFPattern, TuplePatternNode); -}; - -/*! \brief Get index-th field out of a tuple. */ -class TupleGetItemPattern; -class TupleGetItemPatternNode : public DFPatternNode { - public: - /*! \brief The tuple Expression */ - DFPattern tuple; - /*! \brief which value to get */ - int index; - - void VisitAttrs(tvm::AttrVisitor* v) { - v->Visit("tuple", &tuple); - v->Visit("index", &index); - } - - static constexpr const char* _type_key = "relay.dataflow_pattern.TupleGetItemPattern"; - TVM_DECLARE_FINAL_OBJECT_INFO(TupleGetItemPatternNode, DFPatternNode); -}; - -class IfPatternNode : public DFPatternNode { - public: - DFPattern cond, true_branch, false_branch; - - void VisitAttrs(tvm::AttrVisitor* v) { - v->Visit("cond", &cond); - v->Visit("true_branch", &true_branch); - v->Visit("false_branch", &false_branch); - } - - static constexpr const char* _type_key = "relay.dataflow_pattern.IfPattern"; - TVM_DECLARE_FINAL_OBJECT_INFO(IfPatternNode, DFPatternNode); -}; - -class IfPattern : public DFPattern { - public: - TVM_DLL IfPattern(DFPattern cond, DFPattern then_clause, DFPattern else_clause); - TVM_DEFINE_OBJECT_REF_METHODS(IfPattern, DFPattern, IfPatternNode); -}; - -class TupleGetItemPattern : public DFPattern { - public: - TVM_DLL TupleGetItemPattern(DFPattern tuple, int index); - TVM_DEFINE_OBJECT_REF_METHODS(TupleGetItemPattern, DFPattern, TupleGetItemPatternNode); -}; - -class AltPattern; -/*! - * \brief Pattern for Alternate Expressions. - */ -class AltPatternNode : public DFPatternNode { - public: - /*! \brief The left optional pattern. */ - DFPattern left; - /*! \brief The right optional pattern. */ - DFPattern right; - - void VisitAttrs(tvm::AttrVisitor* v) { - v->Visit("left", &left); - v->Visit("right", &right); - } - - static constexpr const char* _type_key = "relay.dataflow_pattern.AltPattern"; - TVM_DECLARE_FINAL_OBJECT_INFO(AltPatternNode, DFPatternNode); -}; - -/*! - * \brief A pattern which matches either of two patterns - */ -class AltPattern : public DFPattern { - public: - TVM_DLL AltPattern(DFPattern left, DFPattern right); - TVM_DEFINE_OBJECT_REF_METHODS(AltPattern, DFPattern, AltPatternNode); -}; - -/*! - * \brief Wildcard Pattern. - */ -class WildcardPatternNode : public DFPatternNode { - public: - void VisitAttrs(tvm::AttrVisitor* v) {} - - /*! \brief If the wildcard is redirected, then pattern is not nullptr, and the wildcard - * redirects to the pattern. */ - Optional pattern{nullptr}; - - static constexpr const char* _type_key = "relay.dataflow_pattern.WildcardPattern"; - TVM_DECLARE_FINAL_OBJECT_INFO(WildcardPatternNode, DFPatternNode); -}; - -/*! - * \brief A pattern which matches anything. - */ -class WildcardPattern : public DFPattern { - public: - TVM_DEFINE_OBJECT_REF_METHODS(WildcardPattern, DFPattern, WildcardPatternNode); - - void redirect_to(DFPattern pat) const; -}; - -class TypePattern; -/*! - * \brief Pattern for Types. - */ -class TypePatternNode : public DFPatternNode { - public: - /*! \brief The pattern. */ - DFPattern pattern; - /*! \brief The type to match */ - Type type; - - void VisitAttrs(tvm::AttrVisitor* v) { - v->Visit("pattern", &pattern); - v->Visit("type", &type); - } - - static constexpr const char* _type_key = "relay.dataflow_pattern.TypePattern"; - TVM_DECLARE_FINAL_OBJECT_INFO(TypePatternNode, DFPatternNode); -}; - -/*! - * \brief A pattern which matches a type in another pattern - */ -class TypePattern : public DFPattern { - public: - TVM_DLL TypePattern(DFPattern pattern, Type type); - TVM_DEFINE_OBJECT_REF_METHODS(TypePattern, DFPattern, TypePatternNode); -}; - -class ShapePattern; -/*! - * \brief Pattern for Shapes. - */ -class ShapePatternNode : public DFPatternNode { - public: - /*! \brief The pattern. */ - DFPattern pattern; - /*! \brief The type to match */ - Array shape; - - void VisitAttrs(tvm::AttrVisitor* v) { - v->Visit("pattern", &pattern); - v->Visit("shape", &shape); - } - - static constexpr const char* _type_key = "relay.dataflow_pattern.ShapePattern"; - TVM_DECLARE_FINAL_OBJECT_INFO(ShapePatternNode, DFPatternNode); -}; - -/*! - * \brief A pattern which matches a type in another pattern - */ -class ShapePattern : public DFPattern { - public: - TVM_DLL ShapePattern(DFPattern pattern, Array type); - TVM_DEFINE_OBJECT_REF_METHODS(ShapePattern, DFPattern, ShapePatternNode); -}; - -class DataTypePattern; -/*! - * \brief Pattern for Types. - */ -class DataTypePatternNode : public DFPatternNode { - public: - /*! \brief The pattern. */ - DFPattern pattern; - /*! \brief The type to match */ - DataType dtype; - - void VisitAttrs(tvm::AttrVisitor* v) { - v->Visit("pattern", &pattern); - v->Visit("dtype", &dtype); - } - - static constexpr const char* _type_key = "relay.dataflow_pattern.DataTypePattern"; - TVM_DECLARE_FINAL_OBJECT_INFO(DataTypePatternNode, DFPatternNode); -}; - -/*! - * \brief A pattern which matches a type in another pattern - */ -class DataTypePattern : public DFPattern { - public: - TVM_DLL DataTypePattern(DFPattern pattern, DataType dtype); - TVM_DEFINE_OBJECT_REF_METHODS(DataTypePattern, DFPattern, DataTypePatternNode); -}; - -class AttrPattern; -/*! - * \brief Pattern for Attributes. - */ -class AttrPatternNode : public DFPatternNode { - public: - /*! \brief The pattern. */ - DFPattern pattern; - /*! \brief The attribute to match */ - DictAttrs attrs; - - void VisitAttrs(tvm::AttrVisitor* v) { - v->Visit("pattern", &pattern); - v->Visit("attrs", &attrs); - } - - static constexpr const char* _type_key = "relay.dataflow_pattern.AttrPattern"; - TVM_DECLARE_FINAL_OBJECT_INFO(AttrPatternNode, DFPatternNode); -}; - -/*! - * \brief A pattern which matches attributes in another pattern - */ -class AttrPattern : public DFPattern { - public: - TVM_DLL AttrPattern(DFPattern pattern, DictAttrs attrs); - TVM_DEFINE_OBJECT_REF_METHODS(AttrPattern, DFPattern, AttrPatternNode); -}; - -class DominatorPattern; -/*! - * \brief Dominated Graph Pattern - * Pattern for fuzzy subgraphs where all outputs of the parent are used finally by the child, and - * every operation between the parent and the child matches the path. - */ -class DominatorPatternNode : public DFPatternNode { - public: - /*! \brief The parent. */ - DFPattern parent; - /*! \brief The path. */ - DFPattern path; - /*! \brief The child. */ - DFPattern child; - - void VisitAttrs(tvm::AttrVisitor* v) { - v->Visit("parent", &parent); - v->Visit("path", &path); - v->Visit("child", &child); - } - - static constexpr const char* _type_key = "relay.dataflow_pattern.DominatorPattern"; - TVM_DECLARE_FINAL_OBJECT_INFO(DominatorPatternNode, DFPatternNode); -}; - -/*! - * \brief A pattern which matches a variable length dominator path - */ -class DominatorPattern : public DFPattern { - public: - TVM_DLL DominatorPattern(DFPattern parent, DFPattern path, DFPattern child); - TVM_DEFINE_OBJECT_REF_METHODS(DominatorPattern, DFPattern, DominatorPatternNode); -}; - -/*! \brief Syntatic Sugar for creating a VarPattern with a name */ -DFPattern IsVar(const String& name); -/*! \brief Syntatic Sugar for creating a ConstantPattern */ -DFPattern IsConstant(); -/*! \brief Syntatic Sugar for creating a WildcardPattern */ -DFPattern IsWildcard(); -/*! \brief Syntatic Sugar for creating a ExprPattern */ -DFPattern IsExpr(const Expr& expr); -/*! \brief Syntatic Sugar for creating a ExprPattern base on an Op*/ -DFPattern IsOp(const String& op_name); -/*! \brief Syntatic Sugar for creating a TuplePattern*/ -DFPattern IsTuple(const Array& fields); -/*! \brief Syntatic Sugar for creating a TupleGetItemPattern*/ -DFPattern IsTupleGetItem(const DFPattern tuple, int index = -1); - -/*! \brief A printer class to print pattern. */ -class DFPatternPrinter : public ReprPrinter { - public: - std::stringstream string_stream{}; - - std::unordered_map, ObjectPtrHash, ObjectPtrEqual> - memo_{}; - /*! \brief Subpatterns that are encountered more than once during printing. If a subpattern has - * already printed, only the pattern ID will be printed in the next encounter of the same pattern. - * This avoids printing a subpattern infinitely many times is the considered pattern involves - * recursion.*/ - std::vector auxiliary_patterns{}; - - DFPatternPrinter(std::ostream& stream) // NOLINT(*) - : ReprPrinter(stream) {} - TVM_DLL void Print(const ObjectRef& node); - using FType = NodeFunctor; - TVM_DLL static FType& vtable(); -}; - -inline std::ostream& operator<<(std::ostream& os, - const DFPattern& n) { // NOLINT(*) - std::stringstream string_stream{}, tmp_stream{}; - DFPatternPrinter printer{tmp_stream}; - printer.Print(n); - string_stream << "Main pattern:" << std::endl; - string_stream << printer.string_stream.str(); - string_stream << std::endl; - string_stream << "Auxiliary patterns:"; - for (const DFPattern& pat : printer.auxiliary_patterns) { - string_stream << std::endl; - string_stream << printer.memo_[pat].second; - } - os << string_stream.str(); - return os; -} - -String PrettyPrint(const DFPattern& pattern); - -} // namespace relay -} // namespace tvm -#endif // TVM_RELAY_DATAFLOW_PATTERN_H_ diff --git a/include/tvm/relay/dataflow_pattern_functor.h b/include/tvm/relay/dataflow_pattern_functor.h deleted file mode 100644 index 490cdc5e3f9d..000000000000 --- a/include/tvm/relay/dataflow_pattern_functor.h +++ /dev/null @@ -1,164 +0,0 @@ -/* - * 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 tvm/relay/dataflow_pattern_functor.h - * \brief A set of passes for operating on pattern graphs. - */ -#ifndef TVM_RELAY_DATAFLOW_PATTERN_FUNCTOR_H_ -#define TVM_RELAY_DATAFLOW_PATTERN_FUNCTOR_H_ - -#include - -#include -#include - -namespace tvm { -namespace relay { - -/*! - * \brief A dynamical functor that dispatches on in the first DFPattern argument. - * - * \tparam FType function signature - * This type is only defined for FType with function signature R(const DFPattern&, - * Args...) - */ -template -class DFPatternFunctor; - -// functions to be overriden. -#define DFPATTERN_FUNCTOR_DEFAULT \ - { return VisitDFPatternDefault_(op, std::forward(args)...); } - -#define RELAY_DFPATTERN_FUNCTOR_DISPATCH(OP) \ - vtable.template set_dispatch([](const ObjectRef& n, TSelf* self, Args... args) { \ - return self->VisitDFPattern_(static_cast(n.get()), std::forward(args)...); \ - }); - -template -class DFPatternFunctor { - private: - using TSelf = DFPatternFunctor; - using FType = tvm::NodeFunctor; - - public: - /*! \brief virtual destructor */ - virtual ~DFPatternFunctor() {} - /*! - * \brief Same as call. - * \param n The expression node. - * \param args Additional arguments. - * \return The result of the call - */ - R operator()(const DFPattern& n, Args... args) { - return VisitDFPattern(n, std::forward(args)...); - } - /*! - * \brief The functor call. - * \param n The expression node. - * \param args Additional arguments. - * \return The result of the call - */ - virtual R VisitDFPattern(const DFPattern& n, Args... args) { - ICHECK(n.defined()); - static FType vtable = InitVTable(); - return vtable(n, this, std::forward(args)...); - } - // Functions that can be overriden by subclass - virtual R VisitDFPattern_(const AltPatternNode* op, Args... args) DFPATTERN_FUNCTOR_DEFAULT; - virtual R VisitDFPattern_(const AttrPatternNode* op, Args... args) DFPATTERN_FUNCTOR_DEFAULT; - virtual R VisitDFPattern_(const CallPatternNode* op, Args... args) DFPATTERN_FUNCTOR_DEFAULT; - virtual R VisitDFPattern_(const ConstantPatternNode* op, Args... args) DFPATTERN_FUNCTOR_DEFAULT; - virtual R VisitDFPattern_(const DataTypePatternNode* op, Args... args) DFPATTERN_FUNCTOR_DEFAULT; - virtual R VisitDFPattern_(const DominatorPatternNode* op, Args... args) DFPATTERN_FUNCTOR_DEFAULT; - virtual R VisitDFPattern_(const ExprPatternNode* op, Args... args) DFPATTERN_FUNCTOR_DEFAULT; - virtual R VisitDFPattern_(const FunctionPatternNode* op, Args... args) DFPATTERN_FUNCTOR_DEFAULT; - virtual R VisitDFPattern_(const IfPatternNode* op, Args... args) DFPATTERN_FUNCTOR_DEFAULT; - virtual R VisitDFPattern_(const LetPatternNode* op, Args... args) DFPATTERN_FUNCTOR_DEFAULT; - virtual R VisitDFPattern_(const ShapePatternNode* op, Args... args) DFPATTERN_FUNCTOR_DEFAULT; - virtual R VisitDFPattern_(const TupleGetItemPatternNode* op, - Args... args) DFPATTERN_FUNCTOR_DEFAULT; - virtual R VisitDFPattern_(const TuplePatternNode* op, Args... args) DFPATTERN_FUNCTOR_DEFAULT; - virtual R VisitDFPattern_(const TypePatternNode* op, Args... args) DFPATTERN_FUNCTOR_DEFAULT; - virtual R VisitDFPattern_(const VarPatternNode* op, Args... args) DFPATTERN_FUNCTOR_DEFAULT; - virtual R VisitDFPattern_(const WildcardPatternNode* op, Args... args) DFPATTERN_FUNCTOR_DEFAULT; - virtual R VisitDFPatternDefault_(const Object* op, Args...) { - LOG(FATAL) << "Do not have a default for " << op->GetTypeKey(); - throw; - } - - private: - // initialize the vtable. - static FType InitVTable() { - FType vtable; - // Set dispatch - RELAY_DFPATTERN_FUNCTOR_DISPATCH(AltPatternNode); - RELAY_DFPATTERN_FUNCTOR_DISPATCH(AttrPatternNode); - RELAY_DFPATTERN_FUNCTOR_DISPATCH(CallPatternNode); - RELAY_DFPATTERN_FUNCTOR_DISPATCH(ConstantPatternNode); - RELAY_DFPATTERN_FUNCTOR_DISPATCH(DataTypePatternNode); - RELAY_DFPATTERN_FUNCTOR_DISPATCH(DominatorPatternNode); - RELAY_DFPATTERN_FUNCTOR_DISPATCH(ExprPatternNode); - RELAY_DFPATTERN_FUNCTOR_DISPATCH(FunctionPatternNode); - RELAY_DFPATTERN_FUNCTOR_DISPATCH(IfPatternNode); - RELAY_DFPATTERN_FUNCTOR_DISPATCH(LetPatternNode); - RELAY_DFPATTERN_FUNCTOR_DISPATCH(ShapePatternNode); - RELAY_DFPATTERN_FUNCTOR_DISPATCH(TupleGetItemPatternNode); - RELAY_DFPATTERN_FUNCTOR_DISPATCH(TuplePatternNode); - RELAY_DFPATTERN_FUNCTOR_DISPATCH(TypePatternNode); - RELAY_DFPATTERN_FUNCTOR_DISPATCH(VarPatternNode); - RELAY_DFPATTERN_FUNCTOR_DISPATCH(WildcardPatternNode); - return vtable; - } -}; - -/*! - * \brief A simple visitor wrapper around DFPatternFunctor. - * Recursively visit the content. - * - * DFPatternVisitor treats the Pattern as dataflow graph,and only visit each Expr node once. - */ -class DFPatternVisitor : public DFPatternFunctor { - public: - void VisitDFPattern(const DFPattern& pattern) override; - void VisitDFPattern_(const AltPatternNode* op) override; - void VisitDFPattern_(const AttrPatternNode* op) override; - void VisitDFPattern_(const CallPatternNode* op) override; - void VisitDFPattern_(const ConstantPatternNode* op) override; - void VisitDFPattern_(const DataTypePatternNode* op) override; - void VisitDFPattern_(const DominatorPatternNode* op) override; - void VisitDFPattern_(const ExprPatternNode* op) override; - void VisitDFPattern_(const FunctionPatternNode* op) override; - void VisitDFPattern_(const IfPatternNode* op) override; - void VisitDFPattern_(const LetPatternNode* op) override; - void VisitDFPattern_(const ShapePatternNode* op) override; - void VisitDFPattern_(const TupleGetItemPatternNode* op) override; - void VisitDFPattern_(const TuplePatternNode* op) override; - void VisitDFPattern_(const TypePatternNode* op) override; - void VisitDFPattern_(const VarPatternNode* op) override; - void VisitDFPattern_(const WildcardPatternNode* op) override; - - protected: - // set of already-visited nodes - std::unordered_set visited_; -}; - -} // namespace relay -} // namespace tvm -#endif // TVM_RELAY_DATAFLOW_PATTERN_FUNCTOR_H_ diff --git a/include/tvm/relay/error.h b/include/tvm/relay/error.h deleted file mode 100644 index abe8278f2f5d..000000000000 --- a/include/tvm/relay/error.h +++ /dev/null @@ -1,181 +0,0 @@ -/* - * 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. - */ -#ifndef TVM_RELAY_ERROR_H_ -#define TVM_RELAY_ERROR_H_ - -#include - -#include -#include -#include -#include - -namespace tvm { -namespace relay { -/*! - * \brief A wrapper around std::stringstream to build error. - *include/tvm/ir/type.h - * Can be consumed by CompileError to construct an error. - * - * \code - * - * void ReportError(const CompileError& err); - * - * void Test(int number) { - * // Use error reporter to construct an error. - * ReportError(ErrorBuilder() << "This is an error number=" << number); - * } - * - * \endcode - */ -struct ErrorBuilder { - public: - template - ErrorBuilder& operator<<(const T& val) { // NOLINT(*) - stream_ << val; - return *this; - } - - private: - std::stringstream stream_; - friend class CompileError; -}; - -/*! - * \brief Custom Error class to be thrown during compilation. - */ -class CompileError : public Error { - public: - /*! \brief Location of the error */ - Span span; - /*! - * \brief construct error from message. - * \param msg The message - */ - explicit CompileError(const std::string& msg) : Error(msg), span(nullptr) {} - /*! - * \brief construct error from error builder. - * \param err The error builder - */ - CompileError(const ErrorBuilder& err) : Error(err.stream_.str()), span(nullptr) {} // NOLINT(*) - /*! - * \brief copy constructor. - * \param other The other ereor. - */ - CompileError(const CompileError& other) : Error(other.what()), span(other.span) {} // NOLINT(*) - /*! - * \brief default constructor. */ - CompileError() : Error(""), span(nullptr) {} -}; - -/*! - * \brief An abstraction around how errors are stored and reported. - * Designed to be opaque to users, so we can support a robust and simpler - * error reporting mode, as well as a more complex mode. - * - * The first mode is the most accurate: we report a Relay error at a specific - * Span, and then render the error message directly against a textual representation - * of the program, highlighting the exact lines in which it occurs. This mode is not - * implemented in this PR and will not work. - * - * The second mode is a general-purpose mode, which attempts to annotate the program's - * textual format with errors. - * - * The final mode represents the old mode, if we report an error that has no span or - * expression, we will default to throwing an exception with a textual representation - * of the error and no indication of where it occurred in the original program. - * - * The latter mode is not ideal, and the goal of the new error reporting machinery is - * to avoid ever reporting errors in this style. - */ -class ErrorReporter { - public: - /*! \brief default constructor. */ - ErrorReporter() : errors_(), node_to_error_() {} - - /*! - * \brief Report a CompileError. - * - * This API is useful for reporting spanned errors. - * - * \param err The error to report. - */ - void Report(const CompileError& err) { - if (!err.span.defined()) { - throw err; - } - - this->errors_.push_back(err); - } - - /*! - * \brief Report an error against a program, using the full program - * error reporting strategy. - * - * This error reporting method requires the global function in which - * to report an error, the expression to report the error on, - * and the error object. - * - * \param global The global function in which the expression is contained. - * \param node The expression or type to report the error at. - * \param err The error message to report. - */ - void ReportAt(const GlobalVar& global, const ObjectRef& node, std::stringstream& err) { - std::string err_msg = err.str(); - this->ReportAt(global, node, CompileError(err_msg)); - } - - /*! - * \brief Report an error against a program, using the full program - * error reporting strategy. - * - * This error reporting method requires the global function in which - * to report an error, the expression to report the error on, - * and the error object. - * - * \param global The global function in which the expression is contained. - * \param node The expression or type to report the error at. - * \param err The error to report. - */ - void ReportAt(const GlobalVar& global, const ObjectRef& node, const CompileError& err); - - /*! - * \brief Render all reported errors and exit the program. - * - * This function should be used after executing a pass to render reported errors. - * - * It will build an error message from the set of errors, depending on the error - * reporting strategy. - * - * \param module The module to report errors on. - * \param use_color Controls whether to colorize the output. - */ - void RenderErrors(const IRModule& module, bool use_color = true); - - inline bool AnyErrors() { return errors_.size() != 0; } - - private: - std::vector errors_; - std::unordered_map, ObjectPtrHash, ObjectPtrEqual> node_to_error_; - std::unordered_map node_to_gv_; -}; - -} // namespace relay -} // namespace tvm -#endif // TVM_RELAY_ERROR_H_ diff --git a/include/tvm/relay/executor.h b/include/tvm/relay/executor.h deleted file mode 100644 index 858ba5cfe198..000000000000 --- a/include/tvm/relay/executor.h +++ /dev/null @@ -1,277 +0,0 @@ -/* - * 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 tvm/relay/executor.h - * \brief Object representation of Executor configuration and registry - */ -#ifndef TVM_RELAY_EXECUTOR_H_ -#define TVM_RELAY_EXECUTOR_H_ - -#include -#include -#include -#include -#include -#include -#include - -#include -#include -#include -#include - -namespace tvm { - -template -class AttrRegistry; - -namespace relay { - -/*! - * \brief Executor information. - * - * This data structure stores the meta-data - * about executors which can be used to pass around information. - * - * \sa Executor - */ -class ExecutorNode : public Object { - public: - /*! \brief name of the Executor */ - String name; - /* \brief Additional attributes storing meta-data about the Executor. */ - DictAttrs attrs; - - /*! - * \brief Should Link Parameters into the module - * \return Whether the Executor is configured to execute modules with linked parameters - */ - Bool ShouldLinkParameters() const { - return name == "aot" || GetAttr("link-params").value_or(Bool(false)); - } - - /*! - * \brief Get an attribute. - * - * \param attr_key The attribute key. - * \param default_value The default value if the key does not exist, defaults to nullptr. - * - * \return The result - * - * \tparam TObjectRef the expected object type. - * \throw Error if the key exists but the value does not match TObjectRef - * - * \code - * - * void GetAttrExample(const Executor& executor) { - * auto value = executor->GetAttr("AttrKey", 0); - * } - * - * \endcode - */ - template - Optional GetAttr( - const std::string& attr_key, - Optional default_value = Optional(nullptr)) const { - return attrs.GetAttr(attr_key, default_value); - } - // variant that uses TObjectRef to enable implicit conversion to default value. - template - Optional GetAttr(const std::string& attr_key, TObjectRef default_value) const { - return GetAttr(attr_key, Optional(default_value)); - } - - void VisitAttrs(AttrVisitor* v) { - v->Visit("name", &name); - v->Visit("attrs", &attrs); - } - - bool SEqualReduce(const ExecutorNode* other, SEqualReducer equal) const { - return name == other->name && equal.DefEqual(attrs, other->attrs); - } - - void SHashReduce(SHashReducer hash_reduce) const { - hash_reduce(name); - hash_reduce(attrs); - } - - static constexpr const char* _type_key = "Executor"; - static constexpr const bool _type_has_method_sequal_reduce = true; - static constexpr const bool _type_has_method_shash_reduce = true; - TVM_DECLARE_FINAL_OBJECT_INFO(ExecutorNode, Object); -}; - -/*! - * \brief Managed reference class to ExecutorNode. - * \sa ExecutorNode - */ -class Executor : public ObjectRef { - public: - /*! - * \brief Create a new Executor object using the registry - * \throws Error if name is not registered - * \param name The name of the executor. - * \param attrs Attributes for the executor. - * \return the new Executor object. - */ - TVM_DLL static Executor Create(String name, Map attrs = {}); - - /*! - * \brief List all registered Executors - * \return the list of Executors - */ - TVM_DLL static Array ListExecutors(); - - /*! - * \brief List all options for a specific Executor - * \param name The name of the Executor - * \return Map of option name to type - */ - TVM_DLL static Map ListExecutorOptions(const String& name); - - /*! \brief specify container node */ - TVM_DEFINE_OBJECT_REF_METHODS(Executor, ObjectRef, ExecutorNode); - TVM_DEFINE_OBJECT_REF_COW_METHOD(ExecutorNode) - - private: - /*! - * \brief Private Constructor - * \param name The executor name - * \param attrs Attributes to apply to this Executor node - */ - TVM_DLL Executor(String name, DictAttrs attrs) { - auto n = make_object(); - n->name = std::move(name); - n->attrs = std::move(attrs); - data_ = std::move(n); - } -}; - -/*! - * \brief Helper structure to register Executors - * \sa TVM_REGISTER_EXECUTOR - */ -class ExecutorRegEntry { - public: - /*! - * \brief Register a valid configuration option and its ValueType for validation - * \param key The configuration key - * \tparam ValueType The value type to be registered - */ - template - inline ExecutorRegEntry& add_attr_option(const String& key); - - /*! - * \brief Register a valid configuration option and its ValueType for validation - * \param key The configuration key - * \param default_value The default value of the key - * \tparam ValueType The value type to be registered - */ - template - inline ExecutorRegEntry& add_attr_option(const String& key, ObjectRef default_value); - - /*! - * \brief Register or get a new entry. - * \param name The name of the operator. - * \return the corresponding entry. - */ - TVM_DLL static ExecutorRegEntry& RegisterOrGet(const String& name); - - private: - /*! \brief Internal storage of value types */ - struct ValueTypeInfo { - std::string type_key; - uint32_t type_index; - }; - std::unordered_map key2vtype_; - /*! \brief A hash table that stores the default value of each attr */ - std::unordered_map key2default_; - - /*! \brief Index used for internal lookup of attribute registry */ - uint32_t index_; - - // the name - std::string name; - - /*! \brief Return the index stored in attr registry */ - uint32_t AttrRegistryIndex() const { return index_; } - /*! \brief Return the name stored in attr registry */ - String AttrRegistryName() const { return name; } - - /*! \brief private constructor */ - explicit ExecutorRegEntry(uint32_t reg_index) : index_(reg_index) {} - - // friend class - template - friend class AttrRegistryMapContainerMap; - template - friend class tvm::AttrRegistry; - friend class Executor; -}; - -template -inline ExecutorRegEntry& ExecutorRegEntry::add_attr_option(const String& key) { - ICHECK(!key2vtype_.count(key)) << "AttributeError: add_attr_option failed because '" << key - << "' has been set once"; - - using ValueNodeType = typename ValueType::ContainerType; - // NOTE: we could further update the function later. - uint32_t value_type_index = ValueNodeType::_GetOrAllocRuntimeTypeIndex(); - - ValueTypeInfo info; - info.type_index = value_type_index; - info.type_key = runtime::Object::TypeIndex2Key(value_type_index); - key2vtype_[key] = info; - return *this; -} - -template -inline ExecutorRegEntry& ExecutorRegEntry::add_attr_option(const String& key, - ObjectRef default_value) { - add_attr_option(key); - key2default_[key] = default_value; - return *this; -} - -// internal macros to make executor entries -#define TVM_EXECUTOR_REGISTER_VAR_DEF \ - static DMLC_ATTRIBUTE_UNUSED ::tvm::relay::ExecutorRegEntry& __make_##Executor - -/*! - * \def TVM_REGISTER_EXECUTOR - * \brief Register a new executor, or set attribute of the corresponding executor. - * - * \param ExecutorName The name of registry - * - * \code - * - * TVM_REGISTER_EXECUTOR("aot") - * .add_attr_option("my_option"); - * .add_attr_option("my_option_default", String("default")); - * - * \endcode - */ -#define TVM_REGISTER_EXECUTOR(ExecutorName) \ - TVM_STR_CONCAT(TVM_EXECUTOR_REGISTER_VAR_DEF, __COUNTER__) = \ - ::tvm::relay::ExecutorRegEntry::RegisterOrGet(ExecutorName) -} // namespace relay -} // namespace tvm - -#endif // TVM_RELAY_EXECUTOR_H_ diff --git a/include/tvm/relay/expr.h b/include/tvm/relay/expr.h deleted file mode 100644 index 854050464d4a..000000000000 --- a/include/tvm/relay/expr.h +++ /dev/null @@ -1,833 +0,0 @@ -/* - * 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 tvm/relay/expr.h - * \brief Relay expression language. - */ -#ifndef TVM_RELAY_EXPR_H_ -#define TVM_RELAY_EXPR_H_ - -#include -#include -#include -#include -#include - -#include -#include -#include -#include - -#include "./base.h" -#include "./type.h" - -namespace tvm { - -/*! - * \brief Returns \p global_var with the given properties. A null property denotes 'no change'. - * Returns \p global_var if all properties are unchanged. Otherwise, returns a copy with the new - * fields. - */ -GlobalVar WithFields(GlobalVar global_var, Optional opt_name_hint = {}, - Optional opt_type = {}, Optional opt_virtual_device = {}, - Optional opt_span = {}); - -namespace relay { - -using Expr = tvm::RelayExpr; -using ExprNode = tvm::RelayExprNode; -using BaseFunc = tvm::BaseFunc; -using BaseFuncNode = tvm::BaseFuncNode; -using GlobalVar = tvm::GlobalVar; -using GlobalVarNode = tvm::GlobalVarNode; - -/*! - * \brief Constant tensor, backed by an NDArray on the cpu(0) device. - * - * \note Scalar constants are represented by rank-0 const tensor. - * Constant folding are handled uniformly via Tensor types. - */ -class Constant; -/*! - * \brief Constant tensor type. - */ -class ConstantNode : public ExprNode { - public: - /*! \brief The data of the tensor */ - runtime::NDArray data; - - /*! \return The corresponding tensor type of the data */ - TensorType tensor_type() const; - - /*! \return Whether it is scalar(rank-0 tensor) */ - bool is_scalar() const { return data->ndim == 0; } - - void VisitAttrs(tvm::AttrVisitor* v) { - v->Visit("data", &data); - v->Visit("virtual_device_", &virtual_device_); - v->Visit("span", &span); - v->Visit("_checked_type_", &checked_type_); - } - - bool SEqualReduce(const ConstantNode* other, SEqualReducer equal) const { - return equal(data, other->data); - } - - void SHashReduce(SHashReducer hash_reduce) const { hash_reduce(data); } - - static constexpr const char* _type_key = "relay.Constant"; - TVM_DECLARE_FINAL_OBJECT_INFO(ConstantNode, ExprNode); -}; - -class Constant : public Expr { - public: - /*! - * \brief The constructor - * \param data The data of the constant tensor. - * \param span The source span of the expression. - */ - TVM_DLL explicit Constant(runtime::NDArray data, Span span = Span()); - - TVM_DEFINE_OBJECT_REF_METHODS(Constant, RelayExpr, ConstantNode); - TVM_DEFINE_OBJECT_REF_COW_METHOD(ConstantNode); -}; - -/*! - * \brief Returns \p constant with the given properties. A null property denotes 'no change'. - * Returns \p constant if all properties are unchanged. Otherwise, returns a copy with the new - * fields. - */ -Constant WithFields(Constant constant, Optional opt_data = {}, - Optional opt_virtual_device = {}, Optional opt_span = {}); - -/*! \brief Tuple of multiple Exprs */ -class Tuple; -/*! \brief Tuple container */ -class TupleNode : public ExprNode { - public: - /*! \brief the fields of the tuple */ - tvm::Array fields; - - void VisitAttrs(tvm::AttrVisitor* v) { - v->Visit("fields", &fields); - v->Visit("virtual_device_", &virtual_device_); - v->Visit("span", &span); - v->Visit("_checked_type_", &checked_type_); - } - - bool SEqualReduce(const TupleNode* other, SEqualReducer equal) const { - // specially handle empty tuple as a constant is not a graph node. - if (fields.size() == other->fields.size() && fields.size() == 0) { - return true; - } else { - equal->MarkGraphNode(); - return equal(fields, other->fields); - } - } - - void SHashReduce(SHashReducer hash_reduce) const { - if (fields.size() != 0) { - hash_reduce->MarkGraphNode(); - hash_reduce(fields); - } - } - - static constexpr const char* _type_key = "relay.Tuple"; - TVM_DECLARE_FINAL_OBJECT_INFO(TupleNode, ExprNode); -}; - -class Tuple : public Expr { - public: - /*! - * \brief The constructor - * \param fields The fields of a tuple. - * \param span The source span of the expression. - */ - TVM_DLL explicit Tuple(tvm::Array fields, Span span = Span()); - - TVM_DEFINE_OBJECT_REF_METHODS(Tuple, RelayExpr, TupleNode); - TVM_DEFINE_OBJECT_REF_COW_METHOD(TupleNode); -}; - -/*! - * \brief Returns \p tuple with the given properties. A null property denotes 'no change'. - * Returns \p tuple if all properties are unchanged. Otherwise, returns a copy with the new - * fields. - */ -Tuple WithFields(Tuple tuple, Optional> opt_fields = Optional>(), - Optional opt_virtual_device = Optional(), - Optional opt_span = Optional()); - -/*! - * \brief Local variables used in the let expression. - * - * Its semantics are similar to tvm.Var node used in TVM's low level - * tensor expression language. - * - * \note Each Var is bind only once and is immutable. - */ -class Var; -/*! \brief Container for Var */ -class VarNode : public ExprNode { - public: - /*! - * \brief The unique identifier of the Var. - * - * vid will be preserved for the same Var during type inference - * and other rewritings, while the VarNode might be recreated - * to attach additional information. - * This property can be used to keep track of parameter Var - * information across passes. - */ - Id vid; - /*! - * \brief type annotaion of the variable. - * This field records user provided type annotation of the Var. - * This field is optional and can be None. - */ - Type type_annotation; - - /*! \return The name hint of the variable */ - const String& name_hint() const { return vid->name_hint; } - - void VisitAttrs(tvm::AttrVisitor* v) { - v->Visit("vid", &vid); - v->Visit("type_annotation", &type_annotation); - v->Visit("virtual_device_", &virtual_device_); - v->Visit("span", &span); - v->Visit("_checked_type_", &checked_type_); - } - - bool SEqualReduce(const VarNode* other, SEqualReducer equal) const { - equal->MarkGraphNode(); - return equal(type_annotation, other->type_annotation) && equal(vid, other->vid) && - equal(virtual_device_, other->virtual_device_); - } - - void SHashReduce(SHashReducer hash_reduce) const { - hash_reduce->MarkGraphNode(); - hash_reduce(type_annotation); - hash_reduce(vid); - } - - static constexpr const char* _type_key = "relay.Var"; - TVM_DECLARE_FINAL_OBJECT_INFO(VarNode, ExprNode); -}; - -class Var : public Expr { - public: - /*! - * \brief The constructor - * \param name_hint The name hint of a variable. - * \param type_annotation The type annotation of a variable. - * \param span The source span of the expression. - */ - TVM_DLL Var(String name_hint, Type type_annotation, Span span = Span()) - : Var(Id(name_hint), type_annotation, span) {} - - /*! - * \brief The constructor - * \param vid The unique id of a variable. - * \param type_annotation The type annotation of a variable. - * \param span The source span of the expression. - */ - TVM_DLL Var(Id vid, Type type_annotation, Span span = Span()); - - /*! - * \brief Return a globally fresh name. Helps with debugging to follow the same - * variable between passes and sub-expressions. - * - * TODO(mbs): Replace with name creation w.r.t. scopes once available as part of - * name gen overhaul. - */ - static Var GenSym(Type type_annotation = {}, Span span = {}); - - TVM_DEFINE_OBJECT_REF_METHODS(Var, RelayExpr, VarNode); - TVM_DEFINE_OBJECT_REF_COW_METHOD(VarNode); -}; - -/*! - * \brief Returns \p var with the given properties. A null property denotes 'no change'. - * Returns \p var if all properties are unchanged. Otherwise, returns a copy with the new - * fields. - */ -Var WithFields(Var var, Optional opt_vid = Optional(), - Optional opt_type_annotation = Optional(), - Optional opt_virtual_device = Optional(), - Optional opt_span = Optional()); - -/*! - * \brief Call corresponds to operator invocation. - * Corresponds to the operator in computational graph terminology. - */ -class Call; -/*! \brief Call container. */ -class CallNode : public ExprNode { - protected: - // CallNode uses own deleter to indirectly call non-recursive destructor - Object::FDeleter saved_deleter_; - static void Deleter_(Object* ptr); - - public: - /*! - * \brief The operator(function) being invoked - * - * - It can be tvm::Op which corresponds to the primitive operators. - * - It can also be user defined functions (Function, GlobalVar, Var). - */ - Expr op; - - /*! \brief The arguments(inputs) of the call */ - tvm::Array args; - - /*! \brief The additional attributes */ - Attrs attrs; - - /*! - * \brief The type arguments passed to polymorphic(template) function. - * - * This is the advance feature that is only used when the function is - * polymorphic. It is safe to be ignored in most cases. For example, in the - * following code, the type_args of addone call is [int]. - * - * \code - * - * template - * T addone(T a) { return a + 1; } - * - * void main() { - * int x = addone(10); - * } - * - * \endcode - */ - tvm::Array type_args; - - void VisitAttrs(tvm::AttrVisitor* v) { - v->Visit("op", &op); - v->Visit("args", &args); - v->Visit("attrs", &attrs); - v->Visit("type_args", &type_args); - v->Visit("virtual_device_", &virtual_device_); - v->Visit("span", &span); - v->Visit("_checked_type_", &checked_type_); - } - - bool SEqualReduce(const CallNode* other, SEqualReducer equal) const { - // skip type_args check for primitive ops. - equal->MarkGraphNode(); - return equal(op, other->op) && equal(args, other->args) && equal(attrs, other->attrs) && - (IsPrimitiveOp(op) || equal(type_args, other->type_args)); - } - - void SHashReduce(SHashReducer hash_reduce) const { - hash_reduce->MarkGraphNode(); - hash_reduce(op); - hash_reduce(args); - hash_reduce(attrs); - if (!IsPrimitiveOp(op)) { - hash_reduce(type_args); - } - } - - static constexpr const char* _type_key = "relay.Call"; - TVM_DECLARE_FINAL_OBJECT_INFO(CallNode, ExprNode); - template - friend class runtime::ObjAllocatorBase; - friend class Call; -}; - -class Call : public Expr { - public: - /*! - * \brief The destructor - */ - ~Call(); - - /*! - * \brief The constructor - * \param op The operator will be invoked. - * \param args The arguments of the call. - * \param attrs The attributes of the call node. - * \param type_args The type arguments passed to a polymorphic function. - * \param span The source span of the expression. - */ - TVM_DLL Call(Expr op, Array args, Attrs attrs = Attrs(), - Array type_args = Array(), Span span = Span()); - - TVM_DEFINE_OBJECT_REF_METHODS(Call, RelayExpr, CallNode); - TVM_DEFINE_OBJECT_REF_COW_METHOD(CallNode); -}; - -/*! - * \brief Returns \p call with the given properties. A null property denotes 'no change'. - * Returns \p call if all properties are unchanged. Otherwise, returns a copy with the new - * fields. - */ -Call WithFields(Call call, Optional opt_op = Optional(), - Optional> opt_args = Optional>(), - Optional opt_attrs = Optional(), - Optional> opt_type_args = Optional>(), - Optional opt_virtual_device = Optional(), - Optional opt_span = Optional()); - -/*! - * \brief Let binding that binds a local var and optionally a type annotation. - * - * \note Let is useful to transform the program to be A-normal form. - * where each of the expression corresponds to a let binding. - * - * For developers who are familar with the computational graph. - * Each of the let can be viewed as a operator node in the computational graph. - * Traversing the list of let bindings is similar to running - * PostDFS-order(topo-order) traversal on the computational graph. - */ -class Let; -/*! \brief A binding of a sub-network. */ -class LetNode : public ExprNode { - protected: - // LetNode uses own deleter to indirectly call non-recursive destructor - Object::FDeleter saved_deleter_; - static void Deleter_(Object* ptr); - - public: - /*! \brief The variable we bind to */ - Var var; - /*! \brief The value we bind var to */ - Expr value; - /*! \brief The body of the let binding */ - Expr body; - - void VisitAttrs(tvm::AttrVisitor* v) { - v->Visit("var", &var); - v->Visit("value", &value); - v->Visit("body", &body); - v->Visit("virtual_device_", &virtual_device_); - v->Visit("span", &span); - v->Visit("_checked_type_", &checked_type_); - } - - bool SEqualReduce(const LetNode* other, SEqualReducer equal) const { - equal->MarkGraphNode(); - return equal.DefEqual(var, other->var) && equal(value, other->value) && - equal(body, other->body); - } - - void SHashReduce(SHashReducer hash_reduce) const { - hash_reduce->MarkGraphNode(); - hash_reduce.DefHash(var); - hash_reduce(value); - hash_reduce(body); - } - - static constexpr const char* _type_key = "relay.Let"; - TVM_DECLARE_FINAL_OBJECT_INFO(LetNode, ExprNode); - template - friend class runtime::ObjAllocatorBase; - friend class Let; -}; - -class Let : public Expr { - public: - /*! - * \brief The destructor - */ - ~Let(); - - /*! - * \brief The constructor - * \param var The variable that is bound to. - * \param value The value used to bind to the variable. - * \param body The body of the let binding. - * \param span The source span of the expression. - */ - TVM_DLL Let(Var var, Expr value, Expr body, Span span = Span()); - - TVM_DEFINE_OBJECT_REF_METHODS(Let, RelayExpr, LetNode); - TVM_DEFINE_OBJECT_REF_COW_METHOD(LetNode); -}; - -/*! - * \brief Returns \p let with the given properties. A null property denotes 'no change'. - * Returns \p let if all properties are unchanged. Otherwise, returns a copy with the new - * fields. - */ -Let WithFields(Let let, Optional opt_var = Optional(), - Optional opt_value = Optional(), - Optional opt_body = Optional(), - Optional opt_virtual_device = Optional(), - Optional opt_span = Optional()); - -/*! - * \brief Condition expression - * - * Unlike traditional statement `if`s, the if evalutes - * to the result of the branch taken. - * - * let x = if (true) { 1 } else { 0 }; // x is 1 - * let y = if (false) { 1 } else { 0 }; // y is 0 - * - * \note This is similar to C's ternary operator. - */ -class If; -/*! \brief container of If */ -class IfNode : public ExprNode { - public: - /*! \brief The condition */ - Expr cond; - /*! \brief The expression evaluated when condition is true. */ - Expr true_branch; - /*! \brief The expression evaluated when condition is false */ - Expr false_branch; - - void VisitAttrs(tvm::AttrVisitor* v) { - v->Visit("cond", &cond); - v->Visit("true_branch", &true_branch); - v->Visit("false_branch", &false_branch); - v->Visit("virtual_device_", &virtual_device_); - v->Visit("span", &span); - v->Visit("_checked_type_", &checked_type_); - } - - bool SEqualReduce(const IfNode* other, SEqualReducer equal) const { - equal->MarkGraphNode(); - return equal(cond, other->cond) && equal(true_branch, other->true_branch) && - equal(false_branch, other->false_branch); - } - - void SHashReduce(SHashReducer hash_reduce) const { - hash_reduce->MarkGraphNode(); - hash_reduce(cond); - hash_reduce(true_branch); - hash_reduce(false_branch); - } - - static constexpr const char* _type_key = "relay.If"; - TVM_DECLARE_FINAL_OBJECT_INFO(IfNode, ExprNode); -}; - -class If : public Expr { - public: - /*! - * \brief The constructor - * \param cond The condition of a if node. - * \param true_branch The fall through branch - * \param false_branch The branch for execution when condition is false. - * \param span The source span of the expression. - */ - TVM_DLL If(Expr cond, Expr true_branch, Expr false_branch, Span span = Span()); - - TVM_DEFINE_OBJECT_REF_METHODS(If, RelayExpr, IfNode); - TVM_DEFINE_OBJECT_REF_COW_METHOD(IfNode); -}; - -/*! - * \brief Returns \p if_expr with the given properties. A null property denotes 'no change'. - * Returns \p if_expr if all properties are unchanged. Otherwise, returns a copy with the new - * fields. - */ -If WithFields(If if_expr, Optional opt_cond = Optional(), - Optional opt_true_branch = Optional(), - Optional opt_false_branch = Optional(), - Optional opt_virtual_device = Optional(), - Optional opt_span = Optional()); - -/*! \brief Get index-th field out of a tuple. */ -class TupleGetItem; -class TupleGetItemNode : public ExprNode { - public: - /*! \brief The tuple Expression */ - Expr tuple; - /*! \brief which value to get */ - int index; - - void VisitAttrs(tvm::AttrVisitor* v) { - v->Visit("tuple_value", &tuple); - v->Visit("index", &index); - v->Visit("virtual_device_", &virtual_device_); - v->Visit("span", &span); - v->Visit("_checked_type_", &checked_type_); - } - - bool SEqualReduce(const TupleGetItemNode* other, SEqualReducer equal) const { - return equal(tuple, other->tuple) && equal(index, other->index); - } - - void SHashReduce(SHashReducer hash_reduce) const { - hash_reduce(tuple); - hash_reduce(index); - } - - static constexpr const char* _type_key = "relay.TupleGetItem"; - TVM_DECLARE_FINAL_OBJECT_INFO(TupleGetItemNode, ExprNode); -}; - -class TupleGetItem : public Expr { - public: - /*! - * \brief The constructor - * \param tuple The tuple to get an element from. - * \param index The index for extracting a value in the tuple. - * \param span The source span of the expression. - */ - TVM_DLL TupleGetItem(Expr tuple, int index, Span span = Span()); - - TVM_DEFINE_OBJECT_REF_METHODS(TupleGetItem, RelayExpr, TupleGetItemNode); - TVM_DEFINE_OBJECT_REF_COW_METHOD(TupleGetItemNode); -}; - -/*! - * \brief Returns \p tuple_get_item with the given properties. A null property denotes 'no change'. - * Returns \p tuple_get_item if all properties are unchanged. Otherwise, returns a copy with the new - * fields. - */ -TupleGetItem WithFields(TupleGetItem tuple_get_item, Optional opt_tuple = Optional(), - Optional opt_index = Optional(), - Optional opt_virtual_device = Optional(), - Optional opt_span = Optional()); - -/*! \brief Create a new Reference out of initial value. */ -class RefCreate; -class RefCreateNode : public ExprNode { - public: - /*! \brief The initial value of the Reference. */ - Expr value; - - void VisitAttrs(tvm::AttrVisitor* v) { - v->Visit("value", &value); - v->Visit("virtual_device_", &virtual_device_); - v->Visit("span", &span); - v->Visit("_checked_type_", &checked_type_); - } - - bool SEqualReduce(const RefCreateNode* other, SEqualReducer equal) const { - equal->MarkGraphNode(); - return equal(value, other->value); - } - - void SHashReduce(SHashReducer hash_reduce) const { - hash_reduce->MarkGraphNode(); - hash_reduce(value); - } - - static constexpr const char* _type_key = "relay.RefCreate"; - TVM_DECLARE_FINAL_OBJECT_INFO(RefCreateNode, ExprNode); -}; - -class RefCreate : public Expr { - public: - /*! - * \brief The constructor - * \param value The initial value of the reference. - * \param span The source span of the expression. - */ - TVM_DLL explicit RefCreate(Expr value, Span span = Span()); - - TVM_DEFINE_OBJECT_REF_METHODS(RefCreate, RelayExpr, RefCreateNode); - TVM_DEFINE_OBJECT_REF_COW_METHOD(RefCreateNode); -}; - -/*! - * \brief Returns \p ref_create with the given properties. A null property denotes 'no change'. - * Returns \p ref_crete if all properties are unchanged. Otherwise, returns a copy with the new - * fields. - */ -RefCreate WithFields(RefCreate ref_create, Optional opt_value = Optional(), - Optional opt_virtual_device = Optional(), - Optional opt_span = Optional()); - -/*! \brief Get value out of Reference. */ -class RefRead; -class RefReadNode : public ExprNode { - public: - /*! \brief The Reference Expression. */ - Expr ref; - - void VisitAttrs(tvm::AttrVisitor* v) { - v->Visit("ref", &ref); - v->Visit("virtual_device_", &virtual_device_); - v->Visit("span", &span); - v->Visit("_checked_type_", &checked_type_); - } - - bool SEqualReduce(const RefReadNode* other, SEqualReducer equal) const { - equal->MarkGraphNode(); - return equal(ref, other->ref); - } - - void SHashReduce(SHashReducer hash_reduce) const { - hash_reduce->MarkGraphNode(); - hash_reduce(ref); - } - - static constexpr const char* _type_key = "relay.RefRead"; - TVM_DECLARE_FINAL_OBJECT_INFO(RefReadNode, ExprNode); -}; - -class RefRead : public Expr { - public: - /*! - * \brief The constructor - * \param ref The reference where to read data. - * \param span The source span of the expression. - */ - TVM_DLL explicit RefRead(Expr ref, Span span = Span()); - - TVM_DEFINE_OBJECT_REF_METHODS(RefRead, RelayExpr, RefReadNode); - TVM_DEFINE_OBJECT_REF_COW_METHOD(RefReadNode); -}; - -/*! - * \brief Returns \p ref_read with the given properties. A null property denotes 'no change'. - * Returns \p ref_read if all properties are unchanged. Otherwise, returns a copy with the new - * fields. - */ -RefRead WithFields(RefRead ref_read, Optional opt_ref = Optional(), - Optional opt_virtual_device = Optional(), - Optional opt_span = Optional()); - -/*! \brief Set value of Reference. The whole expression evaluates to an Empty Tuple. */ -class RefWrite; -class RefWriteNode : public ExprNode { - public: - /*! \brief The Reference Expression. */ - Expr ref; - /*! \brief The value to write into. */ - Expr value; - - void VisitAttrs(tvm::AttrVisitor* v) { - v->Visit("ref", &ref); - v->Visit("value", &value); - v->Visit("virtual_device_", &virtual_device_); - v->Visit("span", &span); - v->Visit("_checked_type_", &checked_type_); - } - - bool SEqualReduce(const RefWriteNode* other, SEqualReducer equal) const { - equal->MarkGraphNode(); - return equal(ref, other->ref) && equal(value, other->value); - } - - void SHashReduce(SHashReducer hash_reduce) const { - hash_reduce->MarkGraphNode(); - hash_reduce(ref); - hash_reduce(value); - } - - static constexpr const char* _type_key = "relay.RefWrite"; - TVM_DECLARE_FINAL_OBJECT_INFO(RefWriteNode, ExprNode); -}; - -class RefWrite : public Expr { - public: - /*! - * \brief The constructor - * \param ref The reference where data is write to. - * \param value The value to write. - * \param span The source span of the expression. - */ - TVM_DLL RefWrite(Expr ref, Expr value, Span span = Span()); - - TVM_DEFINE_OBJECT_REF_METHODS(RefWrite, RelayExpr, RefWriteNode); - TVM_DEFINE_OBJECT_REF_COW_METHOD(RefWriteNode); -}; - -/*! - * \brief Returns \p ref_write with the given properties. A null property denotes 'no change'. - * Returns \p ref_write if all properties are unchanged. Otherwise, returns a copy with the new - * fields. - */ -RefWrite WithFields(RefWrite ref_write, Optional opt_ref = Optional(), - Optional opt_value = Optional(), - Optional opt_virtual_device = Optional(), - Optional opt_span = Optional()); - -/*! - * \brief Base class of the temporary expression. - * - * TempExprs are pass specific expression that can be - * useful to define intermediate result in the - * rewriting pass such as layout or type transformation. - * - * Subclass TempExprNode allows us to pattern match on - * specific kind of TempExpr and use them for expression rewriting. - * - * TempExpr should only be used within a pass, - */ -class TempExprNode : public ExprNode { - public: - /*! \brief virtual destructor */ - virtual ~TempExprNode() {} - /*! - * \brief Convert the expression to a normal(non-temp) Expr. - * \return The corresponding normal(non-temp) expression. - */ - virtual Expr Realize() const = 0; - - static constexpr const char* _type_key = "relay.TempExpr"; - static constexpr const bool _type_has_method_sequal_reduce = false; - static constexpr const bool _type_has_method_shash_reduce = false; - static constexpr const uint32_t _type_child_slots = 0; - TVM_DECLARE_BASE_OBJECT_INFO(TempExprNode, ExprNode); -}; - -class TempExpr : public Expr { - public: - TVM_DEFINE_OBJECT_REF_METHODS(TempExpr, RelayExpr, TempExprNode); -}; - -} // namespace relay - -namespace runtime { - -template <> -template <> -inline ObjectPtr -ObjAllocatorBase::make_object() { - using Derived = SimpleObjAllocator; - using T = relay::LetNode; - using Handler = typename Derived::template Handler; - static_assert(std::is_base_of::value, "make can only be used to create Object"); - T* ptr = Handler::New(static_cast(this)); - ptr->type_index_ = T::RuntimeTypeIndex(); - ptr->saved_deleter_ = Handler::Deleter(); - ptr->deleter_ = relay::LetNode::Deleter_; - return ObjectPtr(ptr); -} - -template <> -template <> -inline ObjectPtr -ObjAllocatorBase::make_object() { - using Derived = SimpleObjAllocator; - using T = relay::CallNode; - using Handler = typename Derived::template Handler; - static_assert(std::is_base_of::value, "make can only be used to create Object"); - T* ptr = Handler::New(static_cast(this)); - ptr->type_index_ = T::RuntimeTypeIndex(); - ptr->saved_deleter_ = Handler::Deleter(); - ptr->deleter_ = relay::CallNode::Deleter_; - return ObjectPtr(ptr); -} - -} // namespace runtime - -} // namespace tvm -#endif // TVM_RELAY_EXPR_H_ diff --git a/include/tvm/relay/expr_functor.h b/include/tvm/relay/expr_functor.h deleted file mode 100644 index 2a295c9da7f9..000000000000 --- a/include/tvm/relay/expr_functor.h +++ /dev/null @@ -1,517 +0,0 @@ -/* - * 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 tvm/relay/expr_functor.h - * \brief A more powerful visitor which enables defining arbitrary function - * signatures with type based dispatch on first argument. - */ -#ifndef TVM_RELAY_EXPR_FUNCTOR_H_ -#define TVM_RELAY_EXPR_FUNCTOR_H_ - -#include -#include -#include -#include -#include -#include - -#include -#include -#include -#include -#include - -namespace tvm { -namespace relay { - -/*! - * \brief A dynamical functor that dispatches on in the first Expr argument. - * You can use this as a more powerful Visitor, since it allows you to - * define function signatures of Visit Function. - * - * \sa tvm/ir_functor.h - * - * \tparam FType function signiture - * This type is only defined for FType with function signature R(const Expr&, - * Args...) - */ -template -class ExprFunctor; - -// functions to be overriden. -#define EXPR_FUNCTOR_DEFAULT \ - { return VisitExprDefault_(op, std::forward(args)...); } - -#define RELAY_EXPR_FUNCTOR_DISPATCH(OP) \ - vtable.template set_dispatch([](const ObjectRef& n, TSelf* self, Args... args) { \ - return self->VisitExpr_(static_cast(n.get()), std::forward(args)...); \ - }); - -template -class ExprFunctor { - private: - using TSelf = ExprFunctor; - using FType = tvm::NodeFunctor; - - public: - /*! \brief the result type of this functor */ - using result_type = R; - /*! \brief virtual destructor */ - virtual ~ExprFunctor() {} - /*! - * \brief Same as call. - * \param n The expression node. - * \param args Additional arguments. - * \return The result of the call - */ - R operator()(const Expr& n, Args... args) { return VisitExpr(n, std::forward(args)...); } - /*! - * \brief The functor call. - * \param n The expression node. - * \param args Additional arguments. - * \return The result of the call - */ - virtual R VisitExpr(const Expr& n, Args... args) { - ICHECK(n.defined()) << "Found null pointer node while traversing AST. The previous pass may " - "have generated invalid data."; - static FType vtable = InitVTable(); - return vtable(n, this, std::forward(args)...); - } - // Functions that can be overriden by subclass - virtual R VisitExpr_(const ConstantNode* op, Args... args) EXPR_FUNCTOR_DEFAULT; - virtual R VisitExpr_(const TupleNode* op, Args... args) EXPR_FUNCTOR_DEFAULT; - virtual R VisitExpr_(const VarNode* op, Args... args) EXPR_FUNCTOR_DEFAULT; - virtual R VisitExpr_(const GlobalVarNode* op, Args... args) EXPR_FUNCTOR_DEFAULT; - virtual R VisitExpr_(const FunctionNode* op, Args... args) EXPR_FUNCTOR_DEFAULT; - virtual R VisitExpr_(const CallNode* op, Args... args) EXPR_FUNCTOR_DEFAULT; - virtual R VisitExpr_(const LetNode* op, Args... args) EXPR_FUNCTOR_DEFAULT; - virtual R VisitExpr_(const IfNode* op, Args... args) EXPR_FUNCTOR_DEFAULT; - virtual R VisitExpr_(const OpNode* op, Args... args) EXPR_FUNCTOR_DEFAULT; - virtual R VisitExpr_(const TupleGetItemNode* op, Args... args) EXPR_FUNCTOR_DEFAULT; - virtual R VisitExpr_(const RefCreateNode* op, Args... args) EXPR_FUNCTOR_DEFAULT; - virtual R VisitExpr_(const RefReadNode* op, Args... args) EXPR_FUNCTOR_DEFAULT; - virtual R VisitExpr_(const RefWriteNode* op, Args... args) EXPR_FUNCTOR_DEFAULT; - virtual R VisitExpr_(const ConstructorNode* op, Args... args) EXPR_FUNCTOR_DEFAULT; - virtual R VisitExpr_(const MatchNode* op, Args... args) EXPR_FUNCTOR_DEFAULT; - virtual R VisitExprDefault_(const Object* op, Args...) { - LOG(FATAL) << "Do not have a default for " << op->GetTypeKey(); - throw; - } - - private: - // initialize the vtable. - static FType InitVTable() { - FType vtable; - // Set dispatch - RELAY_EXPR_FUNCTOR_DISPATCH(ConstantNode); - RELAY_EXPR_FUNCTOR_DISPATCH(TupleNode); - RELAY_EXPR_FUNCTOR_DISPATCH(VarNode); - RELAY_EXPR_FUNCTOR_DISPATCH(GlobalVarNode); - RELAY_EXPR_FUNCTOR_DISPATCH(FunctionNode); - RELAY_EXPR_FUNCTOR_DISPATCH(CallNode); - RELAY_EXPR_FUNCTOR_DISPATCH(LetNode); - RELAY_EXPR_FUNCTOR_DISPATCH(IfNode); - RELAY_EXPR_FUNCTOR_DISPATCH(OpNode); - RELAY_EXPR_FUNCTOR_DISPATCH(TupleGetItemNode); - RELAY_EXPR_FUNCTOR_DISPATCH(RefCreateNode); - RELAY_EXPR_FUNCTOR_DISPATCH(RefReadNode); - RELAY_EXPR_FUNCTOR_DISPATCH(RefWriteNode); - RELAY_EXPR_FUNCTOR_DISPATCH(ConstructorNode); - RELAY_EXPR_FUNCTOR_DISPATCH(MatchNode); - return vtable; - } -}; - -/*! - * \brief A simple visitor wrapper around ExprFunctor. - * Recursively visit the content. - * - * ExprVisitor treats Expr as dataflow graph, - * and only visit each Expr node once. - */ -class ExprVisitor : public ::tvm::relay::ExprFunctor { - public: - void VisitExpr(const Expr& expr) override; - void VisitExpr_(const VarNode* op) override; - void VisitExpr_(const GlobalVarNode* op) override; - void VisitExpr_(const ConstantNode* op) override; - void VisitExpr_(const TupleNode* op) override; - void VisitExpr_(const FunctionNode* op) override; - void VisitExpr_(const CallNode* op) override; - void VisitExpr_(const LetNode* op) override; - void VisitExpr_(const IfNode* op) override; - void VisitExpr_(const OpNode* op) override; - void VisitExpr_(const TupleGetItemNode* op) override; - void VisitExpr_(const RefCreateNode* op) override; - void VisitExpr_(const RefReadNode* op) override; - void VisitExpr_(const RefWriteNode* op) override; - void VisitExpr_(const ConstructorNode* op) override; - void VisitExpr_(const MatchNode* op) override; - virtual void VisitType(const Type& t); - virtual void VisitClause(const Clause& c); - virtual void VisitPattern(const Pattern& c); - virtual void VisitSpan(const Span& span); - - protected: - // Internal visiting counter - std::unordered_map visit_counter_; -}; - -/*! - * \brief A wrapper around ExprFunctor which functionally updates the AST. - * - * ExprMutator treats Expr as dataflow graph, and only Mutate each Expr once. - * The mutated results are memoized in a map and reused so that - * local transformation on the dataflow preserves the graph structure. - */ -class ExprMutator : public ::tvm::relay::ExprFunctor { - public: - /*! - * \brief Mutate is alias for VisitExpr - * \return expr. - */ - Expr Mutate(const Expr& expr) { return this->VisitExpr(expr); } - Expr VisitExpr(const Expr& expr) override; - Expr VisitExpr_(const VarNode* op) override; - Expr VisitExpr_(const ConstantNode* op) override; - Expr VisitExpr_(const GlobalVarNode* op) override; - Expr VisitExpr_(const OpNode* op) override; - Expr VisitExpr_(const TupleNode* op) override; - Expr VisitExpr_(const FunctionNode* op) override; - Expr VisitExpr_(const CallNode* call_node) override; - Expr VisitExpr_(const LetNode* op) override; - Expr VisitExpr_(const IfNode* op) override; - Expr VisitExpr_(const TupleGetItemNode* op) override; - Expr VisitExpr_(const RefCreateNode* op) override; - Expr VisitExpr_(const RefReadNode* op) override; - Expr VisitExpr_(const RefWriteNode* op) override; - Expr VisitExpr_(const ConstructorNode* op) override; - Expr VisitExpr_(const MatchNode* op) override; - - /*! - * \brief Used to visit the types inside of expressions. - * - * Can be overloaded to transform the types in arbitrary - * ways, one way would be to define a sub-class of type - * visitor for types which transform them appropriately. - */ - virtual Type VisitType(const Type& t); - virtual Clause VisitClause(const Clause& c); - virtual Pattern VisitPattern(const Pattern& c); - - protected: - /*! \brief Internal map used for memoization. */ - std::unordered_map memo_; -}; - -/*! - * \brief A wrapper around ExprVisitor which traverses the Dataflow Normal AST. - * - * MixedModeVisitor treats Expr as dataflow graph, and visits in post-DFS order - * - * MixedModeVisitor provides the same recursive API as ExprVisitor, and uses - * recursion to traverse most forms of the IR, but under the hood it expands nested dataflow regions - * of the graph and processes them iteratively to prevent stack overflows - */ -class MixedModeVisitor : public ::tvm::relay::ExprVisitor { - public: - using ::tvm::relay::ExprFunctor::VisitExpr_; - - /*! \brief The constructor of MixedModeVisitor - * \param visit_limit The number of times to allow visitation to a node. Usually 1, ocassionally - * higher (i.e., 2 for dead code elimiation), limited to 10 as a sanity check. - */ - explicit MixedModeVisitor(int visit_limit = 1); - - using ExprVisitor::VisitExpr_; - - /*! - * \brief VisitExpr is finalized to preserve call expansion of dataflow regions - */ - void VisitExpr(const Expr& expr) final; - void VisitExpr_(const CallNode* op) override; - void VisitExpr_(const TupleNode* op) override; - void VisitExpr_(const TupleGetItemNode* op) override; - - protected: - /*! - * \brief A function to apply when reaching a leaf of the graph non-recursively - */ - virtual void VisitLeaf(const Expr& expr); - /*! - * \brief A function to determine if an expression has already been visited or needs to be - * re-visited - */ - virtual bool CheckVisited(const Expr& expr); - /*! - * \brief The max number of times to visit a node - */ - size_t visit_limit_; -}; - -/*! \brief Non-recursive DFS Graph Traversal for Custom Rewriting Passes - * - * MixedModeMutator treats Expr as dataflow graph, and only Rewrites each Expr once. - * The mutated results are memoized in a map and reused so that - * local transformation on the dataflow preserves the graph structure. - * - * MixedModeMutator provides the same recursive API as ExprMutator, and uses - * recursion to traverse most forms of the IR, but under the hood it expands nested dataflow regions - * of the graph and processes them iteratatively to prevent stack overflows - * - * Uses Rewrite_ API of ExprRewriter for a cleaner split between recrusive and non-recursive - * behavior. - */ -class MixedModeMutator : public ::tvm::relay::ExprMutator { - public: - using ::tvm::relay::ExprFunctor::VisitExpr_; - - MixedModeMutator(bool pre = false) : pre_{pre} {}; - Expr VisitExpr(const Expr& expr) final; - - virtual Expr DispatchVisitExpr(const Expr& expr); - Expr VisitExpr_(const TupleNode* op) final { return Rewrite(op); }; - Expr VisitExpr_(const CallNode* call_node) final { return Rewrite(call_node); }; - Expr VisitExpr_(const TupleGetItemNode* op) final { return Rewrite(op); }; - /*! - * \brief Users should override Rewrite_ methods to implement their pass. Rewrite_ functions will - * be able to rewrite the op only with data about the original node `pre` and the same node with - * modified inputs `post` and should not recurse. - * - * \param pre The expression node before rewriting. - * \param post The expression with rewritten inputs. - */ - virtual Expr Rewrite_(const TupleNode* pre, const Expr& post) { return post; } - virtual Expr Rewrite_(const CallNode* pre, const Expr& post) { return post; } - virtual Expr Rewrite_(const TupleGetItemNode* pre, const Expr& post) { return post; } - - protected: - bool pre_; - /*! \brief Implement Rewrite API by calling ExprMutator's VisitExpr_(op) to get a `post` node with - * changed inputs. - */ - template - Expr Rewrite(const T* op) { - Expr post = ExprMutator::VisitExpr_(op); - return Rewrite_(op, post); - } - - virtual void VisitLeaf(const Expr& expr); - virtual bool CheckVisited(const Expr& expr); -}; - -#define RELAY_EXPR_REWRITER_DISPATCH(OP) \ - vtable.template set_dispatch([](const ObjectRef& n, TSelf* self, const Expr& post) { \ - return self->Rewrite_(static_cast(n.get()), post); \ - }); - -#define EXPR_REWRITER_REWRITE_DEFAULT \ - { return post; } - -/*! \brief A non-iterating Expression Rewriter - * - * ExprRewriter provides a Rewrite interface for modifying graphs in Post-DFS order. - * - * The expectation is that ExprRewriter objects will be passed to PostOrderRewrite, which will - * non-recursively unroll the graph and call Rewriting on inputs. It will then pass the original - * node, called `pre`, and a node recreated with any alterned inputs, called `post`, to the - * ExprRewriter. The ExprRewriter can then use the information in those two nodes to do more complex - * graph rewriting. - */ -class ExprRewriter { - private: - using TSelf = ExprRewriter; - using FType = tvm::NodeFunctor; - - public: - /*! \brief virtual destructor */ - virtual ~ExprRewriter() {} - /*! - * \brief Same as call. - * \param pre The expression node before rewriting. - * \param post The expression node with rewritten inputs. - * \return The result of the call - */ - Expr operator()(const Expr& pre, const Expr& post) { return Rewrite(pre, post); } - /*! - * \brief The functor call. - * \param pre The expression node before rewriting. - * \param post The expression node with rewritten inputs. - * \return The result of the call - */ - virtual Expr Rewrite(const Expr& pre, const Expr& post) { - ICHECK(pre.defined()); - static FType vtable = InitVTable(); - return vtable(pre, this, post); - } - // Functions that can be overriden by subclass, should not recurse - virtual Expr Rewrite_(const VarNode* pre, const Expr& post) EXPR_REWRITER_REWRITE_DEFAULT; - virtual Expr Rewrite_(const GlobalVarNode* pre, const Expr& post) EXPR_REWRITER_REWRITE_DEFAULT; - virtual Expr Rewrite_(const ConstantNode* pre, const Expr& post) EXPR_REWRITER_REWRITE_DEFAULT; - virtual Expr Rewrite_(const TupleNode* pre, const Expr& post) EXPR_REWRITER_REWRITE_DEFAULT; - virtual Expr Rewrite_(const FunctionNode* pre, const Expr& post) EXPR_REWRITER_REWRITE_DEFAULT; - virtual Expr Rewrite_(const CallNode* pre, const Expr& post) EXPR_REWRITER_REWRITE_DEFAULT; - virtual Expr Rewrite_(const LetNode* pre, const Expr& post) EXPR_REWRITER_REWRITE_DEFAULT; - virtual Expr Rewrite_(const IfNode* pre, const Expr& post) EXPR_REWRITER_REWRITE_DEFAULT; - virtual Expr Rewrite_(const OpNode* pre, const Expr& post) EXPR_REWRITER_REWRITE_DEFAULT; - virtual Expr Rewrite_(const TupleGetItemNode* pre, - const Expr& post) EXPR_REWRITER_REWRITE_DEFAULT; - virtual Expr Rewrite_(const RefCreateNode* pre, const Expr& post) EXPR_REWRITER_REWRITE_DEFAULT; - virtual Expr Rewrite_(const RefReadNode* pre, const Expr& post) EXPR_REWRITER_REWRITE_DEFAULT; - virtual Expr Rewrite_(const RefWriteNode* pre, const Expr& post) EXPR_REWRITER_REWRITE_DEFAULT; - virtual Expr Rewrite_(const ConstructorNode* pre, const Expr& post) EXPR_REWRITER_REWRITE_DEFAULT; - virtual Expr Rewrite_(const MatchNode* pre, const Expr& post) EXPR_REWRITER_REWRITE_DEFAULT; - - private: - // initialize the vtable. - static FType InitVTable() { - FType vtable; - // Set dispatch - RELAY_EXPR_REWRITER_DISPATCH(ConstantNode); - RELAY_EXPR_REWRITER_DISPATCH(TupleNode); - RELAY_EXPR_REWRITER_DISPATCH(VarNode); - RELAY_EXPR_REWRITER_DISPATCH(GlobalVarNode); - RELAY_EXPR_REWRITER_DISPATCH(FunctionNode); - RELAY_EXPR_REWRITER_DISPATCH(CallNode); - RELAY_EXPR_REWRITER_DISPATCH(LetNode); - RELAY_EXPR_REWRITER_DISPATCH(IfNode); - RELAY_EXPR_REWRITER_DISPATCH(OpNode); - RELAY_EXPR_REWRITER_DISPATCH(TupleGetItemNode); - RELAY_EXPR_REWRITER_DISPATCH(RefCreateNode); - RELAY_EXPR_REWRITER_DISPATCH(RefReadNode); - RELAY_EXPR_REWRITER_DISPATCH(RefWriteNode); - RELAY_EXPR_REWRITER_DISPATCH(ConstructorNode); - RELAY_EXPR_REWRITER_DISPATCH(MatchNode); - return vtable; - } -}; - -/*! \brief Non-recursive DFS Graph Traversal for Custom Rewriting Passes - * - * PostOrderRewrite does a non-recursive traversal of the graph in Post-DFS order and calls the - * ExprRewriter's Rewrite functions on nodes once their inputs are rewritten. At each rewrite call, - * PostOrderRewrite provides the original node and the node with altered inputs for use by the - * ExprRewriter. - */ -Expr PostOrderRewrite(const Expr& expr, ExprRewriter* rewriter); - -/*! - * \brief recursively visit the ir in post DFS order node, apply fvisit - * Each node is guaranteed to be visited only once. - * \param node The ir to be visited. - * \param fvisit The visitor function to be applied. - */ -void PostOrderVisit(const Expr& node, std::function fvisit); - -/*! - * \brief A struct to keep info of traversed expr in ExpandDataflow function - */ -struct v_info { - explicit v_info(Expr node_) : node{node_} {} - v_info(Expr node_, bool children_expanded_) - : node{node_}, children_expanded{children_expanded_} {}; - Expr node{}; - bool children_expanded{false}; -}; - -/*! - * \brief A function to iteratively traverse dataflow regions of a graph - * - * ExpandDataflow manually manages a stack and performs DFS to determine the processing - * order of nodes in an input graph. - * - * By default fexpand_expr implemented in a way that if it finds a dataflow node (Call, Tuple, - * TupleGetItem), it checks if the arguments to that node need to be processed via fcheck_visited. - * If so, the function pushes those arguments to the stack and continues iteratively to process - * the top of the stack. When it finds a node that doesn't match the dataflow types, or a node who's - * inputs have all been processed, it visits the current leaf via fvisit_leaf. - * - * This function should be used internally to other classes to implement mixed-mode traversals. The - * expectation is that fvisit_leaf will perform recursive analysis within mixed-mode traversal if it - * hits a non-dataflow node. - * - * fcheck_visited, fvisit_leaf and fexpand_expr are templated to encourage reusing. - */ -template -void ExpandDataflow(Expr expr, FCheckVisited fcheck_visited, FVisitLeaf fvisit_leaf, - FExpandExpr fexpand_expr) { - std::deque stack; - auto fpush_to_stack = [&fcheck_visited, &stack](const Expr& expr) { - if (!fcheck_visited(expr)) { - stack.emplace_front(v_info(expr)); - } - }; - - fpush_to_stack(expr); - while (stack.size() > 0) { - v_info* front = &stack.front(); - if (fcheck_visited(front->node)) { - stack.pop_front(); - } else if (front->children_expanded) { - fvisit_leaf(front->node); - // TODO(d-smirnov): this is for compatibility with current implementation of MixedModeVisitor - stack.pop_front(); - } else { - front->children_expanded = true; - for (auto e : fexpand_expr(front->node)) { - fpush_to_stack(e); - } - } - } -} - -template -void ExpandDataflow(Expr expr, FCheckVisited fcheck_visited, FVisitLeaf fvisit_leaf) { - auto fexpand_expr = [](const Expr& expr) { - std::vector result; - if (const CallNode* op = expr.as()) { - if (op->op == Op::Get("call_lowered")) { - // Ignore the intermediate tuple since this is purely a calling-convention detail - const auto* tuple_args = op->args[1].as(); - ICHECK(tuple_args) - << "Expected second arg to call_lowered to be a Tuple of input arguments."; - for (auto it = tuple_args->fields.rbegin(); it != tuple_args->fields.rend(); ++it) { - result.push_back(*it); - } - result.push_back(op->args[0]); - } else { - for (auto it = op->args.rbegin(); it != op->args.rend(); ++it) { - result.push_back(*it); - } - } - result.push_back(op->op); - } else if (const TupleNode* op = expr.as()) { - for (auto it = op->fields.rbegin(); it != op->fields.rend(); ++it) { - result.push_back(*it); - } - } else if (const TupleGetItemNode* op = expr.as()) { - result.push_back(op->tuple); - } - return result; - }; - ExpandDataflow(expr, fcheck_visited, fvisit_leaf, fexpand_expr); -} - -void ExpandANormalForm(const LetNode* op, std::function pre_visit, - std::function post_visit); - -} // namespace relay -} // namespace tvm -#endif // TVM_RELAY_EXPR_FUNCTOR_H_ diff --git a/include/tvm/relay/feature.h b/include/tvm/relay/feature.h deleted file mode 100644 index 136dcfa87c68..000000000000 --- a/include/tvm/relay/feature.h +++ /dev/null @@ -1,199 +0,0 @@ -/* - * 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 tvm/relay/feature.h - * \brief Detect features used in Expr/Module. - */ -#ifndef TVM_RELAY_FEATURE_H_ -#define TVM_RELAY_FEATURE_H_ - -#include -#include - -#include -#include - -namespace tvm { -namespace relay { - -/*! \brief Different kinds of relay feature a program might use. */ -enum Feature : int { - fVar = 0, - fGlobalVar = 1, - fConstant = 2, - fTuple = 3, - fTupleGetItem = 4, - fFunction = 5, - fOp = 6, - fCall = 7, - fLet = 8, - fIf = 9, - fRefCreate = 10, - fRefRead = 11, - fRefWrite = 12, - fConstructor = 13, - fMatch = 14, - /*! \brief Whether any non-atom fragment of the program is shared, making the program a graph. */ - fGraph = 15, - /*! \brief Whether there is local fixpoint in the program. */ - fLetRec = 16 -}; - -constexpr size_t feature_count = 17; - -/*! - * \brief A finite set of Feature. - */ -class FeatureSet { - public: - FeatureSet(const FeatureSet&) = default; - /*! \brief A singleton set containing a single Feature. */ - explicit FeatureSet(Feature ft) { bs_.set(static_cast(ft)); } - explicit FeatureSet(const tvm::Array& ft) { - for (Integer i : ft) { - *this += Feature(i.IntValue()); - } - } - explicit operator Array() const { - Array ret; - for (size_t i = 0; i < feature_count; ++i) { - if (bs_[i]) { - ret.push_back(Integer(i)); - } - } - return ret; - } - /*! \brief A set that contain all the Feature. */ - static FeatureSet All() { - FeatureSet fs; - fs.bs_.flip(); - return fs; - } - /*! \brief The empty set. Contain no Feature. */ - static FeatureSet No() { - FeatureSet fs; - return fs; - } - template - FeatureSet& operator+=(const T& rhs) { - bs_ |= FeatureSet(rhs).bs_; - return *this; - } - /*! \brief Set union. */ - template - FeatureSet operator+(const T& rhs) const { - FeatureSet fs(*this); - fs += rhs; - return fs; - } - template - FeatureSet& operator-=(const T& rhs) { - bs_ &= ~(FeatureSet(rhs)).bs_; - return *this; - } - /*! \brief Set difference. */ - template - FeatureSet operator-(const T& rhs) const { - FeatureSet fs(*this); - fs -= rhs; - return fs; - } - /*! - * \brief Is this a subset of rhs? - * - * \param rhs another FeatureSet. - * - * \return true only if this is a subset of rhs. - */ - bool is_subset_of(const FeatureSet& rhs) const { return ((*this) - rhs).bs_.none(); } - - /*! - * \brief return a string representation. - */ - std::string ToString() const; - - private: - std::bitset bs_; - FeatureSet() = default; - explicit FeatureSet(const std::bitset& bs) : bs_(bs) {} -}; - -/*! - * \brief Calculate the feature of the program. - * - * \param expr The expression. - * - * \return The FeatureSet. - */ -FeatureSet DetectFeature(const RelayExpr& expr); - -/*! - * \brief Calculate the feature of the program. - * - * \param mod The module. - * - * \return The FeatureSet. - */ -FeatureSet DetectFeature(const IRModule& mod); - -/*! - * \brief Calculate the feature of the program. - * - * \param expr The expression. - * \param mod The module. - * - * \return The FeatureSet. - */ -inline FeatureSet DetectFeature(const Expr& expr, const IRModule& mod) { - return DetectFeature(expr) + DetectFeature(mod); -} - -/*! - * \brief Check the feature of the program. - * - * \param expr The expression. - * \param fs The feature set of the program. - */ -void CheckFeature(const RelayExpr& expr, const FeatureSet& fs); - -/*! - * \brief Check the feature of the program. - * - * \param mod The module. - * \param fs The feature set of the program. - */ -void CheckFeature(const IRModule& mod, const FeatureSet& fs); - -/*! - * \brief Check the feature of the program. - * - * \param expr The expression. - * \param mod The module. - * \param fs The feature set of the program. - */ -inline void CheckFeature(const RelayExpr& expr, const IRModule& mod, const FeatureSet& fs) { - CheckFeature(expr, fs); - CheckFeature(mod, fs); -} - -} // namespace relay -} // namespace tvm - -#endif // TVM_RELAY_FEATURE_H_ diff --git a/include/tvm/relay/function.h b/include/tvm/relay/function.h deleted file mode 100644 index 798f6d4d2566..000000000000 --- a/include/tvm/relay/function.h +++ /dev/null @@ -1,203 +0,0 @@ -/* - * 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 tvm/relay/function.h - * \brief Relay Function. - */ -#ifndef TVM_RELAY_FUNCTION_H_ -#define TVM_RELAY_FUNCTION_H_ - -#include -#include - -#include - -namespace tvm { -namespace relay { - -/*! - * \brief Relay Function container - * \sa Function - */ -class FunctionNode : public BaseFuncNode { - public: - /*! \brief Function parameters */ - tvm::Array params; - /*! - * \brief - * The expression which represents the computation of the function, - * the expression may reference the parameters, and the type of it - * or sub-expressions may reference the type variables. - */ - Expr body; - /*! \brief User annotated return type of the function. */ - Type ret_type; - /*! - * \brief Type parameters of the function. - * Enables the function to vary its type based on these. - * This corresponds to template paramaters in c++'s terminology. - * - * \note This can be usually empty for non-polymorphic functions. - */ - tvm::Array type_params; - - void VisitAttrs(tvm::AttrVisitor* v) { - v->Visit("params", ¶ms); - v->Visit("body", &body); - v->Visit("ret_type", &ret_type); - v->Visit("type_params", &type_params); - v->Visit("attrs", &attrs); - v->Visit("virtual_device_", &virtual_device_); - v->Visit("span", &span); - v->Visit("_checked_type_", &checked_type_); - } - - bool SEqualReduce(const FunctionNode* other, SEqualReducer equal) const { - // Important to make def equal first. - equal->MarkGraphNode(); - return equal.DefEqual(params, other->params) && - equal.DefEqual(type_params, other->type_params) && equal(ret_type, other->ret_type) && - equal(attrs, other->attrs) && equal(body, other->body); - } - - void SHashReduce(SHashReducer hash_reduce) const { - hash_reduce->MarkGraphNode(); - hash_reduce.DefHash(params); - hash_reduce.DefHash(type_params); - hash_reduce(ret_type); - hash_reduce(attrs); - hash_reduce(body); - } - - /*! - * \brief Return the derived function annotation of this expression. - * - * \return The function type annotation. - * \note The function type annotation can contain IncompleteType. - */ - TVM_DLL FuncType func_type_annotation() const; - - static constexpr const char* _type_key = "relay.Function"; - TVM_DECLARE_FINAL_OBJECT_INFO(FunctionNode, BaseFuncNode); -}; - -/*! - * \brief Managed reference to FunctionNode. - * \sa FunctionNode - */ -class Function : public BaseFunc { - public: - /*! - * \brief Constructor - * \param params The parameters of the function. - * \param body The body of the function. - * \param ret_type The return type of the function. - * \param ty_params The type parameters. - * \param attrs Additional function attributes. - * \param span The span of the function. - */ - TVM_DLL Function(tvm::Array params, Expr body, Type ret_type, tvm::Array ty_params, - tvm::DictAttrs attrs = DictAttrs(), Span span = Span()); - - TVM_DEFINE_OBJECT_REF_METHODS(Function, BaseFunc, FunctionNode); - TVM_DEFINE_OBJECT_REF_COW_METHOD(FunctionNode); -}; - -/*! - * \brief Returns \p function with the given properties. A null property denotes 'no change'. - * Returns \p function if all properties are unchanged. Otherwise, returns a copy with the new - * fields. - */ -Function WithFields(Function function, Optional> opt_params = Optional>(), - Optional opt_body = Optional(), - Optional opt_ret_type = Optional(), - Optional> opt_ty_params = Optional>(), - Optional opt_attrs = Optional(), - Optional opt_virtual_device = Optional(), - Optional opt_span = Optional()); - -/* - * \brief Returns the Relay FunctionNode represented by base_func if it should be optimized, - * otherwise returns nullptr. - * - * This means returns nullptr: - * - For PrimFuncs, since not Relay Functions. - * - For Functions marked for external compilation (with "Compiler"). - * - For Functions marked as already having an external definition (with "ExternalSymbol"). - * - For Functions marked as not to be optimized (with "SkipOptimization"). - * - * TODO(mbs): Audit all enumerations of IRModule::functions to use this or some family of such. - */ -const FunctionNode* AsOptimizableFunctionNode(const BaseFunc& base_func); - -/*! - * \brief namespace of the attributes that can be attached to a relay::Function. - */ -namespace attr { - -/*! - * \brief Mark the function as representing a sub-graph which is to be lowered or compiled as - * a unit. For example, the function may represent a kernel which TVM will lower to a PrimFunc. - * If present should be bound to \p Integer(1). May be accompanied by "Compiler", see below. - * The function body should be considered opaque by Relay, and many passes simply ignore these - * functions. - * - * Type: Integer - */ -constexpr const char* kPrimitive = "Primitive"; - -/*! - * \brief Mark the function as externally implemented, ie bound in a runtime::Module within the - * IRModule's "external_mods" attribute. If present should be bound to \p Integer(1). Generally - * the only attribute when present. - * - * Type: Integer - */ -constexpr const char* kExtern = "Extern"; - -/*! - * \brief Indicates the name of the external codegen 'compiler' that should be used to lower - * or compile the function other than TVM's default lowering pipeline. The name may correspond - * to a TargetKind name. There may be a global function registered under 'relay.ext.{name}'. - * - * Type: String - */ -constexpr const char* kCompiler = "Compiler"; - -/*! \brief Indicate if the function is a closure. */ -constexpr const char* kClosure = "Closure"; -/*! \brief Store a Var to parameter/Constant mapping on a Function. */ -constexpr const char* kParams = "__params__"; -/*! \brief Mark if the function should be avoided being optimized. */ -constexpr const char* kSkipOptimization = "SkipOptimization"; -/*! \brief Treat the function as a composite operator. */ -constexpr const char* kComposite = "Composite"; -/*! \brief Mark the function to be inlined. */ -constexpr const char* kInline = "Inline"; -/*! \brief Indicate the function was created by the Pattern Partitioning Pass. */ -constexpr const char* kPartitionedFromPattern = "PartitionedFromPattern"; -/*! \brief Mark the function as only composed of reshape operations. */ -constexpr const char* kReshapeOnly = "relay.reshape_only"; - -} // namespace attr - -} // namespace relay -} // namespace tvm -#endif // TVM_RELAY_FUNCTION_H_ diff --git a/include/tvm/relay/interpreter.h b/include/tvm/relay/interpreter.h deleted file mode 100644 index f71107258d9a..000000000000 --- a/include/tvm/relay/interpreter.h +++ /dev/null @@ -1,197 +0,0 @@ -/* - * 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 tvm/relay/interpreter.h - * \brief An interpreter for Relay. - * - * This file implements a simple reference interpreter for Relay programs. - * Given a Relay module, and a Relay expression it produces a value. - * - * The interpreter's values are a naive representation of the values that - * can be produced by a Relay program and are exposed via TVM's object - * protocol to Python for introspection and debugging. - * - * The interpreter's intent is to serve as a reference semantics for the Relay IR, - * as well as for debugging and testing. - */ -#ifndef TVM_RELAY_INTERPRETER_H_ -#define TVM_RELAY_INTERPRETER_H_ - -#include -#include -#include -#include -#include - -#include - -namespace tvm { -namespace relay { - -/*! \brief The container type of Closures used by the interpreter. */ -class InterpreterClosureObj : public runtime::ClosureObj { - public: - /*! \brief The set of free variables in the closure. - * - * These are the captured variables which are required for - * evaluation when we call the closure. - */ - tvm::Map env; - /*! \brief The function which implements the closure. - * - * \note May reference the variables contained in the env. - */ - Function func; - - InterpreterClosureObj() {} - - void VisitAttrs(tvm::AttrVisitor* v) { - v->Visit("env", &env); - v->Visit("func", &func); - } - - static constexpr const char* _type_key = "interpreter.Closure"; - TVM_DECLARE_FINAL_OBJECT_INFO(InterpreterClosureObj, runtime::ClosureObj); -}; - -class InterpreterClosure : public runtime::Closure { - public: - TVM_DLL InterpreterClosure(tvm::Map env, Function func); - TVM_DEFINE_OBJECT_REF_METHODS(InterpreterClosure, runtime::Closure, InterpreterClosureObj); -}; - -/*! \brief The container type of RecClosure. */ -class RecClosureObj : public Object { - public: - /*! \brief The closure. */ - InterpreterClosure clos; - /*! \brief variable the closure bind to. */ - Var bind; - - RecClosureObj() {} - - void VisitAttrs(tvm::AttrVisitor* v) { - v->Visit("clos", &clos); - v->Visit("bind", &bind); - } - - static constexpr const char* _type_key = "interpreter.RecClosure"; - TVM_DECLARE_FINAL_OBJECT_INFO(RecClosureObj, Object); -}; - -class RecClosure : public ObjectRef { - public: - TVM_DLL RecClosure(InterpreterClosure clos, Var bind); - TVM_DEFINE_OBJECT_REF_METHODS(RecClosure, ObjectRef, RecClosureObj); -}; - -struct RefValueObj : Object { - mutable ObjectRef value; - - RefValueObj() {} - - void VisitAttrs(tvm::AttrVisitor* v) { v->Visit("value", &value); } - - static constexpr const char* _type_key = "relay.RefValue"; - TVM_DECLARE_FINAL_OBJECT_INFO(RefValueObj, Object); -}; - -class RefValue : public ObjectRef { - public: - TVM_DLL RefValue(ObjectRef val); - TVM_DEFINE_OBJECT_REF_METHODS(RefValue, ObjectRef, RefValueObj); -}; - -struct ConstructorValueObj : Object { - int32_t tag; - - tvm::Array fields; - - /*! \brief Optional field tracking ADT constructor. */ - Constructor constructor; - - void VisitAttrs(tvm::AttrVisitor* v) { - v->Visit("tag", &tag); - v->Visit("fields", &fields); - v->Visit("constructor", &constructor); - } - - static constexpr const char* _type_key = "relay.ConstructorValue"; - TVM_DECLARE_FINAL_OBJECT_INFO(ConstructorValueObj, Object); -}; - -class ConstructorValue : public ObjectRef { - public: - TVM_DLL ConstructorValue(int32_t tag, tvm::Array fields, Constructor construtor = {}); - - TVM_DEFINE_OBJECT_REF_METHODS(ConstructorValue, ObjectRef, ConstructorValueObj); -}; - -/*! - * \brief Returns a packed function over Relay expressions which will evaluate \p expr - * applied to those arguments, where \p expr is w.r.t. the definitions in \p mod. - * - * This function is intended to support the Python 'debug' executor. - * - * The given \p expr should have function type. The given \p mod may be empty or - * undefined if \p expr is self-contained. Relay arguments passed to the result - * packed function must be constants, references, or constructors/tuples over such. - * As much work as possible is done while constructing the result packed function, and - * that function may be reasonably efficiently applied multiple times without redoing - * unnecessary work. - * - * Primitives are lowered and compiled to packed functions for execution on \p device - * with properties given by \p target. All other Relay constructs are interpreted. - * - * The interpreter is intended to be a 'reference' implementation of the Relay semantics - * for testing and interactive use. It is not intended to be particularly efficient. - * - * \param mod A module containing definitions which can be referenced from - * \p expr. May be empty or undefined. - * \param expr An expression of function type to evaluate. May reference definitions from \p mod. - * \param device The device on which all primitives will be executed. - * \param target The compiler target flag for compiling primitives. - * \return A packed function that takes an array of Relay expressions and returns the - * result of applying \p expr to those arguments. - */ -TypedPackedFunc)> EvalFunction(IRModule mod, Expr expr, Device device, - Target target); - -/*! - * \brief Evaluates \p expr and returns its result. - * - * This function is intended to support TVM constant evaluation. - * - * \param expr An expression to evaluate. - * \param type_definitions Global type definitions which \p expr may references. - * \param import_set Already imported external modules. - * \param device The device on which all primitives will be executed. - * \param target The compiler target flag for compiling primitives. - * \param attrs Attributes for the expression to be evaluated with - * @return The object representing the result. - */ -ObjectRef Eval(Expr expr, Map type_definitions, - std::unordered_set import_set, Device device, Target target, - Map attrs = {}); - -} // namespace relay -} // namespace tvm - -#endif // TVM_RELAY_INTERPRETER_H_ diff --git a/include/tvm/relay/op.h b/include/tvm/relay/op.h deleted file mode 100644 index 12845158a22f..000000000000 --- a/include/tvm/relay/op.h +++ /dev/null @@ -1,41 +0,0 @@ -/* - * 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 tvm/relay/op.h - * \brief Primitive operators(builtin intrinsics). - */ -#ifndef TVM_RELAY_OP_H_ -#define TVM_RELAY_OP_H_ - -#include -#include -#include - -namespace tvm { -namespace relay { - -using Op = tvm::Op; -using OpNode = tvm::OpNode; - -#define RELAY_REGISTER_OP(OpName) TVM_REGISTER_OP(OpName) - -} // namespace relay -} // namespace tvm -#endif // TVM_RELAY_OP_H_ diff --git a/include/tvm/relay/op_attr_types.h b/include/tvm/relay/op_attr_types.h deleted file mode 100644 index 97a3d5e2a01f..000000000000 --- a/include/tvm/relay/op_attr_types.h +++ /dev/null @@ -1,234 +0,0 @@ -/* - * 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 tvm/relay/op_attr_types.h - * \brief The Expr and related elements in DataFlow construction. - */ -#ifndef TVM_RELAY_OP_ATTR_TYPES_H_ -#define TVM_RELAY_OP_ATTR_TYPES_H_ - -#include -#include -#include -#include -#include -#include -#include - -#include - -namespace tvm { -namespace relay { - -using tir::BijectiveLayoutNode; -using tir::Layout; -using tir::LayoutAxis; - -/*! \brief operator pattern used in graph fusion */ -enum OpPatternKind { - // Elementwise operation - kElemWise = 0, - // Broadcasting operator, can always map output axis to the input in order. - // for example :code:`out[i, ax1, j, ax2] = input[i, j]`. - // Note that the axis need to be in order so transpose is not a bcast operator. - kBroadcast = 1, - // Injective operator, can always injectively map output axis to a single input axis. - // All injective operator can still be safely fused to injective and reduction. - kInjective = 2, - // Communicative reduction operator. - kCommReduce = 3, - // Complex operation, can still fuse elemwise operations into its output. - // but cannot chain another complex op - kOutEWiseFusable = 4, - // The pattern for tuple nodes. Can fuse into subsequent injective ops, - // but treated specially - kTuple = 7, - // Opaque operation, cannot fuse anything. - kOpaque = 8 -}; - -/*! \brief the operator pattern */ -using TOpPattern = int; - -/*! - * \brief Whether operator is stateful or contain internal state. - * - * All the primitive ops we registered so far are pure. - * This attribute is left for potential future compatible reasons. - * We can always work around the stateful ops by adding an additional - * handle argument and return it. - */ -using TOpIsStateful = bool; - -/*! - * \brief Mark the operator as non-computational. - */ -using TNonComputational = bool; - -/*! - * \brief Mark the operator as reshape op of its first input - * and can be turned into a nop when the first input and output - * shares the same piece of memory. - */ -using TReshapeOp = bool; - -/*! - * \brief Mark the operator whether output shape is data dependent. - */ -using TShapeDataDependent = Array; - -/*! - * \brief Computation description interface. - * - * \note This function have a special convention - * for functions with tuple input/output. - * - * So far we restrict tuple support to the following case: - * - Function which takes a single tuple as input. - * - Function which outputs a single tuple. - * - * In both cases, the tuple is flattened as array. - * - * \param attrs The attribute of the primitive - * \param inputs The input tensors. - * \param out_type The output type information - & these are always placeholders. - * \return The output compute description of the operator. - */ -using FTVMCompute = runtime::TypedPackedFunc( - const Attrs& attrs, const Array& inputs, const Type& out_type)>; - -/*! - * \brief Build the computation schedule for - * op whose root is at current op. - * - * \param attrs The attribute of the node. - * \param outs The output tensors. - * \param target The build target. - * \return schedule The computation schedule. - */ -using FTVMSchedule = runtime::TypedPackedFunc& outs, const Target& target)>; - -/*! - * \brief Generate the strategy of operators. This function is a generic - * function and can be re-defined for different targets. - * - * The function signature of generic function is: - * OpStrategy(const Attrs& attrs, const Array& inputs, - * const Type& out_type, const Target& target) - */ -using FTVMStrategy = GenericFunc; - -/*! - * \brief Alternate the layout of operators or replace the - * operator with other expressions. This function will be invoked - * in AlterOpLayout pass. - * \param attrs The attribute of the original node. - * \param args The input symbols of the original node. - * \param tinfos An array of placeholders, use for getting the inferred shape - * and dtype of the inputs. - * \return new_expr The modified expression. - */ -using FTVMAlterOpLayout = - runtime::TypedPackedFunc& args, - const Array& tinfos, const Type& out_type)>; - -/*! - * \brief Convert the layout of operators or replace the - * operator with other expressions. This function will be invoked - * in ConvertLayout pass. - * \param attrs The attribute of the original node. - * \param inputs The input symbols of the original node. - * \param tinfos An array of placeholders, use for getting the inferred shape - * and dtype of the inputs. - * \param desired_layouts Specify an array of desired layouts for each input. - * For example a conv2d op: Array("NHWC", "OHWI"), this - * specifies the desired layout for data then kernel. - * \return new_expr The modified expression. - */ -using FTVMConvertOpLayout = runtime::TypedPackedFunc& args, const Array& tinfos, - const Array& desired_layouts)>; -/*! - * \brief Legalizes an expression with another expression. This function will be - * invoked in Legalize pass. It is a target-dependent pass. - * \param attrs The attribute of the original node. - * \param args The input symbols of the original node. - * \param arg_types An array of placeholders, use for getting the inferred shape - * and dtype of the inputs. - * \return new_expr The modified expression. - */ -using FTVMLegalize = runtime::TypedPackedFunc& args, - const Array& arg_types)>; - -/*! - * \brief Annotates an expression to indicate if an op should be compiled using - * the given compiler/target. - * \param expr The original expr. - * \return true if this op should be registered to invoke a specific compiler - * for codegen, otherwise, false. - */ -using FTVMAnnotateTarget = runtime::TypedPackedFunc; - -/*! - * \brief Forward rewriting rule for a specific op. - * - * \param ref_call The reference old call type to be rewritten. - * We can make use of the op and type information. - * \param new_args The new arguments (some of them could be TempExpr). - * \param ctx Optional context information about ref_call. - * \return The rewriten result call, can also return nullptr, - * which indicate the rewriter should use the default fallback - * rule that realizes all its input and compose the call. - * - * \note When we register the function, we can register - * a different signature with ctx to be a specific node type. - */ -using FForwardRewrite = runtime::TypedPackedFunc& new_args, const ObjectRef& ctx)>; - -/*! - * \brief Gradient for a specific op. - * - * \param orig_call the original Expr. - * \param output_grad the gradient of the Expr. - * \return the gradient for each parameters. - */ -using FPrimalGradient = - runtime::TypedPackedFunc(const Expr& orig_call, const Expr& output_grad)>; - -/*! - * \brief The codegeneration strategy for dynamic dimensions. - */ -enum AnyCodegenStrategy { - /*! \brief The default strategy of using completely variable dimensions. */ - kVariableDimensions -}; - -/*! \brief A runtime representation of shape. */ -using Shape = Array; - -using FShapeFunc = runtime::TypedPackedFunc( - const Attrs& attrs, const Array& inputs, const Array& out_ndims)>; - -} // namespace relay -} // namespace tvm -#endif // TVM_RELAY_OP_ATTR_TYPES_H_ diff --git a/include/tvm/relay/op_strategy.h b/include/tvm/relay/op_strategy.h deleted file mode 100644 index c5785369f8d5..000000000000 --- a/include/tvm/relay/op_strategy.h +++ /dev/null @@ -1,161 +0,0 @@ -/* - * 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 tvm/relay/op_strategy.h - * \brief The Relay operator Strategy and related data structure. - */ - -#ifndef TVM_RELAY_OP_STRATEGY_H_ -#define TVM_RELAY_OP_STRATEGY_H_ - -#include -#include -#include -#include -#include - -#include - -namespace tvm { -namespace relay { - -/*! - * \brief Operator implementation that includes compute and schedule function. - */ -class OpImplementationNode : public Object { - public: - /*! \brief Compute function */ - FTVMCompute fcompute; - /*! \brief Schedule function */ - FTVMSchedule fschedule; - /*! \brief Name of the implementation */ - String name; - /*! \brief Priority level */ - int plevel; - - void VisitAttrs(tvm::AttrVisitor* v) { - v->Visit("name", &name); - v->Visit("plevel", &plevel); - } - - static constexpr const char* _type_key = "relay.OpImplementation"; - TVM_DECLARE_FINAL_OBJECT_INFO(OpImplementationNode, Object); -}; - -/*! - * \brief Operator implementation class. - */ -class OpImplementation : public ObjectRef { - public: - /*! - * \brief Invoke the operator compute function. - * \param attrs The attribute of the primitive - * \param inputs The input tensors. - * \param out_type The output type information. - * \return The output compute description of the operator. - */ - TVM_DLL Array Compute(const Attrs& attrs, const Array& inputs, - const Type& out_type); - /*! - * \brief Build the computation schedule. - * \param attrs The attribute of the node. - * \param outs The output tensors. - * \param target The build target. - * \return The computation schedule. - */ - TVM_DLL te::Schedule Schedule(const Attrs& attrs, const Array& outs, - const Target& target); - - TVM_DEFINE_OBJECT_REF_METHODS(OpImplementation, ObjectRef, OpImplementationNode); -}; - -/*! - * \brief Specialized implementations for operators under certain conditions. - */ -class OpSpecializationNode : public Object { - public: - /*! \brief List of implementations. */ - Array implementations; - /*! \brief Condition to enable the specialization. - * Could be undefined to represent generic case. */ - te::SpecializedCondition condition; - - void VisitAttrs(tvm::AttrVisitor* v) { - v->Visit("condition", &condition); - v->Visit("implementations", &implementations); - } - - static constexpr const char* _type_key = "relay.OpSpecialization"; - TVM_DECLARE_FINAL_OBJECT_INFO(OpSpecializationNode, ExprNode); -}; - -/*! - * \brief Operator specialization class. - */ -class OpSpecialization : public ObjectRef { - public: - /*! - * \brief Add an implementation. - * \param fcompute Compute function - * \param fschedule Schedule function - * \param name Name of the implementation - * \param plevel Priority level of the implementation - */ - TVM_DLL void AddImplementation(FTVMCompute fcompute, FTVMSchedule fschedule, String name, - int plevel); - - TVM_DEFINE_MUTABLE_OBJECT_REF_METHODS(OpSpecialization, ObjectRef, OpSpecializationNode); -}; - -/*! - * \brief Operator strategy to choose implementation. - */ -class OpStrategyNode : public Object { - public: - /*! \brief List of operator specializations. */ - Array specializations; - - void VisitAttrs(tvm::AttrVisitor* v) { v->Visit("specializations", &specializations); } - - static constexpr const char* _type_key = "relay.OpStrategy"; - TVM_DECLARE_FINAL_OBJECT_INFO(OpStrategyNode, ExprNode); -}; - -/*! - * \brief Operator strategy class. - */ -class OpStrategy : public ObjectRef { - public: - /*! - * \brief Add an implementation. - * \param fcompute Compute function - * \param fschedule Schedule function - * \param name Name of the implementation - * \param plevel Priority level of the implementation - */ - TVM_DLL void AddImplementation(FTVMCompute fcompute, FTVMSchedule fschedule, String name, - int plevel); - - TVM_DEFINE_MUTABLE_OBJECT_REF_METHODS(OpStrategy, ObjectRef, OpStrategyNode); -}; - -} // namespace relay -} // namespace tvm -#endif // TVM_RELAY_OP_STRATEGY_H_ diff --git a/include/tvm/relay/parser.h b/include/tvm/relay/parser.h deleted file mode 100644 index 6e33e7873f60..000000000000 --- a/include/tvm/relay/parser.h +++ /dev/null @@ -1,49 +0,0 @@ -/* - * 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. - */ -#ifndef TVM_RELAY_PARSER_H_ -#define TVM_RELAY_PARSER_H_ - -#include -#include -#include -#include - -#include -#include - -namespace tvm { -namespace relay { - -using MetaTable = Map>; - -IRModule ParseModule(const std::string& file_name, const std::string& file_content, - const Optional& init_module = Optional(), - const MetaTable& init_meta_table = MetaTable()); - -/*! - * \brief This pass pretty-prints mod then parses it back so as to establish spans and sources - * for all Relay sub-expressions. This improves error and debugging diagnostics downstream for - * modules constructed programaticaly rather than textually. - */ -tvm::transform::Pass AnnotateSpans(); - -} // namespace relay -} // namespace tvm - -#endif // TVM_RELAY_PARSER_H_ diff --git a/include/tvm/relay/pattern_functor.h b/include/tvm/relay/pattern_functor.h deleted file mode 100644 index 9d2b6689b2c2..000000000000 --- a/include/tvm/relay/pattern_functor.h +++ /dev/null @@ -1,166 +0,0 @@ -/* - * 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 tvm/relay/pattern_functor.h - * \brief A more powerful visitor on ADT patterns that enables defining - * arbitrary function signatures with type-based dispatch on first argument. - */ -#ifndef TVM_RELAY_PATTERN_FUNCTOR_H_ -#define TVM_RELAY_PATTERN_FUNCTOR_H_ - -#include -#include - -#include -#include -#include - -#include "./adt.h" -#include "./expr.h" -#include "./op.h" - -namespace tvm { -namespace relay { - -/*! - * \brief A dynamical functor on ADT patterns that dispatches on its first argument. - * You can use this as a more powerful visitor, since it allows you to - * define the types of further arguments to VisitPattern. - * - * \sa tvm/ir_functor.h - * - * \tparam FType function signiture - * This type is only defined for FType with function signature R(const Pattern&, - * Args...) - */ -template -class PatternFunctor; - -// functions to be overriden. -#define PATTERN_FUNCTOR_DEFAULT \ - { return VisitPatternDefault_(op, std::forward(args)...); } - -#define RELAY_PATTERN_FUNCTOR_DISPATCH(OP) \ - vtable.template set_dispatch([](const ObjectRef& n, TSelf* self, Args... args) { \ - return self->VisitPattern_(static_cast(n.get()), std::forward(args)...); \ - }); - -template -class PatternFunctor { - private: - using TSelf = PatternFunctor; - using FType = tvm::NodeFunctor; - - public: - /*! \brief the result type of this functor */ - using result_type = R; - /*! \brief virtual destructor */ - virtual ~PatternFunctor() {} - /*! - * \brief Same as call. - * \param n The expression node. - * \param args Additional arguments. - * \return The result of the call - */ - R operator()(const Pattern& n, Args... args) { - return VisitPattern(n, std::forward(args)...); - } - /*! - * \brief The functor call. - * \param n The expression node. - * \param args Additional arguments. - * \return The result of the call - */ - virtual R VisitPattern(const Pattern& n, Args... args) { - ICHECK(n.defined()); - static FType vtable = InitVTable(); - return vtable(n, this, std::forward(args)...); - } - // Functions that can be overriden by subclass - virtual R VisitPattern_(const PatternWildcardNode* op, Args... args) PATTERN_FUNCTOR_DEFAULT; - virtual R VisitPattern_(const PatternVarNode* op, Args... args) PATTERN_FUNCTOR_DEFAULT; - virtual R VisitPattern_(const PatternConstructorNode* op, Args... args) PATTERN_FUNCTOR_DEFAULT; - virtual R VisitPattern_(const PatternTupleNode* op, Args... args) PATTERN_FUNCTOR_DEFAULT; - virtual R VisitPatternDefault_(const Object* op, Args...) { - LOG(FATAL) << "Do not have a default for " << op->GetTypeKey(); - throw; - } - - private: - // initialize the vtable. - static FType InitVTable() { - FType vtable; - // Set dispatch - RELAY_PATTERN_FUNCTOR_DISPATCH(PatternWildcardNode); - RELAY_PATTERN_FUNCTOR_DISPATCH(PatternVarNode); - RELAY_PATTERN_FUNCTOR_DISPATCH(PatternConstructorNode); - RELAY_PATTERN_FUNCTOR_DISPATCH(PatternTupleNode); - return vtable; - } -}; - -/*! \brief A simple visitor wrapper around PatternFunctor. - * - * Exposes two visitors with default traversal strategies, one - * which doesn't compute a result but can mutate internal state, - * and another which functionally builds a new pattern. - */ -class PatternVisitor : public ::tvm::relay::PatternFunctor { - public: - void VisitPattern_(const PatternWildcardNode* op) override; - void VisitPattern_(const PatternVarNode* op) override; - void VisitPattern_(const PatternConstructorNode* op) override; - void VisitPattern_(const PatternTupleNode* op) override; - virtual void VisitType(const Type& t); - virtual void VisitVar(const Var& v); - virtual void VisitConstructor(const Constructor& c); -}; - -/*! \brief A wrapper around ExprFunctor which functionally updates the AST. - * - * ExprMutator uses memoization and self return in order to amortize - * the cost of using functional updates. - */ -class PatternMutator : public ::tvm::relay::PatternFunctor { - public: - Pattern Mutate(const Pattern& pat); - Pattern VisitPattern_(const PatternWildcardNode* op) override; - Pattern VisitPattern_(const PatternVarNode* op) override; - Pattern VisitPattern_(const PatternConstructorNode* op) override; - Pattern VisitPattern_(const PatternTupleNode* op) override; - /*! \brief Used to visit the types inside of patterns. - * - * Can be overloaded to transform the types in arbitrary - * ways, one way would be to define a sub-class of type - * visitor for types which transform them appropriately. - */ - virtual Type VisitType(const Type& t); - /*! \brief Used to visit the vars inside of patterns. */ - virtual Var VisitVar(const Var& v); - /*! \brief Used to visit the vars inside of patterns. */ - virtual Constructor VisitConstructor(const Constructor& c); - - private: - std::unordered_map var_map_; -}; - -} // namespace relay -} // namespace tvm -#endif // TVM_RELAY_PATTERN_FUNCTOR_H_ diff --git a/include/tvm/relay/qnn/attrs.h b/include/tvm/relay/qnn/attrs.h deleted file mode 100644 index 85e008528625..000000000000 --- a/include/tvm/relay/qnn/attrs.h +++ /dev/null @@ -1,133 +0,0 @@ -/* - * 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 tvm/relay/qnn/attrs.h - * \brief Auxiliary attributes for qnn operators. - */ -#ifndef TVM_RELAY_QNN_ATTRS_H_ -#define TVM_RELAY_QNN_ATTRS_H_ - -#include - -#include - -namespace tvm { -namespace relay { -namespace qnn { - -/*! \brief Attribute for requantize operator */ -struct RequantizeAttrs : public tvm::AttrsNode { - int axis; - std::string rounding; - std::string compute_dtype; - DataType out_dtype; - - TVM_DECLARE_ATTRS(RequantizeAttrs, "relay.attrs.RequantizeAttrs") { - TVM_ATTR_FIELD(axis) - .describe( - "The output channel axis for channel wise quantization. Default value is -1," - "which corresponds to the last axis.") - .set_default(-1); - TVM_ATTR_FIELD(rounding).set_default("None").describe( - "Defines the rounding direction when the value is midway between" - "two representable values. There are two supported modes - UPWARD" - "or TONEAREST. Both modes behave exactly same except at the" - "midpoints between the two representable values. At the midpoint," - "UPWARD rounds towards positive infinity (for example -1.5 will be" - "rounded to -1). TONEAREST is the standard rounding where the" - "value is rounded away from zero at midpoints (for example, -1.5" - "rounds to -2). More context can be found at following gblic manual" - "https://www.gnu.org/software/libc/manual/html_node/Rounding.html."); - TVM_ATTR_FIELD(compute_dtype) - .set_default("None") - .describe( - "Specifies the data type used during requantize. Supported " - "options: \"int64\", \"float32\", \"float64\""); - TVM_ATTR_FIELD(out_dtype) - .set_default(NullValue()) - .describe("Output data type, set to explicit type under mixed precision setting"); - } -}; - -/*! \brief Attribute for quantize operator */ -struct QuantizeAttrs : public tvm::AttrsNode { - DataType out_dtype; - int axis; - - TVM_DECLARE_ATTRS(QuantizeAttrs, "relay.attrs.QuantizeAttrs") { - TVM_ATTR_FIELD(out_dtype).describe("Output data type, can be one of [int8 or uint8]."); - TVM_ATTR_FIELD(axis) - .describe( - "The output channel axis for channel wise quantization. Default value is -1," - "which corresponds to the last axis.") - .set_default(-1); - } -}; - -struct SimulatedQuantizeAttrs : public tvm::AttrsNode { - int axis; - - TVM_DECLARE_ATTRS(SimulatedQuantizeAttrs, "relay.attrs.SimulatedQuantizeAttrs") { - TVM_ATTR_FIELD(axis) - .describe( - "The output channel axis for channel wise quantization. Default value is -1," - "which corresponds to the last axis.") - .set_default(-1); - } -}; - -/*! \brief Attribute for dequantize operator */ -struct DequantizeAttrs : public tvm::AttrsNode { - DataType out_dtype; - int axis; - - TVM_DECLARE_ATTRS(DequantizeAttrs, "relay.attrs.DequantizeAttrs") { - TVM_ATTR_FIELD(out_dtype).describe("Output data type, can be one of [float16, float32]."); - TVM_ATTR_FIELD(axis) - .describe( - "The channel axis for channel wise dequantization. Default value is -1," - "which corresponds to the last axis.") - .set_default(-1); - } -}; - -/*! \brief Attribute for broadcast operator */ -struct BroadcastAttrs : public tvm::AttrsNode { - int lhs_axis; - int rhs_axis; - - TVM_DECLARE_ATTRS(BroadcastAttrs, "relay.attrs.BroadcastAttrs") { - TVM_ATTR_FIELD(lhs_axis) - .describe( - "The channel axis for channel wise broadcast. Default value is -1," - "which corresponds to the last axis.") - .set_default(-1); - TVM_ATTR_FIELD(rhs_axis) - .describe( - "The channel axis for channel wise broadcast. Default value is -1," - "which corresponds to the last axis.") - .set_default(-1); - } -}; - -} // namespace qnn -} // namespace relay -} // namespace tvm -#endif // TVM_RELAY_QNN_ATTRS_H_ diff --git a/include/tvm/relay/qnn/transform.h b/include/tvm/relay/qnn/transform.h deleted file mode 100644 index d1f07c924d6b..000000000000 --- a/include/tvm/relay/qnn/transform.h +++ /dev/null @@ -1,60 +0,0 @@ -/* - * 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 tvm/relay/qnn/transform.h - * - * This file implements a pass manager for QNN ops using Relay Pass manager. - */ -#ifndef TVM_RELAY_QNN_TRANSFORM_H_ -#define TVM_RELAY_QNN_TRANSFORM_H_ - -#include -#include - -namespace tvm { -namespace relay { - -using relay::transform::Pass; - -namespace qnn { -namespace transform { - -/*! - * \brief Legalizes a QNN expr. Contains specifically two types of Legalizations. First, - * converts/Lowers an expression containing QNN ops to an expression containing only core Relay ops. - * Each QNN op is lowered to a sequence of exisiting Relay ops. This is a target-independent pass. - * One can register the lowering/transformation function for this op using FTVMQnnCanonicalize - * attr_name for FTVMLegalize op attribute. Second, as opposed to Relay Legalize, this one legalizes - * only QNN ops. One can register a transformation/legalization function for an op by using the - * FTVMQnnLegalize attr_name for FTVMLegalize op attribute. The isolation of QNN and Relay Legalize - * gives us separation of concerns, leading to a better software practice. The legalization can be - * configured to happen per target. - * - * \return The pass. - */ -TVM_DLL Pass Legalize(); - -} // namespace transform - -} // namespace qnn -} // namespace relay -} // namespace tvm - -#endif // TVM_RELAY_QNN_TRANSFORM_H_ diff --git a/include/tvm/relay/runtime.h b/include/tvm/relay/runtime.h deleted file mode 100644 index 10e124bc339b..000000000000 --- a/include/tvm/relay/runtime.h +++ /dev/null @@ -1,276 +0,0 @@ -/* - * 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 tvm/relay/runtime.h - * \brief Object representation of Runtime configuration and registry - */ -#ifndef TVM_RELAY_RUNTIME_H_ -#define TVM_RELAY_RUNTIME_H_ - -#include -#include -#include -#include -#include -#include -#include - -#include -#include -#include -#include - -namespace tvm { - -template -class AttrRegistry; - -namespace relay { - -/*! \brief Value used with Runtime::name to indicate the C++ runtime. */ -static constexpr const char* kTvmRuntimeCpp = "cpp"; - -/*! \brief Value used with Runtime::name to indicate the C runtime. */ -static constexpr const char* kTvmRuntimeCrt = "crt"; - -/*! - * \brief Runtime information. - * - * This data structure stores the meta-data - * about Runtimes which can be used to pass around information. - * - * \sa Runtime - */ -class RuntimeNode : public Object { - public: - /*! \brief name of the Runtime */ - String name; - /* \brief Additional attributes storing meta-data about the Runtime. */ - DictAttrs attrs; - - /*! - * \brief Get an attribute. - * - * \param attr_key The attribute key. - * \param default_value The default value if the key does not exist, defaults to nullptr. - * - * \return The result - * - * \tparam TObjectRef the expected object type. - * \throw Error if the key exists but the value does not match TObjectRef - * - * \code - * - * void GetAttrExample(const Runtime& runtime) { - * auto value = runtime->GetAttr("AttrKey", 0); - * } - * - * \endcode - */ - template - Optional GetAttr( - const std::string& attr_key, - Optional default_value = Optional(nullptr)) const { - return attrs.GetAttr(attr_key, default_value); - } - // variant that uses TObjectRef to enable implicit conversion to default value. - template - Optional GetAttr(const std::string& attr_key, TObjectRef default_value) const { - return GetAttr(attr_key, Optional(default_value)); - } - - void VisitAttrs(AttrVisitor* v) { - v->Visit("name", &name); - v->Visit("attrs", &attrs); - } - - bool SEqualReduce(const RuntimeNode* other, SEqualReducer equal) const { - return name == other->name && equal.DefEqual(attrs, other->attrs); - } - - void SHashReduce(SHashReducer hash_reduce) const { - hash_reduce(name); - hash_reduce(attrs); - } - - static constexpr const char* _type_key = "Runtime"; - static constexpr const bool _type_has_method_sequal_reduce = true; - static constexpr const bool _type_has_method_shash_reduce = true; - TVM_DECLARE_FINAL_OBJECT_INFO(RuntimeNode, Object); -}; - -/*! - * \brief Managed reference class to RuntimeNode. - * \sa RuntimeNode - */ -class Runtime : public ObjectRef { - public: - Runtime() = default; - - /*! - * \brief Create a new Runtime object using the registry - * \throws Error if name is not registered - * \param name The name of the Runtime. - * \param attrs Attributes for the Runtime. - * \return the new Runtime object. - */ - TVM_DLL static Runtime Create(String name, Map attrs = {}); - - /*! - * \brief List all registered Runtimes - * \return the list of Runtimes - */ - TVM_DLL static Array ListRuntimes(); - - /*! - * \brief List all options for a specific Runtime - * \param name The name of the Runtime - * \return Map of option name to type - */ - TVM_DLL static Map ListRuntimeOptions(const String& name); - - /*! \brief specify container node */ - TVM_DEFINE_NOTNULLABLE_OBJECT_REF_METHODS(Runtime, ObjectRef, RuntimeNode); - - private: - /*! - * \brief Private Constructor - * \param name The Runtime name - * \param attrs Attributes to apply to this Runtime node - */ - TVM_DLL Runtime(String name, DictAttrs attrs) { - auto n = make_object(); - n->name = std::move(name); - n->attrs = std::move(attrs); - data_ = std::move(n); - } -}; - -/*! - * \brief Helper structure to register Runtimes - * \sa TVM_REGISTER_Runtime - */ -class RuntimeRegEntry { - public: - /*! - * \brief Register a valid configuration option and its ValueType for validation - * \param key The configuration key - * \tparam ValueType The value type to be registered - */ - template - inline RuntimeRegEntry& add_attr_option(const String& key); - - /*! - * \brief Register a valid configuration option and its ValueType for validation - * \param key The configuration key - * \param default_value The default value of the key - * \tparam ValueType The value type to be registered - */ - template - inline RuntimeRegEntry& add_attr_option(const String& key, ObjectRef default_value); - - /*! - * \brief Register or get a new entry. - * \param name The name of the operator. - * \return the corresponding entry. - */ - TVM_DLL static RuntimeRegEntry& RegisterOrGet(const String& name); - - private: - /*! \brief Internal storage of value types */ - struct ValueTypeInfo { - std::string type_key; - uint32_t type_index; - }; - std::unordered_map key2vtype_; - /*! \brief A hash table that stores the default value of each attr */ - std::unordered_map key2default_; - - /*! \brief Index used for internal lookup of attribute registry */ - uint32_t index_; - - // the name - std::string name; - - /*! \brief Return the index stored in attr registry */ - uint32_t AttrRegistryIndex() const { return index_; } - /*! \brief Return the name stored in attr registry */ - String AttrRegistryName() const { return name; } - - /*! \brief private constructor */ - explicit RuntimeRegEntry(uint32_t reg_index) : index_(reg_index) {} - - // friend class - template - friend class AttrRegistryMapContainerMap; - template - friend class tvm::AttrRegistry; - friend class Runtime; -}; - -template -inline RuntimeRegEntry& RuntimeRegEntry::add_attr_option(const String& key) { - ICHECK(!key2vtype_.count(key)) << "AttributeError: add_attr_option failed because '" << key - << "' has been set once"; - - using ValueNodeType = typename ValueType::ContainerType; - // NOTE: we could further update the function later. - uint32_t value_type_index = ValueNodeType::_GetOrAllocRuntimeTypeIndex(); - - ValueTypeInfo info; - info.type_index = value_type_index; - info.type_key = runtime::Object::TypeIndex2Key(value_type_index); - key2vtype_[key] = info; - return *this; -} - -template -inline RuntimeRegEntry& RuntimeRegEntry::add_attr_option(const String& key, - ObjectRef default_value) { - add_attr_option(key); - key2default_[key] = default_value; - return *this; -} - -// internal macros to make Runtime entries -#define TVM_RUNTIME_REGISTER_VAR_DEF \ - static DMLC_ATTRIBUTE_UNUSED ::tvm::relay::RuntimeRegEntry& __make_##Runtime - -/*! - * \def TVM_REGISTER_RUNTIME - * \brief Register a new Runtime, or set attribute of the corresponding Runtime. - * - * \param RuntimeName The name of registry - * - * \code - * - * TVM_REGISTER_RUNTIME("c") - * .add_attr_option("my_option"); - * .add_attr_option("my_option_default", String("default")); - * - * \endcode - */ -#define TVM_REGISTER_RUNTIME(RuntimeName) \ - TVM_STR_CONCAT(TVM_RUNTIME_REGISTER_VAR_DEF, __COUNTER__) = \ - ::tvm::relay::RuntimeRegEntry::RegisterOrGet(RuntimeName) -} // namespace relay -} // namespace tvm - -#endif // TVM_RELAY_RUNTIME_H_ diff --git a/include/tvm/relay/transform.h b/include/tvm/relay/transform.h deleted file mode 100644 index a767c36d714f..000000000000 --- a/include/tvm/relay/transform.h +++ /dev/null @@ -1,754 +0,0 @@ -/* - * 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 tvm/relay/transform.h - * \brief Relay specific transformation passes. - */ -#ifndef TVM_RELAY_TRANSFORM_H_ -#define TVM_RELAY_TRANSFORM_H_ - -#include -#include -#include -#include -#include -#include -#include -#include -#include - -#include - -namespace tvm { -namespace relay { -namespace transform { - -using Pass = tvm::transform::Pass; -using PassNode = tvm::transform::PassNode; -using PassInfo = tvm::transform::PassInfo; -using PassInfoNode = tvm::transform::PassInfoNode; -using PassContext = tvm::transform::PassContext; -using PassContextNode = tvm::transform::PassContextNode; -using Sequential = tvm::transform::Sequential; -using FTVMRelayToTIR = tvm::transform::Pass; -/*! - * \brief TIRToRuntime conversion specific to a TargetKind - * - * This function is responsible for scanning an IRModule for appropriate Target-specific functions - and generating a Runtime module representing the compiled output - * - * \param ir_module Unified IRModule - * \param target Target to filter on or retrieve arguments from - * \return Runtime Module containing compiled functions - */ -using FTVMTIRToRuntime = tvm::runtime::TypedPackedFunc; - -/*! - * \brief RelayToTIR tvm::transform::Pass specific to a TargetKind - * - * Called before the default lowering passes. - * - * \param mod The module that an optimization pass runs on. - * \param pass_ctx The pass context that can provide information for the optimization. - * - * \return The transformed module. - */ -using FTVMRelayToTIR = tvm::transform::Pass; - -/* - * \brief Create a function pass. - * - * \param pass_func The packed function that contains the optimization. - * \param opt_level The optimization level of the function pass. - * \param name The name of the function pass. - * \param required The list of the passes that the function pass is dependent on. - * - * \return The created function pass. - */ -TVM_DLL Pass CreateFunctionPass( - const runtime::TypedPackedFunc& pass_func, - int opt_level, String name, tvm::Array required, bool traceable = false); - -/*! \brief Remove let-bound expressions which do not effect the program result. - * - * This pass will remove let bindings which are not referenced. If inline_once is True, - * let bindings which are only referenced once will also be inlined. - * - * For example, this pass should turn `let a = 1; 2` into `2`, - * as the value of the expression does not depend on a. - * - * As another example, `let a = 1; a` will be optimized into 1 if inline_once is True. - * - * If ignore_purity is False, possibly side-effecting expressions (such as memory allocation, - * random number generation, reading/writing references, or calls to primitive or external - * functions) are never elided or inlined. This is sound, but ignore_purity can be set to True - * to suppress this check. - * - * The analysis is fairly conservative, for example it assumes all local functions - * may be called more than once, any functions passed as arguments have side effects, - * and so on. - * - * \param inline_once whether or not to inline bindings used exactly once. - * \param ignore_purity whether to ignore whether expressions have side-effects - * - * \return the pass. - */ -TVM_DLL Pass DeadCodeElimination(bool inline_once = false, bool ignore_purity = false); - -/*! - * \brief Convert all expressions of TensorType into GradCell, - * an algebraic data type defined in gradient.rly. - * - * This will delay or decrease memory usage. All calls to - * ones, ones_like, zeros, zeros_like will not immediately instantiate a tensor in memory, - * rather only instantiate if needed. It also defines + and * operation - * between GradCell types which can increase performance when using - * zero-filled or one-filled tensors, which is the case in reverse mode ad. - * - * \return the pass - */ -TVM_DLL Pass LazyGradientInit(); - -/*! - * \brief Fold constant expressions. - * - * Because of backward compatibility reason it skips QNN primitives from folding by default. - * There are some transformation passes like FakeQuantizationToInteger, which requires to keep QNN - * primitives for constant subgraphs. Uncontrolled constant folding of QNN primitives may break - * applicability of FakeQuantizationToInteger. We suggest to use FoldConstant pass with none - * default fold_qnn=True value only when all other QNN sensitive passes were already applied. - * - * \param fold_qnn Whether to fold constants for QNN operations. - * - * \return The pass. - */ -TVM_DLL Pass FoldConstant(bool fold_qnn = false); - -/*! - * \brief Split function with huge number of arguments to smaller pieces. - * - * \param max_function_args Maximum number of function arguments. If it equals 0 then SplitArgs - * shouldn't split the function. - * - * \return The pass. - */ -TVM_DLL Pass SplitArgs(uint64_t max_function_args); - -/*! - * \brief Fuse operations into expr into separate functions. - * - * \param fuse_opt_level Optimization level. If it is -1 it will be inferred from pass context. - * - * \return The pass. - */ -TVM_DLL Pass FuseOps(int fuse_opt_level = -1); - -/*! - * \brief The inverse operation of FuseOps. It transforms a fused program returned by - * FuseOps into the program before FuseOps. (i.e. x == DefuseOps(FuseOps(x))) - * - * \return The pass. - */ -TVM_DLL Pass DefuseOps(); - -/*! - * \brief Rewrite the annotated program. - * - * \param fallback_device The fallback device which is the default device for - * operators without annotation. - * - * \return The pass. - */ -TVM_DLL Pass RewriteAnnotatedOps(int fallback_device); - -/*! - * \brief Turn an expression to Basic Block Normal Form. - * - * We define a block as a group of expressions implied by the scope structure. - * - * Each graph node can only belong to a single block. - * - * For any value that is being used in multiple blocks, it has to be referred - * by a Var which is defined in a block, whose scope is the least common ancestor - * of blocks this value is used. - * - * \return The pass. - */ -TVM_DLL Pass ToBasicBlockNormalForm(); - -/*! - * \brief turn a dataflow graph into Administrative Normal Form, or A-Normal Form (ANF). - * - * It will turn an expression that is in a graph form (with sharing implicit), - * to an expression with explicit sharing (A-Normal Form). - * - * The scope of the root expression is the global scope. - * - * The scope of any non root expression is the least common ancestor of all it's scope. - * - * Values are ordered by post-DFS order in each scope. - * - * \return The pass. - */ -TVM_DLL Pass ToANormalForm(); - -/*! - * \brief ToANormalForm but on incomplete graph. - * - * \param expr the graph. - * - * \return The transformed program. - */ -TVM_DLL Expr ToANormalForm(const Expr& expr); - -/*! - * \brief Turn an expression into continuation passing style(CPS). - * - * CPS mean that every function will, instead of returning the result directly, - * be passed down an extra function (called the continuation) as argument, - * and pass the result to the continuation instead. - * - * Thus, every function call has to be passed an extra argument - * that represent the rest of the computation (Hence the name of continuation). - * - * Similarly, all other compute will be wrapped and call the continuation as well. - * - * \return the pass. - */ -TVM_DLL Pass ToCPS(); - -/*! - * \brief Remove let binding and directly share via pointer instead. - * - * It will remove all let binding, - * and turn all of the variable bound by let into direct pointer reference. - * - * \return the expression in graph normal form. - */ -TVM_DLL Pass ToGraphNormalForm(); - -/*! - * \brief Aggressive constant propagation/constant folding/inlining. - * - * It will do as much computation in compile time as possible. - * It has two benefit: remove runtime overhead, and allow more optimization (typically fusion). - * As a side effect, code size will explode. - * - * \return the optimized expression. - */ -TVM_DLL Pass PartialEval(); - -/*! - * \brief Simplify certain operators during inference. For example, the result - * of a batch norm which is indexed at tuple index 0 will be unpacked into a - * number of simplified operators. - * - * \return The Pass. - */ -TVM_DLL Pass SimplifyInference(); - -/*! - * \brief Replaces non linear activation functions with their fast but approximate counterparts. - * - * \return The Pass. - */ -TVM_DLL Pass FastMath(); - -/*! - * \brief Find Dynamic ops and make them static - * - * Searches the graph for dynamic ops. If the dynamic inputs to those ops are constants, it replaces - * them with static ops and re-performs type inference and constant folding. The pass repeats - * itself until the graph stops changing or we run too many iterations. - * - * \return The pass. - */ -TVM_DLL Pass DynamicToStatic(); - -/*! - * \brief Infer the type of an expression. - * - * The result of type checking is a new expression with unambiguous - * type information filled in, as well as it's checked type field - * populated with the result type. - * - * \return The pass. - */ -TVM_DLL Pass InferType(); - -/*! - * \brief Infer the type of an expression, reusing existing type information. - * - * The result of type checking is a new expression with unambiguous - * type information filled in for the given node only. The local - * version can use existing type information populated throughout - * the expression and assumes this information is correct. The local - * version also avoids examining large amounts of the graph assuming - * type information is filled in properly which makes it much faster if we - * iteratively call type inference. - * - * \return The type of the expression. - */ -TVM_DLL Type InferTypeLocal(const Expr& expr); - -/*! - * \brief Search and eliminate common subexpression. For example, if there are - * two expressions evaluated to an identical value, a single variable is created - * and these two expressions are replaced by this variable. - * - * \param fskip The callback argument that allows to skip certain expressions. - * - * \return The pass. - */ -TVM_DLL Pass EliminateCommonSubexpr(runtime::PackedFunc fskip = nullptr); - -/*! - * \brief Combine parallel 2d convolutions into a single convolution if the - * number of branches of this conv2d operator is not less than - * `min_num_branch`. - * - * \param min_num_branches The minimun number of branches. - * - * \return The pass. - */ -TVM_DLL Pass CombineParallelConv2D(uint64_t min_num_branches = 3); - -/*! - * \brief Combine parallel dense ops into a single batch_matmul if the - * number of branches of this dense operator is not less than - * `min_num_branch`. - * - * \param min_num_branches The minimun number of branches. - * \param to_batch_matmul Whether to combine parallel dense ops to batch matmul. - * If set false, combine dense ops to single dense op. - * - * \return The pass. - */ -TVM_DLL Pass CombineParallelDense(uint64_t min_num_branches = 3, bool to_batch_matmul = true); - -/*! - * \brief Combine parallel batch_matmul ops into a single batch_matmul - * if the number of branches of this dense operator is not less than - * `min_num_branch`. - * - * \param min_num_branches The minimun number of branches. - * - * \return The pass. - */ -TVM_DLL Pass CombineParallelBatchMatmul(uint64_t min_num_branches = 3); - -/*! - * \brief Backward fold axis scaling into weights of conv/dense operators. - * - * \return The pass. - */ -TVM_DLL Pass BackwardFoldScaleAxis(); - -/*! - * \brief Forward fold axis scaling into weights of conv/dense operators. - * - * \return The pass. - */ -TVM_DLL Pass ForwardFoldScaleAxis(); - -/*! - * \brief A sequential pass that executes ForwardFoldScaleAxis and - * BackwardFoldScaleAxis passes. - * - * \return The pass. - */ -TVM_DLL Pass FoldScaleAxis(); - -/*! - * \brief Canonicalize some operators to the simplified operators. For example, - * bias_add can be canonicalized to expand_dims and broadcast_add. - * - * \return The pass. - */ -TVM_DLL Pass CanonicalizeOps(); - -/*! - * \brief Alternate the layouts of operators or replace primitive operators - * with other expressions. - * - * \return The pass. - */ -TVM_DLL Pass AlterOpLayout(); - -/*! - * \brief Do layout rewrite according to the tile structure created by auto-scheduler. - * \return The pass - */ -TVM_DLL Pass AutoSchedulerLayoutRewrite(); - -/*! - * \brief Do layout rewrite according to the tile structure created by meta-schedule. - * \return The pass - */ -TVM_DLL Pass MetaScheduleLayoutRewrite(); - -/*! - * \brief Given a dest layout, this pass transforms the expr such that most of the ops input data - * layout is changed to the dest layout. In ideal situation, there are only 2 layout transforms, one - * at the start and one at the end. - * - * This pass is not a part of relay.build and is expected to be called between framework-relay - * parser and relay.build call. This is very helpful for hardware backends that support/prefer only - * type of data layout. - * - * RFC - https://discuss.tvm.ai/t/layout-conversion-pass/4009 - * - * This pass uses most of the AlterOpLayout and InferCorrectLayout infrastructure. We can define new - * layouts for conv2d ops for now. Most of the other operators try to adapt to their input layout - * using the InferCorrectLayout infrastructure. - * - * \param desired_layouts Specify mapping of op_name to array of desired layouts for each input. - * For example: Map("nn.conv2d", Array("NHWC", "OHWI")), - * this specifies the desired layout for data then kernel for nn.conv2d. - * \return The pass. - */ -TVM_DLL Pass ConvertLayout(const Map>& desired_layouts); - -/*! - * \brief Legalizes an expr with another expression. - * \param legalize_map_attr_name The Op's attr name which corresponds to the legalize rule function. - * One can collect and isolate similar type of legalize transformations using this param. For - * example, transformations that only apply to Dialects can be isolated into a FTVMDialectLegalize - * string. This pass calls only those transformations that have been registered using the supplied - * legalize_map_attr_name. - * - * \return The pass. - */ -TVM_DLL Pass Legalize(const String& legalize_map_attr_name = "FTVMLegalize"); - -/*! - * \brief Canonicalize cast expressions to make operator fusion more efficient. - * - * \return The pass. - */ -TVM_DLL Pass CanonicalizeCast(); - -/*! - * \brief Add abstraction over a constructor or global variable bound to a function. - * - * For example: `square` is transformed to - * `fn (%x: int32) -> int32 { square(x) }`. - * - * See https://en.wikipedia.org/wiki/Lambda_calculus#%CE%B7-conversion - * for more details. - * - * \param expand_constructor Whether to expand constructors. - * \param expand_global_var Whether to expand global variables. - * - * \return The pass. - */ -TVM_DLL Pass EtaExpand(bool expand_constructor, bool expand_global_var); - -/*! - * \brief Partition a Relay program into regions that can be executed on - * different backends. - * - * \return The pass. - */ -TVM_DLL Pass PartitionGraph(); - -/*! - * \brief Inline the global functions marked as `inline` in a given Relay - * IRModule. - * - * \return The pass. - */ -TVM_DLL Pass Inline(); - -/*! - * \brief Remove the unused functions in the Relay IRModule. - * - * \param entry_functions The entry functions used to search the functions that - * are being used. - * - * \return The pass. - */ -TVM_DLL Pass RemoveUnusedFunctions(Array entry_functions); - -/*! - * \brief Simplify the Relay expression. - * - * \return The pass. - */ -TVM_DLL Pass SimplifyExpr(); - -/*! - * \brief Stripped down version of SimplifyExpr which is run after AlterOpLayout. - * - * \return The pass. - */ -TVM_DLL Pass SimplifyExprPostAlterOp(); - -/*! - * \brief Run any custom passes registered under "RelayToTIR" attributes on TargetKinds. - * - * This pass looks for inline, let-bound or global functions which have a "Compiler" attribute. - * If the attribute value corresponds to a TargetKind with a "RelayToTIR" attribute, then the - * 'custom' pass bound to that attribute is run (at most once) on the IRModule as a whole. - * - * If, in addition, the \p config has a Target with a matching TargetKind, that Target is set - * as the 'current' target before the custom pass is executed. In this way it is possible - * for custom passes to pick up target options which may guide how they transform the IRModule. - * (Those targets are referred to as 'extern codegen targets' elsewhere). - * - * A typical custom pass will: - * - Find calls to "Compiler" attributes functions with matching compiler name. - * - Lower those function to TIR PrimFuncs. - * - Bind those functions into the IRModule under the functions' "global_symbol" attribute. - * - Replace all calls to those functions with 'call_lowered' to the matching global. - * Care should be taken to handle multiple calls to the same function. - * See src/relay/backend/contrib/example_target_hooks/relay_to_tir.cc for an example custom pass. - * - * It is also possible (despite the pass and attribute names!) for the custom pass to proceed - * directly to a runtime::Module, which can be attached to the output IRModules "external_mods" - * attribute (taking care not to clobber any existing modules). In this case the flow is as above, - * except: - * - The runtime::Module must contain a binding for each compiled function under their - * "global_symbol" (ie runtime::Module::ImplementsFunction should return true). - * - A Relay Function must be bound (or re-bound) into the result IRModule, again with the same - * "global_symbol", but with only the "Extern" attribute set to Integer(1). The function body - * should be the original function body. In this way we always have a TVM definition matching - * every global function name. - * - * There are many existing runtime::Modules, ranging from source to object to dynamic libaries to - * entirely custom implementations. Some of those may require additional compilation using - * 'export_library' on the final build artifact. - * - * The OutlineCompilerFunctionsWithExistingGlobalSymbols and MarkCompilerFunctionsAsExtern utility - * passes can be used by custom passes to take care of some of the boilerplate. - * - * TODO(mbs): Rename PreLoweringTargetHooks? - * - * \param config All available targets. - * - * \return The pass. - */ -TVM_DLL Pass RelayToTIRTargetHook(CompilationConfig config); - -/*! - * \brief A pass for manifesting explicit memory allocations and rewriting - * specific dialects. - * - * \param cpu_virtual_device VirtualDevice for computations and data which must reside on a CPU, - * such as shapes and shape functions. - * - * \return The pass. - */ -TVM_DLL Pass ManifestAlloc(VirtualDevice cpu_virtual_device); - -/*! - * \brief A pass for manifesting variable lifetimes by inserting kill operations when variables - * become dead. This pass should be run after ManifestAlloc, and should not be run more than once. - * - * \return The pass. - */ -TVM_DLL Pass ManifestLifetimes(); - -/*! - * \brief Uses existing "on_device" and "device_copy" CallNodes to infer the \p VirtualDevice on - * which every Relay sub-expression should run and the result stored. Captures the result of that - * analysis using new "on_device" and "device_copy" CallNodes. - * - * See tvm::relay::transform::{LexicalOnDeviceMixin,DeviceAwareExprVisitor,DeviceAwareExprMutator} - * for help recovering the device for an arbitrary sub-expression in downstream transformations. - * - * \param config Describes the targets and default \p VirtualDevice for all primitive operators and - * host sub-expressions. - * - * \return The pass. - */ -TVM_DLL Pass PlanDevices(CompilationConfig config); - -/*! - * \brief This transform flattens atrous convolution, which corresponds to the sequence of - * operations: "space_to_batch_nd"->"conv2d"->"batch_to_space_nd" and convert them into subgraphs - * with a convolution with the modified "dilation" and recalculated "padding" parameters. - * - * \return The pass. - */ -TVM_DLL Pass FlattenAtrousConv(); - -/*! - * \brief Annotates the minimum required memory of each primitive function callsite by analyzing - * the liveness of the input/output tensors at each function callsite and calculating the total - * amount of memory these tensors require. This is added as a "used_memory" annotation to the - * function in question as a list of the number of bytes for each callsite. In addition, the - * containing function is annotated with an "io_used_memory" annotation which refers to the total - * memory required for the IO tensors. - * - * Note: This pass does not support dynamic shapes, it is the users responsibility to check this - * pass isn't applied where dynamic shapes may be input. - */ -TVM_DLL Pass AnnotateUsedMemory(); - -/*! - * \brief Captures the post-dfs index and dominator post-dfs index of (most) expression nodes in - * their span, in the form "index::". This is useful for - * debugging since a) it helps identify pretty-printed sub-expressions within the overall model - * and b) the indexes are heavily used by Collage for its compact representation of sub-graphs. - * - * Note that Op and Constructor nodes are not changed even though they are assigned an - * post-dfs index. - */ -TVM_DLL Pass CapturePostDfsIndexInSpans(); - -/*! - * \brief Calls device dependent memory scope analysis pass, collects mapping of desirable - * expr->memory_scope and annotates expressions by VirtualDevice with required memory_scope - */ -TVM_DLL Pass AnnotateMemoryScope(); - -/*! - * \brief Removes non-fused reshapes after lowering the graph. - * InferType() cannot be invoked after calling this pass as it removes reshapes from the call - * graph. Many targets only need buffer addresses irrespective of the shapes of them. This makes - * reshapes symbolic once the graph has been lowered. Reshape removal results into smaller code - * size and reduced buffer allocations. It opens up opportunities of operator fusion in the target - * backend. Thus, consequently, it improves the performance of the inference. - */ -TVM_DLL Pass RemoveStandaloneReshapes(); - -} // namespace transform - -/*! - * \brief Bind the free variables to a Relay expression. This is a helper - * function usually called by other pass functions to help optimizations. - * If any free variables are introduced into a function, those are added - * to the functoin parameters. - * Additionally this may change the order of parameters if you map a variable - * to a variable. - * - * \param expr The input expression. - * \param binds The variable to expression map that will be used to help the - * binding. - * - * \return The updated expression. - */ -TVM_DLL Expr Bind(const Expr& expr, const tvm::Map& binds); - -/*! - * \brief Substitute variables with new variables (including function parameters) in a function. - * This is a helper function usually called by other pass functions to help optimizations. - * Expects all values in the bind map to be Vars. - * - * \param func The input function. - * \param binds The variable to expression map that will be used to help the - * binding. - * - * \return The updated expression. - */ -TVM_DLL Function SubstituteBoundVars(const Function& func, const tvm::Map& binds); - -/*! - * \brief Apply rewrite rules to rewrite the expr in post DFS order. This - * function is used as a helper function to rewrtie an expression in a pass. - * - * \param expr The expression. - * \param rewrite_map_attr_name The Op's attr name which corresponds to the rewrite - * rule function. - * \param fcontext Additional callback to provide context argument for each call node. - * \param fmulti_ref_trigger Transformation function to be called when - * an Expr consumed by multiple callers. - * \return The rewritten expression. - */ -TVM_DLL Expr ForwardRewrite(const Expr& expr, const String& rewrite_map_attr_name, - std::function fcontext = nullptr, - std::function fmulti_ref_trigger = nullptr); - -/*! - * \brief Apply rewrite rules to rewrite the expr in post DFS order. This - * function is used as a helper function to rewrtie an expression in a pass. - * - * \param expr The expression. - * \param rewrite_func The rewrite func that will apply to all operators. - * \param fcontext Additional callback to provide context argument for each call node. - * \param fmulti_ref_trigger Transformation function to be called when - * an Expr consumed by multiple callers. - * - * \return The rewritten expression. - */ -TVM_DLL Expr ForwardRewrite(const Expr& expr, const FForwardRewrite& rewrite_func, - std::function fcontext = nullptr, - std::function fmulti_ref_trigger = nullptr); - -/*! - * \brief Rewrite the annotated program. - * - * \param expr The expression. - * \param fallback_device The fallback device which is the default device for - * operators without annotation. - * - * \return The updated program. - */ -TVM_DLL Expr RewriteAnnotatedOps(const Expr& expr, int fallback_device); - -/*! - * \brief Turn an expression into continuation passing style(CPS). - * - * CPS mean that every function will, instead of returning the result directly, - * be passed down an extra function (called the continuation) as argument, - * and pass the result to the continuation instead. - * - * Thus, every function call has to be passed an extra argument - * that represent the rest of the computation (Hence the name of continuation). - * - * Similarly, all other compute will be wrapped and call the continuation as well. - * - * \param f the function. - * \param mod the module. - * - * \return the converted Function. - */ -TVM_DLL Function ToCPS(const Function& f, const IRModule& mod); - -/*! - * \brief Remove the continuation argument of a CPS function. - * - * Note that this only transform the type back into un-CPS form - * when there is no higher order input/output. - * - * \param f the function. - * - * \return the converted Function. - */ -TVM_DLL Function UnCPS(const Function& f); - -/*! - * \brief Deduplicate the bound variables and type variables in the expression. - * - * \param e the expression. - * - * \return the deduplicated expression. - */ -TVM_DLL Expr DeDup(const Expr& e); - -namespace legalize { -TVM_DLL Expr Legalize(const Expr& expr, const std::string& legalize_map_attr_name); -} // namespace legalize - -} // namespace relay -} // namespace tvm - -#endif // TVM_RELAY_TRANSFORM_H_ diff --git a/include/tvm/relay/type.h b/include/tvm/relay/type.h deleted file mode 100644 index a388c82a8d90..000000000000 --- a/include/tvm/relay/type.h +++ /dev/null @@ -1,75 +0,0 @@ -/* - * 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 tvm/relay/type.h - * \brief Relay typed AST nodes. - */ -#ifndef TVM_RELAY_TYPE_H_ -#define TVM_RELAY_TYPE_H_ - -#include -#include -#include -#include -#include -#include -#include - -#include - -#include "base.h" - -namespace tvm { -namespace relay { - -// namespace update for backward compact -// will be removed later. -using AnyNode = tvm::tir::AnyNode; -using Any = tvm::tir::Any; -using Kind = TypeKind; -using Type = tvm::Type; -using TypeNode = tvm::TypeNode; -using TypeVar = tvm::TypeVar; -using TypeVarNode = tvm::TypeVarNode; -using GlobalTypeVar = tvm::GlobalTypeVar; -using GlobalTypeVarNode = tvm::GlobalTypeVarNode; -using TupleType = tvm::TupleType; -using TupleTypeNode = tvm::TupleTypeNode; -using TypeConstraint = tvm::TypeConstraint; -using TypeConstraintNode = tvm::TypeConstraintNode; -using FuncType = tvm::FuncType; -using FuncTypeNode = tvm::FuncTypeNode; -using IncompleteType = tvm::IncompleteType; -using IncompleteTypeNode = tvm::IncompleteTypeNode; -using RelayRefType = tvm::RelayRefType; -using RelayRefTypeNode = tvm::RelayRefTypeNode; -using TensorType = tvm::TensorType; -using TensorTypeNode = tvm::TensorTypeNode; -using TypeCall = tvm::TypeCall; -using TypeCallNode = tvm::TypeCallNode; -using TypeRelation = tvm::TypeRelation; -using TypeRelationNode = tvm::TypeRelationNode; -using TypeRelationFn = tvm::TypeRelationFn; -using TypeReporter = tvm::TypeReporter; -using TypeReporterNode = tvm::TypeReporterNode; - -} // namespace relay -} // namespace tvm -#endif // TVM_RELAY_TYPE_H_ diff --git a/include/tvm/runtime/container/base.h b/include/tvm/runtime/container/base.h index 262041e2ffa7..51d48ae7f23b 100644 --- a/include/tvm/runtime/container/base.h +++ b/include/tvm/runtime/container/base.h @@ -24,7 +24,6 @@ #ifndef TVM_RUNTIME_CONTAINER_BASE_H_ #define TVM_RUNTIME_CONTAINER_BASE_H_ -#include #include #include #include diff --git a/include/tvm/runtime/container/string.h b/include/tvm/runtime/container/string.h index c6382506b355..a7be84de23f9 100644 --- a/include/tvm/runtime/container/string.h +++ b/include/tvm/runtime/container/string.h @@ -25,7 +25,6 @@ #define TVM_RUNTIME_CONTAINER_STRING_H_ #include -#include #include #include #include diff --git a/include/tvm/runtime/metadata.h b/include/tvm/runtime/metadata.h deleted file mode 100644 index f921f3e39c60..000000000000 --- a/include/tvm/runtime/metadata.h +++ /dev/null @@ -1,142 +0,0 @@ -/* - * 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 tvm/runtime/metadata.h - * \brief Defines types which can be used in Metadata. - */ -#ifndef TVM_RUNTIME_METADATA_H_ -#define TVM_RUNTIME_METADATA_H_ - -#include -#include -#include -#include -#include -#include - -#include -#include -#include - -// Version number recorded in emitted artifacts for runtime checking. -#define TVM_METADATA_VERSION 1 - -namespace tvm { -namespace runtime { -namespace metadata { -/*! - * \brief Version of metadata emitted and understood by this compiler/runtime. - * Should be populated into the `version` field of all TVMMetadata. - */ -static const constexpr int64_t kMetadataVersion = TVM_METADATA_VERSION; - -class Metadata; -class TensorInfo; -class ConstantInfoMetadata; - -class MetadataNode : public MetadataBaseNode { - public: - explicit MetadataNode(const struct ::TVMMetadata* data) : data_{data} {} - static constexpr const char* _type_key = "metadata.MetadataNode"; - const char* get_c_struct_name() const override; - inline int64_t version() const { return int64_t(data_->version); } - inline int64_t num_inputs() const { return data_->num_inputs; } - ArrayAccessor inputs(); - inline int64_t num_outputs() const { return data_->num_outputs; } - ArrayAccessor outputs(); - inline int64_t num_workspace_pools() const { return data_->num_workspace_pools; } - ArrayAccessor workspace_pools(); - inline ::tvm::runtime::String mod_name() const { return ::tvm::runtime::String(data_->mod_name); } - const struct ::TVMMetadata* data() const { return data_; } - ArrayAccessor constant_pools(); - inline int64_t num_constant_pools() const { return data_->num_constant_pools; } - TVM_DECLARE_FINAL_OBJECT_INFO(MetadataNode, MetadataBaseNode); - - private: - const struct ::TVMMetadata* data_; -}; - -class Metadata : public MetadataBase { - public: - explicit Metadata(const struct ::TVMMetadata* data); - TVM_DEFINE_MUTABLE_OBJECT_REF_METHODS(Metadata, MetadataBase, MetadataNode); -}; - -class TensorInfoNode : public MetadataBaseNode { - public: - explicit TensorInfoNode(const struct ::TVMTensorInfo* data) : data_{data} {} - static constexpr const char* _type_key = "metadata.TensorInfoNode"; - const char* get_c_struct_name() const override; - inline ::tvm::runtime::String name() const { return ::tvm::runtime::String(data_->name); } - inline int64_t num_shape() const { return data_->num_shape; } - inline ::tvm::support::Span shape() const { - return ::tvm::support::Span(data_->shape, - data_->shape + data_->num_shape); - } - inline ::tvm::runtime::DataType dtype() const { return ::tvm::runtime::DataType(data_->dtype); } - const struct ::TVMTensorInfo* data() const { return data_; } - TVM_DECLARE_FINAL_OBJECT_INFO(TensorInfoNode, MetadataBaseNode); - - private: - const struct ::TVMTensorInfo* data_; -}; - -class TensorInfo : public MetadataBase { - public: - explicit TensorInfo(const struct ::TVMTensorInfo* data); - TVM_DEFINE_MUTABLE_OBJECT_REF_METHODS(TensorInfo, MetadataBase, TensorInfoNode); -}; - -class ConstantInfoMetadataNode : public MetadataBaseNode { - public: - explicit ConstantInfoMetadataNode(const struct ::TVMConstantInfo* data) : data_{data} {} - // This name should match TVMConstantInfo after processing - static constexpr const char* _type_key = "metadata.ConstantInfoNode"; - const char* get_c_struct_name() const override; - inline ::tvm::runtime::String name_hint() const { - return ::tvm::runtime::String(data_->name_hint); - } - inline size_t byte_offset() const { return data_->byte_offset; } - inline ::tvm::runtime::NDArray data() const { - ::tvm::runtime::NDArray ndarray; - if (data_->data_len) { - dmlc::MemoryFixedSizeStream bytes(const_cast(data_->data_bytes), data_->data_len); - ndarray.Load(&bytes); - } - return ndarray; - } - TVM_DECLARE_FINAL_OBJECT_INFO(ConstantInfoMetadataNode, MetadataBaseNode); - - protected: - const struct ::TVMConstantInfo* data_; -}; - -class ConstantInfoMetadata : public MetadataBase { - public: - explicit ConstantInfoMetadata(const struct ::TVMConstantInfo* data); - TVM_DEFINE_MUTABLE_OBJECT_REF_METHODS(ConstantInfoMetadata, MetadataBase, - ConstantInfoMetadataNode); -}; - -} // namespace metadata -} // namespace runtime -} // namespace tvm - -#endif // TVM_RUNTIME_METADATA_H_ diff --git a/include/tvm/runtime/metadata_base.h b/include/tvm/runtime/metadata_base.h deleted file mode 100644 index ca412a3b615c..000000000000 --- a/include/tvm/runtime/metadata_base.h +++ /dev/null @@ -1,220 +0,0 @@ -/* - * 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 tvm/runtime/metadata_base.h - * \brief Defines types which can be used in Metadata. - */ -#ifndef TVM_RUNTIME_METADATA_BASE_H_ -#define TVM_RUNTIME_METADATA_BASE_H_ - -#include -#include -#include -#include -#include - -#include -#include -#include -#include - -namespace tvm { -namespace runtime { -namespace metadata { - -/*! - * \brief Common base class for all Metadata. - * - * This class is used in the visitor classes as a internal check to ensure that verify that all - * parts of the Metadata struct used in codegen are Metadata objects. - */ -class MetadataBaseNode : public ::tvm::runtime::Object { - public: - virtual const char* get_c_struct_name() const = 0; - - static constexpr const char* _type_key = "metadata.MetadataBaseNode"; - TVM_DECLARE_BASE_OBJECT_INFO(MetadataBaseNode, ::tvm::runtime::Object); -}; - -/*! \brief Reference class for the common MetadataBaseNode class. */ -class MetadataBase : public ::tvm::runtime::ObjectRef { - public: - TVM_DEFINE_MUTABLE_OBJECT_REF_METHODS(MetadataBase, ::tvm::runtime::ObjectRef, MetadataBaseNode); -}; - -template -class ArrayAccessor; - -/*! \brief An iterator implementation that lazily instantiates the C++ wrapping Metadata class. */ -template -class ArrayIterator { - public: - ArrayIterator(size_t index, const ArrayAccessor* parent) - : index_{index}, parent_{parent} {} - - inline Ref operator*() { return (*parent_)[index_]; } - - inline ArrayIterator& operator++() { - if (index_ < parent_->size()) { - index_++; - } - - return *this; - } - - inline bool operator==(const ArrayIterator& other) const { - return parent_ == other.parent_ && index_ == other.index_; - } - - inline bool operator!=(const ArrayIterator& other) const { return !operator==(other); } - - private: - size_t index_; - const ArrayAccessor* parent_; -}; - -/*! \brief A span-like class which permits access to Array fields with complex elements. - * These array fields should be accessed from C++ using the Metadata wrapper classes. This class - * lazily instantiates those wrappers as they are accessed. - */ -template -class ArrayAccessor { - public: - using value_type = Ref; - using iterator = ArrayIterator; - using const_iterator = iterator; - - template ::value>::type> - ArrayAccessor(const C* data, size_t num_data) : data_{data}, num_data_{num_data} {} - - inline size_t size() const { return num_data_; } - - inline Ref operator[](size_t index) const { - if (index >= num_data_) { - throw std::runtime_error("Index out of range"); - } - - return Ref(&data_[index]); - } - - inline ArrayIterator begin() const { return ArrayIterator{0, this}; } - - inline ArrayIterator end() const { return ArrayIterator{num_data_, this}; } - - private: - const C* data_; - size_t num_data_; -}; - -/*! \brief A specialization of ArrayAccessor for String. - * This class is needed because the String constructor signature is different from the typical - * Metadata subclass. - */ -template <> -class ArrayAccessor { - public: - using value_type = ::tvm::runtime::String; - using iterator = ArrayIterator; - using const_iterator = iterator; - - ArrayAccessor(const char** data, size_t num_data) : data_{data}, num_data_{num_data} {} - - inline size_t size() const { return num_data_; } - - inline ::tvm::runtime::String operator[](size_t index) const { - if (index >= num_data_) { - throw std::runtime_error("Index out of range"); - } - return ::tvm::runtime::String(data_[index]); - } - - inline ArrayIterator begin() const { - return ArrayIterator{0, this}; - } - - inline ArrayIterator end() const { - return ArrayIterator{num_data_, this}; - } - - private: - const char** data_; - size_t num_data_; -}; - -/*! \brief Enumerates the primitive types which can be part of a Metadata instance. - * - * These are separate from TIR DataType because TIR does not model structs. - */ -enum MetadataKind : uint8_t { - kUint64 = 0, - kInt64 = 1, - kBool = 2, - kString = 3, - kHandle = 4, - kMetadata = 5, -}; - -/*! \brief Container for arrays in the metadata. - * - * Type information is needed when emitting arrays. This container augments the data field with - * the necessary typing information. - */ -class MetadataArrayNode : public MetadataBaseNode { - public: - MetadataArrayNode(Array array, MetadataKind kind, const char* type_key) - : array(::std::move(array)), kind{kind}, type_key{type_key} {} - - const char* get_c_struct_name() const final; - - std::string get_element_c_struct_name() const { - CHECK(kind == MetadataKind::kMetadata) - << "cannot get struct name for MetadataArray with kind=" << kind; - constexpr int prefix_size = sizeof("metadata.") - 1; - constexpr int suffix_size = sizeof("Node") - 1; - std::string type_key_str(type_key); - return std::string("TVM") + - type_key_str.substr(prefix_size, type_key_str.size() - prefix_size - suffix_size); - } - - Array array; - - /*! \brief Describes the storage class of the emitted struct member. */ - MetadataKind kind; - - /*! \brief When `kind` is Metadata, type_key of the MetadataBaseNode used with this array. */ - const char* type_key; - - static constexpr const char* _type_key = "metadata.MetadataArrayNode"; - TVM_DECLARE_BASE_OBJECT_INFO(MetadataArrayNode, MetadataBaseNode); -}; - -/*! \brief Reference class for MetadataArray. */ -class MetadataArray : public MetadataBase { - public: - MetadataArray(Array array, MetadataKind kind, const char* struct_name); - - TVM_DEFINE_MUTABLE_OBJECT_REF_METHODS(MetadataArray, MetadataBase, MetadataArrayNode); -}; - -} // namespace metadata -} // namespace runtime -} // namespace tvm - -#endif // TVM_RUNTIME_METADATA_BASE_H_ diff --git a/include/tvm/runtime/metadata_types.h b/include/tvm/runtime/metadata_types.h deleted file mode 100644 index 5d828843e2b8..000000000000 --- a/include/tvm/runtime/metadata_types.h +++ /dev/null @@ -1,109 +0,0 @@ -/* - * 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. - */ - -// LINT_C_FILE - -/*! - * \file tvm/runtime/metadata_types.h - * \brief Defines types which can be used in metadata here which - * are also shared between C and C++ code bases. - */ -#ifndef TVM_RUNTIME_METADATA_TYPES_H_ -#define TVM_RUNTIME_METADATA_TYPES_H_ - -#include -#include - -#ifdef __cplusplus -extern "C" { -#endif - -/*! - * \brief Top-level metadata structure. Holds all other metadata types. - */ -struct TVMMetadata { - /*! \brief Version identifier for this metadata. */ - int64_t version; - /*! \brief Inputs to the AOT run_model function. - * The order of the elements is the same as in the arguments to run_model. That is to say, - * this array specifies the first `num_inputs` arguments to run_model. - */ - const struct TVMTensorInfo* inputs; - /*! \brief Number of elements in `inputs` array. */ - int64_t num_inputs; - /*! \brief Outputs of the AOT run_model function. - * The order of the elements is the same as in the arguments to run_model. That is to say, - * this array specifies the last `num_outputs` arguments to run_model. - */ - const struct TVMTensorInfo* outputs; - /*! \brief Number of elements in `outputs` array. */ - int64_t num_outputs; - /*! \brief Workspace Memory Pools needed by the AOT main function. - * The order of the elements is the same as in the arguments to run_model. That is to say, - * this array specifies the last `num_workspace_pools` arguments to run_model. - */ - const struct TVMTensorInfo* workspace_pools; - /*! \brief Number of elements in `workspace_pools` array. */ - int64_t num_workspace_pools; - /*! \brief Constant pools needed by the AOT main function. - */ - const struct TVMConstantInfo* constant_pools; - /*! \brief Number of elements in `constant_pools` array. */ - int64_t num_constant_pools; - /*! \brief Name of the model, as passed to tvm.relay.build. */ - const char* mod_name; -}; - -/*! - * \brief Describes one tensor argument to `run_model`. - * NOTE: while TIR allows for other types of arguments, such as scalars, the AOT run_model - * function does not currently accept these. Therefore it's not possible to express those - * in this metadata. A future patch may modify this. - */ -struct TVMTensorInfo { - /*! \brief Name of the tensor, as specified in the Relay program. */ - const char* name; - /*! \brief Shape of the tensor. */ - const int64_t* shape; - /*! \brief Rank of this tensor. */ - int64_t num_shape; - /*! \brief Data type of one element of this tensor. */ - DLDataType dtype; -}; - -/*! - * \brief Describes one constant argument to `run_model`. - * - */ -struct TVMConstantInfo { - /*! \brief Name of the constant */ - const char* name_hint; - /*! \brief Offset in bytes of the constant */ - int64_t byte_offset; - /*! \brief length of the data_bytes field */ - int64_t data_len; - /*! \brief data bytes of serialized NDArray */ - const void* data_bytes; -}; - -#ifdef __cplusplus -} // extern "C" -#endif - -#endif // TVM_RUNTIME_METADATA_TYPES_H_ diff --git a/include/tvm/runtime/object.h b/include/tvm/runtime/object.h index 4483867f3ccb..ea31fcac48e2 100644 --- a/include/tvm/runtime/object.h +++ b/include/tvm/runtime/object.h @@ -339,8 +339,8 @@ class TVM_DLL Object { * \tparam ObjectType The object type * \return The corresponding RefType */ -template -inline RelayRefType GetRef(const ObjectType* ptr); +template +inline ObjectRefType GetRef(const ObjectType* ptr); /*! * \brief Downcast a base reference type to a more specific type. @@ -505,8 +505,8 @@ class ObjectPtr { friend class TVMRetValue; friend class TVMArgValue; friend class TVMMovableArgValue_; - template - friend RelayRefType GetRef(const ObjType* ptr); + template + friend ObjectRefType GetRef(const ObjType* ptr); template friend ObjectPtr GetObjectPtr(ObjType* ptr); }; diff --git a/include/tvm/runtime/vm/bytecode.h b/include/tvm/runtime/vm/bytecode.h deleted file mode 100644 index 637c1e70a79f..000000000000 --- a/include/tvm/runtime/vm/bytecode.h +++ /dev/null @@ -1,415 +0,0 @@ -/* - * 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 tvm/runtime/vm/bytecode.h - * \brief The bytecode for Relay virtual machine. - */ -#ifndef TVM_RUNTIME_VM_BYTECODE_H_ -#define TVM_RUNTIME_VM_BYTECODE_H_ - -#include -#include - -#include -#include - -namespace tvm { -namespace runtime { -namespace vm { - -/*! \brief A register name. */ -using RegName = int64_t; - -/*! \brief An alias for the integer type used ubiquitously - * in the VM. - */ -using Index = int64_t; - -/*! \brief An enumeration of Relay's opcodes. - * - * The opcode is used to implement instruction - * as a tagged union. - */ -enum class Opcode { - Move = 0U, - Ret = 1U, - Invoke = 2U, - InvokeClosure = 3U, - InvokePacked = 4U, - AllocTensor = 5U, - AllocTensorReg = 6U, - AllocADT = 7U, - AllocClosure = 8U, - GetField = 9U, - If = 10U, - LoadConst = 11U, - Goto = 12U, - GetTag = 13U, - LoadConsti = 14U, - Fatal = 15U, - AllocStorage = 16U, - ShapeOf = 17U, - ReshapeTensor = 18U, - DeviceCopy = 19U, - KillRegister = 20U, -}; - -/*! \brief A single virtual machine instruction. - * - * The representation of the instruction is as - * a tagged union. - * - * The first field represents which instruction, - * and by extension which field of the union - * is active. - */ -struct Instruction { - /*! \brief The instruction opcode. */ - Opcode op; - - /*! \brief The destination register. */ - RegName dst; - - union { - struct /* AllocTensor Operands */ { - /*! \brief The storage to allocate from. */ - RegName storage; - /*! \brief The offset into the storage to allocate from. */ - Index offset; - /*! \brief The number of dimensions. */ - uint32_t ndim; - /*! \brief The shape of tensor. */ - int64_t* shape; - /*! \brief The datatype of tensor to be allocated. */ - DLDataType dtype; - } alloc_tensor; - struct /* AllocTensorReg Operands */ { - /*! \brief The storage to allocate from. */ - RegName storage; - /*! \brief The offset into the storage to allocate from. */ - Index offset; - /*! \brief The register to read the shape out of. */ - RegName shape_register; - /*! \brief The datatype of tensor to be allocated. */ - DLDataType dtype; - } alloc_tensor_reg; - struct /* InvokeClosure Operands */ { - /*! \brief The register containing the closure. */ - RegName closure; - /*! \brief The number of arguments to the closure. */ - Index num_closure_args; - /*! \brief The closure arguments as an array. */ - RegName* closure_args; - }; - struct /* Return Operands */ { - /*! \brief The register to return. */ - RegName result; - }; - struct /* Move Operands */ { - /*! \brief The source register for a move operation. */ - RegName from; - }; - struct /* InvokePacked Operands */ { - /*! \brief The index into the packed function table. */ - Index packed_index; - /*! \brief The arity of the packed function. */ - Index arity; - /*! \brief The number of outputs produced by the packed function. */ - Index output_size; - /*! \brief The arguments to pass to the packed function. */ - RegName* packed_args; - }; - struct /* If Operands */ { - /*! \brief The register containing the test value. */ - RegName test; - /*! \brief The register containing the target value. */ - RegName target; - /*! \brief The program counter offset for the true branch. */ - Index true_offset; - /*! \brief The program counter offset for the false branch. */ - Index false_offset; - } if_op; - struct /* Invoke Operands */ { - /*! \brief The function to call. */ - Index func_index; - /*! \brief The number of arguments to the function. */ - Index num_args; - /*! \brief The registers containing the arguments. */ - RegName* invoke_args_registers; - }; - struct /* LoadConst Operands */ { - /* \brief The index into the constant pool. */ - Index const_index; - /*! \brief The index of the device on which the load will be made. */ - Index device_index; - }; - struct /* LoadConsti Operands */ { - /* \brief The index into the constant pool. */ - Index val; - } load_consti; - struct /* Jump Operands */ { - /*! \brief The jump offset. */ - Index pc_offset; - }; - struct /* Proj Operands */ { - /*! \brief The register to project from. */ - RegName object; - /*! \brief The field to read out. */ - Index field_index; - }; - struct /* GetTag Operands */ { - /*! \brief The register to project from. */ - RegName object; - } get_tag; - struct /* AllocADT Operands */ { - // TODO(mbs): Needs a DeviceAndScope. - /*! \brief The datatype's constructor tag. */ - Index constructor_tag; - /*! \brief The number of fields to store in the datatype. */ - Index num_fields; - /*! \brief The fields as an array. */ - RegName* datatype_fields; - }; - struct /* AllocClosure Operands */ { - // TODO(mbs): Needs a DeviceAndScope. - /*! \brief The index into the function table. */ - Index clo_index; - /*! \brief The number of free variables to capture. */ - Index num_freevar; - /*! \brief The free variables as an array. */ - RegName* free_vars; - }; - struct /* AllocStorage Operands */ { - /*! \brief The alignment of the allocation. */ - Index alignment; - /*! \brief The hint of the dtype. */ - DLDataType dtype_hint; - /*! \brief The number of dimensions. */ - uint32_t ndim; - union { - /*! \brief The shape of tensor. */ - int64_t* shape; - /*! \brief The size of the allocation. */ - RegName allocation_size; - }; - /*! \brief The index of the device on which the allocation will be made. */ - Index device_index; - } alloc_storage; - struct /* ShapeOf Operands */ { - RegName tensor; - } shape_of; - struct /* ReshapeTensor Operands */ { - RegName tensor; - RegName newshape; - } reshape_tensor; - struct /* DeviceCopy Operands */ { - RegName src; - /*! \brief The index of the source device to copy from. */ - Index src_device_index; - /*! \brief The index of the destination deviceto copy to. */ - Index dst_device_index; - } device_copy; - }; - - /*! - * \brief Construct a return instruction. - * \param return_reg The register containing the return value. - * \return The return instruction. - */ - static Instruction Ret(RegName return_reg); - /*! - * \brief Construct a fatal instruction. - * \return The fatal instruction. - */ - static Instruction Fatal(); - /*! - * \brief Construct a invoke packed instruction. - * \param packed_index The index of the packed function. - * \param arity The arity of the function. - * \param output_size The number of outputs of the packed function. - * \param args The argument registers. - * \return The invoke packed instruction. - */ - static Instruction InvokePacked(Index packed_index, Index arity, Index output_size, - const std::vector& args); - /*! - * \brief Construct an allocate tensor instruction with constant shape. - * \param storage The storage to allocate out of. - * \param offset The offset to allocate at. - * \param shape The shape of the tensor. - * \param dtype The dtype of the tensor. - * \param dst The destination register. - * \return The allocate tensor instruction. - */ - static Instruction AllocTensor(RegName storage, Index offset, const std::vector& shape, - DLDataType dtype, RegName dst); - /*! - * \brief Construct an allocate tensor instruction with register. - * \param storage The storage to allocate out of. - * \param offset The offset into the storage to allocate from. - * \param shape_register The register containing the shape. - * \param dtype The dtype of the tensor. - * \param dst The destination register. - * \return The allocate tensor instruction. - */ - static Instruction AllocTensorReg(RegName storage, Index offset, RegName shape_register, - DLDataType dtype, RegName dst); - /*! - * \brief Construct an allocate datatype instruction. - * \param tag The datatype tag. - * \param num_fields The number of fields for the datatype. - * \param fields The registers containing the fields. - * \param dst The register name of the destination. - * \return The allocate instruction tensor. - */ - static Instruction AllocADT(Index tag, Index num_fields, const std::vector& fields, - RegName dst); - /*! - * \brief Construct an allocate closure instruction. - * \param func_index The index of the function table. - * \param num_freevar The number of free variables. - * \param free_vars The registers of the free variables. - * \param dst The destination register. - * \return The allocate closure instruction. - */ - static Instruction AllocClosure(Index func_index, Index num_freevar, - const std::vector& free_vars, RegName dst); - /*! - * \brief Construct a get field instruction. - * \param object_reg The register containing the object to project from. - * \param field_index The field to read out of the object. - * \param dst The destination register. - * \return The get field instruction. - */ - static Instruction GetField(RegName object_reg, Index field_index, RegName dst); - /*! - * \brief Construct a get_tag instruction. - * \param object_reg The register containing the object to project from. - * \param dst The destination register. - * \return The get_tag instruction. - */ - static Instruction GetTag(RegName object_reg, RegName dst); - /*! - * \brief Construct an if instruction. - * \param test The register containing the test value. - * \param target The register containing the target value. - * \param true_branch The offset to the true branch. - * \param false_branch The offset to the false branch. - * \return The if instruction. - */ - static Instruction If(RegName test, RegName target, Index true_branch, Index false_branch); - /*! - * \brief Construct a goto instruction. - * \param pc_offset The offset from the current pc. - * \return The goto instruction. - */ - static Instruction Goto(Index pc_offset); - /*! - * \brief Construct an invoke instruction. - * \param func_index The index of the function to invoke. - * \param args The registers containing the arguments. - * \param dst The destination register. - * \return The invoke instruction. - */ - static Instruction Invoke(Index func_index, const std::vector& args, RegName dst); - /*! - * \brief Construct an invoke closure instruction. - * \param closure The register of the closure to invoke. - * \param args The registers containing the arguments. - * \param dst The destination register. - * \return The invoke closure instruction. - */ - static Instruction InvokeClosure(RegName closure, const std::vector& args, RegName dst); - /*! - * \brief Construct a load constant instruction. - * \param const_index The index of the constant. - * \param device_index The index of the device to load on. - * \param dst The destination register. - * \return The load constant instruction. - */ - static Instruction LoadConst(Index const_index, Index device_index, RegName dst); - /*! - * \brief Construct a load_constanti instruction. - * \param val The interger constant value. - * \param dst The destination register. - * \return The load_constanti instruction. - */ - static Instruction LoadConsti(Index val, RegName dst); - /*! - * \brief Construct a move instruction. - * \param src The source register. - * \param dst The destination register. - * \return The move instruction. - */ - static Instruction Move(RegName src, RegName dst); - /*! - * \brief Allocate a storage block. - * \param size The size of the allocation. - * \param alignment The allocation's alignment. - * \param dtype_hint The data type hint for the allocator. - * \param device_index The index of the device to allocate on. - * \param shape The shape of the allocation. - * \param dst The destination to place the storage. - * \return The alloc storage instruction. - */ - static Instruction AllocStorage(RegName size, Index alignment, DLDataType dtype_hint, - Index device_index, const std::vector& shape, - RegName dst); - /*! - * \brief Get the shape of an input tensor. - * \param tensor The input tensor. - * \param dst The destination to store the shape of the given tensor. - * \return The shape of instruction. - */ - static Instruction ShapeOf(RegName tensor, RegName dst); - /*! - * \brief Reshape the tensor given the new shape. - * \param tensor The input tensor. - * \param newshape The shape tensor. - * \param dst The destination to store the output tensor with new shape. - * \return The reshape tensor instruction. - */ - static Instruction ReshapeTensor(RegName tensor, RegName newshape, RegName dst); - /*! - * \brief Copy tensor cross different devices. - * \param src The source register. - * \param src_device_index The index of the device holding the tensor in the source register. - * \param dst_device_index The index of the device to hold the tensor in the destination register. - * \param dst The destination register to store the copied tensor. - * \return The device copy instruction. - */ - static Instruction DeviceCopy(RegName src, Index src_device_index, Index dst_device_index, - RegName dst); - - static Instruction KillRegister(RegName dst); - - Instruction(); - Instruction(const Instruction& instr); - Instruction& operator=(const Instruction& instr); - ~Instruction(); - - friend std::ostream& operator<<(std::ostream& os, const Instruction&); -}; - -} // namespace vm -} // namespace runtime -} // namespace tvm - -#endif // TVM_RUNTIME_VM_BYTECODE_H_ diff --git a/include/tvm/runtime/vm/executable.h b/include/tvm/runtime/vm/executable.h deleted file mode 100644 index 12bb115aa783..000000000000 --- a/include/tvm/runtime/vm/executable.h +++ /dev/null @@ -1,386 +0,0 @@ -/* - * 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 tvm/runtime/vm/executable.h - * \brief The Relay virtual machine executable. - */ -#ifndef TVM_RUNTIME_VM_EXECUTABLE_H_ -#define TVM_RUNTIME_VM_EXECUTABLE_H_ - -#include -#include -#include -#include -#include -#include - -#include -#include -#include -#include -#include - -namespace tvm { -namespace runtime { -namespace vm { - -struct VMFunction; - -/*! - * \brief The executable emitted by the VM compiler. - * - * The executable contains information (e.g. data in different memory regions) - * to run in a virtual machine. - * - * - Global section, containing all globals. - * - Constant section, storing the constant pool. - * - Primitive name section, containing the function name of the primitive ops - * used by the virtual machine. - * - Code section, handling the VM functions and bytecode. - */ -class TVM_DLL Executable : public ModuleNode { - public: - TVM_MODULE_VTABLE_BEGIN("VMExecutable"); - TVM_MODULE_VTABLE_ENTRY("get_lib", &Executable::GetLib); - TVM_MODULE_VTABLE_ENTRY("get_bytecode", &Executable::GetBytecode); - TVM_MODULE_VTABLE_ENTRY("get_constants", &Executable::GetConstants); - TVM_MODULE_VTABLE_ENTRY("get_virtual_devices", &Executable::GetVirtualDevices); - TVM_MODULE_VTABLE_ENTRY("get_primitives", &Executable::GetPrimitives); - TVM_MODULE_VTABLE_ENTRY("get_stats", &Executable::Stats); - TVM_MODULE_VTABLE_ENTRY("save", &Executable::Save); - TVM_MODULE_VTABLE_ENTRY("get_function_arity", &Executable::GetFunctionArity); - TVM_MODULE_VTABLE_ENTRY("get_function_param_name", &Executable::GetFunctionParameterName); - TVM_MODULE_VTABLE_ENTRY("vm_load_executable", &Executable::VMLoadExecutable); - TVM_MODULE_VTABLE_ENTRY("move_late_bound_consts", &Executable::MoveLateBoundConstantsToFile); - TVM_MODULE_VTABLE_ENTRY("get_late_bound_consts", &Executable::GetLateBoundConstants); - TVM_MODULE_VTABLE_ENTRY("load_late_bound_consts", &Executable::LoadLateBoundConstantsFromFile); - TVM_MODULE_VTABLE_ENTRY("load_late_bound_consts_from_map", - &Executable::LoadLateBoundConstantsFromMap); - TVM_MODULE_VTABLE_END(); - - /*! \brief Get the property of the runtime module .*/ - int GetPropertyMask() const final { return ModulePropertyMask::kBinarySerializable; }; - /*! \brief Creates a VM that loads `this` as the executable. */ - Module VMLoadExecutable(); - /*! - * \brief Write the Executable to the binary stream in serialized form. - * - * Late-bound constants (if any) must have already been saved by \p - * MoveLateBoundConstantsToBinary. - * - * \param stream The binary stream to save the executable to. - */ - void SaveToBinary(dmlc::Stream* stream) final; - - /*! - * \brief Write the Executable to the provided path as a file containing its serialized content. - * - * Late-bound constants (if any) must have already been saved by \p - * MoveLateBoundConstantsToBinary. - * - * \param path The path to write the serialized data to. - * \param format The format of the serialized blob. - */ - void SaveToFile(const String& path, const String& format) final; - - /*! - * \brief Serialize the executable into global section, constant section, and - * code section. This object must outlive the returned byte array. - * - * Late-bound constants (if any) must have already been saved by \p - * MoveLateBoundConstantsToBinary. - * - * \return The binary representation of the VM. - */ - TVMByteArray Save(); - - /*! - * \brief Load the saved VM executable. - * - * Late-bound constants (if any) must then be loaded by \p LoadLateBoundConstantsFromBinary. - * - * \param code The bytecode in string. - * \param lib The compiled runtime library. - * - * \return exe The constructed executable. - */ - static runtime::Module Load(const std::string& code, const runtime::Module lib); - - /*! - * \brief Returns the late-bound constants for the executable (if any) as a byte-stream. - * Leaves the executable's late-bound constants map empty. Only constants who's byte - * tensor size is greater than or equal to \p byte_limit are marked as late-bound. \p byte_limit - * may be zero. - * - * Must be called before \p SaveToBinary and friends if late-bound constants are - * desired. Otherwise can be ignore. - */ - void MoveLateBoundConstantsToStream(dmlc::Stream* stream, int64_t byte_limit); - - /*! - * \brief As for \p MoveLateBoundConstantsToStream, but save to file at \p path. - */ - void MoveLateBoundConstantsToFile(const std::string& path, int64_t byte_limit); - - /*! - * \brief Get a map of all constants with larger that byte_limit in size. - */ - Map GetLateBoundConstants(int64_t byte_limit); - - /*! - * \brief Restores the late-bound constants for the executable (if any) from given byte-stream. - * - * Must be called after \p Load but before any other methods if \p MoveLateBoundConstantsToBinary - * was used when saving. Otherwise can be ignored. - */ - void LoadLateBoundConstantsFromStream(dmlc::Stream* stream); - - /*! - * \brief Restores the late-bound constants for the executable (if any) from given map. - * - * Must be called after \p Load but before any other methods if \p MoveLateBoundConstantsToBinary - * was used when saving. Otherwise can be ignored. - */ - void LoadLateBoundConstantsFromMap(Map map); - - /*! - * \brief As for \p LoadLateBoundConstantsFromStream, but load from file at \p path. - */ - void LoadLateBoundConstantsFromFile(const std::string& path); - - /*! - * \brief Get the serialized form of the `functions`. This is - * essentially bytecode serialization. - * - * \return The serialized vm bytecode. - * - * \note The bytecode is in the following format: - * func_name reg_file_size num_instructions - * param1 param2 ... paramM - * instruction1 - * instruction2 - * ... - * instructionN - * - * Each instruction is printed in the following format: - * opcode num_fields field1 ... fieldX # The text format. - * - * Serializing an `Instruction` requires us to deal with the bytecode. Each line - * of the instructions could be serialized as the following format: - * hash, opcode, f1, f2, ..., fX, field with variable length - * 1. hash: the hash of the instruction. This number will be used to help us - * validate if an instruction is well-formed during deserialization. - * 2. opcode: the opcode code of the instruction. - * 3. f1, f2, ..., fX. These fields together represent the fixed fields in - * an instruction, e.g., `from` and `dst` fields of a `Move` instruction. For - * example, `DLDataType` will be unpacked into three fields (code, bits, lanes). - * 4. The rest of the line indicates the field with variable length, e.g., - * the shape of a tensor, the args used by an `InvokPacked` instruction, etc. - * - * The field starting from # is only used for debugging. The serialized code - * doesn't contain it, therefore the deserializer doens't need to handle it. - */ - std::string GetBytecode() const; - - /*! - * \brief Returns a description of all the constants in the executable in human-readable - * format. Intended for debugging and diff-testing. - */ - std::string GetConstants() const; - - /*! - * \brief Returns a description of all the (virtual) devices in the executable in human-readable - * format. Intended for debugging and diff-testing. - */ - std::string GetVirtualDevices() const; - - /*! - * \brief Returns a description of all the 'primitive' (ie PackedFuncs) in the executable in - * human-readable format. These correspond either to PrimFuncs we've compiled locally, or - * functions compiled by a BYOC external codegen. Intended for debugging and diff-testing. - */ - std::string GetPrimitives() const; - - /*! - * \brief Print the detailed statistics of the given code, i.e. number of - * globls and constants, etc. - */ - std::string Stats() const; - - /*! - * \brief Get the `lib` module in an executable. Users have the flexibility to call - * `export_library` from the frontend to save the library to disk. - * - * \return The runtime module that contains the hardware dependent code. - */ - runtime::Module GetLib() const; - - /*! - * \brief Set the `lib` module in an executable. - * - * This allows us to do partial initialization in the case of (de|ser)ialization cases. - * This method also ensures correct initialization of library ensuring we only Import a - * single library. - * - * NB: This also provides some abstraction over how libraries are stored as there are plans - * to iterate on the way runtime::Module works in the backend of the compiler. - */ - void SetLib(const runtime::Module& lib); - - /*! - * \brief Get VMFunction. - * \param func_name The function's name. - * \return VMFunction. - */ - const VMFunction& GetVMFunctionWithName(const std::string& func_name) const; - - /*! - * \brief Get the arity of the VMFunction. - * \param func Function name. - * \return The number of parameters. - */ - int GetFunctionArity(std::string func) const; - - /*! - * \brief Get the parameter name given the function name and parameter index. - * \param func Function name. - * \param index Parameter index. - * \return The parameter name. - */ - std::string GetFunctionParameterName(std::string func, int index) const; - - virtual ~Executable() {} - - /*! - * \brief The (compile-time, virtual) devices corresponding to each device index. - * This vector contains a pair Device and its memory_scope. - */ - std::vector> virtual_devices; - /*! - * \brief The device index corresponding to the 'host' device. That will hold and evaluate - * shape-related data and code. - */ - int host_device_index = -1; - /*! - * \brief The global constant array. - * - * LoadConst instructions indexes are w.r.t. this vector. Late-bound constants are removed - * from this table after saving late-bound constants. - */ - std::vector constants; - /*! - * \brief For each constant index the name of the late-bound constant, or null if constant is - * immediate. Only populated after loading executable but before loading late-bound constants. - */ - std::vector late_bound_constant_names; - - /*! \brief A map from globals (as strings) to their index in the Relay function map. */ - std::unordered_map global_map; - /*! \brief A mapping from the packed function's global name (as string) to the index that - * corresponds to the position of the `packed_funcs` list in a `VirtualMachine` object. - */ - std::unordered_map primitive_map; - /*! \brief The structural hashes of the operators in this function. */ - std::map> op_attrs; - /*! \brief The virtual machine's function table. */ - std::vector functions; - /*! \brief The index of the device holding each constant. */ - std::vector const_device_indexes; - - private: - /*! - * \brief Save the virtual devices - * - * /param strm The output stream. - */ - void SaveVirtualDevicesSection(dmlc::Stream* strm); - - /*! - * \brief Save the globals. - * - * \param strm The output stream. - */ - void SaveGlobalSection(dmlc::Stream* strm); - - /*! - * \brief Save the constant pool. - * - * \param stream The output stream. - */ - void SaveConstantSection(dmlc::Stream* stream); - - /*! - * \brief Load the constant pool. - * - * \param stream The input stream. - */ - void LoadConstantSection(dmlc::Stream* stream); - - /*! - * \brief Save primitive op names. - * - * \param strm The output stream. - */ - void SavePrimitiveOpNames(dmlc::Stream* strm); - - /*! - * \brief Save the vm functions. - * - * \param strm The output stream. - */ - void SaveCodeSection(dmlc::Stream* strm); - - /*! - * \brief Load the virtual devices - * - * /param strm The input stream. - */ - void LoadVirtualDevicesSection(dmlc::Stream* strm); - - /*! - * \brief Load the globals. - * - * \param strm The input stream. - */ - void LoadGlobalSection(dmlc::Stream* strm); - - /*! - * \brief Load primitive op names. - * - * \param strm The input stream. - */ - void LoadPrimitiveOpNames(dmlc::Stream* strm); - - /*! - * \brief Load the vm functions. - * - * \param strm The input stream. - */ - void LoadCodeSection(dmlc::Stream* strm); - - /*! \brief The serialized bytecode. */ - std::string code_; -}; - -} // namespace vm -} // namespace runtime -} // namespace tvm - -#endif // TVM_RUNTIME_VM_EXECUTABLE_H_ diff --git a/include/tvm/runtime/vm/vm.h b/include/tvm/runtime/vm/vm.h deleted file mode 100644 index a5fe91186d99..000000000000 --- a/include/tvm/runtime/vm/vm.h +++ /dev/null @@ -1,476 +0,0 @@ -/* - * 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 tvm/runtime/vm/vm.h - * \brief The Relay virtual machine runtime. - */ -#ifndef TVM_RUNTIME_VM_VM_H_ -#define TVM_RUNTIME_VM_VM_H_ - -#include -#include -#include -#include -#include -#include -#include -#include - -#include -#include -#include -#include -#include - -namespace tvm { -namespace runtime { - -using memory::Allocator; -using memory::AllocatorType; -using memory::MemoryManager; -using memory::Storage; -using memory::StorageObj; - -namespace vm { - -/*! - * \brief An object representing a vm closure. - */ -class VMClosureObj : public ClosureObj { - public: - /*! - * \brief The index into the function list. The function could be any - * function object that is compatible to the VM runtime. - */ - size_t func_index; - /*! \brief The free variables of the closure. */ - std::vector free_vars; - - static constexpr const uint32_t _type_index = TypeIndex::kDynamic; - static constexpr const char* _type_key = "vm.Closure"; - TVM_DECLARE_FINAL_OBJECT_INFO(VMClosureObj, ClosureObj); -}; - -/*! \brief reference to closure. */ -class VMClosure : public Closure { - public: - VMClosure(size_t func_index, std::vector free_vars); - TVM_DEFINE_OBJECT_REF_METHODS(VMClosure, Closure, VMClosureObj); -}; - -/*! - * \brief A representation of a Relay function in the VM. - * - * Contains metadata about the compiled function, as - * well as the compiled VM instructions. - */ -struct VMFunction { - /*! \brief The function's name. */ - std::string name; - /*! \brief The function parameter names. */ - std::vector params; - /*! \brief The instructions representing the function. */ - std::vector instructions; - /*! \brief The size of the frame for this function */ - Index register_file_size = 0; - /*! \brief The indexes for the device holding each function parameter. */ - std::vector param_device_indexes; - - VMFunction(std::string name, std::vector params, - std::vector instructions, Index register_file_size, - std::vector param_device_indexes) - : name(std::move(name)), - params(std::move(params)), - instructions(std::move(instructions)), - register_file_size(register_file_size), - param_device_indexes(std::move(param_device_indexes)) { - ICHECK_EQ(this->params.size(), this->param_device_indexes.size()); - } - - VMFunction() = default; - - friend std::ostream& operator<<(std::ostream& os, const VMFunction&); -}; - -/*! - * \brief A representation of a stack frame. - * - * A stack frame is a record containing the information needed - * to restore the caller's virtual machine state after returning - * from a function call. - */ -struct VMFrame { - /*! \brief The return program counter. */ - Index pc; - /*! \brief The index into the function table, points to the caller. */ - Index func_index; - /*! \brief The number of arguments. */ - Index args; - /*! \brief A pointer into the caller function's instructions. */ - const Instruction* code; - - /*! \brief Statically allocated space for objects */ - std::vector register_file; - - /*! \brief Register in caller's frame to put return value */ - RegName caller_return_register; - - VMFrame(Index pc, Index func_index, Index args, const Instruction* code, Index register_file_size) - : pc(pc), - func_index(func_index), - args(args), - code(code), - register_file(register_file_size), - caller_return_register(0) {} -}; - -/*! - * \brief The virtual machine. - * - * The virtual machine contains all the current execution state, - * as well as the executable. - * - * The goal is to have a single self-contained object, - * enabling one to easily pass around VMs, execute them on - * multiple threads, or serialize them to disk or over the - * wire. - */ -class TVM_DLL VirtualMachine : public runtime::ModuleNode { - public: - /*! - * \brief Get a PackedFunc from module. - * - * The PackedFunc may not be fully initialized, - * there might still be first time running overhead when - * executing the function on certain devices. - * For benchmarking, use prepare to eliminate - * - * \param name the name of the function. - * \param sptr_to_self The shared_ptr that points to this module node. - * - * \return PackedFunc(nullptr) when it is not available. - * - * \note The function will always remain valid. - * If the function needs resource from the module(e.g. late linking), - * it should capture sptr_to_self. - */ - virtual PackedFunc GetFunction(const String& name, const ObjectPtr& sptr_to_self); - - virtual ~VirtualMachine() {} - - const char* type_key() const final { return "VirtualMachine"; } - - VirtualMachine() : frames_(), func_index_(0), code_(nullptr), pc_(0), exec_(nullptr) {} - - /*! - * \brief load the executable for the virtual machine. - * \param exec The executable. - */ - virtual void LoadExecutable(const ObjectPtr& exec); - - /*! \brief Get the property of the runtime module .*/ - int GetPropertyMask() const final { return ModulePropertyMask::kRunnable; } - - protected: - /*! \brief Push a call frame on to the call stack. */ - void PushFrame(Index arg_count, Index ret_pc, const VMFunction& vm_func); - - /*! - * \brief Pop a frame off the call stack. - * \return The number of frames left. - */ - Index PopFrame(); - - /*! - * \brief Write to a VM register. - * \param reg The register to write to. - * \param obj The object to write to. - */ - inline void WriteRegister(RegName reg, const ObjectRef& obj); - - /*! - * \brief Read a VM register. - * \param reg The register to read from. - * \return The read object. - */ - ObjectRef ReadRegister(RegName reg) const; - - /*! - * \brief Read a VM register and cast it to int32_t - * \param reg The register to read from. - * \return The read scalar. - */ - int64_t LoadScalarInt(RegName reg) const; - - /*! - * \brief Invoke a VM function. - * \param func The function. - * \param args The arguments to the function. - * \return The object representing the result. - */ - ObjectRef Invoke(const VMFunction& func, const std::vector& args); - - // TODO(@jroesch): I really would like this to be a global variable. - /*! - * \brief Invoke a VM function by name. - * \param name The function's name. - * \param args The arguments to the function. - * \return The object representing the result. - */ - ObjectRef Invoke(const std::string& name, const std::vector& args); - - /*! - * \brief Invoke a VM function. - * \param func The function. - * \param input_args The input arguments to the function. - * \param output_args The pre-allocated output arguments of the function. - * \return The object(s) representing the result. - */ - ObjectRef Invoke(const VMFunction& func, const std::vector& input_args, - const std::vector& output_args); - - /*! - * \brief Invoke a PackedFunction - * - * \param packed_index The offset of the PackedFunction in all functions. - * \param func The PackedFunction to be invoked. - * \param arg_count The number of arguments to the PackedFunction. - * \param output_size The number of outputs of the PackedFunction. - * \param args Arguments to the PackedFunction. - * - * \note The return value will be stored in the last output_size slots of args. - */ - virtual void InvokePacked(Index packed_index, const PackedFunc& func, Index arg_count, - Index output_size, const std::vector& args); - - /*! - * \brief Initialize the virtual machine for a set of (physical) devices. - * \param physical_devices The set of TVM devices. - * \param alloc_types The allocator types for each device. - */ - void Init(const std::vector& physical_devices, - const std::vector& alloc_types); - - /*! \brief Run VM dispatch loop. */ - void RunLoop(const std::vector& output_tensor_reg_indices = {}); - - /*! \brief Get device from the device list based on a given device index. */ - Device GetDevice(Index device_index) const; - Allocator* GetAllocator(Index device_index) const; - - /*! - * \brief Invoke a global setting up the VM state to execute. - * - * This does not begin execution of the VM. - */ - void InvokeGlobal(const VMFunction& func, const std::vector& args); - - /*! - * \brief Set inputs to a function. - * \param name The function name - * \param args args[offset:] are arguments to the - * function. If the arguments are not of the correct device for the function, - * they will be copied to the device. - * \param offset Starting offset of the arguments in `args`. - */ - void SetInput(std::string name, TVMArgs args, int offset); - - /*! - * \brief Set one input tensor with index or name to a function. - * \param name The function name. - * \param tag index or name of the input tensor . - * \param tensor the input tensor. If the tensor is not of the correct device for the function, - * they will be copied to the device. - */ - void SetOneInput(std::string name, const TVMArgValue& tag, const TVMArgValue& tensor); - - /*! - * \brief Set pre-allocated output tensors to a function. - * It is native implementation of 'set_outputs' python method. - * It is used in scenario when output tensors are allocated outside each invocation. - * Note: it sets set_outputs_enabled_[name] true and fill outputs_[name] - * but after invocation the first is switched off and the second is cleared - * \param name The function name - * \param args outputs to the function. - */ - void SetOutputs(std::string name, TVMArgs args); - - /*! - * \brief Preparation part of Invoke method before RunLoop. - * \param func the function. - * \param args input args - */ - void PrintInfoAndSetInputArgs(const VMFunction& func, const std::vector& args); - - /*! - * \brief Set pre-allocated outputs to register for specified function. - * \param func_name The function's name. - * \param outputs set of output tensors. - */ - void SetOutputTensorsToRegister(const std::string& func_name, - const std::vector& outputs); - - /*! - * \brief Internal hook for profiling the start of an op. - * - * This hook is only called on certain ops that are likely to take a - * significant amount of runtime (normally because they alloc or transfer to - * device). - * - * \param instr Instruction that will be executed after this hook fires - */ - virtual void OpStartHook(Instruction instr); - - /*! - * \brief Internal hook for profiling the end of an op. - */ - virtual void OpStopHook(); - - private: - /*! - * \brief Get index of input tensor from its name. - * \param func_name The function's name. - * \param input_name The input tensor name. - * \return The input tensor index. - */ - int64_t GetInputIndexFromVMFunction(const std::string& func_name, - const std::string& input_name) const; - - /*! - * \brief Get index of input tensor from its name. - * \param params parameter names. - * \param input_name The input tensor name. - * \return The input tensor index. - */ - int64_t GetInputIndexFromName(const std::vector& params, - const std::string& input_name) const; - - /*! - * \brief Check executable exists and get VM function from it. - * \param func_name The function's name. - * \return VM function. - */ - const VMFunction& CheckAndGetVMFunction(const std::string& func_name) const; - - /*! - * \brief Creats inputs_ field, if it exists check its size. - * \param func_name The function's name. - * \param size inputs_ field size. - */ - void CreateInputsOrCheckSize(const std::string& func_name, size_t size); - - /*! - * \brief Set one input tensor with given index to set of input tensors if need copy to given - * device. \param tensors the input tensors set (destination) \param tensor some tensor (not - * necessary DLTensor). \param index The input tensor index. \param dev device to copy if need. - */ - void SetInputTensorWithIndex(std::vector& tensors, // NOLINT(*) - const TVMArgValue& tensor, int index, Device dev); - - /*! - * \brief Convert tensor from TVMArgValue to ObjectRef. - * DLTensor and NDArray types are supported. - * \param tensor given arg value containing tensor. - * \return tensor in ObjectRef format - */ - ObjectRef TensorFromTVMArgValueToObjectRef(const TVMArgValue& tensor) const; - - /*! - * \brief Get index of outputs in register_file from func code - * \return result register index - */ - Index GetResultRegisterIndex() const; - - /*! - * \brief Calculate the index of operation which destination is result - * \param res_index is the index of op returning result - */ - void CalculatePreResultOpIndex(Index res_index); - - /*! - * \brief Get indices from register_file for output tensors. - * It helps to replace output tensors allocated in RunLoop by - * tensors pre-allocated outside. Scenario is when `set_output` is used - * \return indices from register_file for output tensors. - */ - std::vector GetOutputTensorRegIndices(); - - /*! - * \brief Write new allocated tensor to register_file of frame. - * \param instr current instruction containing shape and storage info. - */ - void WriteAllocatedTensor(const Instruction& instr); - - /*! - * \brief 'set_outputs_enabled' is assumed true for using this method. - * It is expected that result register has already contained tensor from outside, - * new memory is not allocated and write, but expected shape and data type are checked. - * For other register WriteAllocatedTensor method is used. - * \param instr current instruction containing shape and storage info. - */ - void WriteAllocatedTensorFromOutside(const Instruction& instr); - - bool FindIndex(const std::vector& indices, Index val) const; - - protected: - /*! \brief The virtual machine's packed function table. */ - std::vector packed_funcs_; - /*! \brief The current stack of call frames. */ - std::vector frames_; - /*! \brief The fuction table index of the current function. */ - Index func_index_; - /*! \brief The current pointer to the code section. */ - const Instruction* code_; - /*! \brief The virtual machine PC. */ - Index pc_; - /*! \brief The special return register. */ - ObjectRef return_register_; - /*! \brief The executable the VM will operate on. */ - ObjectPtr exec_; - /*! \brief The function name to inputs mapping. */ - std::unordered_map> inputs_; - /*! \brief The function name to flag enabling scenario with set outputs. */ - std::unordered_map set_outputs_enabled_; - /*! \brief The index of operation which destination is result. */ - Index preresult_op_index_ = -1; - /*! \brief The function name to indices of output tensors in register file. */ - std::unordered_map> output_tensor_reg_indices_; - /*! \brief The function name to pre-allocated outputs mapping. */ - std::unordered_map> outputs_; - /*! - * \brief The "physical" devices the VM can execute primitives on. All "device indexes" - * are w.r.t. this vector. Each entry in this vector must match the corresponding entry - * in the executable's "virtual" devices vector. - */ - std::vector devices_; - /*! \brief The cached memory allocators, one per device. */ - std::vector allocators_; - /*! - * \brief The constant pool for runtime. It caches the device dependent - * object to avoid rellocation of constants during inference. - */ - std::vector const_pool_; -}; - -} // namespace vm -} // namespace runtime -} // namespace tvm - -#endif // TVM_RUNTIME_VM_VM_H_ diff --git a/include/tvm/target/compilation_config.h b/include/tvm/target/compilation_config.h deleted file mode 100644 index eab34de1fb9a..000000000000 --- a/include/tvm/target/compilation_config.h +++ /dev/null @@ -1,205 +0,0 @@ -/* - * 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 tvm/target/compilation_config.h - * \brief A helper class to collect all the targets in canonical form necessary for compilation. - */ - -#ifndef TVM_TARGET_COMPILATION_CONFIG_H_ -#define TVM_TARGET_COMPILATION_CONFIG_H_ - -#include - -#include - -namespace tvm { - -/*! - * \brief Gathers the \p Targets and distinguished \p VirtualDevices in canonical form needed to - * compile a Relay module for execution over possibly heterogeneous devices. Centralizes the - * validation and canonicalization logic needed to transition from targets supplied by the Python - * APIs to a single internal representation. Also holds a cache of canonical \p VirtualDevices - * so that structural equal virtual devices have pointer equal canonical virtual devices. - * - * The construction of \p CompilationConfig is idempotent, in that given the same \p PassContext - * \p ctx and an arbitrary \p Array \p raw_targets: - * - * \code - * CompilationConfig(ctxt, raw_targets) - * is structurally equal to - * CompilationConfig(ctxt, CompilationConfig(ctxt, raw_targets)->primitive_targets) - * \endcode - * - * TODO(mbs): This is subject to change as we rework compilation options in general. This class - * is probably better called a 'CompositeTarget', and may be better made a sub-class of Target or - * some other common-target-root class. - */ -class CompilationConfigNode : public Object { - public: - /*! - * \brief The host target. Used for 'scalar' data and code (such as shapes and shape - * functions) and residual Relay expressions and data (such as conditionals and ADTs). - * Each \p primitive_target below will have this exact target object as its 'host'. - * - * Note that it is possible for a \p Target used for primitive operations to be structurally - * equal to the host \p Target (up to the \p host field.) However the \p Target objects will - * be distinct, and can be used as keys within a \p Map without collision. - */ - Target host_target; - - /*! - * \brief Vector of all available \p Targets for partitioning or compiling primitive tensor - * operators (kernels). May contain a \p Target for the same device type as for the - * \p host_target, however the \p host_target should be used for all host computations and data. - * Each \p Target will have \p host_target as its 'host'. - * - * Primitive targets must be unique by their kind name. In this way the - * \p FindPrimitiveTargetForKind method will find the unique target for the given kind name. - * This method is used when transitioning from an external codegen "Compiler" attribute value - * to the external codegen target representing that compiler. - * - * It is possible to have multiple primitive targets for the same device type. However given - * primitive targets left and right where: - * - left appears before right in the array - * - left->GetTargetDeviceType() == right->GetTargetDeviceType() - * then: - * - right.IsExternalCodegenFor(left) must be true - * In this way the \p FindPrimitiveTargetForDeviceOrFail method will find the 'most general' - * target for the requested device type. This method is used when transitioning from a device - * constraint to the target needed to compile for that device. - * - * In the homogeneous case primitive_targets will have just one entry, which will be pointer equal - * to optional_homogeneous_target. - * - * In the homogenous case where the 'host' is the same device as used for compiling kernels it - * is *not* the case that optional_homogenous_target == host_target. This is because all - * primitive always have their host field set to the host_target. Ie, it is valid to have: - * \code - * host_target=Target("llvm") - * optional_homogenous_target=Target("llvm", host=host_target) - * \endcode - */ - Array primitive_targets; - - /*! - * \brief \p VirtualDevice for primitive operators which are not otherwise constrained to a - * particular device. Used by the PlanDevices pass to determine a virtual device for every - * sub-expression. - */ - VirtualDevice default_primitive_virtual_device = VirtualDevice::FullyUnconstrained(); - - /*! \brief VirtualDevice for the host. */ - VirtualDevice host_virtual_device = VirtualDevice::FullyUnconstrained(); - - /*! - * \brief If defined then compile and/or run in 'homogenous execution mode'. In this mode all - * primitives are compiled for this target only. - * - * This is to support legacy passes which have not been adapted to heterogeneous execution and - * rely on an implicit global \p Target to be in scope. - * - * TODO(mbs): Remove once all passes are 'heterogeneous aware'. - */ - Target optional_homogeneous_target; - - void VisitAttrs(AttrVisitor* v); - - /*! - * \brief Returns the unique \p Target to use for \p device_type. Fail if no such target exists. - * - * This will be the first primitive target with matching device type. - */ - Target FindPrimitiveTargetForDeviceOrFail(DLDeviceType device_type) const; - - /*! - * \brief Returns the unique \p Target to use for \p kind_name. Returns null if none such. - */ - Optional FindPrimitiveTargetForKind(const std::string& kind_name) const; - - /*! - * \brief Returns a \p Target structurally equal to \p target, however prefer a structually equal - * known host or primitive target if the configuration has one. - */ - Target CanonicalTarget(const Target& target) const; - - /*! - * \brief Returns a \p VirtualDevice which is structurally equal to \p virtual_device on all its - * constrained fields, however: - * - If \p virtual_device has a device type but not a target, fill in a target using - * \p FindPrimitiveTargetOrFail. This is the one place we allow targets to be defaulted - * from device types alone. - * - If \p virtual_device has a target, also canonicalize it using \p CanonicalTarget. - * The returned object will be unique for the adjusted virtual device w.r.t. all other - * \p VirtualDevices returned by this method. - * - * We call the result the 'canonical' \p VirtualDevice. Two canonical \p VirtualDevices are - * structurally equal if and only if they are pointer equal. In this way we can build maps - * from virtual devices using just pointer equality. - */ - VirtualDevice CanonicalVirtualDevice(const VirtualDevice& virtual_device) const; - - static constexpr const char* _type_key = "CompilationConfig"; - TVM_DECLARE_FINAL_OBJECT_INFO(CompilationConfigNode, Object) - - private: - /*! - * \brief Sets the primitive targets, the host target, the default primitive virtual device, and - * the host virtual device given: - * - the vector of 'raw' targets (in any order) supplied by one of the TVM entry points. - * - any "relay.fallback_device_type" attribute on \p pass_ctx. - * - whether the LLVM backend is available. - * Will look for a suitable host target in the given primitive targets, but if none found may - * reuse a raw target or create a default CPU target. - */ - void Init(const transform::PassContext& pass_ctx, const Array& raw_targets); - - /*! - * \brief Returns a freshly constructed CPU \p Target. - */ - static Target MakeDefaultCPUTarget(); - - /*! - * \brief A cache of constructed virtual devices. - */ - mutable VirtualDeviceCache virtual_device_cache_; - - friend class CompilationConfig; -}; - -/*! - * \brief Managed reference class to \p CompilationConfig - * - * \sa CompilationConfig - */ -class CompilationConfig : public ObjectRef { - public: - /*! - * \brief Constructs the compilation config given the settings in \p pass_ctx and supplied - * \p raw_targets. See \p CompilationConfigNode::Init for details. - */ - TVM_DLL CompilationConfig(const transform::PassContext& pass_ctx, - const Array& raw_targets); - - TVM_DEFINE_OBJECT_REF_METHODS(CompilationConfig, ObjectRef, CompilationConfigNode); -}; - -} // namespace tvm - -#endif // TVM_TARGET_COMPILATION_CONFIG_H_ diff --git a/include/tvm/target/target.h b/include/tvm/target/target.h index 4c1d1fc1f3d2..6f22803b6844 100644 --- a/include/tvm/target/target.h +++ b/include/tvm/target/target.h @@ -238,30 +238,6 @@ class Target : public ObjectRef { /*! \return The target with the host stripped out */ Target WithoutHost() const; - /*! - * \brief Returns true if \p this target represents an external codegen. If so, - * \p this->kind->name can be used as the "Compiler" attribute on partitioned functions, - * and can be used to retrieve a partitioning pattern table using - * \p get_pattern_table. - */ - bool IsExternalCodegen() const; - - /*! - * \brief Returns true if \p this target represents an external codegen which is compatible - * with \p that target. In particular: - * - \p this has a true ::tvm::attr::kIsExternalCodegen attribute - * - \p that does not have a true ::tvm::attr::kIsExternalCodegen attribute - * - \p this and \p that have the same GetTargetDeviceType() - * - * After partitioning, the external codegen compilation path may use \p that to guide it's - * compilation to a \p runtime::Module. Given \p this, an appropriate \p that can be - * found using \p CompilationConfig::FindPrimitiveTargetOrFail(this->GetTargetDeviceType()). - * - * The \p CollagePartition pass uses this method to guide it's search over candidate partitions - * using external codegen. - */ - bool IsExternalCodegenFor(const Target& that) const; - private: Target(TargetKind kind, Optional host, String tag, Array keys, Map attrs); diff --git a/include/tvm/target/target_kind.h b/include/tvm/target/target_kind.h index 6b3b9c31a645..f398736c7394 100644 --- a/include/tvm/target/target_kind.h +++ b/include/tvm/target/target_kind.h @@ -387,36 +387,6 @@ inline TargetKindRegEntry& TargetKindRegEntry::set_name() { #define TVM_TARGET_KIND_REGISTER_VAR_DEF \ static DMLC_ATTRIBUTE_UNUSED ::tvm::TargetKindRegEntry& __make_##TargetKind -namespace attr { -// -// Distinguished TargetKind attribute names. -// - -/*! - * \brief A \p TargetKind attribute of type \p Bool. If true, then the target kind name also - * corresponds to an external codegen 'compiler' name. That name may be used: - * - To retrieve partitioning rules using \p get_partition_table. - * - To attach to Relay Functions under the \p attr::kCompiler attribute to indicate - * the function is to be compiled by the external codegen path. - * - * The \p CollagePartition pass uses this attribute to guide it's search over candidate partitions - * using external codegen. - * - * See also \p Target::IsExternalCodegenFor - */ -constexpr const char* kIsExternalCodegen = "is_external_codegen"; - -/*! - * \brief A \p TargetKind attribute of type \p FTVMRelayToTIR. If set, then the target kind name - * also corresponds to an external codegen 'compiler' name, and the bound value is a \p Pass - * to apply before the TVM lowering. - * - * See also \p Target::IsExternalCodegenFor - */ -constexpr const char* kRelayToTIR = "RelayToTIR"; - -} // namespace attr - /*! * \def TVM_REGISTER_TARGET_KIND * \brief Register a new target kind, or set attribute of the corresponding target kind. diff --git a/include/tvm/tir/transform.h b/include/tvm/tir/transform.h index a8d93bf898c4..b03b4d3a12a1 100644 --- a/include/tvm/tir/transform.h +++ b/include/tvm/tir/transform.h @@ -603,13 +603,6 @@ TVM_DLL Pass LowerAsyncDMA(); */ TVM_DLL Pass CommonSubexprElimTIR(bool enable_cse_tir = true, bool identify_equiv_terms = false); -/*! - * \brief Add TIR-printer output as debug information to all ops in the module - * \return The pass. - */ - -TVM_DLL Pass InstallDebugSpans(); - /*! * \brief Unify all the thread bindings for "blockIdx.x/y/z", "threadIdx.x/y/z", and * "vthread.x/y/z". Before the unification, two vars that are bound to a thread axis (e.g., diff --git a/include/tvm/tir/usmp/algo/greedy.h b/include/tvm/tir/usmp/algo/greedy.h deleted file mode 100644 index 8f0ed873593e..000000000000 --- a/include/tvm/tir/usmp/algo/greedy.h +++ /dev/null @@ -1,85 +0,0 @@ -/* - * 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 include/tvm/tir/usmp/algo/greedy.h - * \brief This header file contains helper methods used in greedy algorithms - * for planning memory for USMP - */ -#pragma once -#include -#include -#include -#include -#include -#include - -#include -#include - -namespace tvm { -namespace tir { -namespace usmp { -namespace algo { - -/*! - * \brief This is the base class for Greedy Algorithms where the sorting - * is specialized in the extended classes based on the greedy criteria. - */ -class GreedyBase { - public: - GreedyBase() {} - /*! - * \brief This function should be implemented by the extended classes to sort the BufferInfo - * objects based on a criteria and then calling PostSortAllocation. - */ - virtual Map PlanMemory(const Array& buffer_info_arr) = 0; - - protected: - /*! - * \brief Rounds up the offset to satisfy the alignement requirement - */ - size_t round_up_to_byte_alignment(const size_t& non_aligned_byte_offset, - const int& byte_alignment); - - /*! - * \brief A helper function check whether a offset is valid given the constraints - */ - bool IsValidPlacement(const PoolInfo& candidate_pool, const size_t& next_offset, - const size_t& size_bytes); - - /*! - * \brief Selects a pool for placement in the given set of ordered pool candidates - */ - PoolInfo SelectPlacementPool( - const BufferInfo& buf_info, - const std::unordered_map& pool_offsets); - - /*! - * \brief This is the base allocation function that works on sorted BufferInfo objects based - * on the greedy heuristic. The sorting algorithm has to be called before calling this. - */ - Map PostSortAllocation( - const std::vector& buffer_info_vec); -}; - -} // namespace algo -} // namespace usmp -} // namespace tir -} // namespace tvm diff --git a/include/tvm/tir/usmp/algorithms.h b/include/tvm/tir/usmp/algorithms.h deleted file mode 100644 index 54431b59d21c..000000000000 --- a/include/tvm/tir/usmp/algorithms.h +++ /dev/null @@ -1,84 +0,0 @@ -/* - * 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 tir/usmp/algorithms.h - * \brief The memory planning algorithm for USMP - */ - -#ifndef TVM_TIR_USMP_ALGORITHMS_H_ -#define TVM_TIR_USMP_ALGORITHMS_H_ - -#include - -namespace tvm { -namespace tir { -namespace usmp { -namespace algo { - -/*! - * \brief The Greedy-by-Size algorithm to plan memory - * - * This will perform a greedy algorithm in deciding the offsets - * within provided Pools, using the size of the buffer. - * - * \return A Map of BufferInfo objects and their associated PoolAllocation - */ -Map GreedyBySize(const Array& buffer_info_arr, - const Integer& memory_pressure); - -/*! - * \brief The Greedy-by-Conflicts algorithm to plan memory - * - * This will perform a greedy algorithm in deciding the offsets - * within provided Pools, using the number of liveness conflicts of the buffer. - * - * \return A Map of BufferInfo objects and their associated PoolAllocation - */ -Map GreedyByConflicts(const Array& buffer_info_arr, - const Integer& memory_pressure); -/*! - *\brief The Hill-Climb algoritm to plan memory - * - * This will perform an attempt to utilize probabalistic approach to memory - * allocation. Typically better than greedy family, but quite slow due to large - * number of iterations. - * - * \return A Map of BufferInfo objects and their associated PoolAllocation - */ -Map HillClimb(const Array& buffer_info_arr, - const Integer& memory_pressure); - -/*! - * \brief The Hill-Climb algorithm to plan memory - * - * This will perform a hill climbing algorithm in deciding the offsets - * within provided Pools. - * - * \return A Map of BufferInfo objects and their associated PoolAllocation - */ -Map HillClimb(const Array& buffer_info_arr, - const Integer& memory_pressure); - -} // namespace algo -} // namespace usmp -} // namespace tir -} // namespace tvm - -#endif // TVM_TIR_USMP_ALGORITHMS_H_ diff --git a/include/tvm/tir/usmp/analysis.h b/include/tvm/tir/usmp/analysis.h deleted file mode 100644 index a24851d33182..000000000000 --- a/include/tvm/tir/usmp/analysis.h +++ /dev/null @@ -1,49 +0,0 @@ -/* - * 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 tir/usmp/analysis.h - * \brief The analysis passes for TIR-based Unified Static Memory Planner - */ - -#ifndef TVM_TIR_USMP_ANALYSIS_H_ -#define TVM_TIR_USMP_ANALYSIS_H_ - -#include -#include - -namespace tvm { -namespace tir { -namespace usmp { - -/*! - * \brief Extract BufferInfo objects from a TIR IRModule - * - * This pass would extract the buffer information of allocate nodes - * including liveness conflict with other buffer info objects. - * - * \return A Map of BufferInfo objects and their associated Stmts - */ -BufferInfoAnalysis ExtractBufferInfo(const PrimFunc& main_func, const IRModule& mod); - -} // namespace usmp -} // namespace tir -} // namespace tvm - -#endif // TVM_TIR_USMP_ANALYSIS_H_ diff --git a/include/tvm/tir/usmp/transform.h b/include/tvm/tir/usmp/transform.h deleted file mode 100644 index ccb684463f18..000000000000 --- a/include/tvm/tir/usmp/transform.h +++ /dev/null @@ -1,75 +0,0 @@ -/* - * 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 tir/usmp/transform.h - * \brief The transform passes for TIR-based Unified Static Memory Planner - */ - -#ifndef TVM_TIR_USMP_TRANSFORM_H_ -#define TVM_TIR_USMP_TRANSFORM_H_ - -#include - -namespace tvm { -namespace tir { -namespace usmp { -namespace transform { - -using Pass = tvm::transform::Pass; - -/*! - * \brief Convert the analyzed PoolAllocation to offsets from pool variables - * - * This pass would convert the main function to accept pool variables as an input - * that get passed onto the operator PrimFuncs. Furthermore, the static allocations - * will be converted to offsets within the pool variable. - * - * \return the pass - */ -TVM_DLL Pass ConvertPoolAllocationsToOffsets(const Map& pool_allocations, - Bool emit_tvmscript_printable = Bool(false)); - -/*! - * \brief Assign PoolInfo objects to tir.allocate nodes depending on the PrimFunc's target - * - * This pass would assign default PoolInfo objects to allocate nodes that are not otherwise - * annotated, depending on pool info supplied for each target. - * - * \return the pass - */ -TVM_DLL Pass AssignPoolInfo(); - -/*! - * \brief This pass creates Allocate nodes for I/O tensors - * - * If the user wants to place the I/O tensors in the workspace, this pass is required to be - * run. In doing so, it will create Allocate nodes for I/O tensors to be planned, and be removed - * from function arguments. - * - * \return the pass - */ -TVM_DLL Pass CreateAllocatesForIO(); - -} // namespace transform -} // namespace usmp -} // namespace tir -} // namespace tvm - -#endif // TVM_TIR_USMP_TRANSFORM_H_ diff --git a/include/tvm/tir/usmp/utils.h b/include/tvm/tir/usmp/utils.h deleted file mode 100644 index a67350a2bb13..000000000000 --- a/include/tvm/tir/usmp/utils.h +++ /dev/null @@ -1,326 +0,0 @@ -/* - * 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 tir/usmp/utils.h - * \brief Utilities for Unified Static Memory Planner - */ - -#ifndef TVM_TIR_USMP_UTILS_H_ -#define TVM_TIR_USMP_UTILS_H_ - -#include -#include -#include -#include -#include -#include - -namespace tvm { - -/*! - * \brief PassContext option to enable the USMP - */ -constexpr const char* kUSMPEnableOption = "tir.usmp.enable"; -/*! - * \brief PassContext option to select the memory planning algorithm in USMP - */ -constexpr const char* kUSMPAlgorithmOption = "tir.usmp.algorithm"; -/*! - * \brief PassContext option to enable placing I/O tensors in the workspace - */ -constexpr const char* kUSMPUseWorkspaceIO = "tir.usmp.use_workspace_io"; -/*! - * \brief PassContext option to specify a custom memory planning algorithm in USMP. - * The algorithm should be provided as registered PackedFunc with the name tir.usmp.algorithm.NAME - */ -constexpr const char* kUSMPCustomAlgorithmOption = "tir.usmp.custom_algorithm"; - -namespace tir { -namespace usmp { -/*! - * \brief A special kind to distinguish between I/O tensors to the model - * and intermediate tensors of the model - */ -enum class BufferInfoKind { kIntermediate = 0, kInput = 1, kOutput = 2 }; - -/*! - * \brief Describes an abstract memory buffer that will get allocated inside a pool. - * The actual memory buffer in represented by PoolAllocationNode after static memory planning. - * - * See also for relay-level counterparts: - * relay::StorageToken (graph_plan_memory.cc) - * relay::backend::StorageInfoNode (relay/backend/utils.h) - * Region (python/tvm/relay/transform/memory_plan.py) - */ -struct BufferInfoNode : public Object { - /*! \brief The name of the buffer var */ - String name_hint; - /*! \brief The size in terms of bytes */ - Integer size_bytes; - /*! \brief The pool candidates that this buffer can get pooled to*/ - Array pool_candidates; - /*! \brief The byte alignment required for buffers that will placed within the pool */ - Integer alignment; - /*! \brief The liveness conflicting other buffer info objects */ - Array conflicts; - /*! \brief Whether BufferInfo object retains info about IO tensors or intermediaries */ - BufferInfoKind kind; - - void VisitAttrs(tvm::AttrVisitor* v) { - v->Visit("name_hint", &name_hint); - v->Visit("size_bytes", &size_bytes); - v->Visit("pool_candidates", &pool_candidates); - v->Visit("alignment", &alignment); - v->Visit("conflicts", &conflicts); - v->Visit("kind", &kind); - } - - bool SEqualReduce(const BufferInfoNode* other, SEqualReducer equal) const { - return equal(name_hint, other->name_hint) && equal(size_bytes, other->size_bytes) && - equal(pool_candidates, other->pool_candidates) && equal(alignment, other->alignment) && - equal(conflicts, other->conflicts) && equal(kind, other->kind); - } - - void SHashReduce(SHashReducer hash_reduce) const { - hash_reduce(name_hint); - hash_reduce(size_bytes); - hash_reduce(alignment); - hash_reduce(conflicts); - hash_reduce(pool_candidates); - hash_reduce(kind); - } - /*! - * \brief Set the liveness conflicts of this BufferInfo - * - * \param conflicting_buffer_info_objs An array of BufferInfo that conflicts in liveness - */ - TVM_DLL void SetConflicts(Array conflicting_buffer_info_objs); - - static constexpr const char* _type_key = "tir.usmp.BufferInfo"; - TVM_DECLARE_FINAL_OBJECT_INFO(BufferInfoNode, Object); -}; - -class BufferInfo : public ObjectRef { - public: - TVM_DLL BufferInfo(String name_hint, Integer size_bytes, Array pool_candidates, - Integer alignment = runtime::kDefaultWorkspaceAlignment, - BufferInfoKind kind = BufferInfoKind::kIntermediate); - TVM_DEFINE_MUTABLE_OBJECT_REF_METHODS(BufferInfo, ObjectRef, BufferInfoNode); -}; - -/*! - * \brief This is a composite node that is produced by extract_buffer_info - * analysis pass that contains useful global information that could be useful - * for memory planning algorithms. - */ -struct BufferInfoAnalysisNode : public Object { - /*! \brief The BufferInfo object and its associated TIR statement */ - Map buffer_info_stmts; - /*! \brief This represent maximum amount of memory being used at - * any point of time in the inference. This value is largely the - * best allocation an algorithm could achieve. Due to - * the complexities of conflict graphs, it would not be feasible - * to achieve this value, practically. However, it can be useful - * for iterative algorithms to know this value to define termination - * criteria.*/ - Integer memory_pressure; - - void VisitAttrs(tvm::AttrVisitor* v) { - v->Visit("buffer_info_stmts", &buffer_info_stmts); - v->Visit("memory_pressure", &memory_pressure); - } - - bool SEqualReduce(const BufferInfoAnalysisNode* other, SEqualReducer equal) const { - return equal(buffer_info_stmts, other->buffer_info_stmts) && - equal(memory_pressure, other->memory_pressure); - } - - void SHashReduce(SHashReducer hash_reduce) const { - hash_reduce(buffer_info_stmts); - hash_reduce(memory_pressure); - } -}; - -class BufferInfoAnalysis : public ObjectRef { - public: - TVM_DLL BufferInfoAnalysis(Map buffer_info_stmts, Integer memory_pressure); - TVM_DEFINE_MUTABLE_OBJECT_REF_METHODS(BufferInfoAnalysis, ObjectRef, BufferInfoAnalysisNode); -}; - -/*! - * \brief The pool allocation produced after the USMP algorithm - */ -struct PoolAllocationNode : public Object { - /*! \brief The assigned WorkspacePoolInfo or ConstantPoolInfo object */ - PoolInfo pool_info; - /*! \brief The byte offset within the pool*/ - Integer byte_offset; - - void VisitAttrs(tvm::AttrVisitor* v) { - v->Visit("pool_info", &pool_info); - v->Visit("byte_offset", &byte_offset); - } - - bool SEqualReduce(const PoolAllocationNode* other, SEqualReducer equal) const { - return equal(pool_info, other->pool_info) && equal(byte_offset, other->byte_offset); - } - - void SHashReduce(SHashReducer hash_reduce) const { - hash_reduce(pool_info); - hash_reduce(byte_offset); - } - - static constexpr const char* _type_key = "tir.usmp.PoolAllocation"; - TVM_DECLARE_FINAL_OBJECT_INFO(PoolAllocationNode, Object); -}; - -class PoolAllocation : public ObjectRef { - public: - TVM_DLL PoolAllocation(PoolInfo pool_info, Integer byte_offset); - TVM_DEFINE_MUTABLE_OBJECT_REF_METHODS(PoolAllocation, ObjectRef, PoolAllocationNode); -}; - -/*! - * \brief This object contains information post-allocation for PoolInfo objects - */ -struct AllocatedPoolInfoNode : public Object { - /*! \brief The assigned PoolInfo object */ - PoolInfo pool_info; - /*! \brief The allocated size into this pool */ - Integer allocated_size; - /*! \brief An optional associated pool Var index of PrimFunc params*/ - Optional pool_var_idx; - - void VisitAttrs(tvm::AttrVisitor* v) { - v->Visit("pool_info", &pool_info); - v->Visit("allocated_size", &allocated_size); - v->Visit("pool_var_idx", &pool_var_idx); - } - - bool SEqualReduce(const AllocatedPoolInfoNode* other, SEqualReducer equal) const { - return equal(pool_info, other->pool_info) && equal(allocated_size, other->allocated_size) && - equal(pool_var_idx, other->pool_var_idx); - } - - void SHashReduce(SHashReducer hash_reduce) const { - hash_reduce(pool_info); - hash_reduce(allocated_size); - hash_reduce(pool_var_idx); - } - - static constexpr const char* _type_key = "ir.AllocatedPoolInfo"; - TVM_DECLARE_FINAL_OBJECT_INFO(AllocatedPoolInfoNode, Object); -}; - -class AllocatedPoolInfo : public ObjectRef { - public: - TVM_DLL AllocatedPoolInfo(PoolInfo pool_info, Integer allocated_size, - Integer pool_var_idx = Integer()); - TVM_DEFINE_MUTABLE_OBJECT_REF_METHODS(AllocatedPoolInfo, ObjectRef, AllocatedPoolInfoNode); -}; - -/*! - * \brief Convert the IR-bound BufferInfo map to an array of BufferInfo - * - * \param buffer_info_map IR-bound BufferInfo map - */ -Array ConvertToArrayOfBufferInfo(const Map& buffer_info_map); - -/*! - * \brief Calculate workspace required to execute a IRModule with main expressed in TIR - * - * \param mod the IRModule with TIR-based main function - */ -Integer CalculateModuleWorkspaceSize(const IRModule& mod); - -/*! - * \brief The allocate node attribute to indicate candidate memory pools. - * This needs to be kept in sync with CANDIDATE_MEMORY_POOL_ATTR in - * python/tvm/tir/usmp/utils.py. - */ -static constexpr const char* kPoolCandidatesAllocateAttr = "candidate_memory_pools"; - -/*! - * \brief The allocate node attribute to indicate it is being used to hold - * an input tensor, that needs to be initialized with. - */ -static constexpr const char* kInputTensorAllocate = "input_tensor"; - -/*! - * \brief The allocate node attribute to indicate it is being used to hold - * an output tensor. - */ -static constexpr const char* kOutputTensorAllocate = "output_tensor"; - -/*! - * \brief Calculate the size of the extents in bytes - * - * \param op the allocate node - */ -Integer CalculateExtentsSize(const AllocateNode* op); - -/*! - * \brief Calculate the size of the extents in bytes - * - * \param op the allocate const node - */ -Integer CalculateExtentsSize(const AllocateConstNode* op); - -/*! - * \brief Joins the Stmt nodes with PoolAllocation objects - * - * \param buffer_info_to_stmt the map of BufferInfo objects to Stmt nodes - * \param buffer_info_to_pool_allocation the map of BufferInfo objects to PoolAllocation objects - */ -Map AssignStmtPoolAllocations( - const Map& buffer_info_to_stmt, - const Map& buffer_info_to_pool_allocation); - -/*! - * \brief Obtains I/O tensor names to their PoolAllocation objects - * - * \param buffer_info_to_pool_allocation the map of BufferInfo objects to PoolAllocation objects - * - * This function will obtain pool allocations for I/O tensors if that had been planned - */ -Map GetIOPoolAllocations( - const Map& buffer_info_to_pool_allocation); - -} // namespace usmp -} // namespace tir - -namespace attr { -/*! - * \brief This is a BaseFunc attribute to indicate which input var represent - * a PoolInfo Object in the form of a Map. - */ -static constexpr const char* kPoolArgs = "pool_args"; - -/*! - * \brief This is a IRModule attribute that contains I/O Tensor names to pool - * allocations. - */ -static constexpr const char* kIOTensorPoolAllocations = "io_tensor_pool_allocations"; - -} // namespace attr - -} // namespace tvm - -#endif // TVM_TIR_USMP_UTILS_H_ diff --git a/jvm/core/src/main/java/org/apache/tvm/contrib/GraphExecutor.java b/jvm/core/src/main/java/org/apache/tvm/contrib/GraphExecutor.java deleted file mode 100644 index 30b2fb1acafb..000000000000 --- a/jvm/core/src/main/java/org/apache/tvm/contrib/GraphExecutor.java +++ /dev/null @@ -1,82 +0,0 @@ -/* - * 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. - */ - -package org.apache.tvm.contrib; - -import org.apache.tvm.Device; -import org.apache.tvm.Function; -import org.apache.tvm.Module; -import org.apache.tvm.TVMValue; -import org.apache.tvm.rpc.RPC; -import org.apache.tvm.rpc.RPCSession; -import org.apache.tvm.rpc.TVMRemoteDevice; - -import java.lang.reflect.Field; -import java.lang.reflect.InvocationTargetException; -import java.lang.reflect.Method; - -public class GraphExecutor { - /** - * Create a runtime executor module given a graph and module. - * @param graphJson The graph deployed in json format output by compiler. - * @param libmod The module of the corresponding function. - * @param dev The local or remote device to deploy the module. - * @return Runtime graph module that can be used to execute the graph. - */ - public static GraphModule create(String graphJson, Module libmod, Device dev) { - Function fcreate = Function.getFunction("tvm.graph_executor.create"); - if (fcreate == null) { - throw new RuntimeException("Cannot find global function tvm.graph_executor.create." - + "Did you compile tvm_runtime with correct version?"); - } - Module graphModule = fcreate.pushArg(graphJson) - .pushArg(libmod).pushArg(dev.deviceType).pushArg(dev.deviceId) - .invoke().asModule(); - - return new GraphModule(graphModule, dev); - } - - private static Object reflectionGetField(Object obj, String fieldName) { - try { - Field field = obj.getClass().getDeclaredField(fieldName); - field.setAccessible(true); - return field.get(obj); - } catch (NoSuchFieldException e) { - throw new RuntimeException(e); - } catch (IllegalAccessException e) { - throw new RuntimeException(e); - } - } - - private static Object reflectionStaticCall(Class clazz, String methodName, Object ... args) { - Class[] types = new Class[args.length]; - for (int i = 0; i < args.length; ++i) { - types[i] = args[i].getClass(); - } - try { - Method method = clazz.getDeclaredMethod(methodName, types); - method.setAccessible(true); - return method.invoke(null, args); - } catch (NoSuchMethodException e) { - throw new RuntimeException(e); - } catch (IllegalAccessException e) { - throw new RuntimeException(e); - } catch (InvocationTargetException e) { - throw new RuntimeException(e); - } - } -} diff --git a/jvm/core/src/main/java/org/apache/tvm/contrib/GraphModule.java b/jvm/core/src/main/java/org/apache/tvm/contrib/GraphModule.java deleted file mode 100644 index 0a0bc7efc46d..000000000000 --- a/jvm/core/src/main/java/org/apache/tvm/contrib/GraphModule.java +++ /dev/null @@ -1,189 +0,0 @@ -/* - * 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. - */ - -package org.apache.tvm.contrib; - -import org.apache.tvm.Device; -import org.apache.tvm.Function; -import org.apache.tvm.Module; -import org.apache.tvm.NDArray; - -/** - * Wrapper runtime module. - * This is a thin wrapper of the underlying TVM module. - * you can also directly call set_input, run, and get_output - * of underlying module functions. - */ -public class GraphModule { - private Module module; - private Device device; - - private Function fsetInput; - private Function frun; - private Function fgetOutput; - private Function fgetInput; - private Function fdebugGetOutput; - private Function floadParams; - - public GraphModule(Module module, Device dev) { - this.module = module; - this.device = dev; - fsetInput = module.getFunction("set_input"); - frun = module.getFunction("run"); - fgetInput = module.getFunction("get_input"); - fgetOutput = module.getFunction("get_output"); - try { - fdebugGetOutput = module.getFunction("debug_get_output"); - } catch (IllegalArgumentException ignored) { - // ignore - } - floadParams = module.getFunction("load_params"); - } - - /** - * Release the GraphModule. - *

- * We highly recommend you to do this manually since the GC strategy is lazy. - *

- */ - public void release() { - fsetInput.release(); - frun.release(); - fgetInput.release(); - fgetOutput.release(); - if (fdebugGetOutput != null) { - fdebugGetOutput.release(); - } - floadParams.release(); - module.release(); - } - - /** - * Set inputs to the module. - * @param key The input key. - * @param value The input value - * @return self. - */ - public GraphModule setInput(String key, NDArray value) { - NDArray input = value; - if (!value.device().equals(device)) { - input = NDArray.empty(value.shape(), device); - value.copyTo(input); - } - fsetInput.pushArg(key).pushArg(input).invoke(); - return this; - } - - /** - * Set inputs to the module. - * @param key The input key. - * @param value The input value. - * @return self. - */ - public GraphModule setInput(int key, NDArray value) { - NDArray input = value; - if (!value.device().equals(device)) { - input = NDArray.empty(value.shape(), device); - value.copyTo(input); - } - fsetInput.pushArg(key).pushArg(input).invoke(); - return this; - } - - /** - * Run forward execution of the graph. - * @return self. - */ - public GraphModule run() { - frun.invoke(); - return this; - } - - /** - * Get index-th input to out. - * @param index The input index. - * @param out The output array container. - * @return out. - */ - public NDArray getInput(int index, NDArray out) { - fgetInput.pushArg(index).pushArg(out).invoke(); - return out; - } - - /** - * Get index-th output to out. - * @param index The output index. - * @param out The output array container. - * @return out. - */ - public NDArray getOutput(int index, NDArray out) { - fgetOutput.pushArg(index).pushArg(out).invoke(); - return out; - } - - /** - * Run graph up to node and get the output to out. - * @param node The node name. - * @param out The output array container. - * @return out. - */ - public NDArray debugGetOutput(String node, NDArray out) { - if (fdebugGetOutput != null) { - fdebugGetOutput.pushArg(node).pushArg(out).invoke(); - } else { - throw new RuntimeException("Please compile runtime with USE_PROFILER = ON"); - } - return out; - } - - /** - * Run graph up to node and get the output to out. - * @param node The node index. - * @param out The output array container. - * @return out. - */ - public NDArray debugGetOutput(int node, NDArray out) { - if (fdebugGetOutput != null) { - fdebugGetOutput.pushArg(node).pushArg(out).invoke(); - } else { - throw new RuntimeException("Please compile runtime with USE_PROFILER = ON"); - } - return out; - } - - /** - * Load parameters from serialized byte array of parameter dict. - * @param params The serialized parameter. - * @return self. - */ - public GraphModule loadParams(byte[] params) { - floadParams.pushArg(params).invoke(); - return this; - } - - /** - * Get internal module function. - * @param key The key to the module. - * @return The function. - * @throws IllegalArgumentException if function does not exist. - */ - public Function getFunction(String key) { - return module.getFunction(key); - } -} diff --git a/jvm/core/src/test/java/org/apache/tvm/contrib/GraphExecutorTest.java b/jvm/core/src/test/java/org/apache/tvm/contrib/GraphExecutorTest.java deleted file mode 100644 index f79253e487e1..000000000000 --- a/jvm/core/src/test/java/org/apache/tvm/contrib/GraphExecutorTest.java +++ /dev/null @@ -1,117 +0,0 @@ -/* - * 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. - */ - -package org.apache.tvm.contrib; - -import org.apache.tvm.Module; -import org.apache.tvm.NDArray; -import org.apache.tvm.Device; -import org.apache.tvm.TestUtils; -import org.apache.tvm.rpc.Client; -import org.apache.tvm.rpc.RPCSession; -import org.apache.tvm.rpc.Server; -import org.junit.BeforeClass; -import org.junit.Test; -import org.slf4j.Logger; -import org.slf4j.LoggerFactory; - -import java.io.File; -import java.io.IOException; -import java.util.Scanner; - -import static org.junit.Assert.assertArrayEquals; - -public class GraphExecutorTest { - private final Logger logger = LoggerFactory.getLogger(GraphExecutor.class); - private static String loadingDir; - - @BeforeClass - public static void beforeClass() { - loadingDir = System.getProperty("test.tempdir"); - } - - @Test - public void test_add_one_local() throws IOException { - Module libmod = Module.load(loadingDir + File.separator + "graph_addone_lib.so"); - String graphJson = new Scanner(new File( - loadingDir + File.separator + "graph_addone.json")) - .useDelimiter("\\Z").next(); - - Device dev = Device.cpu(); - GraphModule graph = GraphExecutor.create(graphJson, libmod, dev); - - long[] shape = new long[]{4}; - NDArray arr = NDArray.empty(shape, dev); - arr.copyFrom(new float[]{1f, 2f, 3f, 4f}); - - NDArray out = NDArray.empty(shape, dev); - - graph.setInput("x", arr).run(); - graph.getOutput(0, out); - - assertArrayEquals(new float[]{2f, 3f, 4f, 5f}, out.asFloatArray(), 1e-3f); - - arr.release(); - out.release(); - graph.release(); - } - - @Test - public void test_add_one_remote() throws IOException { - if (!Module.enabled("rpc")) { - logger.warn("RPC is not enabled. Skip."); - return; - } - - String libPath = loadingDir + File.separator + "graph_addone_lib.so"; - String graphJson = new Scanner(new File( - loadingDir + File.separator + "graph_addone.json")) - .useDelimiter("\\Z").next(); - - TestUtils.RefInt port = new TestUtils.RefInt(); - Server server = null; - try { - server = TestUtils.startServer(port); - RPCSession remote = Client.connect("127.0.0.1", port.value); - Device dev = remote.cpu(); - - remote.upload(new File(libPath)); - Module mlib = remote.loadModule("graph_addone_lib.so"); - - GraphModule graph = GraphExecutor.create(graphJson, mlib, dev); - - long[] shape = new long[]{4}; - NDArray arr = NDArray.empty(shape, dev); - arr.copyFrom(new float[]{1f, 2f, 3f, 4f}); - - NDArray out = NDArray.empty(shape, dev); - - graph.setInput("x", arr).run(); - graph.getOutput(0, out); - - assertArrayEquals(new float[]{2f, 3f, 4f, 5f}, out.asFloatArray(), 1e-3f); - - arr.release(); - out.release(); - graph.release(); - } finally { - if (server != null) { - server.terminate(); - } - } - } -} diff --git a/jvm/native/linux-x86_64/pom.xml b/jvm/native/linux-x86_64/pom.xml index 5aacc7d5f617..ff7499d45ca8 100644 --- a/jvm/native/linux-x86_64/pom.xml +++ b/jvm/native/linux-x86_64/pom.xml @@ -114,7 +114,7 @@ under the License. - -std=c++0x + -std=c++17 -I../../../include diff --git a/jvm/native/osx-x86_64/pom.xml b/jvm/native/osx-x86_64/pom.xml index e035a76a5647..df3406ced8a6 100644 --- a/jvm/native/osx-x86_64/pom.xml +++ b/jvm/native/osx-x86_64/pom.xml @@ -115,7 +115,7 @@ under the License. - -std=c++0x + -std=c++17 -I../../../include diff --git a/jvm/native/src/main/native/org_apache_tvm_native_c_api.cc b/jvm/native/src/main/native/org_apache_tvm_native_c_api.cc index c039508b4b7f..77bc8d636098 100644 --- a/jvm/native/src/main/native/org_apache_tvm_native_c_api.cc +++ b/jvm/native/src/main/native/org_apache_tvm_native_c_api.cc @@ -25,9 +25,9 @@ #include "tvm_runtime.h" #else #include -#include #include #include +#include #endif #include #include diff --git a/python/tvm/__init__.py b/python/tvm/__init__.py index ab11f33cc035..a622df496959 100644 --- a/python/tvm/__init__.py +++ b/python/tvm/__init__.py @@ -43,12 +43,6 @@ from .ir import transform from .ir import instrument from .ir import container -from .ir import PoolInfo -from .ir import WorkspacePoolInfo -from .ir import ConstantPoolInfo -from .ir import PoolInfoProperties -from .ir import WorkspaceMemoryPools -from .ir import ConstantMemoryPools from . import ir # tvm.tir diff --git a/python/tvm/contrib/debugger/__init__.py b/python/tvm/contrib/debugger/__init__.py deleted file mode 100644 index 13a83393a912..000000000000 --- a/python/tvm/contrib/debugger/__init__.py +++ /dev/null @@ -1,16 +0,0 @@ -# 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. diff --git a/python/tvm/contrib/debugger/debug_executor.py b/python/tvm/contrib/debugger/debug_executor.py deleted file mode 100644 index b0bd46c123b7..000000000000 --- a/python/tvm/contrib/debugger/debug_executor.py +++ /dev/null @@ -1,525 +0,0 @@ -# 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. -"""Graph debug runtime executes TVM debug packed functions.""" - -import logging -import json -import os -import shutil -import struct -import tempfile - -import tvm._ffi -from tvm._ffi.base import string_types -from tvm.contrib import graph_executor -from tvm.runtime.module import BenchmarkResult - -from ...runtime.profiling import Report -from . import debug_result - -_DUMP_ROOT_PREFIX = "tvmdbg_" -_DUMP_PATH_PREFIX = "_tvmdbg_" - - -def create(graph_json_str, libmod, device, dump_root=None): - """Create a runtime executor module given a graph and module. - - Parameters - ---------- - graph_json_str : str - The graph to be deployed in json format output by graph compiler. - The graph can contain operator(tvm_op) that points to the name - of PackedFunc in the libmod. - - libmod : tvm.Module - The module of the corresponding function. - - device : Device - The device to deploy the module, can be local or remote. - - dump_root : str - To select which folder the outputs should be kept. - None will make a temp folder in /tmp/tvmdbg and does the dumping - Returns - ------- - graph_module : GraphModuleDebug - Debug Runtime graph module that can be used to execute the graph. - """ - assert isinstance(graph_json_str, string_types) - - try: - dev, num_rpc_dev, device_type_id = graph_executor.get_device(libmod, device) - if num_rpc_dev == len(dev): - fcreate = dev[0]._rpc_sess.get_function("tvm.graph_executor_debug.create") - else: - fcreate = tvm._ffi.get_global_func("tvm.graph_executor_debug.create") - except ValueError: - raise ValueError( - "Please set '(USE_PROFILER ON)' in " "config.cmake and rebuild TVM to enable debug mode" - ) - func_obj = fcreate(graph_json_str, libmod, *device_type_id) - gmod = GraphModuleDebug(func_obj, dev, graph_json_str, dump_root) - - # Automatically set params if they can be extracted from the libmod - try: - params = libmod["get_graph_params"]() - if isinstance(params, tvm.ir.container.Map): - gmod.set_input(**params) - except (AttributeError, tvm.error.RPCError): - # Params can not be extracted from the libmod and must be set somewhere else manually - # Do not set params during RPC communication - pass - - return gmod - - -class GraphModuleDebug(graph_executor.GraphModule): - """Graph debug runtime module. - - This is a debug wrapper over the TVM runtime. - Runtime interfaces are wrapped with debug functionalities. - Manage the debug framework to format the debug data and - trigger the user interfaces. - - Parameters - ---------- - module : Module - The internal tvm module that holds the actual graph functions. - - device : Device - The device that this module is under. - - graph_json_str : str or graph class - Content of graph json file in string format - - dump_root : str - To select which folder the outputs should be kept. - None will make a temp folder in /tmp/tvmdbg and does the dumping - """ - - def __init__(self, module, device, graph_json_str, dump_root): - self._dump_root = dump_root - self._dump_path = None - self._run_individual = module["run_individual"] - self._run_individual_node = module["run_individual_node"] - self._debug_get_output = module["debug_get_output"] - self._execute_node = module["execute_node"] - self._debug_run_ext_compiler = module["debug_run_ext_compiler"] - self._get_node_output = module["get_node_output"] - self._profile = module["profile"] - self._profile_rpc = module["profile_rpc"] - graph_executor.GraphModule.__init__(self, module) - self._create_debug_env(graph_json_str, device) - - def _format_device(self, device): - return str(device[0]).upper().replace("(", ":").replace(")", "") - - def _ensure_dir(self, directory): - """Create a directory if not exists - - Parameters - ---------- - - directory : str - File path to create - """ - if not os.path.exists(directory): - os.makedirs(directory, 0o700) - - def _get_dump_path(self, device): - """Make the graph and tensor dump folder and return the path. - - Parameters - ---------- - device : Device - The device that this module is under. - - Returns - ------- - path : str - Directory path where the graph and node outputs will be stored. - """ - # save to file - folder_name = _DUMP_PATH_PREFIX + "device_" - folder_name = folder_name + device.replace(":", "_") - path = os.path.join(self._dump_root, folder_name) - self._ensure_dir(path) - return path - - def _remove_dump_root(self): - if os.path.isdir(self._dump_root): - shutil.rmtree(self._dump_root) - - def _create_debug_env(self, graph_json, device): - """Create UI wrapper framework to handle multiple UI frontends for tvmdbg - - Parameters - ---------- - graph_json : json format - json formatted NNVM graph contain list of each node's name, shape and type. - - nodes_list : list - List of all the nodes presented in the graph - - device : Device - The device that this module is under. - """ - # make the dump folder if not given - if not self._dump_root: - self._dump_root = tempfile.mkdtemp(prefix=_DUMP_ROOT_PREFIX) - - # format the device - device = self._format_device(device) - - # updates the dumping directories - self._dump_path = self._get_dump_path(device) - - # init the debug dumping environment - self.debug_datum = debug_result.DebugResult(graph_json, self._dump_path) - - def _execute_next_node(self, node_index, output_index): - """Execute node assuming all previous nodes has been executed. - Return the output of this node. - - Parameters - ---------- - node_index : int - The node index - output_index: int - The node output index - Return - ------ - output_tensors : Array - Array of output tensors - """ - output_tensors = self._execute_next_node_get_output(node_index, output_index) - return output_tensors - - def _run_per_layer(self): - """Execute up to each node and each debug output will be - copied to the buffer. - - """ - output_tensors = [] - for i, node in enumerate(self.debug_datum.get_graph_nodes()): - self._execute_node(i) - num_outputs = self.debug_datum.get_graph_node_output_num(node) - for j in range(num_outputs): - logging.info( - "running node=%d, output_ind=%d, with node_name: %s", i, j, node["name"] - ) - output_tensors.append(self._get_node_output(i, j)) - self.debug_datum.update_output_tensors(output_tensors) - - def _run_external_debug(self): - ext_trace = self._debug_run_ext_compiler() - ext_json = json.loads(ext_trace) - for op in ext_json: - ext_debug = tvm.get_global_func("runtime.ext.debug." + op["compiler"], True) - if isinstance(ext_debug, tvm.runtime.packed_func.PackedFunc): - ext_debug(op["op"], op["dump"], self._dump_path) - - def _run_debug( - self, - number, - repeat, - min_repeat_ms, - limit_zero_time_iterations, - cooldown_interval_ms, - repeats_to_cooldown, - ): - """Execute the node specified with index will be executed. - Each debug output will be copied to the buffer - Time consumed for each execution will be set as debug output. - """ - # Get timing. - self.debug_datum._time_list = self.run_individual( - number=number, - repeat=repeat, - min_repeat_ms=min_repeat_ms, - limit_zero_time_iterations=limit_zero_time_iterations, - cooldown_interval_ms=cooldown_interval_ms, - repeats_to_cooldown=repeats_to_cooldown, - ) - - # Get outputs. - self._run_per_layer() - - # Run external compiler debug if supported - self._run_external_debug() - - def debug_get_output(self, node, out=None): - """Run graph up to node and get the output to out - - Parameters - ---------- - node : int / str - The node index or name - - out : NDArray - The output array container - """ - if isinstance(node, str): - node_index = None - for i, graph_node in enumerate(self.debug_datum.get_graph_nodes()): - if graph_node["name"] == node: - node_index = i - break - else: - raise AttributeError(f"Could not find a node named {node} in this graph.") - elif isinstance(node, int): - node_index = node - else: - raise RuntimeError("Require node index or name only.") - if out: - self._debug_get_output(node_index, out) - return out - return self._debug_get_output(node_index) - - # pylint: disable=arguments-differ - def run( - self, - number=10, - repeat=1, - min_repeat_ms=1, - limit_zero_time_iterations=100, - cooldown_interval_ms=0, - repeats_to_cooldown=1, - sort_by_time=True, - **input_dict, - ): - """Run forward execution of the graph with debug - - Parameters - ---------- - number: int, optional - The number of times to run this function for taking average. - We call these runs as one `repeat` of measurement. - - repeat: int, optional - The number of times to repeat the measurement. - In total, the function will be invoked (1 + number x repeat) times, - where the first one is warm up and will be discarded. - The returned result contains `repeat` costs, - each of which is an average of `number` costs. - - min_repeat_ms: int, optional - The minimum duration of one `repeat` in milliseconds. - By default, one `repeat` contains `number` runs. If this parameter is set, - the parameters `number` will be dynamically adjusted to meet the - minimum duration requirement of one `repeat`. - i.e., When the run time of one `repeat` falls below this time, the `number` parameter - will be automatically increased. - - limit_zero_time_iterations: int, optional - The maximum number of repeats when measured time is equal to 0. - It helps to avoid hanging during measurements. - - cooldown_interval_ms: int, optional - The cooldown interval in milliseconds between the number of repeats defined by - `repeats_to_cooldown`. - - repeats_to_cooldown: int, optional - The number of repeats before the cooldown is activated. - - sort_by_time: bool, optional - Whether to sort the debug output by time. - - input_dict : dict of str to NDArray - List of input values to be feed to - """ - if input_dict: - self.set_input(**input_dict) - - # Step 1. Execute the graph - self._run_debug( - number=number, - repeat=repeat, - min_repeat_ms=min_repeat_ms, - limit_zero_time_iterations=limit_zero_time_iterations, - cooldown_interval_ms=cooldown_interval_ms, - repeats_to_cooldown=repeats_to_cooldown, - ) - # Step 2. Dump the output tensors to the dump folder - self.debug_datum.dump_output_tensor() - # Step 3. Dump the Chrome trace to the dump folder - self.debug_datum.dump_chrome_trace() - # Step 4. Display the collected information - self.debug_datum.display_debug_result(sort_by_time) - - def run_individual( - self, - number, - repeat=1, - min_repeat_ms=0, - limit_zero_time_iterations=100, - cooldown_interval_ms=0, - repeats_to_cooldown=1, - ): - """Run each operation in the graph and get the time per op for all ops. - - number: int - The number of times to run this function for taking average. - We call these runs as one `repeat` of measurement. - - repeat: int, optional - The number of times to repeat the measurement. - In total, the function will be invoked (1 + number x repeat) times, - where the first one is warm up and will be discarded. - The returned result contains `repeat` costs, - each of which is an average of `number` costs. - - min_repeat_ms: int, optional - The minimum duration of one `repeat` in milliseconds. - By default, one `repeat` contains `number` runs. If this parameter is set, - the parameters `number` will be dynamically adjusted to meet the - minimum duration requirement of one `repeat`. - i.e., When the run time of one `repeat` falls below this time, the `number` parameter - will be automatically increased. - - limit_zero_time_iterations: int, optional - The maximum number of repeats when measured time is equal to 0. - It helps to avoid hanging during measurements. - - cooldown_interval_ms: int, optional - The cooldown interval in milliseconds between the number of repeats defined by - `repeats_to_cooldown`. - - repeats_to_cooldown: int, optional - The number of repeats before the cooldown is activated. - - Returns - ------- - A 2-dimensional array where the dimensions are: the index of the operation and - the repeat of the measurement. - """ - res = self._run_individual( - number, - repeat, - min_repeat_ms, - limit_zero_time_iterations, - cooldown_interval_ms, - repeats_to_cooldown, - ) - results = [] - offset = 0 - format_size = "@q" - (nodes_count,) = struct.unpack_from(format_size, res, offset) - offset += struct.calcsize(format_size) - format_data = "@" + repeat * "d" - for _ in range(0, nodes_count): - ret = struct.unpack_from(format_data, res, offset) - offset += struct.calcsize(format_data) - results.append([*ret]) - return results - - def run_individual_node( - self, - index, - number=10, - repeat=1, - min_repeat_ms=0, - limit_zero_time_iterations=100, - cooldown_interval_ms=0, - repeats_to_cooldown=1, - ): - """Benchmark a single node in the serialized graph. - - This does not do any data transfers and uses arrays already on the device. - - Parameters - ---------- - index : int - The index of the node, see `self.debug_datum.get_graph_nodes` - - number: int - The number of times to run this function for taking average. - We call these runs as one `repeat` of measurement. - - repeat: int, optional - The number of times to repeat the measurement. - In total, the function will be invoked (1 + number x repeat) times, - where the first one is warm up and will be discarded. - The returned result contains `repeat` costs, - each of which is an average of `number` costs. - - min_repeat_ms : int, optional - The minimum duration of one `repeat` in milliseconds. - By default, one `repeat` contains `number` runs. If this parameter is set, - the parameters `number` will be dynamically adjusted to meet the - minimum duration requirement of one `repeat`. - i.e., When the run time of one `repeat` falls below this time, the `number` parameter - will be automatically increased. - - limit_zero_time_iterations: int, optional - The maximum number of repeats when measured time is equal to 0. - It helps to avoid hanging during measurements. - - cooldown_interval_ms: int, optional - The cooldown interval in milliseconds between the number of repeats defined by - `repeats_to_cooldown`. - - repeats_to_cooldown: int, optional - The number of repeats before the cooldown is activated. - - Returns - ------- - A module BenchmarkResult - """ - # Results are returned as serialized strings which we deserialize - res = self._run_individual_node( - index, - number, - repeat, - min_repeat_ms, - limit_zero_time_iterations, - cooldown_interval_ms, - repeats_to_cooldown, - ) - fmt = "@" + ("d" * repeat) - results = struct.unpack(fmt, res) - return BenchmarkResult(list(results)) - - def profile(self, collectors=None, **input_dict): - """Run forward execution of the graph and collect overall and per-op - performance metrics. - - Parameters - ---------- - collectors : Optional[Sequence[MetricCollector]] - Extra metrics to collect. If profiling over RPC, collectors must be `None`. - - input_dict : dict of str to NDArray - List of input values to be feed to - - Return - ------ - timing_results : str - Per-operator and whole graph timing results in a table format. - """ - if input_dict: - self.set_input(**input_dict) - - if self.module.type_key == "rpc": - # We cannot serialize MetricCollectors over RPC - assert collectors is None, "Profiling with collectors is not supported over RPC" - return Report.from_json(self._profile_rpc()) - return self._profile(collectors) - - def exit(self): - """Exits the dump folder and all its contents""" - self._remove_dump_root() diff --git a/python/tvm/contrib/debugger/debug_result.py b/python/tvm/contrib/debugger/debug_result.py deleted file mode 100644 index 946afd8a0be3..000000000000 --- a/python/tvm/contrib/debugger/debug_result.py +++ /dev/null @@ -1,301 +0,0 @@ -# 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=pointless-exception-statement, unnecessary-list-index-lookup -"""Graph debug results dumping class.""" -import collections -import json -import os - -import numpy as np -import tvm - -GRAPH_DUMP_FILE_NAME = "_tvmdbg_graph_dump.json" -CHROME_TRACE_FILE_NAME = "_tvmdbg_execution_trace.json" - -ChromeTraceEvent = collections.namedtuple("ChromeTraceEvent", ["ts", "tid", "pid", "name", "ph"]) - - -class DebugResult(object): - """Graph debug data module. - - Data dump module manage all the debug data formatting. - Output data and input graphs are formatted and dumped to file. - Frontend read these data and graph for visualization. - - Parameters - ---------- - graph_json : str - The graph to be deployed in json format output by graph compiler. Each operator (tvm_op) - in the graph will have a one to one mapping with the symbol in libmod which is used - to construct a "PackedFunc" . - - dump_path : str - Output data path is read/provided from frontend - """ - - def __init__(self, graph_json, dump_path): - self._dump_path = dump_path - self._output_tensor_list = [] - self._time_list = [] - json_obj = self._parse_graph(graph_json) - # dump the json information - self._dump_graph_json(json_obj) - - def _parse_graph(self, graph_json): - """Parse and extract the JSON graph and update the nodes, shapes and dltype. - - Parameters - ---------- - graph_json : str or graph class - The graph to be deployed in json format output by JSON graph. - """ - json_obj = json.loads(graph_json) - self._nodes_list = json_obj["nodes"] - self._shapes_list = json_obj["attrs"]["shape"] - self._dtype_list = json_obj["attrs"]["dltype"] - self._update_graph_json() - return json_obj - - def _update_graph_json(self): - """update the nodes_list with name, shape and data type, - for temporarily storing the output. - """ - eid = 0 - for node in self._nodes_list: - input_list = [] - if node["op"] == "null": - node["attrs"] = {} - node["op"] = "param" - num_outputs = 1 - elif node["op"] == "tvm_op": - for input_node in node["inputs"]: - input_list.append(self._nodes_list[input_node[0]]["name"]) - node["op"] = node["attrs"]["func_name"] - num_outputs = int(node["attrs"]["num_outputs"]) - else: - raise ValueError("") - node["inputs"] = input_list - dtype = str("type: " + self._dtype_list[1][eid]) - node["attrs"].update({"T": dtype}) - node["shape"] = self._shapes_list[1][eid] - eid += num_outputs - - def _cleanup_tensors(self): - """Remove the tensor dump file (graph wont be removed)""" - for filename in os.listdir(self._dump_path): - if os.path.isfile(filename) and not filename.endswith(".json"): - os.remove(filename) - - def get_graph_nodes(self): - """Return the nodes list""" - return self._nodes_list - - def get_graph_node_shapes(self): - """Return the nodes shapes list""" - return self._shapes_list - - def get_graph_node_output_num(self, node): - """Return the number of outputs of a node""" - return 1 if node["op"] == "param" else int(node["attrs"]["num_outputs"]) - - def get_graph_node_dtypes(self): - """Return the nodes dtype list""" - return self._dtype_list - - def get_output_tensors(self): - """Get the output tensors of each operation in numpy format""" - eid = 0 - output_tensors = {} - for i, node in enumerate(self._nodes_list): - num_outputs = self.get_graph_node_output_num(node) - for j in range(num_outputs): - - # the node name is not unique, so we need a consistent - # indexing based on the list ordering in the nodes - key = f"{node['name']}____topo-index:{i}____output-num:{j}" - output_tensors[key] = self._output_tensor_list[eid] - eid += 1 - return output_tensors - - def update_output_tensors(self, tensors): - """Update output tensors list - - Parameters - ---------- - tensors : list[NDArray] - """ - if not isinstance(tensors, list): - AttributeError("tensors with incorrect type.") - - for output_array in tensors: - self._output_tensor_list.append(output_array) - - def dump_output_tensor(self): - """Dump the outputs to a temporary folder, the tensors are in numpy format""" - # cleanup existing tensors before dumping - self._cleanup_tensors() - output_tensors = self.get_output_tensors() - - np_tensors = {} - for key, val in output_tensors.items(): - np_tensors[key] = val.asnumpy() - np.savez(os.path.join(self._dump_path, "output_tensors.npz"), **np_tensors) - with open(os.path.join(self._dump_path, "output_tensors.params"), "wb") as param_f: - param_f.write(save_tensors(output_tensors)) - - def dump_chrome_trace(self): - """Dump the trace to the Chrome trace.json format.""" - - def s_to_us(t): - return t * 10**6 - - starting_times = np.zeros(len(self._time_list) + 1) - starting_times[1:] = np.cumsum([np.mean(times) for times in self._time_list]) - - def node_to_events(node, times, starting_time): - return [ - ChromeTraceEvent( - ts=s_to_us(starting_time), - tid=1, - pid=1, - ph="B", - name=node["name"], - ), - ChromeTraceEvent( - # Use start + duration instead of end to ensure precise timings. - ts=s_to_us(np.mean(times) + starting_time), - tid=1, - pid=1, - ph="E", - name=node["name"], - ), - ] - - events = [ - e - for (node, times, starting_time) in zip( - self._nodes_list, self._time_list, starting_times - ) - for e in node_to_events(node, times, starting_time) - ] - result = dict(displayTimeUnit="ns", traceEvents=[e._asdict() for e in events]) - - with open(os.path.join(self._dump_path, CHROME_TRACE_FILE_NAME), "w") as trace_f: - json.dump(result, trace_f) - - def _dump_graph_json(self, graph): - """Dump json formatted graph. - - Parameters - ---------- - graph : json format - json formatted JSON graph contain list of each node's - name, shape and type. - """ - graph_dump_file_name = GRAPH_DUMP_FILE_NAME - with open(os.path.join(self._dump_path, graph_dump_file_name), "w") as outfile: - json.dump(graph, outfile, indent=4, sort_keys=False) - - def get_debug_result(self, sort_by_time=True): - """Return the debugger result""" - header = [ - "Node Name", - "Ops", - "Time(us)", - "Time(%)", - "Shape", - "Inputs", - "Outputs", - "Measurements(us)", - ] - lines = [ - "---------", - "---", - "--------", - "-------", - "-----", - "------", - "-------", - "----------------", - ] - eid = 0 - data = [] - total_time = sum([np.mean(time) for time in self._time_list]) - for node, time in zip(self._nodes_list, self._time_list): - time_mean = np.mean(time) - num_outputs = self.get_graph_node_output_num(node) - for j in range(num_outputs): - op = node["op"] - if node["op"] == "param": - eid += 1 - continue - name = node["name"] - shape = str(self._output_tensor_list[eid].shape) - time_us = round(time_mean * 1e6, 3) - time_percent = round(((time_mean / total_time) * 100), 3) - inputs = str(node["attrs"]["num_inputs"]) - outputs = str(node["attrs"]["num_outputs"]) - measurements = str([round(repeat_data * 1e6, 3) for repeat_data in time]) - node_data = [name, op, time_us, time_percent, shape, inputs, outputs, measurements] - data.append(node_data) - eid += 1 - - if sort_by_time: - # Sort on the basis of execution time. Prints the most expensive ops in the start. - data = sorted(data, key=lambda x: x[2], reverse=True) - # Insert a row for total time at the end. - rounded_total_time_us = round(total_time * 1e6, 3) - data.append(["Total_time", "-", rounded_total_time_us, "-", "-", "-", "-", "-", "-"]) - - fmt = "" - for i, _ in enumerate(header): - max_len = len(header[i]) - for j, _ in enumerate(data): - item_len = len(str(data[j][i])) - if item_len > max_len: - max_len = item_len - fmt = fmt + "{:<" + str(max_len + 2) + "}" - log = [fmt.format(*header)] - log.append(fmt.format(*lines)) - for row in data: - log.append(fmt.format(*row)) - return "\n".join(log) - - def display_debug_result(self, sort_by_time=True): - """Displays the debugger result""" - print(self.get_debug_result(sort_by_time)) - - -def save_tensors(params): - """Save parameter dictionary to binary bytes. - - The result binary bytes can be loaded by the - GraphModule with API "load_params". - - Parameters - ---------- - params : dict of str to NDArray - The parameter dictionary. - - Returns - ------- - param_bytes: bytearray - Serialized parameters. - """ - _save_tensors = tvm.get_global_func("tvm.relay._save_param_dict") - - return _save_tensors(params) diff --git a/python/tvm/contrib/debugger/debug_runtime.py b/python/tvm/contrib/debugger/debug_runtime.py deleted file mode 100644 index ebd903b47570..000000000000 --- a/python/tvm/contrib/debugger/debug_runtime.py +++ /dev/null @@ -1,29 +0,0 @@ -# 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. -"""Deprecated Python API for DebugExecutor.""" - -import warnings - -from . import debug_executor - - -def create(*args, **kwargs): - warnings.warn( - "This function has been moved to tvm.contrib.graph_executor and will be removed " - "in the next TVM release" - ) - return debug_executor.create(*args, **kwargs) diff --git a/python/tvm/contrib/graph_runtime.py b/python/tvm/contrib/graph_runtime.py deleted file mode 100644 index f8ecfdd70a5b..000000000000 --- a/python/tvm/contrib/graph_runtime.py +++ /dev/null @@ -1,29 +0,0 @@ -# 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. -"""Deprecated Python API for GraphExecutor.""" - -import warnings - -from . import graph_executor - - -def create(*args, **kwargs): - warnings.warn( - "This function has been moved to tvm.contrib.graph_executor and will be removed " - "in the next TVM release" - ) - return graph_executor.create(*args, **kwargs) diff --git a/python/tvm/contrib/hexagon/hexagon_profiler.py b/python/tvm/contrib/hexagon/hexagon_profiler.py index c1ab2dd8aea9..59711d28b760 100644 --- a/python/tvm/contrib/hexagon/hexagon_profiler.py +++ b/python/tvm/contrib/hexagon/hexagon_profiler.py @@ -20,11 +20,8 @@ import os import subprocess -import typing from tvm.ir.transform import PassContext from tvm.contrib.hexagon.profiling.process_lwp_data import process_lwp_output -from tvm.relay.backend.executor_factory import ExecutorFactoryModule -from tvm.driver.build_module import OperatorModule from tvm.contrib import utils @@ -34,7 +31,7 @@ class HexagonProfiler: def __init__( self, dso_binary: str, - module: typing.Union[ExecutorFactoryModule, OperatorModule], + module, hexagon_server_process, enable_debug, ): @@ -42,10 +39,7 @@ def __init__( # Save test .so to process profiling data self._temp_dir = utils.tempdir(keep_for_debug=enable_debug) self._dso_binary_path = self._temp_dir.relpath(dso_binary) - if isinstance(module, OperatorModule): - module.save(self._dso_binary_path) - else: - module.get_lib().save(self._dso_binary_path) + module.save(self._dso_binary_path) self._android_serial_number = os.environ.get("ANDROID_SERIAL_NUMBER") self._remote_path = "" diff --git a/python/tvm/ir/__init__.py b/python/tvm/ir/__init__.py index e7376f4c1f0d..3e893099f454 100644 --- a/python/tvm/ir/__init__.py +++ b/python/tvm/ir/__init__.py @@ -36,30 +36,15 @@ from .expr import BaseExpr, GlobalVar, PrimExpr, Range, RelayExpr from .function import BaseFunc, CallingConv from .global_info import GlobalInfo, DummyGlobalInfo, VDevice -from .memory_pools import ( - ConstantMemoryPools, - ConstantPoolInfo, - PoolInfo, - PoolInfoProperties, - WorkspaceMemoryPools, - WorkspacePoolInfo, -) from .module import IRModule from .op import Op, register_intrin_lowering, register_op_attr from .tensor_type import TensorType from .type import ( FuncType, - GlobalTypeVar, - IncompleteType, PointerType, PrimType, - RelayRefType, TupleType, Type, - TypeConstraint, - TypeKind, - TypeVar, ) -from .type_relation import TypeCall, TypeRelation from . import analysis diff --git a/python/tvm/ir/memory_pools.py b/python/tvm/ir/memory_pools.py deleted file mode 100644 index 37b903c7bb00..000000000000 --- a/python/tvm/ir/memory_pools.py +++ /dev/null @@ -1,268 +0,0 @@ -# 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. -"""Objects for Memory Pools to be used within the compilation""" - -from typing import Optional, List - -from tvm._ffi import register_object -from tvm.runtime import Object -from tvm.runtime import NDArray -from . import _ffi_api - - -@register_object("ir.PoolInfo") -class PoolInfo(Object): - """PoolInfo object holds information related to memory pools - where the statically sized allocate nodes will pooled into. - This is a base class for WorkspacePoolInfo and ConstantPoolInfo. - """ - - def __init__(self): - pass - - -@register_object("ir.PoolInfoProperties") -class PoolInfoProperties(Object): - """PoolInfo object holds information related to memory pools - where the statically sized allocate nodes will pooled into. - - Parameters - ---------- - size_hint_bytes : Optional[int] - The expected size hint to be used by the allocator. - The default value would be -1 which means the pool - is not size restricted. - - clock_frequency_hz : Optional[int] - The clock frequency that the memory pool runs at in Hz. - If not specified/known, this will default to -1 indicating - it hasn't been defined. - - read_bandwidth_bytes_per_cycle : Optional[int] - The read bandwidth of the memory pool in bytes/cycle. - If not specified/known, this will default to -1 indicating - it hasn't been defined. - - write_bandwidth_bytes_per_cycle : Optional[int] - The write bandwidth of the memory pool in bytes/cycle. - If not specified/known, this will default to -1 indicating - it hasn't been defined. - - read_latency_cycles : Optional[int] - The read latency of the memory pool in cycles. - If not specified/known, this will default to 0. - - write_latency_cycles : Optional[int] - The write latency of the memory pool in cycles. - If not specified/known, this will default to 0. - - target_burst_bytes : Optional[Union[Dict[Target, int], None]] - The burst length of the memory pool in bytes per target. - If not specified/known for a given target, a burst length - of 1 byte will be assumed. - - """ - - def __init__( - self, - size_hint_bytes: Optional[int] = -1, - clock_frequency_hz: Optional[int] = -1, - read_bandwidth_bytes_per_cycle: Optional[int] = -1, - write_bandwidth_bytes_per_cycle: Optional[int] = -1, - read_latency_cycles: Optional[int] = 0, - write_latency_cycles: Optional[int] = 0, - target_burst_bytes=None, - ): - if not target_burst_bytes: - target_burst_bytes = dict() - - self.__init_handle_by_constructor__( - _ffi_api.PoolInfoProperties, # type: ignore # pylint: disable=no-member - size_hint_bytes, - clock_frequency_hz, - read_bandwidth_bytes_per_cycle, - write_bandwidth_bytes_per_cycle, - read_latency_cycles, - write_latency_cycles, - target_burst_bytes, - ) - - -@register_object("ir.ConstantInfo") -class ConstantInfo(Object): - """ConstantInfo object hold information on a constant pool. - - Parameters - ---------- - name_hint : str - Name of the constant. - byte_offset : int - The byte_offset of the constant. - data : NDArray - The data of the constant. - """ - - def __init__( - self, - name_hint: str, - byte_offset: int, - data: NDArray, - ): - self.__init_handle_by_constructor__( - _ffi_api.ConstantInfo, # type: ignore # pylint: disable=no-member - name_hint, - byte_offset, - data, - ) - - -@register_object("ir.WorkspacePoolInfo") -class WorkspacePoolInfo(PoolInfo): - """WorkspacePoolInfo object holds information related to RW memory pools - where the statically sized allocate nodes will pooled into. - - Parameters - ---------- - pool_name : str - The name of the memory pool - - targets : list[Target] - A list of targets which could access the pool - - pool_info_properties : PoolInfoProperties - The properties of the pool. - """ - - def __init__( - self, - pool_name: str, - targets, - pool_info_properties=None, - ): - super().__init__() - - if pool_info_properties is None: - pool_info_properties = PoolInfoProperties() - - self.__init_handle_by_constructor__( - _ffi_api.WorkspacePoolInfo, # type: ignore # pylint: disable=no-member - pool_name, - targets, - pool_info_properties, - ) - - -@register_object("ir.ConstantPoolInfo") -class ConstantPoolInfo(PoolInfo): - """ConstantPoolInfo object holds information related to RO memory pools - where the statically sized allocate nodes are pooled into. - - Parameters - ---------- - pool_name : str - The name of the memory pool - - targets : list[Target] - describes which targets could access the pool - - pool_info_properties : PoolInfoProperties - The properties of the pool. - """ - - def __init__( - self, - pool_name: str, - targets, # list[Target] - constant_info_arr=None, # list[ConstantInfo] - pool_info_properties=None, - ): - super().__init__() - - if constant_info_arr is None: - constant_info_arr = [] - if pool_info_properties is None: - pool_info_properties = PoolInfoProperties() - self.__init_handle_by_constructor__( - _ffi_api.ConstantPoolInfo, # type: ignore # pylint: disable=no-member - pool_name, - targets, - constant_info_arr, - pool_info_properties, - ) - - -@register_object("ir.WorkspaceMemoryPools") -class WorkspaceMemoryPools(Object): - """This object contains a list of WorkspacePoolInfo objects to be used as - workspace memory in the compilation - - Parameters - ---------- - pools : List[WorkspacePoolInfo] - The list of ConstantPoolInfo objects to be used with the compilation - """ - - def __init__( - self, - pools: List[WorkspacePoolInfo], - ): - self.__init_handle_by_constructor__( - _ffi_api.WorkspaceMemoryPools, pools # type: ignore # pylint: disable=no-member - ) - - -@register_object("ir.ConstantMemoryPools") -class ConstantMemoryPools(Object): - """This object contains a list of ConstantPoolInfo objects to be used as - read-only memory in the compilation - - Parameters - ---------- - pools : List[ConstantPoolInfo] - The list of ConstantPoolInfo objects to be used with the compilation - """ - - def __init__( - self, - pools: List[ConstantPoolInfo], - ): - self.__init_handle_by_constructor__( - _ffi_api.ConstantMemoryPools, pools # type: ignore # pylint: disable=no-member - ) - - -@register_object("ir.AllocatedPoolInfo") -class AllocatedPoolInfo(Object): - """Allocate memory in a given pool. - - Parameters - ---------- - pool : PoolInfo - The pool in which to allocate memory. - allocated_size : int - The size of memory to allocate. - """ - - def __init__( - self, - pool: PoolInfo, - allocated_size: int, - pool_var_idx: int = 0, - ): - self.__init_handle_by_constructor__( - _ffi_api.AllocatedPoolInfo, pool, allocated_size, pool_var_idx # type: ignore # pylint: disable=no-member - ) diff --git a/python/tvm/ir/module.py b/python/tvm/ir/module.py index 9ee87c5224e6..ed7ea4439957 100644 --- a/python/tvm/ir/module.py +++ b/python/tvm/ir/module.py @@ -27,7 +27,6 @@ from . import _ffi_api from . import expr as _expr -from . import type as _ty from .attrs import DictAttrs from .base import Node @@ -44,7 +43,7 @@ class IRModule(Node, Scriptable): Map of global var to BaseFunc """ - def __init__(self, functions=None, type_definitions=None, attrs=None, global_infos=None): + def __init__(self, functions=None, attrs=None, global_infos=None): if functions is None: functions = {} elif isinstance(functions, dict): @@ -56,17 +55,6 @@ def __init__(self, functions=None, type_definitions=None, attrs=None, global_inf raise TypeError("Expect functions to be Dict[GlobalVar, Function]") mapped_funcs[k] = v functions = mapped_funcs - if type_definitions is None: - type_definitions = {} - elif isinstance(type_definitions, dict): - mapped_type_defs = {} - for k, v in type_definitions.items(): - if isinstance(k, string_types): - k = _ty.GlobalTypeVar(k) - if not isinstance(k, _ty.GlobalTypeVar): - raise TypeError("Expect type_definitions to be Dict[GlobalTypeVar, Type]") - mapped_type_defs[k] = v - type_definitions = mapped_type_defs attrs = None if not attrs else attrs if attrs is not None: @@ -76,7 +64,6 @@ def __init__(self, functions=None, type_definitions=None, attrs=None, global_inf self.__init_handle_by_constructor__( _ffi_api.IRModule, functions, - type_definitions, attrs, global_infos, ) @@ -117,11 +104,6 @@ def _add(self, var, val, update=True): else: var = _expr.GlobalVar(var) _ffi_api.Module_Add(self, var, val, update) - else: - assert isinstance(val, _ty.Type) - if isinstance(var, string_types): - var = _ty.GlobalTypeVar(var) - _ffi_api.Module_AddDef(self, var, val, update) def __getitem__(self, var): """Lookup a global definition by name or by variable. @@ -138,9 +120,8 @@ def __getitem__(self, var): """ if isinstance(var, string_types): return _ffi_api.Module_Lookup_str(self, var) - if isinstance(var, _expr.GlobalVar): - return _ffi_api.Module_Lookup(self, var) - return _ffi_api.Module_LookupDef(self, var) + assert isinstance(var, _expr.GlobalVar) + return _ffi_api.Module_Lookup(self, var) def __delitem__(self, var: Union[str, _expr.GlobalVar]): _ffi_api.Module_Remove(self, var) @@ -244,61 +225,8 @@ def replace_global_vars( """ return _ffi_api.Module_ReplaceGlobalVars(self, replacements) - def get_global_type_vars(self): - """Collect all global type vars defined in this module. - - Returns - ------- - global_type_vars: Array[GlobalTypeVar] - An array of global type vars. - """ - return _ffi_api.Module_GetGlobalTypeVars(self) - - def get_global_type_var(self, name): - """Get a global type variable in the function by name. - - Parameters - ---------- - name: str - The name of the global type variable. - - Returns - ------- - global_type_var: GlobalTypeVar - The global variable mapped to :code:`name`. - - Raises - ------ - tvm.error.TVMError if we cannot find corresponding global type var. - """ - return _ffi_api.Module_GetGlobalTypeVar(self, name) - - def get_constructor(self, tag): - """Look up an ADT constructor by tag. - - Parameters - ---------- - tag: int - The tag for a constructor. - - Returns - ------- - constructor: Constructor - The constructor associated with the given tag, - - Raises - ------ - tvm.error.TVMError if the corresponding constructor cannot be found. - """ - return _ffi_api.Module_LookupTag(self, tag) - - def get_type(self, name): - ty_var = self.get_global_type_var(name) - ty_data = self.type_definitions[ty_var] - return tuple([ty_var] + list(ty_data.constructors)) - @staticmethod - def from_expr(expr, functions=None, type_defs=None): + def from_expr(expr, functions=None): """Construct a module from a standalone expression. Parameters @@ -309,9 +237,6 @@ def from_expr(expr, functions=None, type_defs=None): global_funcs: Optional[dict] Map of global vars to function definitions - type_defs: Optional[dict] - Map of global type vars to type definitions - Returns ------- mod: Module @@ -320,16 +245,7 @@ def from_expr(expr, functions=None, type_defs=None): (wrapped in a function if necessary) """ funcs = functions if functions is not None else {} - defs = type_defs if type_defs is not None else {} - return _ffi_api.Module_FromExpr(expr, funcs, defs) - - def _import(self, file_to_import): - return _ffi_api.Module_Import(self, file_to_import) - - def import_from_std(self, file_to_import): - # TODO(@jroesch): clean up prelude - _ffi_api.Module_ImportFromStd(self, file_to_import) - return tvm.relay.transform.InferType()(self) + return _ffi_api.Module_FromExpr(expr, funcs) def get_attr(self, attr_key): """Get the IRModule attribute. diff --git a/python/tvm/ir/op.py b/python/tvm/ir/op.py index 3ab5bb55c051..dae97f114b6e 100644 --- a/python/tvm/ir/op.py +++ b/python/tvm/ir/op.py @@ -101,34 +101,6 @@ def reset_attr(self, attr_name): """ _ffi_api.OpResetAttr(self, attr_name) - def add_type_rel(self, rel_name, type_rel_func=None): - """Attach the type function corresponding to the return type. - - Parameters - ---------- - rel_name : str - The type relation name to register. - - type_rel_func : Optional[function (args: List[Type], attrs: Attrs) -> Type] - The backing relation function which can solve an arbitrary relation on variables. - Differences with type_rel_func in C++: - - 1) When type_rel_func is not None - - a) OpAddTypeRel on C++ side will adjust type_rel_func with TypeReporter to - calling convention of relay type system. - - b) type_rel_func returns output argument's type, return None means can't - infer output's type. - - c) only support single output operators for now, the last argument is output tensor. - - 2) when type_rel_func is None, will call predefined type_rel_funcs in relay - according to ``tvm.relay.type_relation.`` + rel_name. - - """ - _ffi_api.OpAddTypeRel(self, rel_name, type_rel_func) - def add_argument(self, name, type, description): # pylint: disable=redefined-builtin """Add arguments information to the function. diff --git a/python/tvm/ir/type.py b/python/tvm/ir/type.py index c83cef3f6cea..8ecf1ffb4a0f 100644 --- a/python/tvm/ir/type.py +++ b/python/tvm/ir/type.py @@ -15,8 +15,6 @@ # specific language governing permissions and limitations # under the License. """Unified type system in the project.""" -from enum import IntEnum - import tvm import tvm._ffi from tvm.runtime import Scriptable @@ -40,17 +38,6 @@ def same_as(self, other): return super().__eq__(other) -class TypeKind(IntEnum): - """Possible kinds of TypeVars.""" - - Type = 0 - ShapeVar = 1 - BaseType = 2 - Constraint = 4 - AdtHandle = 5 - TypeData = 6 - - @tvm._ffi.register_object("PrimType") class PrimType(Type): """Primitive data type in the low level IR @@ -82,82 +69,6 @@ def __init__(self, element_type, storage_scope=""): self.__init_handle_by_constructor__(_ffi_api.PointerType, element_type, storage_scope) -@tvm._ffi.register_object("TypeVar") -class TypeVar(Type): - """Type parameter in functions. - - A type variable represents a type placeholder which will - be filled in later on. This allows the user to write - functions which are generic over types. - - Parameters - ---------- - name_hint: str - The name of the type variable. This name only acts as a hint, and - is not used for equality. - - kind : Optional[TypeKind] - The kind of the type parameter. - """ - - def __init__(self, name_hint, kind=TypeKind.Type): - self.__init_handle_by_constructor__(_ffi_api.TypeVar, name_hint, kind) - - def __call__(self, *args): - """Create a type call from this type. - - Parameters - ---------- - args: List[Type] - The arguments to the type call. - - Returns - ------- - call: Type - The result type call. - """ - # pylint: disable=import-outside-toplevel - from .type_relation import TypeCall - - return TypeCall(self, args) - - -@tvm._ffi.register_object("GlobalTypeVar") -class GlobalTypeVar(Type): - """A global type variable that is used for defining new types or type aliases. - - Parameters - ---------- - name_hint: str - The name of the type variable. This name only acts as a hint, and - is not used for equality. - - kind : Optional[TypeKind] - The kind of the type parameter. - """ - - def __init__(self, name_hint, kind=TypeKind.AdtHandle): - self.__init_handle_by_constructor__(_ffi_api.GlobalTypeVar, name_hint, kind) - - def __call__(self, *args): - """Create a type call from this type. - - Parameters - ---------- - args: List[Type] - The arguments to the type call. - - Returns - ------- - call: Type - The result type call. - """ - # pylint: disable=import-outside-toplevel - from .type_relation import TypeCall - - return TypeCall(self, args) - - @tvm._ffi.register_object("TupleType") class TupleType(Type): """The type of tuple values. @@ -172,11 +83,6 @@ def __init__(self, fields): self.__init_handle_by_constructor__(_ffi_api.TupleType, fields) -@tvm._ffi.register_object("TypeConstraint") -class TypeConstraint(Type): - """Abstract class representing a type constraint.""" - - @tvm._ffi.register_object("FuncType") class FuncType(Type): """Function type. @@ -186,9 +92,6 @@ class FuncType(Type): a set of type constraints which we omit for the time being, a sequence of argument types, and a return type. - We can informally write them as: - `forall (type_params), (arg_types) -> ret_type where type_constraints` - Parameters ---------- arg_types : List[tvm.relay.Type] @@ -196,45 +99,11 @@ class FuncType(Type): ret_type : tvm.relay.Type The return type. - - type_params : Optional[List[tvm.relay.TypeVar]] - The type parameters - - type_constraints : Optional[List[tvm.relay.TypeConstraint]] - The type constraints. """ - def __init__(self, arg_types, ret_type, type_params=None, type_constraints=None): - if type_params is None: - type_params = [] - if type_constraints is None: - type_constraints = [] + def __init__(self, arg_types, ret_type): self.__init_handle_by_constructor__( - _ffi_api.FuncType, arg_types, ret_type, type_params, type_constraints + _ffi_api.FuncType, + arg_types, + ret_type, ) - - -@tvm._ffi.register_object("IncompleteType") -class IncompleteType(Type): - """Incomplete type during type inference. - - kind : Optional[TypeKind] - The kind of the incomplete type. - """ - - def __init__(self, kind=TypeKind.Type): - self.__init_handle_by_constructor__(_ffi_api.IncompleteType, kind) - - -@tvm._ffi.register_object("relay.RefType") -class RelayRefType(Type): - """Reference Type in relay. - - Parameters - ---------- - value: Type - The value type. - """ - - def __init__(self, value): - self.__init_handle_by_constructor__(_ffi_api.RelayRefType, value) diff --git a/python/tvm/meta_schedule/testing/torchbench/__init__.py b/python/tvm/meta_schedule/testing/torchbench/__init__.py deleted file mode 100644 index 13a83393a912..000000000000 --- a/python/tvm/meta_schedule/testing/torchbench/__init__.py +++ /dev/null @@ -1,16 +0,0 @@ -# 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. diff --git a/python/tvm/meta_schedule/testing/torchbench/run.py b/python/tvm/meta_schedule/testing/torchbench/run.py deleted file mode 100644 index cd50d180446f..000000000000 --- a/python/tvm/meta_schedule/testing/torchbench/run.py +++ /dev/null @@ -1,791 +0,0 @@ -# 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=unused-variable -""" -This script is for benchmarking TVM performance on models from TorchBench. -It uses the TorchDynamo as the frontend to ingest models into TVM, and it also -leverages the benchmark util from TorchDynamo. - -TorchDynamo (https://github.com/pytorch/torchdynamo) and TorchBench -(https://github.com/pytorch/benchmark) need to be in the parent directory of TVM. -We need a local clone of these repos because torchbench and the benchmark runner -in TorchDynamo isn't designed to be used as a Python package. - -To setup the environment, run the following commands in the parent directory of TVM and with -the appropriate Python environment: -```bash -# torchdynamo requires nightly pytorch. If it fails to find the specified version, try -# installing the latest nightly pytorch. -pip3 install --pre \ - --extra-index-url https://download.pytorch.org/whl/nightly/cu116 \ - torch==1.13.0.dev20220926 \ - torchvision==0.14.0.dev20220926 \ - torchtext==0.14.0.dev20220926 - -git clone https://github.com/pytorch/torchdynamo -pushd torchdynamo -git checkout c537639f9712621dc04ca09908796dbbe86c354b -pip install -e . -popd - -sudo apt install git-lfs # git lfs is used for TorchBench -git clone https://github.com/pytorch/benchmark -pushd benchmark -python install.py --continue_on_fail # fambench_xlmr might fail to install -popd -``` - -To run a benchmark, the script can be run under 'tune' mode by -```bash -python python/tvm/meta_schedule/testing/torchbench/run.py \ - --mode tune \ - --model resnet50 \ - --target "nvidia/geforce-rtx-3070" \ - --work-dir /path/to/work/dir/ \ - --num-trials 20000 \ - --rpc-host \ - --rpc-port \ - --rpc-key \ -``` - -All available target tags (like nvidia/geforce-rtx-3070) can be found at -https://github.com/apache/tvm/blob/main/src/target/tag.cc - -Then the script can be run under 'eval' mode to actual benchmark the performance, -using the tuning database under the work directory. This can be executed on a different -machine than the one executes tuning (the database json files need to be inside -of the work directory). -```bash -python python/tvm/meta_schedule/testing/torchbench/run.py \ - --mode eval \ - --model resnet50 \ - --target "nvidia/geforce-rtx-3070" \ - --work-dir /path/to/work/dir/ \ - --num-trials 0 -``` - -Alternatively, both tuning and evaluation can be done in a single run on the same machine, -by -```bash -python python/tvm/meta_schedule/testing/torchbench/run.py \ - --mode all \ - --model resnet50 \ - --target "llvm -num-cores 6" \ - --work-dir /path/to/work/dir/ \ - --num-trials 0 -``` -""" -# pylint: disable=logging-format-interpolation - -import argparse -import contextlib -import logging -import os -import pickle -import sys -import warnings -from collections import defaultdict -from enum import Enum -from typing import Callable, Dict, List, Tuple - -import numpy as np # type: ignore -import torch # type: ignore -from scipy.stats import ttest_ind # type: ignore - -import tvm -import tvm.relay -from tvm import meta_schedule as ms -from tvm._ffi import get_global_func -from tvm.contrib.graph_executor import GraphModule -from tvm.meta_schedule.testing.torchbench.utils import ( - DisallowedOperator, - load_torchdynamo_benchmark_runner, - same, - timed, -) -from tvm.runtime.vm import VirtualMachine -from tvm.support import describe - -# Needs to be imported after the .utils is executed -import torchdynamo # type: ignore # isort: skip, pylint: disable=wrong-import-order - - -class RunMode(Enum): - """ - The running mode of this script. Available values are: - - extract: Only import the model and extract tuning tasks from it. - - tune: Only tune the tasks and create the tuning database. - - eval: Only benchmark model using pre-existing tuning database. - - all: Run both tuning and benchmark - """ - - ALL = "all" - EXTRACT = "extract" - TUNE = "tune" - EVAL = "eval" - - @property - def should_extract(self): - """ - Returns whether it should extract tuning tasks. - """ - return self in (RunMode.ALL, RunMode.EXTRACT) - - @property - def should_tune(self): - """ - Returns whether it should tune the tasks. - """ - return self in (RunMode.ALL, RunMode.TUNE) - - @property - def should_eval(self): - """ - Returns whether it should actually benchmark the model. - """ - return self in (RunMode.ALL, RunMode.EVAL) - - -class ResultComparisonMetric(Enum): - """ - This changes how it compares the results with the expected value during - accuracy check. - - cosine: Use the cosine similarity. It should be greater than 0.99. - - allclose-1e-4: Use the max elementwise absolute difference. It should be less than 1e-4. - """ - - COSINE = "cosine" - ALLCLOSE = "allclose-1e-4" - - -def parse_args(): - """ - Parse arguments - """ - args = argparse.ArgumentParser() - - args.add_argument( - "--mode", - type=RunMode, - default=RunMode.ALL, - help=RunMode.__doc__, - ) - args.add_argument( - "--batch-size", - type=int, - default=None, - help="The batch size of model input. Use TorchBench's default value if not specified.", - ) - args.add_argument( - "--result-metric", - type=ResultComparisonMetric, - default=ResultComparisonMetric.ALLCLOSE, - help=ResultComparisonMetric.__doc__, - ) - args.add_argument( - "--benchmark-repeat", - type=int, - default=10, - help="The number of times to repeat the benchmark measurement.", - ) - args.add_argument( - "--benchmark-warmup-rounds", - type=int, - default=5, - help="The number of rounds to warmup before starting to measure the performance.", - ) - args.add_argument( - "--disallowed-op", - type=str, - default="all", - help=DisallowedOperator.__doc__, - ) - - # Model selection - args.add_argument( - "--model", - type=str, - required=True, - help=""" - The name of model to run. It should a directory name under - https://github.com/pytorch/benchmark/tree/main/torchbenchmark/models. - """, - ) - args.add_argument( - "--float32", - action="store_true", - help=""" - Cast model and inputs to fp32 - """, - ) - - # Tuning-related config - args.add_argument( - "--target", - type=tvm.target.Target, - required=True, - help="The target to tune and run benchmark for.", - ) - args.add_argument( - "--work-dir", - type=str, - required=True, - help=""" - The working directory to save intermediate results and store databases for compilation. - """, - ) - args.add_argument( - "--strategy", - type=str, - default="evolutionary", - help="The search strategy used by MetaSchdule.", - ) - args.add_argument( - "--num-trials", - type=int, - required=True, - help="The max number of trials to run MetaSchedule.", - ) - args.add_argument( - "--max-trials-per-task", - type=int, - default=None, - help=""" - The max number of trials to run per task extracted in MetaSchedule. - By default it's the same as --num-trials. - """, - ) - args.add_argument( - "--backend", - type=str, - choices=["graph", "vm"], - default="graph", - help="The backend to use for relay compilation(graph / vm).", - ) - # TODO(@yelite): Add a layout arg to transform the network after - # ingesting into Relay and before feeding into MetaSchedule. - - # Evaluator-related config - args.add_argument( - "--number", - type=int, - default=3, - help="The number of times to run the model for taking average in a single measurement.", - ) - args.add_argument( - "--repeat", - type=int, - default=1, - help="The number of times to repeat the measurement.", - ) - args.add_argument( - "--min-repeat-ms", - type=int, - default=100, - help=""" - Minimum repeat time in ms. The number of runs will be increased if the actual - repeat time is lowered than this. - """, - ) - args.add_argument( - "--adaptive-training", - action="store_true", - help="Whether to use adaptive training for cost model.", - ) - args.add_argument( - "--cpu-flush", - action="store_true", - help="Whether to perform CPU cache flush.", - ) - - # RPC-related args - args.add_argument( - "--rpc-host", - type=str, - help="Host of the RPC Tracker for tuning. Use LocalRunner if not provided", - ) - args.add_argument( - "--rpc-port", - type=int, - help="Port of the RPC Tracker for tuning", - ) - args.add_argument( - "--rpc-key", - type=str, - help="Key of the RPC Tracker for tuning", - ) - - parsed = args.parse_args() - - if parsed.disallowed_op == "all": - disallowed_op = set(DisallowedOperator) - else: - disallowed_op = {DisallowedOperator(v) for v in parsed.disallowed_op.split(",")} - parsed.disallowed_op = disallowed_op - - # Trim all args, otherwise it confuses the arg parser of timm_efficientdet - sys.argv = sys.argv[:1] - - return parsed - - -logging.basicConfig( - format="%(asctime)s.%(msecs)03d %(levelname)s %(message)s", - datefmt="%Y-%m-%d %H:%M:%S", -) -logging.getLogger("tvm.meta_schedule").setLevel(logging.DEBUG) -ARGS = parse_args() -IS_CUDA = ARGS.target.kind.name == "cuda" - -logger = logging.getLogger(__name__) # pylint: disable=invalid-name -logger.setLevel(logging.INFO) - - -runner = load_torchdynamo_benchmark_runner( # pylint: disable=invalid-name - IS_CUDA, - cosine_similarity=ARGS.result_metric == ResultComparisonMetric.COSINE, - float32=ARGS.float32, - disallowed_operators=ARGS.disallowed_op, -) - - -def get_meta_schedule_runner() -> ms.runner.PyRunner: - """ - Get the Runner for MetaSchedule. - - It returns RPCRunner if --rpc-host is given, otherwise it returns LocalRunner - """ - if ARGS.rpc_host is not None: - assert ARGS.rpc_port is not None, "Missing rpc_port" - assert ARGS.rpc_key is not None, "Missing rpc_key" - return ms.runner.RPCRunner( - rpc_config=ms.runner.RPCConfig( - tracker_host=ARGS.rpc_host, - tracker_port=ARGS.rpc_port, - tracker_key=ARGS.rpc_key, - session_timeout_sec=600, - ), - evaluator_config=ms.runner.EvaluatorConfig( - number=ARGS.number, - repeat=ARGS.repeat, - min_repeat_ms=ARGS.min_repeat_ms, - enable_cpu_cache_flush=ARGS.cpu_flush, - ), - alloc_repeat=1, - ) - else: - warnings.warn("Falling back to MetaSchedule LocalRunner because --rpc-host isn't provided.") - return ms.runner.LocalRunner() - - -def get_graph_executor_forward( - graph_executor_factory: tvm.runtime.Module, device: tvm.runtime.Device -) -> Callable: - """ - Get the forward function for graph executor, in order to integrate with TorchDynamo. - """ - - # It has to lazily import this package, loading the C++ PyTorch integration - # after the transformers package is imported when loading model. Otherwise - # there will be segfault caused by the protobuf library. - import tvm.contrib.torch # pylint: disable=import-outside-toplevel, unused-import, redefined-outer-name - - save_runtime_mod = get_global_func("tvmtorch.save_runtime_mod", allow_missing=True) - if save_runtime_mod is None: - warnings.warn( - "C++ PyTorch TVM integration is missing. Fallback to Python forward function." - "Build TVM with 'USE_PT_TVMDSOOP' to enable the C++ custom operator" - ) - mod = GraphModule(graph_executor_factory["default"](device)) - - def forward(*args): - if IS_CUDA: - torch.cuda.synchronize() - args = tuple(arg.detach().contiguous() for arg in args) - for idx, arg in enumerate(args, 0): - mod.set_input( - f"inp_{idx}", - tvm.nd.from_dlpack(arg), - ) - mod.run() - device.sync() - result = [torch.from_dlpack(mod.get_output(i)) for i in range(mod.get_num_outputs())] - return result - - return forward - else: - save_runtime_mod(graph_executor_factory.module) - module = torch.classes.tvm_torch.GraphExecutorFactoryWrapper() - - def forward(*args): # type: ignore # isort: skip, pylint: disable=function-redefined - return module.forward(args) - - return forward - - -def get_vm_forward(virtual_machine: VirtualMachine, device: tvm.runtime.Device) -> Callable: - """ - Get the forward function for VM, in order to integrate with TorchDynamo. - """ - - def forward(*args): - if IS_CUDA: - torch.cuda.synchronize() - args = tuple(tvm.nd.from_dlpack(arg.detach().contiguous()) for arg in args) - result = virtual_machine.invoke("main", *args) - device.sync() - - if isinstance(result, tvm.nd.NDArray): - result = [result] - return [torch.from_dlpack(m) for m in result] - - return forward - - -def should_skip_subgraph(graph_module: torch.fx.GraphModule) -> bool: - """ - Returns whether it should skip optimizing the input graph module. - The graph could be empyt or only containing nodes calling function - for side effect. - """ - graph = graph_module.graph - - inputs = [n for n in graph.nodes if n.op == "placeholder"] - outputs = [n for n in graph.nodes if n.op == "output"] - - return len(inputs) == 0 and all(output.args == ((),) for output in outputs) - - -def create_tvm_task_collection_backend() -> Tuple[Callable, List[ms.ExtractedTask]]: - """ - This torchdynamo backend only collects the extracted tasks from MetaSchedule. - It doesn't tune the model. - """ - - subgraph_idx = 0 - subgraphs_dir = os.path.join(ARGS.work_dir, "subgraphs") - os.makedirs(subgraphs_dir, exist_ok=True) - - collected_tasks = [] - task_index: Dict[int, List[ms.ExtractedTask]] = defaultdict(list) - - def collect_task(task): - task_hash = tvm.ir.structural_hash(task.dispatched[0]) - - for duplicate_task in task_index[task_hash]: - if tvm.ir.structural_equal(duplicate_task.dispatched[0], task.dispatched[0]): - duplicate_task.weight += task.weight - return - - task_index[task_hash].append(task) - collected_tasks.append(task) - - def backend(graph_module, example_inputs): - nonlocal subgraph_idx - - torch.save(graph_module, os.path.join(subgraphs_dir, f"graph_module_{subgraph_idx}")) - torch.save(example_inputs, os.path.join(subgraphs_dir, f"example_inputs_{subgraph_idx}")) - - if should_skip_subgraph(graph_module): - return graph_module.forward - - jit_mod = torch.jit.trace(graph_module, example_inputs) - shape_list = [(f"inp_{idx}", i.shape) for idx, i in enumerate(example_inputs)] - ir_mod, params = tvm.relay.frontend.from_pytorch(jit_mod, shape_list) - - extracted_tasks = ms.relay_integration.extract_tasks( - mod=ir_mod, - target=ARGS.target, - params=params, - ) - old_tasks_count = len(collected_tasks) - for task in extracted_tasks: - collect_task(task) - logger.info( - "Extracted %d tasks from graph %d, with %d new tasks", - len(extracted_tasks), - subgraph_idx, - len(collected_tasks) - old_tasks_count, - ) - - subgraph_idx += 1 - - return graph_module.forward - - return backend, collected_tasks - - -def create_tvm_compilation_backend(database: ms.database.Database) -> Callable: - """ - This torchdynamo backend compiles the model using history best record from the - MetaSchedule database. - """ - - def backend(graph_module, example_inputs): - if should_skip_subgraph(graph_module): - return graph_module.forward - - jit_mod = torch.jit.trace(graph_module, example_inputs) - shape_list = [(f"inp_{idx}", i.shape) for idx, i in enumerate(example_inputs)] - ir_mod, params = tvm.relay.frontend.from_pytorch(jit_mod, shape_list) - - lib = ms.relay_integration.compile_relay( - database=database, - mod=ir_mod, - target=ARGS.target, - params=params, - backend=ARGS.backend, - ) - device = tvm.cuda(0) if IS_CUDA else tvm.cpu(0) - - if ARGS.backend == "graph": - return get_graph_executor_forward(lib, device) - elif ARGS.backend == "vm": - vm = VirtualMachine(lib, device) # pylint: disable=invalid-name - return get_vm_forward(vm, device) - else: - raise RuntimeError(f"Unknown backend {ARGS.backend}") - - return backend - - -def format_time(seconds: float) -> str: - """ - Format elapsed time based on its value. - """ - if seconds > 1: - return f"{seconds:.3g}s" - else: - return f"{seconds * 1000:.3g}ms" - - -def is_output_correct(output: torch.Tensor, expected: torch.Tensor) -> bool: - """ - Check whether the output is correct. - """ - comparison_metric = ARGS.result_metric - if comparison_metric == ResultComparisonMetric.COSINE: - return same(expected, output, cosine_similarity=True) - elif comparison_metric == ResultComparisonMetric.ALLCLOSE: - return same(expected, output, tol=1e-4) - else: - raise RuntimeError(f"Unknown comparison metric {comparison_metric}") - - -def inspect_output_error(output, expected): - """ - Inpsect the error between the actual output and expected output. - """ - if not isinstance(output, torch.Tensor): - logger.info( - f"Unsupported type for error inspection: {type(output).__name__}." - f"Please manually check output.pt" - ) - return - output = output.cpu().float() - expected = expected.cpu().float() - - abs_error = (output - expected).abs() - rel_error = (abs_error / expected).abs() - - def format_error_table(error, bins) -> str: - bin_tensor = torch.as_tensor([float(b) for b in bins], dtype=error.dtype) - error_hist = torch.histogram(error, bin_tensor).hist.int() - return "\n".join(f"< {b}\t{e}" for e, b in zip(error_hist, bins[1:])) - - abs_error_bins = [ - "-1e10", - "0", - "1e-8", - "1e-6", - "1e-5", - "1e-4", - "1e-3", - "1e-2", - "1e-1", - "1", - "1e10", - ] - rel_error_bins = [ - "-1e10", - "0", - "1e-4", - "1e-3", - "1e-2", - "1e-1", - "1", - "1e1", - "1e2", - "1e3", - "1e100", - ] - - large_rel_error_idx = rel_error > 1 - abs_error_with_large_rel_error = abs_error[large_rel_error_idx] - - logger.error(f"Expected (PyTorch eager): {expected}") - logger.error(f"Actual (Optimized): {output}") - logger.error(f"Absolute Error\n{format_error_table(abs_error, abs_error_bins)}") - logger.error(f"Relative Error\n{format_error_table(rel_error, rel_error_bins)}") - logger.error( - f"Max absolute error for position with large relative error (> 1):" - f"{abs_error_with_large_rel_error.max()}" - ) - - -def performance_experiment( - model_iter_fn: Callable, - model: torch.nn.Module, - example_inputs: Tuple[torch.Tensor], -) -> str: - """ - Performs the actual benchmarking - Simplified from https://github.com/pytorch/torchdynamo/blob/c537639f9712621dc04ca09908796dbbe86c354b/benchmarks/common.py#L494 pylint: disable=line-too-long - """ - timings = np.zeros((ARGS.benchmark_repeat, 2), np.float64) - if IS_CUDA: - torch.cuda.empty_cache() - - is_correct = True - - frozen_model_iter_fn = torchdynamo.run(model_iter_fn) - - for _ in range(ARGS.benchmark_warmup_rounds): - frozen_model_iter_fn(model, example_inputs) - model_iter_fn(model, example_inputs) - - for rep in range(ARGS.benchmark_repeat): - # interleave the runs to handle frequency scaling and load changes - timings[rep, 0], expected_output = timed( - model, model_iter_fn, example_inputs, return_result=True - ) - timings[rep, 1], actual_output = timed( - model, frozen_model_iter_fn, example_inputs, return_result=True - ) - is_correct = is_correct and is_output_correct(expected_output, actual_output) - - pvalue = ttest_ind(timings[:, 0], timings[:, 1]).pvalue - median = np.median(timings, axis=0) - speedup = median[0] / median[1] - logger.info( - f"eager:{format_time(median[0])} " - f"optimized:{format_time(median[1])} " - f"speedup:{speedup:.3f}x p:{pvalue:.3f}" - ) - torch.save(actual_output, os.path.join(ARGS.work_dir, "output.pt")) - torch.save(expected_output, os.path.join(ARGS.work_dir, "expected.pt")) - if not is_correct: - logger.error("Result is incorrect.") - inspect_output_error(actual_output, expected_output) - - return "" - - -def get_torch_device_type(target: tvm.target.Target) -> str: - if target.kind.name == "llvm": - return "cpu" - elif target.kind.name == "cuda": - return "cuda" - else: - raise RuntimeError(f"Unsupported target {target}") - - -def main(): - """ - Entry point of the benchmark - """ - describe() - - meta_schedule_work_dir = os.path.join(ARGS.work_dir, "meta_schedule") - os.makedirs(meta_schedule_work_dir, exist_ok=True) - - database = ms.database.JSONDatabase(work_dir=meta_schedule_work_dir) - if not ARGS.mode.should_tune: - if len(database) == 0: - raise RuntimeError( - "Script is running in eval mode while the tuning database is empty. " - "Please tune the model first." - ) - - if IS_CUDA and ARGS.cpu_flush: - warnings.warn( - "Benchmark is running on CUDA, while --cpu-flush is turned on. " - "This flag will have no effect on CUDA." - ) - ARGS.cpu_flush = False - - try: - logger.info(f"Loading model with batch size: {ARGS.batch_size}") - _, name, model, example_inputs, batch_size = runner.load_model( - get_torch_device_type(ARGS.target), - ARGS.model, - batch_size=ARGS.batch_size, - ) - model, example_inputs = runner.maybe_cast(model, example_inputs) - logger.info(f"Got model with batch size: {batch_size}") - except NotImplementedError: - logger.exception(f"{ARGS.model} failed to load") - raise - - with contextlib.ExitStack() as stack: - profiler = stack.enter_context(ms.Profiler()) - stack.enter_context(torch.no_grad()) - - tasks_path = os.path.join(ARGS.work_dir, "extracted_tasks") - - if ARGS.mode.should_extract: - task_collect_backend, extracted_tasks = create_tvm_task_collection_backend() - task_collect_ctx = torchdynamo.optimize(task_collect_backend) - task_collect_ctx(runner.model_iter_fn)(model, example_inputs) - with open(tasks_path, "wb") as f: - pickle.dump(extracted_tasks, f) - else: - with open(tasks_path, "rb") as f: - extracted_tasks = pickle.load(f) - - if ARGS.mode.should_tune: - tasks, task_weights = ms.relay_integration.extracted_tasks_to_tune_contexts( - extracted_tasks=extracted_tasks, - work_dir=ARGS.work_dir, - strategy=ARGS.strategy, - ) - database = ms.tune.tune_tasks( - tasks=tasks, - task_weights=task_weights, - work_dir=ARGS.work_dir, - max_trials_global=ARGS.num_trials, - max_trials_per_task=ARGS.max_trials_per_task, - runner=get_meta_schedule_runner(), # type: ignore - database=database, - cost_model=ms.cost_model.XGBModel( # type: ignore - extractor=ms.feature_extractor.PerStoreFeature(), - adaptive_training=ARGS.adaptive_training, - ), - ) - - if ARGS.mode.should_eval: - torchdynamo.reset() - model_compile_ctx = torchdynamo.optimize(create_tvm_compilation_backend(database)) - model_compile_ctx(runner.model_iter_fn)(model, example_inputs) - with torch.no_grad(): - performance_experiment(runner.model_iter_fn, model, example_inputs) - - print(profiler.table()) - - -if __name__ == "__main__": - main() diff --git a/python/tvm/meta_schedule/testing/torchbench/utils.py b/python/tvm/meta_schedule/testing/torchbench/utils.py deleted file mode 100644 index 7094a282403f..000000000000 --- a/python/tvm/meta_schedule/testing/torchbench/utils.py +++ /dev/null @@ -1,167 +0,0 @@ -# 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, pointless-exception-statement -""" -Helper functions for running TorchBench through the benchmark functions -from TorchDynamo. -""" - -import functools -import os -import sys -from dataclasses import dataclass -from enum import Enum -from typing import Set - -import torch # type: ignore - - -class DisallowedOperator(Enum): - """ - The operators to disallow in the fx graph produced by TorchDynamo. - This is to workaround the limitation in TVM's PyTorch frontend. - - - inplace_copy: aten::copy_ as inplace assign A[...] = ..., or method call A.copy_(...) - - einsum: torch.functional.einsum - - multihead_attention: torch.nn.MultiheadAttention - - as_stride: Tensor.as_stride - """ - - INPLACE_COPY = "inplace_copy" - EINSUM = "einsum" - MULTIHEAD_ATTENTION = "multihead_attention" - AS_STRIDE = "as_stride" - - -def find_torchdynamo() -> str: - """ - Find the directory of TorchDynamo repo. - - It can't directly import the benchmark runner in TorchDynamo - becuase it isn't designed to be used as a Python package. - """ - candidates = [ - "torchdynamo", - "../torchdynamo", - "../../torchdynamo", - ] - for library_dir in candidates: - if os.path.exists(f"{library_dir}/benchmarks"): - return library_dir - - raise RuntimeError( - """ - Cannot find directory for torchdynamo. - You need to clone https://github.com/pytorch/torchdynamo to the parent directory of cwd. - """ - ) - - -DYNAMO_DIR = find_torchdynamo() -sys.path.insert( - 0, DYNAMO_DIR -) # opacus_cifar10 depends on opacus, which installs a package called 'benchmarks' -sys.path.append(f"{DYNAMO_DIR}/benchmarks") - -# pylint: disable=wrong-import-position, unused-import -import torchdynamo # type: ignore -from benchmarks.common import same, timed # type: ignore -from torchbench import TorchBenchmarkRunner # type: ignore - -# pylint: disable=wrong-import-position, unused-import - - -def _disallow_operators(disallowed_ops: Set[DisallowedOperator]): - """ - Disallow certain operators in the fx graph produced by TorchDynamo. - There are two ways to disallow operator in TorchDynamo, - 1. Use the disallow_in_graph API, which only applies to free function call. - 2. Patch the TensorVariable class, which applies to method call on torch.Tensor. - """ - disallowed_tensor_methods: Set[str] = set() - - if DisallowedOperator.INPLACE_COPY in disallowed_ops: - torchdynamo.disallow_in_graph(torch.Tensor.copy_) - disallowed_tensor_methods.update({"copy_", "__setitem__"}) - - if DisallowedOperator.EINSUM in disallowed_ops: - torchdynamo.disallow_in_graph(torch.functional.einsum) - - if DisallowedOperator.MULTIHEAD_ATTENTION in disallowed_ops: - torchdynamo.disallow_in_graph(torch.nn.MultiheadAttention) - - if DisallowedOperator.AS_STRIDE in disallowed_ops: - disallowed_tensor_methods.add("as_stride") - - tensor_variable_cls = torchdynamo.variables.tensor.TensorVariable - old_call_method = tensor_variable_cls.call_method - - @functools.wraps(old_call_method) - def call_method(self, translator, name, args, kwargs): - if name in disallowed_tensor_methods: - raise torchdynamo.exc.Unsupported(f"Tensor.{name} not supported by TVM.") - return old_call_method(self, translator, name, args, kwargs) - - tensor_variable_cls.call_method = call_method - - -def load_torchdynamo_benchmark_runner( - is_cuda: bool, - cosine_similarity: bool = False, - float32: bool = False, - disallowed_operators: Set[DisallowedOperator] = None, -) -> TorchBenchmarkRunner: - """ - Load the benchmark runner from TorchDynamo. - """ - - @dataclass - class RunnerArgs: - """ - This class simulates the parsed args required by the benchmark code from TorchDynamo. - """ - - ci: bool = False # Whether runs in CI mode. pylint: disable=invalid-name - training: bool = False # Whether it benchmarks training workload. - use_eval_mode: bool = True # Whether the model should be in eval mode. - dynamic_shapes: bool = False # Whether runs the model in dynamic shape mode. - float16: bool = False # Whether to cast model and inputs to float16 - float32: bool = False # Whether to cast model and inputs to float32 - - accuracy: bool = False # Whether to perform a accuracy test - performance: bool = True # Whether to perform a performance test - - cosine: bool = False # Whether to use consine similarity to check if output is correct. - - args = RunnerArgs(cosine=cosine_similarity, float32=float32) - - runner = TorchBenchmarkRunner() - runner.args = args - runner.model_iter_fn = runner.forward_pass - - if disallowed_operators: - _disallow_operators(disallowed_operators) - - if is_cuda: - # pylint: disable=import-outside-toplevel - import benchmarks.common # type: ignore - - # pylint: enable=import-outside-toplevel - - benchmarks.common.synchronize = torch.cuda.synchronize - - return runner diff --git a/python/tvm/target/__init__.py b/python/tvm/target/__init__.py index beaddf03c7fa..1bb883e840cc 100644 --- a/python/tvm/target/__init__.py +++ b/python/tvm/target/__init__.py @@ -74,7 +74,5 @@ from .virtual_device import VirtualDevice from .compilation_config import make_compilation_config from .tag import list_tags -from .generic_func import GenericFunc -from .generic_func import generic_func, get_native_generic_func, override_native_generic_func from . import datatype from . import codegen diff --git a/python/tvm/target/generic_func.py b/python/tvm/target/generic_func.py deleted file mode 100644 index 7b6f916bd975..000000000000 --- a/python/tvm/target/generic_func.py +++ /dev/null @@ -1,304 +0,0 @@ -# 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. -"""Generic function.""" - -import tvm._ffi - -try: - from decorator import decorate -except ImportError: - # Allow decorator to be missing in runtime - if not tvm._ffi.base._RUNTIME_ONLY: - raise - -from tvm.runtime import Object -from .target import Target -from . import _ffi_api - - -@tvm._ffi.register_object -class GenericFunc(Object): - """GenericFunc node reference. This represents a generic function - that may be specialized for different targets. When this object is - called, a specialization is chosen based on the current target. - - Note - ---- - Do not construct an instance of this object, it should only ever be - used as a return value from calling into C++. - """ - - def __call__(self, *args): - return _ffi_api.GenericFuncCallFunc(self, *args) - - def set_default(self, func, allow_override=False): - """Set the default function to be used if no specializations match - the current target. - - Parameters - ---------- - func : function - The default function - - allow_override : bool - Whether to allow the current default to be overridden - """ - _ffi_api.GenericFuncSetDefault(self, func, allow_override) - - def register(self, func, key_list, allow_override=False): - """Register a specialization for this GenericFunc. - - Parameters - ---------- - func : function - The function to be registered. - - key : str or list of str - The key to be registered. - - allow_override : bool, optional - Whether to allow existing keys to be overridden. - """ - key_list = [key_list] if isinstance(key_list, str) else key_list - _ffi_api.GenericFuncRegisterFunc(self, func, key_list, allow_override) - - def get_packed_func(self): - """Get the packed function specified for the current target. - - Returns - ------- - func : PackedFunc - The function specified for the current target. Return the default - function if no specializations match the current target. - """ - return _ffi_api.GenericFuncGetPackedFunc(self) - - -def get_native_generic_func(name): - """Get a generic function from the global registry. If no - function is registered under the given name, a new generic - function is created. - - Parameters - ---------- - name : string - The name of the generic function to get - - Returns - ------- - func : GenericFunc - The generic function for the given name - """ - return _ffi_api.GenericFuncGetGlobal(name) - - -def override_native_generic_func(func_name): - """Override a generic function defined in C++ - - Generic function allows registration of further functions - that can be dispatched on current target context. - If no registered dispatch is matched, the fdefault will be called. - - Parameters - ---------- - func_name : string - The name of the generic func to be overridden - - Returns - ------- - fgeneric : function - A wrapped generic function. - - Example - ------- - .. code-block:: python - - import tvm - # wrap function as target generic - @tvm.target.override_native_generic_func("my_func") - def my_func(a): - return a + 1 - # register specialization of my_func under target cuda - @my_func.register("cuda") - def my_func_cuda(a): - return a + 2 - # displays 3, because my_func is called - print(my_func(2)) - # displays 4, because my_func_cuda is called - with tvm.target.cuda(): - print(my_func(2)) - """ - generic_func_node = get_native_generic_func(func_name) - - def fdecorate(fdefault): - """Wrap a target generic function, overriding the previous - default that was set for the generic function. - - Parameters - ---------- - fdefault : function - The default function. - - Returns - ------- - fgeneric : function - A wrapped generic function. - - """ - generic_func_node.set_default(fdefault, allow_override=True) - - def register(key, func=None, override=True): - """Register function to be the dispatch function. - - Parameters - ---------- - key : str or list of str - The key to be registered. - - func : function - The function to be registered. - - override : bool, optional - Whether override existing registration. - - Returns - ------- - The register function is necessary. - """ - - def _do_reg(myf): - generic_func_node.register(myf, key, override) - return myf - - if func: - return _do_reg(func) - return _do_reg - - def dispatch_func(func, *args, **kwargs): - # pylint: disable=unused-argument - """The wrapped dispath function""" - if kwargs: - raise RuntimeError( - "Keyword arguments cannot be used when invoking generic_func %s" % func_name - ) - return generic_func_node(*args) - - fresult = decorate(fdefault, dispatch_func) - fresult.fdefault = fdefault - fresult.register = register - fresult.generic_func_node = generic_func_node - return fresult - - return fdecorate - - -def generic_func(fdefault): - """Wrap a target generic function. - - Generic function allows registration of further functions - that can be dispatched on current target context. - If no registered dispatch is matched, the fdefault will be called. - - Parameters - ---------- - fdefault : function - The default function. - - Returns - ------- - fgeneric : function - A wrapped generic function. - - Example - ------- - .. code-block:: python - - import tvm - # wrap function as target generic - @tvm.target.generic_func - def my_func(a): - return a + 1 - # register specialization of my_func under target cuda - @my_func.register("cuda") - def my_func_cuda(a): - return a + 2 - # displays 3, because my_func is called - print(my_func(2)) - # displays 4, because my_func_cuda is called - with tvm.target.cuda(): - print(my_func(2)) - """ - dispatch_dict = {} - func_name = fdefault.__name__ - - def register(key, func=None, override=False): - """Register function to be the dispatch function. - - Parameters - ---------- - key : str or list of str - The key to be registered. - - func : function - The function to be registered. - - override : bool - Whether override existing registration. - - Returns - ------- - The register function is necessary. - """ - - def _do_reg(myf): - key_list = [key] if isinstance(key, str) else key - for k in key_list: - if k in dispatch_dict and not override: - raise ValueError("Key is already registered for %s" % func_name) - dispatch_dict[k] = myf - return myf - - if func: - return _do_reg(func) - return _do_reg - - def dispatch_func(func, *args, **kwargs): - """The wrapped dispatch function""" - target = Target.current() - if target is None: - return func(*args, **kwargs) - for k in target.keys: - if k in dispatch_dict: - return dispatch_dict[k](*args, **kwargs) - return func(*args, **kwargs) - - def get_packed_func(): - """The wrapped to get dispatched function""" - target = Target.current() - if target is None: - return fdefault - for k in target.keys: - if k in dispatch_dict: - return dispatch_dict[k] - return fdefault - - fdecorate = decorate(fdefault, dispatch_func) - fdecorate.register = register - fdecorate.fdefault = fdefault - fdecorate.dispatch_dict = dispatch_dict - fdecorate.get_packed_func = get_packed_func - return fdecorate diff --git a/python/tvm/tir/__init__.py b/python/tvm/tir/__init__.py index b4172de77d8f..1d7352f66527 100644 --- a/python/tvm/tir/__init__.py +++ b/python/tvm/tir/__init__.py @@ -108,4 +108,3 @@ from . import transform from . import analysis from . import stmt_functor -from . import usmp diff --git a/python/tvm/tir/analysis/analysis.py b/python/tvm/tir/analysis/analysis.py index 67eb7471d22d..e98c176dd093 100644 --- a/python/tvm/tir/analysis/analysis.py +++ b/python/tvm/tir/analysis/analysis.py @@ -164,44 +164,6 @@ def get_block_read_write_region( return _ffi_api.GetBlockReadWriteRegion(block, buffer_var_map) # type: ignore -def calculate_workspace_bytes(func: PrimFunc, workspace_byte_alignment: int) -> int: - """Calculate the workspace size in bytes needed by the TIR allocates inside the TIR - PrimFunc. - - Parameters - ---------- - func: tvm.tir.PrimFunc - The function to be detected. - workspace_byte_alignment : int - The byte alignment required for each tensor - - Returns - ------- - result : int - Workspace size in bytes. - """ - return _ffi_api.calculate_workspace_bytes(func, workspace_byte_alignment) # type: ignore - - -def calculate_constant_bytes(func: PrimFunc, constant_byte_alignment: int) -> int: - """Calculate the constant size in bytes needed by the TIR allocates inside the TIR - PrimFunc. - - Parameters - ---------- - func: tvm.tir.PrimFunc - The function to be detected. - constant_byte_alignment : int - The byte alignment required for each tensor - - Returns - ------- - result : int - Workspace size in bytes. - """ - return _ffi_api.calculate_constant_bytes(func, constant_byte_alignment) # type: ignore - - def calculate_allocated_bytes( func_or_mod: Union[PrimFunc, IRModule] ) -> Union[Dict[str, int], Dict[str, Dict[str, int]]]: @@ -314,41 +276,6 @@ def get_prim_func_arg_and_result_memory_constraints( ) -def apply_prim_func_arg_and_result_memory_constraints( - func: PrimFunc, relay_func_type: Object, arg_and_result_memory_scopes: List[str] -) -> PrimFunc: - """Returns func written to capture the memory (aka storage) scope constraints - for each of the func's parameters given by arg_and_result_memory_scopes. However, - arg_and_result_memory_scopes should be w.r.t. the func's representation as a Relay - Function of relay_func_type before lowering and conversion to DPS. - - Visible for testing. - - CAUTION: This is experimental. The resulting PrimFunc may not have fully accounted - for all new memory scopes. - - Parameters - ---------- - func: tvm.tir.PrimFunc - The function to retrieve constraints from. - - relay_func_type: tvm.relay.FuncType - The type of the Relay Function from which the func was derived. - - arg_and_result_memory_scopes: Array[AnyStr] - Memory constraints for funcs args and result in Relay form. The empty string denotes - 'no constraint'. - - Returns - ------- - result: tvm.tir.PrimFunc - The rewritten func. - """ - return _ffi_api.ApplyPrimFuncArgAndResultMemoryConstraints( # type: ignore # pylint: disable=no-member - func, relay_func_type, arg_and_result_memory_scopes - ) - - def verify_well_formed(obj: Union[PrimFunc, IRModule], assert_mode: bool = True) -> bool: """Verify if the given TIR is well-formed. The verification includes: - Check if expressions not contain vars that is defined outside the block. diff --git a/python/tvm/tir/transform/transform.py b/python/tvm/tir/transform/transform.py index d8531401d49d..b08659e1c712 100644 --- a/python/tvm/tir/transform/transform.py +++ b/python/tvm/tir/transform/transform.py @@ -1171,18 +1171,6 @@ def InstrumentProfileIntrinsics(): return _ffi_api.InstrumentProfileIntrinsics() # type: ignore -def InstallDebugSpans(): - """Add line information from the TIR printer as spans on each statement and - expression. - - Returns - ------- - fpass : tvm.transform.Pass - The result pass - """ - return _ffi_api.InstallDebugSpans() # type: ignore - - def DefaultGPUSchedule(): """The pass sets default thread bindings for PrimFuncs, including symbolic shape functions, allowing their build and execution on GPU devices. It examines all the blocks within the diff --git a/python/tvm/tir/usmp/__init__.py b/python/tvm/tir/usmp/__init__.py deleted file mode 100644 index 514727d52e2e..000000000000 --- a/python/tvm/tir/usmp/__init__.py +++ /dev/null @@ -1,22 +0,0 @@ -# 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=unused-import, redefined-builtin -"""Namespace for Unified Static Memory Planner""" - -from . import analysis -from . import transform -from .utils import BufferInfo diff --git a/python/tvm/tir/usmp/_ffi_api.py b/python/tvm/tir/usmp/_ffi_api.py deleted file mode 100644 index 5899ef0c86ea..000000000000 --- a/python/tvm/tir/usmp/_ffi_api.py +++ /dev/null @@ -1,21 +0,0 @@ -# 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. -"""FFI APIs for tvm.tir.usmp""" -import tvm._ffi - - -tvm._ffi._init_api("tir.usmp", __name__) diff --git a/python/tvm/tir/usmp/analysis/__init__.py b/python/tvm/tir/usmp/analysis/__init__.py deleted file mode 100644 index 756e8c7204c5..000000000000 --- a/python/tvm/tir/usmp/analysis/__init__.py +++ /dev/null @@ -1,20 +0,0 @@ -# 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=unused-import, redefined-builtin -"""Namespace for Unified Static Memory Planner""" - -from .analysis import extract_buffer_info diff --git a/python/tvm/tir/usmp/analysis/_ffi_api.py b/python/tvm/tir/usmp/analysis/_ffi_api.py deleted file mode 100644 index 36973f19905c..000000000000 --- a/python/tvm/tir/usmp/analysis/_ffi_api.py +++ /dev/null @@ -1,21 +0,0 @@ -# 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. -"""FFI APIs for tvm.tir.usmp.analysis""" -import tvm._ffi - - -tvm._ffi._init_api("tir.usmp.analysis", __name__) diff --git a/python/tvm/tir/usmp/analysis/analysis.py b/python/tvm/tir/usmp/analysis/analysis.py deleted file mode 100644 index ff70355a967b..000000000000 --- a/python/tvm/tir/usmp/analysis/analysis.py +++ /dev/null @@ -1,39 +0,0 @@ -# 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. -"""USMP Analysis Python API for passes""" -# pylint: disable=invalid-name -from . import _ffi_api -from ...function import PrimFunc -from ....ir.module import IRModule - - -def extract_buffer_info(main_func: PrimFunc, mod: IRModule): - """Convert Parallel For Loop to Serial. - - Parameters - ---------- - main_func: tvm.tir.PrimFunc - The main function containing calls to operator PrimFuncs. - mod : tvm.ir.IRModule - The full IRModule containing all PrimFuncs - - Returns - ------- - Map - extracted buffer info objects - """ - return _ffi_api.extract_buffer_info(main_func, mod) diff --git a/python/tvm/tir/usmp/transform/__init__.py b/python/tvm/tir/usmp/transform/__init__.py deleted file mode 100644 index 1a9d83328f8d..000000000000 --- a/python/tvm/tir/usmp/transform/__init__.py +++ /dev/null @@ -1,20 +0,0 @@ -# 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=unused-import, redefined-builtin -"""Namespace for Unified Static Memory Planner""" - -from .transform import convert_pool_allocations_to_offsets diff --git a/python/tvm/tir/usmp/transform/_ffi_api.py b/python/tvm/tir/usmp/transform/_ffi_api.py deleted file mode 100644 index 7973ca5b0da0..000000000000 --- a/python/tvm/tir/usmp/transform/_ffi_api.py +++ /dev/null @@ -1,21 +0,0 @@ -# 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. -"""FFI APIs for tvm.tir.usmp.analysis""" -import tvm._ffi - - -tvm._ffi._init_api("tir.usmp.transform", __name__) diff --git a/python/tvm/tir/usmp/transform/transform.py b/python/tvm/tir/usmp/transform/transform.py deleted file mode 100644 index f472172cf36f..000000000000 --- a/python/tvm/tir/usmp/transform/transform.py +++ /dev/null @@ -1,46 +0,0 @@ -# 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. -"""USMP Transform Python API for passes""" -# pylint: disable=invalid-name - -from typing import Dict - -import tvm -from tvm.tir import Stmt -from tvm.tir.usmp.utils import PoolAllocation -from . import _ffi_api - - -def convert_pool_allocations_to_offsets( - pool_allocations: Dict[Stmt, PoolAllocation], emit_tvmscript_printable: bool = False -) -> tvm.transform.Pass: - """Convert pool allocations to Load nodes with offsets from pools. - - Parameters - ---------- - pool_allocations : Dict[Stmt, PoolAllocation] - Allocate or AllocateConst node to pool allocation mapping - emit_tvmscript_printable : bool - A toggle to emit TVMScript printable IRModule for unit tests - removing all attributes that should be attached for integration - - Returns - ------- - ret: tvm.transform.Pass - The registered pass that converts the allocations to offsets. - """ - return _ffi_api.ConvertPoolAllocationsToOffsets(pool_allocations, emit_tvmscript_printable) diff --git a/python/tvm/tir/usmp/utils.py b/python/tvm/tir/usmp/utils.py deleted file mode 100644 index 834783414c00..000000000000 --- a/python/tvm/tir/usmp/utils.py +++ /dev/null @@ -1,105 +0,0 @@ -# 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. -"""USMP Utilities and Data Structures""" -# pylint: disable=invalid-name - -from typing import Optional, List - -import tvm -from tvm._ffi import register_object -from tvm.runtime import Object -from . import _ffi_api -from ...ir.memory_pools import PoolInfo - - -# The allocate node attribute to indicate candidate memory pools. -# This needs to be kept in sync with CANDIDATE_MEMORY_POOL_ATTR in -# include/tvm/tir/usmp/utils.h -CANDIDATE_MEMORY_POOL_ATTR = "candidate_memory_pools" - - -def use_workspace_io_is_enabled() -> bool: - """ - Check whether placing I/O tensors in the workspace is enabled. - """ - ctx = tvm.transform.PassContext.current() - return bool(ctx.config.get("tir.usmp.use_workspace_io", False)) - - -@register_object("tir.usmp.BufferInfo") -class BufferInfo(Object): - """BufferInfo object holds information related to buffers - that are associated with tir.allocates and tir.allocate_consts - that will be used with USMP - - Parameters - ---------- - name_hint : str - The name associated with the buffer (derived from TIR) - - size_bytes : int - The size in bytes - - pool_candidates : List[PoolInfo] - The list of candidates pools this buffer could be placed - - alignment : Optional[int] - The byte alignment required in the workspace memory - - """ - - def __init__( - self, - name_hint: str, - size_bytes: int, - pool_candidates: List[PoolInfo], - alignment: Optional[int] = None, - ): - self.__init_handle_by_constructor__( - _ffi_api.BufferInfo, # type: ignore # pylint: disable=no-member - name_hint, - size_bytes, - pool_candidates, - alignment, - ) - - def set_conflicts(self, conflicts: list): - """Sets the conflicting array of buffer info objects""" - _ffi_api.BufferInfoSetConflicts(self, conflicts) - - -@register_object("tir.usmp.PoolAllocation") -class PoolAllocation(Object): - """PoolAllocation object holds information related to an allocation - that indicates an offset in a pool - - Parameters - ---------- - pool_info : PoolInfo - The PoolInfo to which this allocation corresponds to - - byte_offset : int - The offset in the pool where the allocate node should be placed - - """ - - def __init__(self, pool_info: PoolInfo, byte_offset: int): - self.__init_handle_by_constructor__( - _ffi_api.PoolAllocation, # type: ignore # pylint: disable=no-member - pool_info, - byte_offset, - ) diff --git a/python/tvm/topi/math.py b/python/tvm/topi/math.py index 63a1e48c2bab..865d62683c8b 100644 --- a/python/tvm/topi/math.py +++ b/python/tvm/topi/math.py @@ -94,28 +94,6 @@ def erf(x): return te.compute(x.shape, lambda *i: te.erf(x(*i))) -@tvm.target.generic_func -def erf_legalize(attrs, inputs, types): - """Legalizes ERF op. - - Parameters - ---------- - attrs : tvm.ir.Attrs - Attributes of current convolution - inputs : list of tvm.relay.Expr - The args of the Relay expr to be legalized - types : list of types - List of input and output types - - Returns - ------- - result : tvm.relay.Expr - The legalized expr. - """ - # Note changed by default. - return None - - @tvm.te.tag_scope(tag=tag.ELEMWISE) def tanh(x): """Take hyperbolic tanh of input x. diff --git a/python/tvm/topi/nn/bitserial_conv2d.py b/python/tvm/topi/nn/bitserial_conv2d.py index 87acd4e3602c..eb8f19cbfc67 100644 --- a/python/tvm/topi/nn/bitserial_conv2d.py +++ b/python/tvm/topi/nn/bitserial_conv2d.py @@ -272,25 +272,3 @@ def _conv(nn, yy, xx, ff): ) return conv - - -@tvm.target.generic_func -def bitserial_conv2d_legalize(attrs, inputs, types): - """Legalizes Bitserial Conv2D op. - - Parameters - ---------- - attrs : tvm.ir.Attrs - Attributes of current convolution - inputs : list of tvm.relay.Expr - The args of the Relay expr to be legalized - types : list of types - List of input and output types - - Returns - ------- - result : tvm.relay.Expr - The legalized expr - """ - # not to change by default - return None diff --git a/python/tvm/topi/nn/conv2d.py b/python/tvm/topi/nn/conv2d.py index e145add5f01b..2965944d61c8 100644 --- a/python/tvm/topi/nn/conv2d.py +++ b/python/tvm/topi/nn/conv2d.py @@ -98,94 +98,6 @@ def conv2d( return conv(input, filter, strides, padding, dilation, 1, data_layout, kernel_layout, out_dtype) -@tvm.target.generic_func -def conv2d_legalize(attrs, inputs, types): - """Legalizes Conv2D op. - - Parameters - ---------- - attrs : tvm.ir.Attrs - Attributes of current convolution - inputs : list of tvm.relay.Expr - The args of the Relay expr to be legalized - types : list of types - List of input and output types - - Returns - ------- - result : tvm.relay.Expr - The legalized expr - """ - # not to change by default - return None - - -@tvm.target.generic_func -def conv2d_alter_layout(attrs, inputs, tinfos, out_type): - """Change Conv2D layout. - - Parameters - ---------- - attrs : tvm.ir.Attrs - Attributes of current convolution - inputs : tvm.relay.Expr - Grouped input symbols - tinfos : list - Input shape and dtype - out_type: type - The output type - - Note - ---- - Unlike other TOPI functions, this function operates on both graph level and operator level. - """ - # not to change by default - return None - - -@tvm.target.generic_func -def conv2d_transpose_alter_layout(attrs, inputs, tinfos, out_type): - """Change Conv2D_Transpose layout. - - Parameters - ---------- - attrs : tvm.ir.Attrs - Attributes of current convolution - inputs : tvm.relay.Expr - Grouped input symbols - tinfos : list - Input shape and dtype - out_type: type - The output type - - Note - ---- - Unlike other TOPI functions, this function operates on both graph level and operator level. - """ - # not to change by default - return None - - -@tvm.target.generic_func -def conv2d_infer_layout(workload, cfg): - """Infer input/output shapes and layouts from a workload and cfg. - - Parameters - ---------- - workload : tuple - conv2d workload - - cfg : tuple - tvm.autotvm config - - Returns - ------- - Output : [tuple of tuple and str, tuple of tuple and str] - Input shapes and layouts, and output shapes and layouts - """ - raise ValueError("missing register for topi.nn.conv2d_infer_layout") - - def _get_workload(data, kernel, stride, padding, dilation, out_dtype, data_layout="NCHW"): """Get the workload structure.""" if data_layout == "NCHW": @@ -953,7 +865,6 @@ def unpack_NCHWc_to_nchw(packed_out, out_dtype): return unpacked_out -@tvm.target.generic_func def conv2d_winograd_nhwc( data, weight, @@ -1010,7 +921,6 @@ def conv2d_winograd_nhwc( ) -@tvm.target.generic_func def conv2d_winograd_nchw( data, weight, diff --git a/python/tvm/topi/nn/conv3d.py b/python/tvm/topi/nn/conv3d.py index 1897484dc8cd..e3774f67b336 100644 --- a/python/tvm/topi/nn/conv3d.py +++ b/python/tvm/topi/nn/conv3d.py @@ -17,7 +17,6 @@ # pylint: disable=invalid-name, unused-variable, too-many-locals # pylint: disable=unused-argument, redefined-builtin, no-else-return """Conv3D operators""" -import tvm from tvm import te from ..utils import get_const_tuple @@ -167,26 +166,3 @@ def conv3d_winograd_weight_transform(kernel, tile_size): ), name="transform_weight", ) - - -@tvm.target.generic_func -def conv3d_alter_layout(attrs, inputs, tinfos, out_type): - """Change Conv3D layout. - - Parameters - ---------- - attrs : tvm.ir.Attrs - Attributes of current convolution - inputs : tvm.relay.Expr - Grouped input symbols - tinfos : list - Input shape and dtype - out_type: type - The output type - - Note - ---- - Unlike other TOPI functions, this function operates on both graph level and operator level. - """ - # not to change by default - return None diff --git a/python/tvm/topi/nn/dense.py b/python/tvm/topi/nn/dense.py index 5df1674627c3..449045f8c9b2 100644 --- a/python/tvm/topi/nn/dense.py +++ b/python/tvm/topi/nn/dense.py @@ -163,29 +163,6 @@ def compute(*indices): return mat -@tvm.target.generic_func -def matmul_legalize(attrs, inputs, types): - """Legalizes matmul op. - - Parameters - ---------- - attrs : tvm.ir.Attrs - Attributes of current matmul - inputs : list of tvm.relay.Expr - The args of the Relay expr to be legalized - types : list of types - List of input and output types - - Returns - ------- - result : tvm.relay.Expr - The legalized expr - """ - # not to change by default - # pylint: disable=unused-argument - return None - - def dense( data, weight, @@ -236,29 +213,6 @@ def dense( ) -@tvm.target.generic_func -def dense_legalize(attrs, inputs, types): - """Legalizes dense op. - - Parameters - ---------- - attrs : tvm.ir.Attrs - Attributes of current dense - inputs : list of tvm.relay.Expr - The args of the Relay expr to be legalized - types : list of types - List of input and output types - - Returns - ------- - result : tvm.relay.Expr - The legalized expr - """ - # not to change by default - # pylint: disable=unused-argument - return None - - def dense_pack(data, weight, bias=None, out_dtype=None): """The default implementation of dense_pack in topi. diff --git a/python/tvm/topi/nn/depthwise_conv2d.py b/python/tvm/topi/nn/depthwise_conv2d.py index ad1e4a55177f..b0ea7e051ac4 100644 --- a/python/tvm/topi/nn/depthwise_conv2d.py +++ b/python/tvm/topi/nn/depthwise_conv2d.py @@ -457,23 +457,3 @@ def depthwise_conv2d_NCHWc( 5-D with shape [batch, out_channel_chunk, out_height, out_width, out_channel_block] """ raise ValueError("missing register for topi.nn.depthwise_conv2d_NCHWc") - - -@tvm.target.generic_func -def depthwise_conv2d_infer_layout(workload, cfg): - """Infer input/output shapes and layouts from a workload and cfg. - - Parameters - ---------- - workload : tuple - conv2d workload - - cfg : tuple - tvm.autotvm config - - Returns - ------- - Output : [tuple of tuple and str, tuple of tuple and str] - Input shapes and layouts, and output shapes and layouts - """ - raise ValueError("missing register for topi.nn.depthwise_conv2d_infer_layout") diff --git a/python/tvm/topi/nn/qnn.py b/python/tvm/topi/nn/qnn.py index 98bbb7ebe50f..7b790d6cbbf4 100644 --- a/python/tvm/topi/nn/qnn.py +++ b/python/tvm/topi/nn/qnn.py @@ -62,6 +62,7 @@ def simulated_quantize(data, out_dtype, output_scale=None, output_zero_point=Non The channel axis for quantization. Default value is -1 which corresponds to the last axis. """ + # When disabled, just pass through the input values. def _compute_pass_through(value, *indices): return value[indices] @@ -151,6 +152,7 @@ def simulated_dequantize(data, in_dtype, input_scale=None, input_zero_point=None The channel axis for quantization. Default value is -1 which corresponds to the last axis. """ + # When disabled simply return the input tensor. def _compute_pass_through(value, *indices): return value[indices] @@ -188,112 +190,3 @@ def _dispatch_sim_dequantize(value): return intn_value return te.compute(data.shape, lambda *indices: _dispatch_sim_dequantize(data)[indices]) - - -@tvm.target.generic_func -def qnn_conv2d_alter_layout(_attrs, _inputs, _tinfos, _out_type): - """Change qnn.conv2d layout. - - Parameters - ---------- - attrs : tvm.ir.Attrs - Attributes of current convolution - inputs : tvm.relay.Expr - Grouped input symbols - tinfos : list - Input shape and dtype - out_type: type - The output type - - Note - ---- - Unlike other TOPI functions, this function operates on both graph level and operator level. - """ - return None - - -@tvm.target.generic_func -def bias_add_legalize(_attrs, _inputs, _tinfos): - """Legalize bias_add layout. - - Bias add is not a QNN-specific function, but this generic exists so that empty channels can - be excised from quantized conv2d operators and folded into bias adds. - - Parameters - ---------- - attrs : tvm.ir.Attrs - Attributes of current convolution - inputs : tvm.relay.Expr - Grouped input symbols - tinfos : list - Input shape and dtype - - """ - return None - - -@tvm.target.generic_func -def add_alter_layout(_attrs, _inputs, _tinfos, _out_type): - """Change add layout. - - Add is not a QNN-specific function, but this generic exists so that bias add operations can be - fused with input zero point add optimizations, which only happens if the previous operator is - quantized. - - Parameters - ---------- - attrs : tvm.ir.Attrs - Attributes of current convolution - inputs : tvm.relay.Expr - Grouped input symbols - tinfos : list - Input shape and dtype - out_type: type - The output type - - Note - ---- - Unlike other TOPI functions, this function operates on both graph level and operator level. - """ - return None - - -@tvm.target.generic_func -def qnn_requantize_alter_layout(_attrs, _inputs, _tinfos, _out_type): - """Change requantize layout. - - Parameters - ---------- - attrs : tvm.ir.Attrs - Attributes of current convolution - inputs : tvm.relay.Expr - Grouped input symbols - tinfos : list - Input shape and dtype - out_type: type - The output type - - Note - ---- - Unlike other TOPI functions, this function operates on both graph level and operator level. - """ - return None - - -@tvm.target.generic_func -def qnn_dense_alter_layout(_attrs, _inputs, _tinfos, _out_type): - """Change qnn.dense layout. - Not to change by default - - Parameters - ---------- - attrs : tvm.ir.Attrs - Attributes of current dense op - inputs : tvm.relay.Expr - Grouped input symbols - tinfos : list - Input shape and dtype - out_type: type - The output type - """ - return None diff --git a/python/tvm/topi/transform.py b/python/tvm/topi/transform.py index c1f5bce94870..6101c2e57d21 100644 --- a/python/tvm/topi/transform.py +++ b/python/tvm/topi/transform.py @@ -473,28 +473,6 @@ def take(a, indices, axis=None, batch_dims=0, mode="clip"): return cpp.take(a, indices, int(batch_dims), int(axis), mode) -@tvm.target.generic_func -def take_legalize(attrs, inputs, types): - """Legalizes dyn.topk op. - - Parameters - ---------- - attrs : tvm.ir.Attrs - Attributes of current op - inputs : list of tvm.relay.Expr - The args of the Relay expr to be legalized - types : list of types - List of input and output types - Returns - ------- - result : tvm.relay.Expr - The legalized expr - """ - if tvm.relay.ty.is_dynamic(types[0]): - return tvm.relay.take(tvm.relay.annotation.stop_fusion(inputs[0]), inputs[1], **attrs) - return None - - def gather(data, axis, indices): """Gather values along given axis from given indices. diff --git a/rust/tvm-rt/Cargo.toml b/rust/tvm-rt/Cargo.toml index e813c6941921..789c15a6be80 100644 --- a/rust/tvm-rt/Cargo.toml +++ b/rust/tvm-rt/Cargo.toml @@ -46,10 +46,7 @@ use-rpc = ["tvm-sys/use-rpc"] use-threads = ["tvm-sys/use-threads"] use-llvm = ["tvm-sys/use-llvm"] use-stackvm-runtime = ["tvm-sys/use-stackvm-runtime"] -use-graph-runtime = ["tvm-sys/use-graph-runtime"] -use-graph-runtime-debug = ["tvm-sys/use-graph-runtime-debug"] use-openmp = ["tvm-sys/use-openmp"] -use-relay-debug = ["tvm-sys/use-relay-debug"] use-rtti = ["tvm-sys/use-rtti"] use-mscv-mt = ["tvm-sys/use-mscv-mt"] use-install-dev = ["tvm-sys/use-install-dev"] diff --git a/rust/tvm-rt/src/graph_rt.rs b/rust/tvm-rt/src/graph_rt.rs deleted file mode 100644 index 53f3210aa742..000000000000 --- a/rust/tvm-rt/src/graph_rt.rs +++ /dev/null @@ -1,108 +0,0 @@ -/* - * 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. - */ - -use std::convert::TryInto; - -use crate::Function; -use crate::{function::Result, ByteArray, Device, Module, NDArray}; - -/// An instance of the C++ graph executor. -/// -/// An efficient and light weight runtime for static deep learning models. -pub struct GraphRt { - /// The backing graph executor module which exposes a set of packed functions - /// which can be invoked by a client. - /// - /// In the graph executor module, it exposes create, load_params, set_input, get_output, and run. - module: Module, -} - -impl GraphRt { - /// Create a graph executor directly from a runtime module. - pub fn from_module(module: Module, dev: Device) -> Result { - let default: Box Result> = - module.get_function("default", false)?.into(); - - Ok(Self { - module: default(dev)?, - }) - } - - /// Create a graph executor from the deprecated graph, lib, dev triple. - pub fn create_from_parts(graph: &str, lib: Module, dev: Device) -> Result { - let runtime_create_fn = Function::get("tvm.graph_executor.create").unwrap(); - - let runtime_create_fn_ret = runtime_create_fn.invoke(vec![ - graph.into(), - (&lib).into(), - (&dev.device_type).into(), - // NOTE you must pass the device id in as i32 because that's what TVM expects - (dev.device_id as i32).into(), - ]); - - let graph_executor_module: Module = runtime_create_fn_ret?.try_into()?; - Ok(Self { - module: graph_executor_module, - }) - } - - /// Load the parameters of the model into the runtime. - pub fn load_params

(&mut self, params: P) -> Result<()> - where - P: Into, - { - let load_param_fn = self.module.get_function("load_params", false)?; - - let params: ByteArray = params.into(); - - load_param_fn.invoke(vec![(¶ms).into()])?; - - Ok(()) - } - - /// Set the input with name `name` with the value of `input`. - pub fn set_input(&mut self, name: &str, input: NDArray) -> Result<()> { - let ref set_input_fn = self.module.get_function("set_input", false)?; - - set_input_fn.invoke(vec![name.into(), (&input).into()])?; - Ok(()) - } - - /// Run the graph module, once setting parameters and inputs. - pub fn run(&mut self) -> Result<()> { - let ref run_fn = self.module.get_function("run", false)?; - - // execute the run function. Note that it has no argument - run_fn.invoke(vec![])?; - Ok(()) - } - - /// Extract the ith output from the graph executor and returns it. - pub fn get_output(&mut self, i: i64) -> Result { - let get_output_fn = self.module.get_function("get_output", false)?; - get_output_fn.invoke(vec![i.into()])?.try_into() - } - - /// Extract the ith output from the graph executor and write the results into output. - pub fn get_output_into(&mut self, i: i64, output: NDArray) -> Result<()> { - let get_output_fn = self.module.get_function("get_output", false)?; - get_output_fn.invoke(vec![i.into(), (&output).into()])?; - Ok(()) - } -} diff --git a/rust/tvm-rt/src/lib.rs b/rust/tvm-rt/src/lib.rs index 3b7d066e7b78..0e9936318ded 100644 --- a/rust/tvm-rt/src/lib.rs +++ b/rust/tvm-rt/src/lib.rs @@ -53,7 +53,6 @@ pub mod array; pub mod device; pub mod errors; pub mod function; -pub mod graph_rt; pub mod map; pub mod module; pub mod ndarray; diff --git a/rust/tvm-sys/Cargo.toml b/rust/tvm-sys/Cargo.toml index e31ae66881dc..70daf3b388e4 100644 --- a/rust/tvm-sys/Cargo.toml +++ b/rust/tvm-sys/Cargo.toml @@ -39,10 +39,7 @@ use-rpc = [] use-threads = [] use-llvm = [] use-stackvm-runtime = [] -use-graph-runtime = [] -use-graph-runtime-debug = [] use-openmp = [] -use-relay-debug = [] use-rtti = [] use-mscv-mt = [] use-install-dev = [] diff --git a/src/arith/domain_touched.cc b/src/arith/domain_touched.cc index d2c5d79a0960..8c7c33bcc3ee 100644 --- a/src/arith/domain_touched.cc +++ b/src/arith/domain_touched.cc @@ -135,9 +135,9 @@ Region DomainTouched(const Stmt& stmt, const Buffer& buffer, bool consider_loads return BufferTouchedDomain(stmt).FindUnion(buffer, consider_loads, consider_stores); } -Map DomainTouchedAccessMap(const PrimFunc& func) { +Map> DomainTouchedAccessMap(const PrimFunc& func) { auto buffer_access_map = BufferTouchedDomain(func->body).GetAccessedBufferRegions(); - Map ret; + Map> ret; auto& buffer_map = func->buffer_map; for (auto& var : func->params) { auto& buffer = buffer_map[var]; @@ -153,11 +153,11 @@ Map DomainTouchedAccessMap(const PrimFunc& func) { combined.push_back(Array(touch)); } - std::vector fields; + runtime::Array fields; fields.push_back(loads); fields.push_back(stores); fields.push_back(combined); - ret.Set(buffer, runtime::ADT::Tuple(fields)); + ret.Set(buffer, fields); } return ret; } diff --git a/src/auto_scheduler/auto_schedule.cc b/src/auto_scheduler/auto_schedule.cc deleted file mode 100755 index 41aa49c77193..000000000000 --- a/src/auto_scheduler/auto_schedule.cc +++ /dev/null @@ -1,85 +0,0 @@ -/* - * 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 auto_scheduler/auto_schedule.cc - * \brief The user interface and tuning options of the TVM auto-scheduler. - */ - -#include -#include - -#include "utils.h" - -namespace tvm { -namespace auto_scheduler { - -TVM_REGISTER_NODE_TYPE(TuningOptionsNode); - -TuningOptions::TuningOptions(int num_measure_trials, int early_stopping, int num_measures_per_round, - int verbose, ProgramBuilder builder, ProgramRunner runner, - Optional> measure_callbacks) { - auto node = make_object(); - node->num_measure_trials = num_measure_trials; - node->early_stopping = early_stopping; - node->num_measures_per_round = num_measures_per_round; - node->verbose = verbose; - node->builder = std::move(builder); - node->runner = std::move(runner); - node->measure_callbacks = std::move(measure_callbacks); - data_ = std::move(node); -} - -std::pair> AutoSchedule(SearchPolicy search_policy, - TuningOptions tuning_options) { - // Create a ProgramMeasurer to handle the schedule build and performance measure - ProgramMeasurer measurer = - ProgramMeasurer(tuning_options->builder, tuning_options->runner, - tuning_options->measure_callbacks, tuning_options->verbose); - // Search for the best schedule - State state = - search_policy->Search(tuning_options->num_measure_trials, tuning_options->early_stopping, - tuning_options->num_measures_per_round, measurer); - if (state.defined()) { - return search_policy->search_task->compute_dag.ApplySteps(state->transform_steps); - } else { - StdCout(tuning_options->verbose) - << "No valid state found in this search round. Check if it has traversed all of the " - << "search space." << std::endl; - // Return the default schedule - return {te::Schedule(search_policy->search_task->compute_dag->ops), - search_policy->search_task->compute_dag->tensors}; - } -} - -TVM_REGISTER_GLOBAL("auto_scheduler.TuningOptions") - .set_body_typed([](int num_measure_trials, int early_stopping, int num_measures_per_round, - int verbose, ProgramBuilder builder, ProgramRunner runner, - Optional> measure_callbacks) { - return TuningOptions(num_measure_trials, early_stopping, num_measures_per_round, verbose, - builder, runner, measure_callbacks); - }); - -TVM_REGISTER_GLOBAL("auto_scheduler.AutoSchedule") - .set_body_typed([](SearchPolicy search_policy, TuningOptions tuning_options) { - auto [sch, return_tensors] = AutoSchedule(search_policy, tuning_options); - return Array{sch, return_tensors}; - }); -} // namespace auto_scheduler -} // namespace tvm diff --git a/src/auto_scheduler/compute_dag.cc b/src/auto_scheduler/compute_dag.cc deleted file mode 100644 index 82e439cddbc2..000000000000 --- a/src/auto_scheduler/compute_dag.cc +++ /dev/null @@ -1,1542 +0,0 @@ -/* - * 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 auto_scheduler/compute_dag.cc - * \brief Compute declaration graph and its related analysis tools. - */ - -#include -#include -#include -#include -#include -#include -#include -#include -#include -#include -#include -#include - -#include -#include -#include -#include -#include -#include - -#include "../arith/pattern_match.h" -#include "../relay/transforms/auto_scheduler_layout_rewrite.h" -#include "search_policy/utils.h" -#include "utils.h" - -namespace tvm { -namespace auto_scheduler { - -using namespace tvm::tir; - -template -using OperationMap = AccessAnalyzerNode::OperationMap; -using OperationSet = std::unordered_set; - -TVM_REGISTER_NODE_TYPE(ComputeDAGNode); - -// Topo-sort ops from tensors according to their read-write relations. -Array TopoSortOps(const Array& tensors) { - std::unordered_map degree; - std::unordered_map> edge_set; - std::unordered_map priority; - std::unordered_set visited; - - // traverse to build edge_set and count degree - std::vector stack; - stack.reserve(tensors.size()); - for (const auto& x : tensors) { - stack.push_back(x->op.operator->()); - } - - int ct = 0; - while (!stack.empty()) { - const te::OperationNode* op = stack.back(); - stack.pop_back(); - if (visited.count(op)) { - continue; - } - - priority[op] = ct; - ct++; - visited.insert(op); - - if (op->IsInstance()) { - degree[op] = 0; - } else if (auto cop = GetRef(op).as()) { - const Array& input_tensors = cop->InputTensors(); - degree[op] = input_tensors.size(); - for (const auto& ten : input_tensors) { - edge_set[ten->op.operator->()].push_back(op); - stack.push_back(ten->op.operator->()); - } - } else { - LOG(FATAL) << "Unsupported op " << GetRef(op); - } - } - - // topo sort - Array ops; - - using Item = std::pair; - auto cmp = [](const Item& left, const Item& right) { return left.second < right.second; }; - std::priority_queue, decltype(cmp)> queue(cmp); - for (const auto& iter : degree) { - if (iter.second == 0) { - queue.push(Item(iter.first, priority[iter.first])); - } - } - - ops.reserve(degree.size()); - while (!queue.empty()) { - Item item = queue.top(); - queue.pop(); - ops.push_back(GetRef(item.first)); - for (const auto& dst : edge_set[item.first]) { - degree[dst] -= 1; - if (degree[dst] == 0) { - queue.push(Item(dst, priority[dst])); - } - } - } - - return ops; -} - -// Extract all tensor accesses in an expr -class ReadAccessExtractor : public StmtExprVisitor { - public: - void Extract(PrimExpr expr) { this->VisitExpr(expr); } - - void VisitExpr_(const CallNode* op) final { - if (op->op.same_as(builtin::if_then_else())) { - has_branch = true; - } - StmtExprVisitor::VisitExpr_(op); - } - - void VisitExpr_(const ProducerLoadNode* op) final { - read_access[Downcast(op->producer)->op].emplace_back(op->indices.begin(), - op->indices.end()); - StmtExprVisitor::VisitExpr_(op); - } - - void VisitStmt_(const IfThenElseNode* op) final { - has_branch = true; - StmtExprVisitor::VisitStmt_(op); - } - - void VisitExpr_(const SelectNode* op) final { - has_branch = true; - StmtExprVisitor::VisitExpr_(op); - } - - // All read accesses to all operations - // The innermost vector stores multi-dimensional indices. - // The middle vector stores possible multiple accesses - OperationMap>> read_access; - // Whether this expression has branch - bool has_branch{false}; -}; - -// Returns whether the expr equals to the var with an optional const shift -bool IsConstShiftEqual(const Var& var, const PrimExpr& expr) { - arith::PVar x; - arith::PVar c; - - if (((x + c).Match(expr) || (x - c).Match(expr) || (c + x).Match(expr) || x.Match(expr)) && - x.Eval().same_as(var)) { - return true; - } - return false; -} - -// Return whether the access to an operation is a simple access -// (i.e. all index is just a variable with an optional constant shift) -// For example, A[i][j], A[i+1][j] are simple accesses but A[i][j+i] is not. -bool IsSimpleAccess(const te::Operation& op, const std::vector& indices, - bool* axis_missing, bool* axis_duplicated, bool* same_order) { - auto cop = op.as(); - if (cop == nullptr) { - return false; - } - - std::vector index_to_var_idx; - std::vector var_idx_ct(cop->axis.size(), 0); - - for (const auto& expr : indices) { - if (!is_const_int(expr)) { - bool found = false; - for (size_t i = 0; i < cop->axis.size(); ++i) { - if (IsConstShiftEqual(cop->axis[i]->var, expr)) { - index_to_var_idx.push_back(i); - var_idx_ct[i]++; - found = true; - break; - } - } - if (!found) { - return false; - } - } - } - - *axis_missing = false; // Some axes are missing - *axis_duplicated = false; // Some axes appear more than once - *same_order = true; // The axis order is the same as op->axis - for (int ct : var_idx_ct) { - if (ct == 0) { - *axis_missing = true; - } else if (ct > 1) { - *axis_duplicated = true; - } - } - for (size_t i = 1; i < index_to_var_idx.size(); ++i) { - if (index_to_var_idx[i] < index_to_var_idx[i - 1]) { - *same_order = false; - break; - } - } - - return true; -} - -// Gather all VarNodes in an expr -void GatherVars(const PrimExpr& expr, std::unordered_set* vars) { - PostOrderVisit(expr, [&vars](const ObjectRef& node) { - if (const VarNode* op = node.as()) { - vars->insert(op); - } - }); -} - -// Check whether an expr has expensive operations (e.g. exp) -bool HasExpensiveOp(const PrimExpr& expr) { - bool found = false; - PostOrderVisit(expr, [&found](const ObjectRef& node) { - if (const CallNode* op = node.as()) { - if (op->op.as()->name == "tir.exp") { - found = true; - } - } - }); - return found; -} - -AccessAnalyzer::AccessAnalyzer(const Array& tensors) { - auto node = make_object(); - OperationMap has_branch; - - // Get all ops in topological order - node->ops_topo_order = TopoSortOps(tensors); - - arith::Analyzer analyzer; - - // Build read & write access map - for (const auto& op : node->ops_topo_order) { - if (op->IsInstance()) { - node->read_from[op] = OperationMap>>(); - } else if (auto cop = op.as()) { - ReadAccessExtractor extractor; - for (const auto& exp : cop->body) { - extractor.Extract(exp); - } - - // read_by and read_from map - for (const auto& iter : extractor.read_access) { - std::vector>& accesses = node->read_by[iter.first][op]; - accesses.insert(accesses.begin(), iter.second.begin(), iter.second.end()); - } - - node->read_from[op] = std::move(extractor.read_access); - has_branch[op] = extractor.has_branch; - - // compute number of common outer iterators - for (const auto& pair : node->read_from[op]) { - const te::Operation& producer = pair.first; - const std::vector>& access_list = pair.second; - const Array& output_shape = op->output_shape(0); - const Array& producer_shape = producer->output_shape(0); - - int n_common; - for (n_common = 0; - n_common < static_cast(std::min(output_shape.size(), producer_shape.size())); - n_common++) { - if (!is_zero(analyzer.Simplify(output_shape[n_common] - producer_shape[n_common]))) { - break; - } - - bool injective = true; - for (const auto& access : access_list) { - if (!IsConstShiftEqual(cop->axis[n_common]->var, access[n_common])) { - injective = false; - break; - } - } - - if (!injective) { - break; - } - } - - node->num_common_outer_iterators[op][producer] = n_common; - node->num_common_outer_iterators[producer][op] = n_common; - } - } else { - LOG(FATAL) << "Invalid op: " << op; - } - } - - // Do some static analysis on ComputeOps - for (const auto& op : node->ops_topo_order) { - if (op->IsInstance()) { - node->is_simple_access[op] = true; - node->needs_multi_level_tiling[op] = false; - node->is_strictly_inlineable[op] = false; - node->is_output[op] = false; - } else if (auto cop = op.as()) { - // check whether this op is element-wise and strict-inlineable - bool is_simple_access = true; - bool is_strictly_inlineable = true; - - bool axis_missing, axis_duplicated, same_order; - for (const auto& pair : node->read_from[op]) { - const std::vector>& access_list = pair.second; - for (const auto& access : access_list) { - if (!auto_scheduler::IsSimpleAccess(op, access, &axis_missing, &axis_duplicated, - &same_order)) { - is_simple_access = false; - is_strictly_inlineable = false; - break; - } - if (!same_order || axis_duplicated) { - // do not strictly inline transpose - is_strictly_inlineable = false; - } - } - if (!is_simple_access) { - break; - } - } - - // don't strictly inline expensive op (e.g. exp) - bool has_expensive_op = false; - for (const auto& expr : cop->body) { - has_expensive_op |= HasExpensiveOp(expr); - } - if (has_expensive_op || has_branch[op]) { - is_strictly_inlineable = false; - } - - // constant tensor is strict-inlineable - if (node->read_from[op].empty()) { - is_strictly_inlineable = true; - } - - node->is_simple_access[op] = is_simple_access; - node->is_strictly_inlineable[op] = is_strictly_inlineable; - - // check whether the op needs multi-level tiling - bool needs_multi_level_tiling = false; - int n_missing = 0; - - for (const auto& pair : node->read_from[op]) { - const std::vector>& access_list = pair.second; - std::unordered_set vars; - for (const std::vector& access : access_list) { - for (const PrimExpr& expr : access) { - GatherVars(expr, &vars); - } - } - - for (const auto& axis : cop->axis) { - if (GetIntImm(axis->dom->extent) > 1 && vars.count(axis->var.get()) == 0) { - n_missing++; - break; - } - } - - if (n_missing >= 2 || (n_missing >= 1 && !cop->reduce_axis.empty())) { - needs_multi_level_tiling = true; - break; - } - } - - // do not perform multi-level tiling on "fake reduction" with const tensors - if (op->attrs.count(SearchPolicyKey::simplify_const_tensor_indices)) { - needs_multi_level_tiling = false; - } - - node->needs_multi_level_tiling[op] = needs_multi_level_tiling; - - // check whether the op is output - node->is_output[op] = node->read_by[op].empty(); - } else { - LOG(FATAL) << "Invalid op" << op; - } - } - - data_ = std::move(node); -} - -bool AccessAnalyzer::NeedsMultiLevelTiling(const te::Operation& op) const { - return operator->()->needs_multi_level_tiling.at(op); -} - -bool AccessAnalyzer::IsOutput(const te::Operation& op) const { - return operator->()->is_output.at(op); -} - -bool AccessAnalyzer::IsSimpleAccess(const te::Operation& op) const { - return operator->()->is_simple_access.at(op); -} - -bool AccessAnalyzer::IsStrictlyInlineable(const te::Operation& op) const { - return operator->()->is_strictly_inlineable.at(op); -} - -OperationSet AccessAnalyzer::GetConsumers(const State& state, const te::Operation& op) const { - OperationSet inlined_ops; - for (const auto& stage : state->stages) { - if (stage->compute_at == ComputeAtKind::kInlined) { - inlined_ops.insert(stage->op); - } - } - - OperationSet consumers; - std::function collect; - collect = [this, &collect, &inlined_ops, &consumers](const te::Operation& op) { - for (const auto& iter : operator->()->read_by.at(op)) { - if (inlined_ops.count(iter.first)) { - collect(iter.first); - } else { - consumers.insert(iter.first); - } - } - }; - - collect(op); - return consumers; -} - -OperationSet AccessAnalyzer::GetDirectProducers(const te::Operation& op) const { - OperationSet producers; - for (const auto& iter : operator->()->read_from.at(op)) { - producers.insert(iter.first); - } - return producers; -} - -OperationSet AccessAnalyzer::GetProducers(const State& state, const te::Operation& op) const { - OperationSet inlined_ops; - for (const auto& stage : state->stages) { - if (stage->compute_at == ComputeAtKind::kInlined) { - inlined_ops.insert(stage->op); - } - } - - OperationSet producers; - std::function collect; - collect = [this, &collect, &inlined_ops, &producers](const te::Operation& op) { - for (const auto& iter : operator->()->read_from.at(op)) { - if (inlined_ops.count(iter.first)) { - collect(iter.first); - } else { - producers.insert(iter.first); - } - } - }; - - collect(op); - return producers; -} - -int AccessAnalyzer::GetNumCommonOuterIterator(const te::Operation& op, - const te::Operation& target_op) const { - int ret = INT32_MAX; - bool meet = false; - - std::function traverse; - traverse = [this, &traverse, &target_op, &ret, &meet](const te::Operation& cur_op, int cur_num) { - if (cur_op == target_op) { - ret = std::min(ret, cur_num); - meet = true; - return; - } - - for (const auto& iter : operator->()->read_by.at(cur_op)) { - traverse( - iter.first, - std::min(cur_num, operator->()->num_common_outer_iterators.at(cur_op).at(iter.first))); - } - }; - - traverse(op, op->output_shape(0).size()); - return meet ? ret : 0; -} - -bool AccessAnalyzer::ElementWiseMatch(const te::Operation& op, - const te::Operation& target_op) const { - te::Operation cur_op = op; - while (cur_op != target_op) { - const AccessAnalyzerNode::OperationMap>>& map = - operator->()->read_by.at(cur_op); - - if (map.size() != 1) { - return false; - } - te::Operation next_op = map.begin()->first; - - // Check condition 1: They have the same output size - auto p_cur = cur_op.as(); - auto p_next = next_op.as(); - if (p_cur == nullptr || p_next == nullptr) { - return false; - } - - Array output_shape = p_cur->output_shape(0); - for (int i = 1; i < p_cur->num_outputs(); ++i) { - if (!IntArrayEqual(p_cur->output_shape(i), output_shape)) { - return false; - } - } - for (int i = 0; i < p_next->num_outputs(); ++i) { - if (!IntArrayEqual(p_next->output_shape(i), output_shape)) { - return false; - } - } - - // Check condition 2: The read is elementwise - const std::vector> reads = map.begin()->second; - bool is_simple_access, axis_missing, axis_duplicated, same_order; - for (const auto& read : reads) { - is_simple_access = auto_scheduler::IsSimpleAccess(next_op, read, &axis_missing, - &axis_duplicated, &same_order); - if (!is_simple_access || axis_missing || axis_duplicated || !same_order) { - return false; - } - } - - cur_op = std::move(next_op); - } - return true; -} - -// Estimate the number of float operations in an expression -class FlopEstimator : public ExprFunctor { - public: - double EstimateFlop(const Array& ops) { - double ret = 0; - for (const auto& op : ops) { - if (auto pop = op.as()) { - if (pop->attrs.count("FLOP")) { - // Use user-provided FLOP - ObjectRef annotation = pop->attrs["FLOP"]; - auto value = [&]() -> int64_t { - if (auto runtime_int = annotation.as()) { - return runtime_int->value; - } else if (auto int_imm = annotation.as()) { - return int_imm->value; - } else { - LOG(FATAL) << "FLOP annotation must be an integer, " - << "but was an object of type " << annotation->GetTypeKey(); - } - }(); - - ret += value; - } else { - // Estimate by parsing the compute body - double num_element = AxisLengthProd(pop->axis); - if (num_element == -1) { - fail_ = true; - break; - } - cur_type_code_ = pop->output_dtype(0).code(); - double op_per_element = 0; - for (const auto& x : pop->body) { - op_per_element += VisitExpr(x); - } - ret += num_element * op_per_element; - } - } else if (op->IsInstance()) { - {} // do nothing - } else { - LOG(FATAL) << "Invalid op type " << op; - } - } - - return fail_ ? -1 : ret; - } - - double VisitExpr_(const ReduceNode* op) final { - uint64_t num_iter = 1; - for (const auto& x : op->axis) { - if (auto imm = x->dom->extent.as()) { - num_iter *= imm->value; - } else { - fail_ = true; - num_iter = -1; - } - } - double body_flop = 0; - for (size_t i = 0; i < op->combiner->result.size(); ++i) { - body_flop += VisitExpr(op->combiner->result[i]); - body_flop += VisitExpr(op->source[i]); - } - return num_iter * body_flop; - } - - double VisitExpr_(const FloatImmNode* op) final { return 0.0; } - double VisitExpr_(const IntImmNode* op) final { return 0.0; } - double VisitExpr_(const ProducerLoadNode* op) final { return 0.0; } - - double VisitExpr_(const CastNode* op) final { return VisitExpr(op->value); } - double VisitExpr_(const VarNode* op) final { return 0.0; } - - double VisitExpr_(const SelectNode* op) final { - return VisitExpr(op->condition) + - std::max(VisitExpr(op->true_value), VisitExpr(op->false_value)); - } - -// Index calculations (e.g., the "i + j" expression in A[i + j]) are not counted in FLOPS. -#define VisitBinary(Node) \ - double VisitExpr_(const Node* op) final { \ - double base = 1.0; \ - if ((op->a->dtype.code() != cur_type_code_) && (op->b->dtype.code() != cur_type_code_)) { \ - base = 0.0; \ - } \ - return base + VisitExpr(op->a) + VisitExpr(op->b); \ - } - -#define VisitUnary(Node) \ - double VisitExpr_(const Node* op) final { \ - double base = op->dtype.code() == cur_type_code_ ? 1.0 : 0.0; \ - return base + VisitExpr(op->a); \ - } - - VisitBinary(AddNode); - VisitBinary(SubNode); - VisitBinary(MulNode); - VisitBinary(DivNode); - VisitBinary(ModNode); - VisitBinary(FloorDivNode); - VisitBinary(FloorModNode); - VisitBinary(MaxNode); - VisitBinary(MinNode); - VisitBinary(EQNode); - VisitBinary(NENode); - VisitBinary(LTNode); - VisitBinary(LENode); - VisitBinary(GTNode); - VisitBinary(GENode); - VisitBinary(AndNode); - VisitBinary(OrNode); - VisitUnary(NotNode); - - double VisitExpr_(const CallNode* op) final { - double ret = 0.0; - for (const auto& x : op->args) { - ret += VisitExpr(x); - } - return ret; - } - - double VisitExprDefault_(const Object* op) final { - fail_ = true; - return -1.0; - } - - private: - bool fail_{false}; - int cur_type_code_; -}; - -void CheckComputeValidity(const te::Schedule& sch) { - // Check the validity of a compute definition: - // The name of each iterator should be unique. - for (auto stage : sch->stages) { - if (stage->op->IsInstance()) { - std::unordered_set names; - for (const auto& x : stage->leaf_iter_vars) { - ICHECK(!names.count(x->var->name_hint)) - << "Find duplicated iterator names in the compute definition: " << x->var->name_hint - << ". Please use different names for different iterators."; - names.insert(x->var->name_hint); - } - } - } -} - -ComputeDAG::ComputeDAG(Array tensors) { - auto node = make_object(); - node->tensors = std::move(tensors); - node->access_analyzer = AccessAnalyzer(node->tensors); - - Array out_ops; - for (const auto& op : node->access_analyzer->ops_topo_order) { - if (node->access_analyzer.IsOutput(op)) { - out_ops.push_back(op); - } - } - te::Schedule sch = te::create_schedule(out_ops); - for (auto stage : sch->stages) { - node->ops.push_back(stage->op); - } - - // Make sure it is a valid compute definition - CheckComputeValidity(sch); - - node->flop_ct = FlopEstimator().EstimateFlop(node->ops); - node->init_state = State(node->ops); - data_ = std::move(node); -} - -ComputeDAG::ComputeDAG(const te::Schedule& sch) { - auto node = make_object(); - - // Make sure it is a valid compute definition - CheckComputeValidity(sch); - - // Initialize ops. Here we enforce the order of ops and stages are consistent - for (auto stage : sch->stages) { - node->ops.push_back(stage->op); - } - - // Collect input and output tensors - Array tensors; - for (auto stage : sch->stages) { - if (stage->op->IsInstance() || stage->is_output) { - for (auto i = 0; i < stage->op->num_outputs(); ++i) { - tensors.push_back(stage->op.output(i)); - } - } - } - node->tensors = std::move(tensors); - node->access_analyzer = AccessAnalyzer(node->tensors); - node->flop_ct = FlopEstimator().EstimateFlop(node->ops); - node->init_state = State(node->ops); - data_ = std::move(node); -} - -class IndexRewriter : public StmtExprMutator { - public: - IndexRewriter(const te::Operation& placeholder_op, const std::string& new_layout) - : placeholder_op_(placeholder_op) { - ParseKernelLayout(new_layout, &new_shape_, &new_names_); - } - - PrimExpr Rewrite(PrimExpr expr) { return this->VisitExpr(expr); } - - PrimExpr VisitExpr_(const ProducerLoadNode* op) final { - te::Tensor t = Downcast(op->producer); - if (t->op == placeholder_op_) { - std::unordered_map name_to_arg; - for (const auto& arg : op->indices) { - std::string axis_name; - if (const auto* int_imm = arg.as()) { - ICHECK_EQ(int_imm->value, 0); - axis_name = "IntImm"; - } else { - axis_name = AxisBaseName(CleanName(Downcast(arg)->name_hint)); - ICHECK_EQ(name_to_arg.count(axis_name), 0); - name_to_arg[axis_name] = arg; - } - } - - std::unordered_map div_factors; - std::vector r_new_args; - for (int i = new_names_.size() - 1; i >= 0; --i) { - auto ori_iter_name = new_names_[i]; - auto name_it = name_to_arg.find(ori_iter_name); - ICHECK(name_it != name_to_arg.end()); - PrimExpr ori_arg = name_it->second; - - PrimExpr mod_factor = new_shape_[i]; - - PrimExpr div_factor = 1; - if (div_factors.count(ori_iter_name)) { - div_factor = div_factors[ori_iter_name]; - } - div_factors[ori_iter_name] = div_factor * new_shape_[i]; - - PrimExpr new_arg = indexmod(indexdiv(ori_arg, div_factor), mod_factor); - - r_new_args.push_back(new_arg); - } - - Array new_args(std::make_move_iterator(r_new_args.rbegin()), - std::make_move_iterator(r_new_args.rend())); - return ProducerLoad(op->producer, new_args); - } - return GetRef(op); - } - - private: - const te::Operation& placeholder_op_; - Array new_shape_; - std::vector new_names_; -}; - -std::string GetOrigLayout(std::set* placeholder_axis_names, const te::Operation& op, - const te::Tensor& placeholder) { - ReadAccessExtractor extractor; - for (const auto& exp : op.as()->body) { - extractor.Extract(exp); - } - - std::ostringstream os; - uint32_t i = 0; - const auto& placeholder_op = placeholder->op; - ICHECK_GT(extractor.read_access.count(placeholder_op), 0); - for (const auto& ev : extractor.read_access[placeholder_op]) { - for (const auto& e : ev) { - std::string axis_name; - if (const auto* int_imm = e.as()) { - ICHECK_EQ(int_imm->value, 0); - axis_name = "IntImm"; - } else { - axis_name = AxisBaseName(CleanName(Downcast(e)->name_hint)); - } - - placeholder_axis_names->insert(axis_name); - os << placeholder->shape[i++] << axis_name; - } - } - - ICHECK_EQ(placeholder_axis_names->size(), placeholder->shape.size()); - std::string orig_layout = os.str(); - os.str(""); - ::tvm::relay::AutoSchedulerLayoutRewriter::global_ori_layouts_queue.push_back(orig_layout); - return orig_layout; -} - -std::string GetNewLayout(const State& state, const int stage_id, const Stage& stage, - const te::Operation& op, const te::Tensor& placeholder, - const std::set& placeholder_axis_names) { - std::ostringstream os; - Array stage_iters; - - auto attach_it = state->attach_map->stage_to_attach_iter.find(stage_id); - int attach_pos = -1; - size_t iters_before_attach = 0; - if (attach_it != state->attach_map->stage_to_attach_iter.end()) { - auto attach = attach_it->second; - const auto& attach_stage = state->stages[attach.first]; - attach_pos = attach.second; - stage_iters.insert(stage_iters.end(), attach_stage->iters.begin(), - attach_stage->iters.begin() + attach_pos + 1); - } - - stage_iters.insert(stage_iters.end(), stage->iters.begin(), stage->iters.end()); - - std::vector iters; - for (size_t i = 0; i < stage_iters.size(); ++i) { - const auto& iter = stage_iters[i]; - if (iter->orig_iters.empty()) { - iters.push_back(iter); - } else { - for (const Iterator& ori_iter : iter->orig_iters) { - iters.push_back(ori_iter); - } - } - if (static_cast(i) == attach_pos) { - iters_before_attach = iters.size(); - } - } - - std::vector new_names; - std::vector new_axis_names; - for (const Iterator& iter : iters) { - std::set ori_iter_names; - ExtractOriginalIterators(iter->name, &ori_iter_names); - // fused iters have been replaced with iter->orig_iters. - // So there should be only one ori iter name extracted from iter->name. - ICHECK_EQ(ori_iter_names.size(), 1); - auto ori_iter_name = AxisBaseName(*ori_iter_names.begin()); - new_axis_names.push_back(ori_iter_name); - } - for (size_t i = 0; i < new_axis_names.size(); ++i) { - auto iter = iters[i]; - std::string ori_iter_name; - if (i < iters_before_attach) { - ori_iter_name = new_axis_names[i + iters_before_attach]; - } else { - ori_iter_name = new_axis_names[i]; - } - if (placeholder_axis_names.count(ori_iter_name)) { - PrimExpr extent; - if (iter->range.defined()) { - extent = iter->range->extent; - } else { - // This iter is simplified by InferBound, so it must have a length of one. - extent = 1; - } - os << extent << ori_iter_name; - new_names.push_back(ori_iter_name); - } - } - std::string new_layout = os.str(); - os.str(""); - ::tvm::relay::AutoSchedulerLayoutRewriter::global_new_layouts_queue.push_back(new_layout); - return new_layout; -} - -ComputeDAG ComputeDAG::RewriteLayout(Array* transform_steps, - LayoutRewriteOption layout_rewrite) const { - CHECK(layout_rewrite != LayoutRewriteOption::NoRewrite) - << "Call ComputeDAG::RewriteLayout with NoRewrite."; - ComputeDAG new_dag = *this; - ComputeDAGNode* p_dag = new_dag.CopyOnWrite(); - - auto node = make_object(); - node->transform_steps = *transform_steps; - node->concrete = true; - const State& state = InferBound(State(node)); - - OperationSet handled_ops; - for (size_t stage_id = 0; stage_id < state->stages.size(); stage_id++) { - const auto& stage = state->stages[stage_id]; - - const te::Operation& op = stage->op; - if (!op->IsInstance()) { - continue; - } - const Map& attrs = op->attrs; - if (attrs.count(layout_free_placeholders_key) == 0) { - continue; - } - const ObjectRef& attr_value = attrs[layout_free_placeholders_key]; - for (const auto& placeholder : Downcast>(attr_value)) { - const auto& placeholder_op = placeholder->op; - - // Check whether this placeholder has already been handled - if (handled_ops.count(placeholder_op)) { - continue; - } - // Skip the op that is not direct consumer of this placeholder. - // This is usually caused by cache read/write. - bool direct_consumer = false; - for (auto& t : op->InputTensors()) { - if (t->op == placeholder_op) { - direct_consumer = true; - break; - } - } - if (!direct_consumer) { - continue; - } - handled_ops.insert(placeholder_op); - - // Process original layout - std::set placeholder_axis_names; - std::string origin_layout = GetOrigLayout(&placeholder_axis_names, op, placeholder); - Array origin_shape; - std::vector origin_axes; - ParseKernelLayout(origin_layout, &origin_shape, &origin_axes); - - // Process new layout - std::string new_layout = - GetNewLayout(state, stage_id, stage, op, placeholder, placeholder_axis_names); - Array new_shape; - std::vector new_axes; - ParseKernelLayout(new_layout, &new_shape, &new_axes); - - // Process op updates - te::Operation new_op_to_update; - if (layout_rewrite == LayoutRewriteOption::RewriteForPreTransformed) { - // Create new placeholder - new_op_to_update = te::PlaceholderOp(placeholder_op->name, new_shape, - placeholder_op.as()->dtype); - } else if (layout_rewrite == LayoutRewriteOption::InsertTransformStage) { - // Process index strides - std::unordered_map axes_stride; - for (const auto& i : origin_axes) { - axes_stride[i] = Integer(1); - } - Array new_stride(new_shape.size(), PrimExpr()); - PrimExpr temp = Integer(1); - for (int i = new_shape.size() - 1; i >= 0; i--) { - new_stride.Set(i, axes_stride[new_axes[i]]); - axes_stride[new_axes[i]] *= new_shape[i]; - } - - // Add an extra layout transform stage - const auto& layout_transform_tensor = te::compute( - new_shape, - [&new_stride, &placeholder_op, &origin_shape, &new_shape, &origin_axes, - &new_axes](const tvm::runtime::Array& indices) -> tvm::PrimExpr { - Array access_indices; - for (size_t indice_index = 0; indice_index < origin_shape.size(); indice_index++) { - PrimExpr temp = Integer(0); - for (size_t i = 0; i < new_shape.size(); i++) { - if (origin_axes[indice_index].compare(new_axes[i]) == 0) { - temp += indices[i] * new_stride[i]; - } - } - access_indices.push_back(temp); - } - return placeholder_op.output(0)(access_indices); - }, - "auto_scheduler_layout_transform"); - new_op_to_update = layout_transform_tensor->op; - - // Update the transform steps - for (size_t i = 0; i < transform_steps->size(); i++) { - Step step = (*transform_steps)[i]; - if (step->stage_id >= static_cast(stage_id)) { - step.CopyOnWrite()->stage_id++; - } - if (step->IsInstance()) { - auto compute_at_step = tvm::Downcast(step); - if (compute_at_step->target_stage_id >= static_cast(stage_id)) { - dynamic_cast(compute_at_step.CopyOnWrite())->target_stage_id++; - } - transform_steps->Set(i, std::move(compute_at_step)); - } else { - transform_steps->Set(i, std::move(step)); - } - } - - // Add schedule for the new added transform stage - Array to_fuse; - - if (new_shape.size() >= 5) { - to_fuse.push_back(0); - to_fuse.push_back(1); - to_fuse.push_back(2); - transform_steps->push_back(FuseStep(stage_id, to_fuse)); - } else if (new_shape.size() >= 3) { - to_fuse.push_back(0); - to_fuse.push_back(1); - transform_steps->push_back(FuseStep(stage_id, to_fuse)); - } - transform_steps->push_back(AnnotationStep(stage_id, 0, IteratorAnnotation::kParallel)); - } - - te::Operation new_compute_op, original_compute_op; - Array new_body; - IndexRewriter index_rewriter(placeholder_op, new_layout); - for (const auto& op : p_dag->ops) { - if (auto* pop = op.as()) { - bool need_update = false; - for (auto& t : op->InputTensors()) { - if (t->op == placeholder_op) { - need_update = true; - break; - } - } - if (need_update) { - for (const auto& body : pop->body) { - new_body.push_back(index_rewriter.Rewrite(body)); - } - original_compute_op = op; - CHECK(!new_compute_op.defined()); - auto new_attrs = pop->attrs; - new_attrs.Set("ori_placeholder_layout", tvm::String(origin_layout)); - new_attrs.Set("new_placeholder_layout", tvm::String(new_layout)); - new_compute_op = te::ComputeOp(pop->name, pop->tag, new_attrs, pop->axis, new_body); - } - } - } - - // construct the map from original_op to new_op - std::unordered_map updated_ops; - - Array original_ops = p_dag->ops; - p_dag->ops.clear(); - for (size_t i = 0; i < original_ops.size(); ++i) { - const auto& original_op = original_ops[i]; - if (original_op == placeholder_op) { - if (layout_rewrite == LayoutRewriteOption::InsertTransformStage) { - p_dag->ops.push_back(placeholder_op); - } - p_dag->ops.push_back(new_op_to_update); - updated_ops[placeholder_op] = new_op_to_update; - } else if (original_op == original_compute_op) { - p_dag->ops.push_back(new_compute_op); - updated_ops[original_compute_op] = new_compute_op; - } else { - p_dag->ops.push_back(original_op); - } - } - - ArrayNode* pops = p_dag->ops.CopyOnWrite(); - // Because ops is sorted in topo-order, only do one pass linear scan here. - for (size_t i = 0; i < pops->size(); ++i) { - const auto& original_op = Downcast(pops->at(i)); - if (auto* pop = original_op.as()) { - if (original_op == new_op_to_update) { - continue; - } - auto inputs = pop->InputTensors(); - std::unordered_map rmap; - for (auto input : inputs) { - auto it = updated_ops.find(input->op); - te::Operation new_op; - while (it != updated_ops.end()) { - new_op = it->second; - it = updated_ops.find(new_op); - } - if (new_op.defined()) { - int index = input->value_index; - rmap[input] = new_op.output(index); - } - } - if (!rmap.empty()) { - te::Operation new_op = pop->ReplaceInputs(original_op, rmap); - updated_ops[original_op] = new_op; - pops->SetItem(i, new_op); - } - } - } - - Array old_tensors = p_dag->tensors; - ArrayNode* p_tensors = p_dag->tensors.CopyOnWrite(); - for (size_t i = 0; i < old_tensors.size(); ++i) { - const auto& old_tensor = old_tensors[i]; - if (layout_rewrite != LayoutRewriteOption::RewriteForPreTransformed && - old_tensor->op->IsInstance()) { - continue; - } - auto it = updated_ops.find(old_tensor->op); - te::Operation new_op; - while (it != updated_ops.end()) { - new_op = it->second; - it = updated_ops.find(new_op); - } - if (new_op.defined()) { - auto index = old_tensor->value_index; - p_tensors->SetItem(i, new_op.output(index)); - } - } - } // end for placeholder - } // end for stage - p_dag->access_analyzer = AccessAnalyzer(p_dag->tensors); - - Array out_ops; - for (const auto& op : p_dag->access_analyzer->ops_topo_order) { - if (p_dag->access_analyzer.IsOutput(op)) { - out_ops.push_back(op); - } - } - - p_dag->ops.clear(); - te::Schedule sch = te::create_schedule(out_ops); - for (auto stage : sch->stages) { - p_dag->ops.push_back(stage->op); - } - p_dag->flop_ct = FlopEstimator().EstimateFlop(p_dag->ops); - p_dag->init_state = State(p_dag->ops); - - return new_dag; -} - -// Return whether a DAG has placeholders that are marked as "layout free". -bool HasLayoutFreeTensors(const ComputeDAG& dag) { - for (const auto& op : dag->ops) { - if (!op->IsInstance()) { - continue; - } - if (op->attrs.count(ComputeDAG::layout_free_placeholders_key)) { - return true; - } - } - - return false; -} - -std::pair> ComputeDAG::ApplySteps( - const Array& transform_steps, Array* stages, StageToAxesMap* stage_to_axes, - LayoutRewriteOption layout_rewrite) const { - if (layout_rewrite != LayoutRewriteOption::NoRewrite && HasLayoutFreeTensors(*this) && - !transform_steps.empty()) { - Array steps = transform_steps; - const auto& dag = RewriteLayout(&steps, layout_rewrite); - return dag.ApplySteps(steps); - } - - // Temporal object to be used if the input pointer is nullptr - Array temp_stages; - StageToAxesMap temp_stage_to_axes; - if (stages == nullptr) { - stages = &temp_stages; - } - if (stage_to_axes == nullptr) { - stage_to_axes = &temp_stage_to_axes; - } - Array out_ops; - for (const auto& op : operator->()->ops) { - if (operator->()->access_analyzer.IsOutput(op)) { - out_ops.push_back(op); - } - } - - // Create the initial schedule - te::Schedule schedule = te::create_schedule(out_ops); - - // init axes - for (const auto& x : operator->()->ops) { - const te::Stage& stage = schedule[x]; - stages->push_back(stage); - UpdateStageToAxesMap(stage, stage_to_axes); - } - - // Apply the history steps to TVM schedule - // Call each step's ApplyToSchedule method - for (const auto& step : transform_steps) { - StepApplyToSchedule(step, stages, stage_to_axes, &schedule, transform_steps); - } - - return std::make_pair(schedule, operator->()->tensors); -} - -String ComputeDAG::PrintStepsAsPython(const Array& transform_steps) const { - Array stages; - StageToAxesMap stage_to_axes; - Array out_ops; - for (const auto& op : operator->()->ops) { - if (operator->()->access_analyzer.IsOutput(op)) { - out_ops.push_back(op); - } - } - // Create the initial schedule - te::Schedule schedule = te::create_schedule(out_ops); - - // init axes - for (const auto& x : operator->()->ops) { - const te::Stage& stage = schedule[x]; - stages.push_back(stage); - UpdateStageToAxesMap(stage, &stage_to_axes); - } - - std::stringstream ss; - for (const auto& stage : stages) { - if (stage->op->IsInstance()) { - auto op_name = CleanName(stage->op->name); - - for (size_t i = 0; i < stage->leaf_iter_vars.size(); ++i) { - ss << CleanName(stage->leaf_iter_vars[i]->var->name_hint, op_name); - if (i != stage->leaf_iter_vars.size() - 1) { - ss << ", "; - } - } - ss << " = " - << "tuple(" << op_name << ".op.axis)" - << " + " - << "tuple(" << op_name << ".op.reduce_axis)\n"; - } - } - // Call each step's PrintAsPythonAPI method - for (const auto& step : transform_steps) { - ss << StepPrintAsPythonAPI(step, &stages, &stage_to_axes, &schedule, transform_steps); - } - - return ss.str(); -} - -String ComputeDAG::PrintDAG(bool simple_mode) const { - std::stringstream ss; - - for (const auto& op : operator->()->ops) { - if (op->IsInstance()) { - ss << op->name << " = PLACEHOLDER "; - if (!simple_mode) { - ss << op.output(0)->shape; - } - ss << "\n"; - } else if (auto pop = op.as()) { - for (size_t k = 0; k < pop->body.size(); ++k) { - ss << op->name << "("; - for (size_t i = 0; i < pop->axis.size(); i++) { - ss << pop->axis[i]->var->name_hint; - if (i != pop->axis.size() - 1) { - ss << ", "; - } - } - ss << ")"; - if (pop->body.size() > 1) { - ss << ".v" << k; - } - if (auto p_reduce = pop->body[k].as()) { - ICHECK_LT(k, p_reduce->combiner->result.size()); - PrimExpr combiner = p_reduce->combiner->result[k]; - if (combiner->IsInstance()) { - ss << " += " << AsLegacyRepr(p_reduce->source[0]) << "\n"; - } else if (combiner->IsInstance()) { - ss << " max= " << AsLegacyRepr(p_reduce->source[0]) << "\n"; - } else if (combiner->IsInstance()) { - ss << " min= " << AsLegacyRepr(p_reduce->source[0]) << "\n"; - } else if (combiner->IsInstance()) { - const auto& select = combiner.as(); - ss << " select(" << AsLegacyRepr(select->condition) // - << ", " << AsLegacyRepr(select->true_value) // - << ", " << AsLegacyRepr(select->false_value) // - << ")= (" << AsLegacyRepr(p_reduce->source[0]) // - << ',' << AsLegacyRepr(p_reduce->source[1]) // - << ")\n"; - } else { - ss << "reduce" << AsLegacyRepr(combiner) << "\n"; - } - } else { - auto call = pop->body[k].as(); - if (simple_mode && call) { - ss << " = " << AsLegacyRepr(call->op) << "\n"; - } else { - ss << " = " << AsLegacyRepr(pop->body[k]) << "\n"; - } - } - } - } else { - LOG(FATAL) << "Invalid op"; - } - } - return String(ss.str()); -} - -State ComputeDAG::InferBound(const State& state) const { - ICHECK(state->concrete) << "Only concrete state can be processed to get bound info."; - - State ret_state; - StateNode* pstate; - - if (state->stages.empty()) { - // If the input state is incomplete with empty operation stage - // create a new state from init_state and update it first - ret_state = operator->()->init_state; - pstate = ret_state.CopyOnWrite(); - pstate->transform_steps = state->transform_steps; - for (const auto& step : pstate->transform_steps) { - StepApplyToState(step, &ret_state, *this); - } - } else { - ret_state = state; - pstate = ret_state.CopyOnWrite(); - } - - Array stages; - StageToAxesMap stage_to_axes; - // Replay steps to tvm::Schedule - auto [sch, tensors] = ApplySteps(pstate->transform_steps, &stages, &stage_to_axes); - (void)tensors; // https://gcc.gnu.org/bugzilla/show_bug.cgi?id=81767 - sch = sch.normalize_for_feature_extraction(); - // Get bound information from TVM schedule - Map bounds = te::InferBound(sch); - - // Update the state bound information - for (size_t i = 0; i < pstate->stages.size(); ++i) { - const Stage& stage = pstate->stages[i]; - - if (stage->compute_at == ComputeAtKind::kInlined) { - continue; - } - - Array new_iters; - new_iters.reserve(stage->iters.size()); - // Get bound information from schedule - // the StageToAxesMap is used to find the corresponding IterVar in TVM schedule result - for (size_t j = 0; j < stage->iters.size(); ++j) { - const Iterator& iter = stage->iters[j]; - const IterVar& axis = stage_to_axes.at(stages[i])[j]; - - auto find_res = bounds.find(axis); - if (find_res != bounds.end()) { - new_iters.push_back(Iterator(iter->name, (*find_res).second, iter->iter_kind, - iter->annotation, &iter->orig_iters)); - } else { - LOG(FATAL) << "Infer bound fails"; - } - } - - pstate->stages.Set( - i, Stage(stage->op, stage->op_type, new_iters, stage->compute_at, stage->attrs)); - } - - return ret_state; -} - -Array ComputeDAG::InferBound(const Array& states) const { - Array out_states(states.size(), State()); - - support::parallel_for(0, states.size(), [this, &states, &out_states](int i) { - try { - out_states.Set(i, (states[i].defined()) ? this->InferBound(states[i]) : states[i]); - } catch (Error& e) { - LOG(WARNING) << "InferBound fails on the state:\n" - << states[i] << "\n" - << "with: " << e.what() << std::endl; - } - }); - - return out_states; -} - -ComputeDAG ComputeDAG::ReplayAndGetDAG(const Array& transform_steps) const { - auto [sch, old_tensors] = ApplySteps(transform_steps); - (void)old_tensors; // https://gcc.gnu.org/bugzilla/show_bug.cgi?id=81767 - return ComputeDAG(sch); -} - -TVM_STATIC_IR_FUNCTOR(ReprPrinter, vtable) - .set_dispatch([](const ObjectRef& ref, ReprPrinter* p) { - auto* node = static_cast(ref.get()); - for (const auto& op : node->ops_topo_order) { - p->stream << op << std::endl; - p->stream << "is_simple_access:\t" << node->is_simple_access.at(op) << "\t\t"; - p->stream << "needs_multi_level_tiling:\t" << node->needs_multi_level_tiling.at(op) - << std::endl; - p->stream << "is_strictly_inlinable:\t" << node->is_strictly_inlineable.at(op) << "\t"; - p->stream << "is_output:\t" << node->is_output.at(op) << std::endl; - p->stream << "Read from:\t"; - for (const auto& pair : node->read_from.at(op)) { - for (const auto& index : pair.second) { - p->stream << pair.first->name << Array(index) << ", "; - } - } - p->stream << std::endl; - p->stream << "Read by:\t"; - for (const auto& pair : node->read_by.at(op)) { - for (const auto& index : pair.second) { - p->stream << pair.first->name << Array(index) << ", "; - } - } - p->stream << std::endl; - p->stream << Chars('=', 50) << std::endl; - } - - AccessAnalyzer ana = GetRef(node); - p->stream << "ElementwiseMatch: \n"; - for (size_t i = 0; i < node->ops_topo_order.size(); ++i) { - for (size_t j = 0; j < node->ops_topo_order.size(); ++j) { - if (i == j) { - continue; - } - if (ana.ElementWiseMatch(node->ops_topo_order[i], node->ops_topo_order[j])) { - p->stream << node->ops_topo_order[i]->name << " -> " << node->ops_topo_order[j]->name - << std::endl; - } - } - } - p->stream << Chars('=', 50) << std::endl; - - p->stream << "NumCommonOuterIterators: \n"; - for (const auto& src_pair : node->num_common_outer_iterators) { - for (const auto& dst_pair : src_pair.second) { - p->stream << src_pair.first->name << " " << dst_pair.first->name << " " << dst_pair.second - << std::endl; - } - } - p->stream << Chars('=', 50) << std::endl; - }); - -TVM_STATIC_IR_FUNCTOR(ReprPrinter, vtable) - .set_dispatch([](const ObjectRef& ref, ReprPrinter* p) { - auto* node = static_cast(ref.get()); - auto dag = GetRef(node); - auto dag_str = dag.PrintDAG(); - p->stream << dag_str; - }); - -Array GetShapeFromRewrittenLayout(String rewritten_layout, Array axis_names) { - Array shape; - std::vector extracted_names; - topi::parse_auto_scheduler_layout(rewritten_layout, &shape, &extracted_names); - - Array ret(axis_names.size(), 1); - - size_t ct = 0; - for (size_t i = 0; i < axis_names.size(); ++i) { - for (size_t j = 0; j < extracted_names.size(); ++j) { - if (axis_names[i] == extracted_names[j]) { - ret.Set(i, ret[i] * shape[j]); - ct++; - } - } - } - - CHECK_EQ(ct, extracted_names.size()) << "The number or names of axes do not match"; - - return ret; -} - -TVM_REGISTER_GLOBAL("auto_scheduler.ComputeDAG") - .set_body_typed([](Optional> tensors, Optional sch) { - if (sch) { - return ComputeDAG(sch.value()); - } - ICHECK(tensors) << "Both tensors and schedule are null"; - return ComputeDAG(tensors.value()); - }); - -TVM_REGISTER_GLOBAL("auto_scheduler.ComputeDAGApplyStepsFromState") - .set_body_typed([](const ComputeDAG& dag, const State& state, int layout_rewrite) { - auto [sch, return_tensors] = dag.ApplySteps(state->transform_steps, nullptr, nullptr, - static_cast(layout_rewrite)); - return Array{sch, return_tensors}; - }); - -TVM_REGISTER_GLOBAL("auto_scheduler.ComputeDAGPrintPythonCodeFromState") - .set_body_typed([](const ComputeDAG& dag, const State& state) { - return dag.PrintStepsAsPython(state->transform_steps); - }); - -TVM_REGISTER_GLOBAL("auto_scheduler.ComputeDAGPrintDAG") - .set_body_typed([](const ComputeDAG& dag, bool simple_mode) { - return dag.PrintDAG(simple_mode); - }); - -TVM_REGISTER_GLOBAL("auto_scheduler.ComputeDAGInferBoundFromState") - .set_body_typed([](const ComputeDAG& dag, const State& state) { - return dag.InferBound(state); - }); - -TVM_REGISTER_GLOBAL("auto_scheduler.ComputeDAGRewriteLayoutFromState") - .set_body_typed([](const ComputeDAG& dag, const State& state) { - Array* transform_steps = const_cast*>(&state->transform_steps); - return dag.RewriteLayout(transform_steps, LayoutRewriteOption::RewriteForPreTransformed); - }); - -TVM_REGISTER_GLOBAL("auto_scheduler.RewriteIndexForNewLayout") - .set_body_typed([](const te::Operation& placeholder_op, const std::string& new_layout, - const PrimExpr& body) { - IndexRewriter index_rewriter(placeholder_op, new_layout); - return index_rewriter.Rewrite(body); - }); - -TVM_REGISTER_GLOBAL("auto_scheduler.RewriteTensorShape") - .set_body_typed([](te::Tensor tensor, Array new_shape) -> void { - ICHECK(tensor->op->IsInstance()); - te::PlaceholderOpNode* op = - const_cast(tensor->op.as()); - te::TensorNode* t = const_cast(tensor.get()); - op->shape = new_shape; - t->shape = new_shape; - }); - -TVM_REGISTER_GLOBAL("auto_scheduler.GetShapeFromRewrittenLayout") - .set_body_typed(GetShapeFromRewrittenLayout); - -} // namespace auto_scheduler -} // namespace tvm diff --git a/src/auto_scheduler/cost_model.cc b/src/auto_scheduler/cost_model.cc deleted file mode 100755 index 4ed5ca2bfbe8..000000000000 --- a/src/auto_scheduler/cost_model.cc +++ /dev/null @@ -1,173 +0,0 @@ -/* - * 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 auto_scheduler/cost_model.cc - * \brief Cost models that estimate the performance of programs - */ - -#include - -namespace tvm { -namespace auto_scheduler { - -TVM_REGISTER_OBJECT_TYPE(CostModelNode); -TVM_REGISTER_OBJECT_TYPE(RandomModelNode); -TVM_REGISTER_OBJECT_TYPE(PythonBasedModelNode); - -RandomModel::RandomModel() { - ObjectPtr node = make_object(); - const auto* f = runtime::Registry::Get("auto_scheduler.cost_model.random_fill_float"); - ICHECK(f != nullptr); - node->random_number_func = reinterpret_cast*>(f); - data_ = std::move(node); -} - -void RandomModelNode::Update(const Array& inputs, - const Array& results) {} - -void RandomModelNode::Predict(const SearchTask& task, const Array& states, - std::vector* scores) { - scores->resize(states.size()); - (*random_number_func)(states.size(), static_cast(scores->data())); -} - -PythonBasedModel::PythonBasedModel(PackedFunc update_func, PackedFunc predict_func, - PackedFunc predict_stage_func) { - auto node = make_object(); - node->update_func = std::move(update_func); - node->predict_func = std::move(predict_func); - node->predict_stage_func = std::move(predict_stage_func); - data_ = std::move(node); -} - -void PythonBasedModelNode::Update(const Array& inputs, - const Array& results) { - update_func(inputs, results); -} - -void PythonBasedModelNode::Predict(const SearchTask& task, const Array& states, - std::vector* scores) { - scores->resize(states.size()); - predict_func(task, states, static_cast(scores->data())); -} - -void PythonBasedModelNode::PredictStages(const SearchTask& task, const Array& states, - std::vector* state_scores, - std::vector>* stage_scores) { - size_t n_states = states.size(); - size_t n_stages = task->compute_dag->init_state->stages.size(); - std::vector flatten_scores; - // Allocate sufficient spaces. - flatten_scores.resize(n_states * n_stages * 2); - predict_stage_func(task, states, static_cast(flatten_scores.data())); - - /* For faster data copy between c++ and python, the python part returns scores in a - * single flatten array using a packed format. The c++ part then unpacks the flatten array. - * - * The packed format is: - * { - * float scores[N]; // scores[i] is the score for states[i]. - * int n_stage_0; // the number of stages in states[0] - * float stage_scores_0[[n_stage_0] // the scores for all stages in states[0] - * int n_stage_1; // the number of stages in states[1] - * float stage_scores_1[n_stage_1]; // the scores for all stages in states[1] - * ... - * int n_stage_i; // the number of stages in states[i] - * float stage_scores_1[n_stage_i]; // the scores for all stages in states[i] - * ... // until i == N - 1 - * } - * To implement this format, we also store int as float, so we can store all numbers - * into a single float array. - */ - - // Unpack flatten scores. - state_scores->clear(); - stage_scores->clear(); - - // Score of each states. - for (size_t i = 0; i < n_states; ++i) { - state_scores->push_back(flatten_scores[i]); - } - - // Score of each stage in each states. - size_t idx = n_states; - for (size_t i = 0; i < n_states; ++i) { - ICHECK_LE(idx, flatten_scores.size()); - - // Number of scored stages of this state. - int s_length = static_cast(flatten_scores[idx++]); - - if (s_length > 0) { - std::vector scores; - int offset = 0; - - if ((*state_scores)[i] > -INFINITY) { - // If the score is valid. Copy scored stages and assign 0 to placeholder - // and inlined stages. If the score is 0, meaning this state failed to - // be lowered. Just bypass to update offset. - for (const Stage& stage : states[i]->stages) { - if (stage->op_type == StageKind::kPlaceholder) { - scores.push_back(0); - continue; - } - if (stage->compute_at == ComputeAtKind::kInlined) { - scores.push_back(0); - continue; - } - scores.push_back(flatten_scores[idx + offset]); - offset++; - } - ICHECK_EQ(offset, s_length); - stage_scores->push_back(std::move(scores)); - } - idx += s_length; - } else { - // Cost model does not provide any stage score details. - stage_scores->push_back({}); - } - } -} - -TVM_REGISTER_GLOBAL("auto_scheduler.RandomModel").set_body_typed([]() { return RandomModel(); }); - -TVM_REGISTER_GLOBAL("auto_scheduler.PythonBasedModel") - .set_body_typed([](PackedFunc update_func, PackedFunc predict_func, - PackedFunc predict_stage_func) { - return PythonBasedModel(update_func, predict_func, predict_stage_func); - }); - -TVM_REGISTER_GLOBAL("auto_scheduler.CostModelUpdate") - .set_body_typed([](CostModel model, Array inputs, Array results) { - model->Update(inputs, results); - }); - -TVM_REGISTER_GLOBAL("auto_scheduler.CostModelPredict") - .set_body_typed([](CostModel model, SearchTask task, Array states) { - std::vector scores; - model->Predict(task, states, &scores); - Array ret; - for (auto x : scores) { - ret.push_back(FloatImm(DataType::Float(32), x)); - } - return ret; - }); - -} // namespace auto_scheduler -} // namespace tvm diff --git a/src/auto_scheduler/feature.cc b/src/auto_scheduler/feature.cc deleted file mode 100644 index 09255b5da539..000000000000 --- a/src/auto_scheduler/feature.cc +++ /dev/null @@ -1,1761 +0,0 @@ -/* - * 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 auto_scheduler/feature.cc - * \brief Feature extraction for the cost model - */ - -#include -#include -#include -#include -#include -#include -#include -#include -#include -#include -#include -#include -#include -#include - -#include -#include -#include -#include -#include -#include - -#include "search_policy/utils.h" -#include "utils.h" - -namespace tvm { -namespace auto_scheduler { - -using namespace tvm::tir; -using arith::Analyzer; -using arith::ConstIntBound; - -template -using BufferMap = std::unordered_map; - -// The number of samples to extract for arithmetic intensity curves -static const int ARITH_INTENSITY_CURVE_SAMPLE_N = 10; - -// Annotation position encoding -enum class AnnotationPosType : int { - kPosNone = 0, // Does not have this kind of annotation - kPosInnerSpatial = 1, // The annotated iterator is the innermost spatial iterator - kPosMiddleSpatial = 2, // The annotated iterator is a middle spatial iterator - kPosOuterSpatial = 3, // The annotated iterator is the outermost spatial iterator - kPosInnerReduce = 4, // The annotated iterator is the innermost reduce iterator - kPosMiddleReduce = 5, // The annotated iterator is a middle reduce iterator - kPosOuterReduce = 6, // The annotated iterator is the outermost reduce iterator - kPosMixed = 7 // The annotated iterator is a mixed space and reduce iterator -}; - -// Buffer access type -enum class BufferAccessType : int { kRead = 0, kWrite = 1, kReadWrite = 2, kUnknownRW = 3 }; - -// Accesses to a buffer -struct BufferAccess { - // data reuse type - BufferAccessType acc_type{BufferAccessType::kUnknownRW}; - // Use a two-dimensional array to store multiple multi-dimensional accesses. - // The innermost vector stores the multi-dimensional indices of one access. - std::vector> indices; -}; - -// Data reuse type -enum class ReuseType : int { kLoopMultipleRead = 0, kSerialMultipleReadWrite = 1, kNoReuse = 2 }; - -// Feature for an access of a buffer -struct BufferAccessFeature { - std::string buffer_name; // The name of the buffer - BufferAccessType acc_type; // The type of the access - float bytes; // The touched memory in bytes - float unique_bytes; // The touched unique memory in bytes - float lines; // The number of touched cache lines - float unique_lines; // The number touched unique cache lines - ReuseType reuse_type; // Tye type of data reuse - float reuse_dis_iter; // The reuse distance in iterator number - float reuse_dis_bytes; // The reuse distance in total touched bytes - float reuse_ct; // The reuse ratio - float bytes_d_reuse_ct; // bytes / reuse_ct - float unique_bytes_d_reuse_ct; // unique_bytes / reuse_ct - float lines_d_reuse_ct; // lines / reuse_ct - float unique_lines_d_reuse_ct; // unique_lines / reuse_ct - float stride; // The stride in access -}; - -// Feature set of a BufferStore statement -struct FeatureSet { - // Group 1: Computation related features - float float_mad; // The number of float MAD (Multiply–add) ops - float float_addsub; // The number of float add and sub ops - float float_mul; // The number of float multiply ops - float float_divmod; // The number of float div and mod ops - float float_cmp; // The number of float comparison ops - float float_math_func; // The number of float math func calls - float float_other_func; // The number of other float func calls - float int_mad; // The number of integer MAD (Multiply–add) ops - float int_addsub; // The number of integer add and sub ops - float int_mul; // The number of float multiply ops - float int_divmod; // The number of float div and mod ops - float int_cmp; // The number of float comparison ops - float int_math_func; // The number of float math func calls - float int_other_func; // The number of other float func calls - float bool_op; // The number of bool ops - float select_op; // The number of select ops - float vec_num; // The number of vectorized iterators - float vec_prod; // The product of the lengths of vectorized iterators - float vec_len; // The length of the innermost vectorized iterator - AnnotationPosType vec_type; // The type of vectorization position - float unroll_num; // The number of unrolled iterators - float unroll_prod; // The product of the lengths of vectorized iterators - float unroll_len; // The length of the innermost unrolled iterator - AnnotationPosType unroll_type; // The type of unroll position - float parallel_num; // The number of paralleled iterators - float parallel_prod; // The product of the lengths of paralleled iterators - float parallel_len; // The length of the innermost paralleled iterators - AnnotationPosType parallel_type; // The type of parallel position - float is_gpu; // Whether it is a GPU task - float blockIdx_x_len; // The length of blockIdx.x - float blockIdx_y_len; // The length of blockIdx.y - float blockIdx_z_len; // The length of blockIdx.z - float threadIdx_x_len; // The length of threadIdx.x - float threadIdx_y_len; // The length of threadIdx.y - float threadIdx_z_len; // The length of threadIdx.z - float vthread_len; // The length of virtual thread - - // Group 2: Buffer access related features (per buffer) - std::vector access_feas; - - // Group 3: Arithmetic intensity related features - float arith_intensity_curve[ARITH_INTENSITY_CURVE_SAMPLE_N]; // points sampled from the - // arithmetic intensity curve - - // Group 4: Allocation related features - float alloc_size; // The size of allocated buffer in bytes - float alloc_outer_prod; // The product of lengths of loops outside the scope of the allocation - float alloc_inner_prod; // The product of lengths of loops inside the score of the allocation - float alloc_prod; // alloc_outer_prod * alloc_inner_prod - - // Group 5: Outer scope related features - float outer_prod; // The product of lengths of outer loops - float num_loops; // The number of outer loops - float auto_unroll_max_step; // The value of pragma "auto_unroll_max_step" -}; - -// Return whether a var is in an expr -bool VarInExpr(const Var& var, const PrimExpr& expr) { - bool find = false; - - PostOrderVisit(expr, [&find, &var](const ObjectRef& node) { - if (find) { - return; - } - - if (const VarNode* op = node.as()) { - if (op == var.get()) { - find = true; - } - } - }); - - return find; -} - -// Get position encoding for annotation -AnnotationPosType GetAnnotationPosEncoding(const Var& var, const Array& spatial_args, - const Array& axis, - const Array& reduce_axis) { - // Try to match spatial args first - size_t find_i = 0; - size_t find_ct = 0; - for (size_t i = 0; i < spatial_args.size(); ++i) { - if (VarInExpr(var, spatial_args[i])) { - find_i = i; - find_ct += 1; - } - } - - if (find_ct == 0) { - // If it is not found in spacial args, then it is a reduce iterator. - // Use name to match - const std::string& var_name = var->name_hint; - for (size_t i = 0; i < reduce_axis.size(); ++i) { - if (var_name.find(reduce_axis[i]->var->name_hint) != std::string::npos) { - find_i = i; - find_ct++; - } - } - if (find_ct >= 1) { - if (find_i == 0) { - return AnnotationPosType::kPosInnerReduce; - } else if (find_i == reduce_axis.size() - 1) { - return AnnotationPosType::kPosOuterReduce; - } else { - return AnnotationPosType::kPosMiddleReduce; - } - } else { - // If the axis is not found in both spatial args and reduce axis, - // then this stage must compute_at somewhere under this axis and this axis is simplified out - // We assume it is an outer spatial - return AnnotationPosType::kPosOuterSpatial; - } - } else if (find_ct == 1) { - if (find_i == spatial_args.size() - 1) { - return AnnotationPosType::kPosInnerSpatial; - } else if (find_i == 0) { - return AnnotationPosType::kPosOuterSpatial; - } else { - return AnnotationPosType::kPosMiddleSpatial; - } - } else { - return AnnotationPosType::kPosMixed; - } -} - -// Return the maximum extent of a for loop -int64_t GetLoopExtent(const ForNode* node, const Analyzer& ana) { - int64_t bound = ana.const_int_bound(node->extent)->max_value; - if (bound == ConstIntBound::kPosInf) { - return 1; // Analyzer could not determine a valid bound, use 1 instead. - } else { - return bound; - } -} - -// Count math ops in an expr -class MathOpCounter : public StmtExprVisitor { - public: -#define VisitBinary(Type, float_ct, int_ct) \ - void VisitExpr_(const Type* op) final { \ - if (op->a.dtype().is_float() || op->a.dtype().is_bfloat16()) { \ - float_ct += op->a.dtype().lanes(); \ - } else { \ - int_ct += op->a.dtype().lanes(); \ - } \ - StmtExprVisitor::VisitExpr_(op); \ - } - - VisitBinary(AddNode, float_addsub, int_addsub); - VisitBinary(SubNode, float_addsub, int_addsub); - VisitBinary(MulNode, float_mul, int_mul); - VisitBinary(DivNode, float_divmod, int_divmod); - VisitBinary(ModNode, float_divmod, int_divmod); - VisitBinary(FloorDivNode, float_divmod, int_divmod); - VisitBinary(FloorModNode, float_divmod, int_divmod); - VisitBinary(MaxNode, float_cmp, int_cmp); - VisitBinary(MinNode, float_cmp, int_cmp); - VisitBinary(EQNode, float_cmp, int_cmp); - VisitBinary(NENode, float_cmp, int_cmp); - VisitBinary(LTNode, float_cmp, int_cmp); - VisitBinary(LENode, float_cmp, int_cmp); - VisitBinary(GTNode, float_cmp, int_cmp); - VisitBinary(GENode, float_cmp, int_cmp); - -#undef VisitBinary - - void VisitExpr_(const AndNode* op) final { - bool_op++; - StmtExprVisitor::VisitExpr_(op); - } - void VisitExpr_(const OrNode* op) final { - bool_op++; - StmtExprVisitor::VisitExpr_(op); - } - void VisitExpr_(const NotNode* op) final { - bool_op++; - StmtExprVisitor::VisitExpr_(op); - } - void VisitExpr_(const SelectNode* op) final { - select_op++; - StmtExprVisitor::VisitExpr_(op); - } - - void VisitExpr_(const CallNode* op) final { - auto* pop = op->op.as(); - ICHECK(pop != nullptr); - auto effect_kind = op_call_effect_[GetRef(pop)]; - bool is_pure = - effect_kind == CallEffectKind::kPure || effect_kind == CallEffectKind::kExprAnnotation; - - if (is_pure) { - if (op->dtype.is_float() || op->dtype.is_bfloat16()) { - float_math_func++; - } else { - int_math_func++; - } - } else { - if (op->dtype.is_float() || op->dtype.is_bfloat16()) { - float_other_func++; - } else { - int_other_func++; - } - } - StmtExprVisitor::VisitExpr_(op); - } - - // todo(merrymercy): Detect MAD (Multiply–add) - size_t float_mad{0}; // The number of float MAD (Multiply–add) ops - size_t float_addsub{0}; // The number of float add and sub ops - size_t float_mul{0}; // The number of float multiply ops - size_t float_divmod{0}; // The number of float div and mod ops - size_t float_cmp{0}; // The number of float comparison ops - size_t float_math_func{0}; // The number of float math func calls - size_t float_other_func{0}; // The number of other float func calls - size_t int_mad{0}; // The number of integer MAD (Multiply–add) ops - size_t int_addsub{0}; // The number of integer add and sub ops - size_t int_mul{0}; // The number of float multiply ops - size_t int_divmod{0}; // The number of float div and mod ops - size_t int_cmp{0}; // The number of float comparison ops - size_t int_math_func{0}; // The number of float math func calls - size_t int_other_func{0}; // The number of other float func calls - size_t bool_op{0}; // The number of bool ops - size_t select_op{0}; // The number of select ops - - OpAttrMap op_call_effect_ = Op::GetAttrMap("TCallEffectKind"); -}; - -// Extract all buffer accesses in an expr -class BufferAccessExtractor : public StmtExprVisitor { - public: - void ExtractReads(const PrimExpr& expr) { this->VisitExpr(expr); } - - void InsertAccess(const Var& buf, BufferAccessType acc_type, const Array& indices) { - BufferAccess& acc = buf_accesses[buf]; - acc.acc_type = acc_type; - acc.indices.push_back(std::vector(indices.begin(), indices.end())); - } - - void VisitExpr_(const BufferLoadNode* op) final { - AddAccess(op->buffer->data, op->indices); - StmtExprVisitor::VisitExpr_(op); - } - - void AddAccess(const Var& buffer, const Array& indices) { - BufferAccess& acc = buf_accesses[buffer]; - switch (acc.acc_type) { - case BufferAccessType::kRead: - break; - case BufferAccessType::kWrite: - acc.acc_type = BufferAccessType::kReadWrite; - break; - case BufferAccessType::kReadWrite: - break; - case BufferAccessType::kUnknownRW: - default: - acc.acc_type = BufferAccessType::kRead; - break; - } - - if (acc.acc_type != BufferAccessType::kReadWrite) { - // If a buffer is both read and written, in the tvm DSL, it must be a update, - // so the indices should be the same. Then we can skip appending indices for it. - // Otherwise we do the following. - buf_accesses[buffer].indices.push_back(std::vector(indices.begin(), indices.end())); - } - } - - BufferMap buf_accesses; -}; - -// Compute the coefficient for an loop iterator in an expression -// Note: we use an approximation strategy to find coefficient. -// Hopefully, it is faster than DetectLinearEquation and can handle more cases (non-linear) -class CoefficientExtractor : public StmtExprVisitor { - public: - void VisitExpr_(const MulNode* node) final { - StmtExprVisitor::VisitExpr_(node); - if (visited_var) { - if (!visited_add) { - if (auto a = node->a.as()) { - visited_mul = true; - stride = a->value; - } else if (auto b = node->b.as()) { - visited_mul = true; - stride = b->value; - } - } - } - } - - void VisitExpr_(const AddNode* node) final { - StmtExprVisitor::VisitExpr_(node); - if (visited_var) { - if (!visited_mul) { - visited_add = true; - stride = 1; - } - } - } - - void VisitExpr_(const VarNode* node) final { - if (node == var_) { - visited_var = true; - // This is a magic default stride in case our approximation strategy fails - stride = 2; - } - } - - int ExtractCoefficient(const PrimExpr& expr, const VarNode* var) { - visited_var = visited_mul = visited_add = false; - var_ = var; - - this->VisitExpr(expr); - - if (visited_var && !visited_mul && !visited_add) { - return 1; - } else { - return stride; - } - } - - bool visited_var{false}; - bool visited_mul{false}; - bool visited_add{false}; - int stride{0}; - - private: - const VarNode* var_{nullptr}; -}; - -// Compute stride for the accesses to a buffer -int64_t ComputeStride(const std::vector>& indices, - const std::vector& shape, const VarNode* stride_var) { - // Use stride of 1 for 0-dimensional buffers. 0-dim buffers has a single - // index access, so we have to check here. - if (shape.size() == 0) { - return 1; - } - int64_t min_stride = std::numeric_limits::max(); - bool find = false; - CoefficientExtractor extractor; - - for (const auto& index : indices) { - int64_t shape_stride = 1; - for (int i = static_cast(index.size()) - 1; i >= 0; i--) { - int coefficient = extractor.ExtractCoefficient(index[i], stride_var); - if (extractor.visited_var) { - find = true; - min_stride = std::min(min_stride, std::abs(coefficient) * shape_stride); - break; - } - shape_stride *= shape[i]; - } - } - - return find ? min_stride : 0; -} - -// Compute touched bytes and cache lines for accesses to a buffer -void ComputeRegion(const std::vector>& indices, arith::Analyzer* ana, - std::vector* region) { - region->clear(); - - if (indices.empty()) { - return; - } - - region->reserve(indices[0].size()); - - if (indices.size() == 1) { - for (const auto& index : indices[0]) { - ConstIntBound bound = ana->const_int_bound(index); - region->push_back(bound->max_value - bound->min_value + 1); - } - } else { - // future(lmzheng): implement a more accurate IntSet? - for (size_t i = 0; i < indices[0].size(); ++i) { - int64_t minimum = ConstIntBound::kPosInf, maximum = ConstIntBound::kNegInf; - for (size_t j = 0; j < indices.size(); ++j) { - ConstIntBound bound = ana->const_int_bound(indices[j][i]); - - minimum = std::min(minimum, bound->min_value); - maximum = std::max(maximum, bound->max_value); - } - region->push_back(maximum - minimum + 1); - } - } -} - -// Compute reuse distance and reuse ratio for accesses to a buffer -// return values: reuse_type, reuse_dis_iter, reuse_dis_bytes, reuse_ct -std::tuple ComputeReuse( - const Var& buf, const std::vector>& indices, - const std::vector& for_loop_stack, - const std::unordered_map>>>& - for_touch_regions, - const Analyzer& ana) { - float reuse_dis_iter = 1.0f; - float reuse_dis_bytes = -1.0f; - - for (int i = static_cast(for_loop_stack.size()) - 1; i >= 0; --i) { - const ForNode* cur_for = for_loop_stack[i]; - bool find = false; - - for (size_t j = 0; j < indices.size(); j++) { - for (size_t k = 0; k < indices[j].size(); k++) { - if (VarInExpr(cur_for->loop_var, indices[j][k])) { - find = true; - break; - } - } - if (find) { - break; - } - } - - int64_t extent = GetLoopExtent(for_loop_stack[i], ana); - if (find) { - // accumulate/update reuse distance - reuse_dis_iter *= extent; - reuse_dis_bytes = 0.0f; - for (const auto& iter : for_touch_regions.at(cur_for)) { - for (const auto& access : iter.second) { - reuse_dis_bytes += std::get<1>(access) * std::get<2>(access); - } - } - } else { - // Have LoopMultipleRead reuse - if (reuse_dis_bytes < 0) { - // For the reuse in the innermost axis, the above code won't be executed. - // So we compute bytes here - reuse_dis_bytes = 0.0f; - for (const auto& iter : for_touch_regions.at(cur_for)) { - for (const auto& access : iter.second) { - reuse_dis_bytes += 1 * std::get<2>(access); - } - } - } - return std::make_tuple(ReuseType::kLoopMultipleRead, reuse_dis_iter, reuse_dis_bytes, extent); - } - - const BufferMap>>& buffer_map = - for_touch_regions.at(cur_for); - - int serial_reuse = static_cast(buffer_map.at(buf).size()) - 1; - if (serial_reuse > 0) { - int64_t extent = GetLoopExtent(cur_for, ana); - - // Have SerialMultipleReadWrite reuse - reuse_dis_iter = std::numeric_limits::max(); - for (const auto& acc_info : buffer_map.at(buf)) { - reuse_dis_iter = std::min(reuse_dis_iter, static_cast(std::get<1>(acc_info))); - } - - reuse_dis_bytes = 0.0f; - for (const auto& iter : for_touch_regions.at(cur_for)) { - for (const auto& access : iter.second) { - reuse_dis_bytes += std::get<1>(access) * std::get<2>(access); - } - } - - return std::make_tuple(ReuseType::kSerialMultipleReadWrite, reuse_dis_iter / extent, - reuse_dis_bytes / extent, serial_reuse); - } - } - - return std::make_tuple(ReuseType::kNoReuse, 0, 0, 0); -} - -// Extract features for every BufferStore statement -// -// This visitor assumes that loop bounds do no depend on data or on parent loop -// bounds. For example, `for i in .. { for j in range(i, ..) }` would result in -// inaccurate features. This visitor also does not take conditionals into -// consideration when creating features. Each branch of the conditional is -// taken at the same time. -class PerStoreFeatureExtractor : public StmtExprVisitor { - public: - explicit PerStoreFeatureExtractor(int cache_line_size, const Map& existing_buffers) - : cache_line_size_(cache_line_size) { - for (const auto& buffer : existing_buffers) { - buffer_shapes[buffer.first] = buffer.second->shape; - buffer_dtypes[buffer.first] = buffer.second->dtype; - // Also need to add a reference from the buffers internal variable. This - // is usually how buffers are referenced within the body of a PrimFunc - buffer_shapes[buffer.second->data] = buffer.second->shape; - buffer_dtypes[buffer.second->data] = buffer.second->dtype; - } - } - - void VisitStmt_(const AttrStmtNode* node) final { - if (node->attr_key == tir::attr::thread_extent || node->attr_key == tir::attr::virtual_thread) { - const Var& var = node->node.as()->var; - int extent = GetIntImm(node->value); - - int* plen = nullptr; - - const std::string& name = var.get()->name_hint; - if (node->attr_key == tir::attr::thread_extent) { - if (name == "blockIdx.x") { - plen = &blockIdx_x_len_; - } else if (name == "blockIdx.y") { - plen = &block_idx_y_len_; - } else if (name == "blockIdx.z") { - plen = &block_idx_z_len_; - } else if (name == "threadIdx.x") { - plen = &threadIdx_x_len_; - } else if (name == "threadIdx.y") { - plen = &thread_idx_y_len_; - } else if (name == "threadIdx.z") { - plen = &thread_idx_z_len_; - } else { - LOG(FATAL) << "invalid thread itervar " + name; - } - } else { - plen = &vthread_len_; - } - - int extent_before = *plen; - if (node->attr_key == tir::attr::thread_extent) { - *plen = extent; - } else { - *plen *= extent; - } - - is_gpu_ = true; - - // make a fake for node for blockIdx.x or threadIdx.x - Stmt fake_for_node = For(var, 0, extent, ForKind::kParallel, node->body); - - outer_loop_prod_ *= extent; - for_loop_stack_.push_back(fake_for_node.as()); - variable_definition_stack_.push_back({}); - StmtExprVisitor::VisitStmt_(node); - variable_definition_stack_.pop_back(); - for_loop_stack_.pop_back(); - outer_loop_prod_ /= extent; - - *plen = extent_before; - } else if (node->attr_key == "pragma_auto_unroll_max_step") { - int value = GetIntImm(node->value); - - int16_t old_value = cur_auto_unroll_max_step_; - cur_auto_unroll_max_step_ = value; - StmtExprVisitor::VisitStmt_(node); - cur_auto_unroll_max_step_ = old_value; - } else { - StmtExprVisitor::VisitStmt_(node); - } - } - - void VisitStmt_(const ForNode* node) final { - ana_.Bind(node->loop_var, Range::FromMinExtent(node->min, node->extent)); - int64_t loop_extent = GetLoopExtent(node, ana_); - - if (node->kind == ForKind::kVectorized) { - vec_for_stack_.push_back(node); - } else if (node->kind == ForKind::kUnrolled) { - unroll_for_stack_.push_back(node); - } else if (node->kind == ForKind::kParallel) { - parallel_for_stack_.push_back(node); - } - - outer_loop_prod_ *= loop_extent; - for_loop_stack_.push_back(node); - variable_definition_stack_.push_back({}); - StmtExprVisitor::VisitStmt_(node); - variable_definition_stack_.pop_back(); - for_loop_stack_.pop_back(); - outer_loop_prod_ /= loop_extent; - - if (node->kind == ForKind::kVectorized) { - vec_for_stack_.pop_back(); - } else if (node->kind == ForKind::kUnrolled) { - unroll_for_stack_.pop_back(); - } else if (node->kind == ForKind::kParallel) { - parallel_for_stack_.pop_back(); - } - } - - void VisitExpr_(const BufferLoadNode* node) final { - // Store buffer shape/dtype. It may already be stored. - buffer_shapes[node->buffer->data] = node->buffer->shape; - buffer_dtypes[node->buffer->data] = node->buffer->dtype; - StmtExprVisitor::VisitExpr_(node); - } - - void VisitStmt_(const BufferStoreNode* node) final { - // Store buffer shape/dtype. It may already be stored. - buffer_shapes[node->buffer->data] = node->buffer->shape; - buffer_dtypes[node->buffer->data] = node->buffer->dtype; - - MathOpCounter math_op_counter; - math_op_counter(node->value); - std::vector mem_bytes_list; - std::vector compute_ops_list; - double cur_compute_ops; - - // Group 1: Computation related features - ExtractComputationFeature(node->buffer->data, node->indices, math_op_counter); - - // Group 2: Buffer access related features (per buffer) - ExtractBufferAccessFeature(node->buffer->data, node->indices, node->value, math_op_counter, - &cur_compute_ops, &compute_ops_list, &mem_bytes_list); - - // Group 3: Arithmetic intensity related features - ExtractArithmeticIntensityFeature(node->buffer->data, cur_compute_ops, compute_ops_list, - mem_bytes_list); - - // Group 4: Allocation related features - ExtractOuterScopeFeature(node->buffer->data); - } - - void VisitStmt_(const BufferRealizeNode* node) final { - // Store buffer shape/dtype. It may already be stored. - buffer_shapes[node->buffer->data] = node->buffer->shape; - buffer_dtypes[node->buffer->data] = node->buffer->dtype; - StmtExprVisitor::VisitStmt_(node); - - // Group 5: Outer scope related features - ExtractAllocationFeature(node); - } - - void VisitStmt_(const AllocateNode* node) final { - buffer_dtypes[node->buffer_var] = node->dtype; - buffer_shapes[node->buffer_var] = node->extents; - StmtExprVisitor::VisitStmt_(node); - - // Group 5: Outer scope related features - ExtractAllocationFeature(node); - } - - void VisitStmt_(const LetStmtNode* node) final { - // TODO(tkonolige): add arithmetic counts from this statement to counts of inner stores. - ana_.Bind(node->var, node->value); - ICHECK(variable_definition_stack_.size() > 0) - << "Variable definition outside of a for loop is not handled by feature extraction"; - variable_definition_stack_.back().push_back(std::make_tuple(node->var, node->value)); - StmtExprVisitor::VisitStmt_(node); - } - - // Extract computation related features (group 1) - void ExtractComputationFeature(const Var& buffer, const Array& indices, - const MathOpCounter& math_op_counter) { - FeatureSet& fea = buffer_features[buffer]; - - // Computation related features - fea.float_mad += outer_loop_prod_ * math_op_counter.float_mad; - fea.float_addsub += outer_loop_prod_ * math_op_counter.float_addsub; - fea.float_mul += outer_loop_prod_ * math_op_counter.float_mul; - fea.float_divmod += outer_loop_prod_ * math_op_counter.float_divmod; - fea.float_cmp += outer_loop_prod_ * math_op_counter.float_cmp; - fea.float_math_func += outer_loop_prod_ * math_op_counter.float_math_func; - fea.float_other_func += outer_loop_prod_ * math_op_counter.float_other_func; - fea.int_mad += outer_loop_prod_ * math_op_counter.int_mad; - fea.int_addsub += outer_loop_prod_ * math_op_counter.int_addsub; - fea.int_mul += outer_loop_prod_ * math_op_counter.int_mul; - fea.int_divmod += outer_loop_prod_ * math_op_counter.int_divmod; - fea.int_math_func += outer_loop_prod_ * math_op_counter.int_math_func; - fea.int_cmp += outer_loop_prod_ * math_op_counter.int_cmp; - fea.int_other_func += outer_loop_prod_ * math_op_counter.int_other_func; - fea.bool_op += outer_loop_prod_ * math_op_counter.bool_op; - fea.select_op += outer_loop_prod_ * math_op_counter.select_op; - - fea.vec_len = fea.unroll_len = fea.parallel_len = 0.0f; - fea.vec_type = fea.unroll_type = fea.parallel_type = AnnotationPosType::kPosNone; - - fea.vec_num = vec_for_stack_.size(); - if (!vec_for_stack_.empty()) { - fea.vec_len = GetLoopExtent(vec_for_stack_.back(), ana_); - fea.vec_prod = 1.0; - for (const ForNode* pfor : vec_for_stack_) { - fea.vec_prod *= GetLoopExtent(pfor, ana_); - } - fea.vec_type = AnnotationPosType::kPosMixed; - // todo(merrymercy): this feature requires operation (tvm.compute) information - // GetAnnotationPosEncoding(vec_for_stack_.back()->loop_var, - // node->args, pcompute->axis, pcompute->reduce_axis); - } - - fea.unroll_num = unroll_for_stack_.size(); - if (!unroll_for_stack_.empty()) { - fea.unroll_len = GetLoopExtent(unroll_for_stack_.back(), ana_); - fea.unroll_prod = 1.0; - for (const ForNode* pfor : unroll_for_stack_) { - fea.unroll_prod *= GetLoopExtent(pfor, ana_); - } - fea.unroll_type = AnnotationPosType::kPosMixed; - // GetAnnotationPosEncoding(unroll_for_stack_.back()->loop_var, - // node->args, pcompute->axis, pcompute->reduce_axis); - } - - fea.parallel_num = parallel_for_stack_.size(); - if (!parallel_for_stack_.empty()) { - fea.parallel_len = GetLoopExtent(parallel_for_stack_.back(), ana_); - fea.parallel_prod = 1.0; - for (const ForNode* pfor : parallel_for_stack_) { - fea.parallel_prod *= GetLoopExtent(pfor, ana_); - } - fea.parallel_type = AnnotationPosType::kPosMixed; - // GetAnnotationPosEncoding(parallel_for_stack_.back()->loop_var, - // node->args, pcompute->axis, pcompute->reduce_axis); - } - - // GPU threads - fea.is_gpu = is_gpu_; - fea.blockIdx_x_len = blockIdx_x_len_; - fea.blockIdx_y_len = block_idx_y_len_; - fea.blockIdx_z_len = block_idx_z_len_; - fea.threadIdx_x_len = threadIdx_x_len_; - fea.threadIdx_y_len = thread_idx_y_len_; - fea.threadIdx_z_len = thread_idx_z_len_; - fea.vthread_len = vthread_len_; - } - - // Extract buffer access related features (group 2) - void ExtractBufferAccessFeature(const Var& buffer, const Array& indices, - const PrimExpr& value, const MathOpCounter& math_op_counter, - double* cur_compute_ops, std::vector* compute_ops_list, - std::vector* mem_bytes_list) { - FeatureSet& fea = buffer_features[buffer]; - - // Extract all buffer accesses - std::vector acc_feas; - BufferAccessExtractor buf_extractor; - buf_extractor.InsertAccess(buffer, BufferAccessType::kWrite, indices); - buf_extractor.ExtractReads(value); - - mem_bytes_list->reserve(for_loop_stack_.size()); - compute_ops_list->reserve(for_loop_stack_.size()); - - *cur_compute_ops = math_op_counter.float_mad + math_op_counter.float_addsub + - math_op_counter.float_mul + math_op_counter.float_divmod + - math_op_counter.float_cmp + math_op_counter.float_math_func + - math_op_counter.float_other_func; - - ICHECK_EQ(for_loop_stack_.size(), variable_definition_stack_.size()) - << "variable_definition_stack_ should mirror for_loop_stack_ in size"; - std::vector tmp_region; - for (int i = static_cast(for_loop_stack_.size()) - 1; i >= 0; i--) { - const ForNode* p_for = for_loop_stack_[i]; - - // Construct a local analyzer context which contains definitions (for and - // let) from innermost loops up to and including `i`. For loop variable - // definitions in loops more outer than `i` are set to 1 so that we can - // get per-loop-iteration features. Note that we add these definitions - // from outermost to innermost because inner definitions may depend on - // outer ones. - Analyzer local_analyzer; - for (int j = 0; j < i; j++) { - local_analyzer.Bind(for_loop_stack_.at(j)->loop_var, - Range::FromMinExtent(for_loop_stack_.at(j)->min, 1)); - } - for (int j = i; j < static_cast(for_loop_stack_.size()); j++) { - local_analyzer.Bind( - for_loop_stack_.at(j)->loop_var, - Range::FromMinExtent(for_loop_stack_.at(j)->min, for_loop_stack_.at(j)->extent)); - for (auto definition : variable_definition_stack_.at(j)) { - local_analyzer.Bind(std::get<0>(definition), std::get<1>(definition)); - } - } - - // Note, here we do overwrite. - // So if there are multiple BufferStoreNode, the last one will overwrite the first few. - // e.g. The update part in gemm will overwrite the init part. - BufferMap>>& buffer_regions_map = - for_touch_regions_[p_for]; - - int64_t mem_bytes = 0; - for (const auto& x : buf_extractor.buf_accesses) { - const Var& t = x.first; - const BufferAccess& acc = x.second; - - ComputeRegion(acc.indices, &local_analyzer, &tmp_region); - int64_t touched_size = ElementProduct(tmp_region); - touched_size = std::max(0, touched_size); - buffer_regions_map[t].push_back( - std::make_tuple(acc.acc_type, touched_size, buffer_dtypes.at(t).bytes())); - mem_bytes += touched_size * buffer_dtypes.at(t).bytes(); - } - - mem_bytes_list->push_back(mem_bytes); - *cur_compute_ops *= GetLoopExtent(for_loop_stack_[i], local_analyzer); - compute_ops_list->push_back(*cur_compute_ops); - } - - // Buffer access related features (per buffer) - for (const auto& x : buf_extractor.buf_accesses) { - const Var& t = x.first; - const BufferAccess& acc = x.second; - - std::vector int_shape; - for (const auto& dim : buffer_shapes.at(t)) { - int_shape.push_back(GetIntImm(dim)); - } - - size_t ele_bytes = buffer_dtypes.at(t).bytes(); - - // calculate bytes - float bytes = outer_loop_prod_ * ele_bytes; - float unique_bytes; - - // calculate cache lines - int64_t stride; - float lines; - float unique_lines; - - if (for_loop_stack_.empty()) { - unique_bytes = ele_bytes; - stride = 0; - lines = 1.0f; - unique_lines = 1.0f; - } else { - unique_bytes = static_cast( - std::get<1>(for_touch_regions_[for_loop_stack_.front()][t].front())) * - ele_bytes; - - stride = 0; - int64_t reduce_ratio = 1; - - int i; - for (i = static_cast(for_loop_stack_.size()) - 1; i >= 0; i--) { - stride = ComputeStride(acc.indices, int_shape, for_loop_stack_[i]->loop_var.get()); - if (stride != 0) { - break; - } - reduce_ratio *= GetLoopExtent(for_loop_stack_.back(), ana_); - } - - lines = outer_loop_prod_ / reduce_ratio * - std::min(1.0f, 1.0f * stride * ele_bytes / cache_line_size_); - lines = std::max(lines, 1.0f); - - // convert `stride` back to the stride of the innermost iterator - stride = (i == static_cast(for_loop_stack_.size()) - 1 ? stride : 0); - - float n_continuous = ele_bytes; - for (int i = std::min(static_cast(tmp_region.size()) - 1, - static_cast(int_shape.size()) - 1); - i >= 0; i--) { - if (tmp_region[i] == int_shape[i]) { - n_continuous *= tmp_region[i]; - break; - } - } - unique_lines = unique_bytes / std::min(n_continuous, static_cast(cache_line_size_)); - unique_lines = std::max(unique_lines, 1.0f); - } - - auto [reuse_type, reuse_dis_iter, reuse_dis_bytes, reuse_ct] = - ComputeReuse(t, acc.indices, for_loop_stack_, for_touch_regions_, ana_); - - acc_feas.emplace_back(); - BufferAccessFeature& acc_fea = acc_feas.back(); - - // TODO(tkonolige): save buffer names and use those instead? - acc_fea.buffer_name = t->name_hint; - acc_fea.acc_type = acc.acc_type; - acc_fea.stride = stride; - acc_fea.bytes = bytes; - acc_fea.unique_bytes = unique_bytes; - acc_fea.lines = lines; - acc_fea.unique_lines = unique_lines; - acc_fea.reuse_type = reuse_type; - acc_fea.reuse_dis_iter = reuse_dis_iter; - acc_fea.reuse_dis_bytes = reuse_dis_bytes; - acc_fea.reuse_ct = reuse_ct; - if (acc_fea.reuse_ct > 0.5) { - acc_fea.bytes_d_reuse_ct = bytes / reuse_ct; - acc_fea.unique_bytes_d_reuse_ct = unique_bytes / reuse_ct; - acc_fea.lines_d_reuse_ct = lines / reuse_ct; - acc_fea.unique_lines_d_reuse_ct = unique_lines / reuse_ct; - } else { - // no reuse, multiply by a magic number '2' - acc_fea.bytes_d_reuse_ct = bytes * 2; - acc_fea.unique_bytes_d_reuse_ct = unique_bytes * 2; - acc_fea.lines_d_reuse_ct = lines * 2; - acc_fea.unique_lines_d_reuse_ct = unique_lines * 2; - } - } - - fea.access_feas = acc_feas; - } - - // Extract arithmetic intensity related feature (group 3) - void ExtractArithmeticIntensityFeature(const Var& buffer, double cur_compute_ops, - const std::vector& compute_ops_list, - const std::vector& mem_bytes_list) { - FeatureSet& fea = buffer_features[buffer]; - - // Compute arithmetic intensity curve (y axis : arithmetic intensity, x axis : flops). - // We use piecewise linear interpolation to fit this curve. - int pt = 0; - if (cur_compute_ops <= 0 || compute_ops_list.empty()) { - std::fill(fea.arith_intensity_curve, - fea.arith_intensity_curve + ARITH_INTENSITY_CURVE_SAMPLE_N, 0.0); - } else { - for (size_t i = 0; i < ARITH_INTENSITY_CURVE_SAMPLE_N; ++i) { - float cur_compute_ops = compute_ops_list.back() * (i + 1) / ARITH_INTENSITY_CURVE_SAMPLE_N; - while (compute_ops_list[pt] < cur_compute_ops - 1e-4) { - pt++; - } - ICHECK_LT(pt, compute_ops_list.size()); - - float value; - if (pt == 0) { - value = compute_ops_list[pt] / mem_bytes_list[pt]; - } else { - float base = compute_ops_list[pt - 1] / mem_bytes_list[pt - 1]; - float slope = (compute_ops_list[pt] / mem_bytes_list[pt] - - compute_ops_list[pt - 1] / mem_bytes_list[pt - 1]) / - (compute_ops_list[pt] - compute_ops_list[pt - 1]); - value = base + slope * (cur_compute_ops - compute_ops_list[pt - 1]); - } - fea.arith_intensity_curve[i] = value; - } - } - } - - // Extract allocation related features (group 4) - void ExtractAllocationFeature(const BufferRealizeNode* node) { - FeatureSet& fea = buffer_features[node->buffer->data]; - - float allocation_size = 1.0f; - for (const auto& x : node->bounds) { - allocation_size *= GetIntImm(x->extent); - } - // allocation feature - fea.alloc_size = allocation_size * node->buffer->dtype.bytes(); - fea.alloc_prod = allocation_size * outer_loop_prod_; - fea.alloc_outer_prod = outer_loop_prod_; - fea.alloc_inner_prod = fea.outer_prod / outer_loop_prod_; - } - - void ExtractAllocationFeature(const AllocateNode* node) { - FeatureSet& fea = buffer_features[node->buffer_var]; - - float allocation_size = 1.0f; - for (const auto& x : node->extents) { - // TODO(tkonolige): will not handle dynamic shape - allocation_size *= GetIntImm(x); - } - // allocation feature - fea.alloc_size = allocation_size * node->dtype.bytes(); - fea.alloc_prod = allocation_size * outer_loop_prod_; - fea.alloc_outer_prod = outer_loop_prod_; - fea.alloc_inner_prod = fea.outer_prod / outer_loop_prod_; - } - - // Extract outer scope related features (group 5) - void ExtractOuterScopeFeature(const Var& buffer) { - FeatureSet& fea = buffer_features[buffer]; - - fea.outer_prod = outer_loop_prod_; - fea.num_loops = for_loop_stack_.size(); - fea.auto_unroll_max_step = cur_auto_unroll_max_step_; - } - - // Stores FeatureSet for every buffer - BufferMap buffer_features; - - private: - // The shared arithmetic analyzer - Analyzer ana_; - - // The product of outer loop - float outer_loop_prod_ = 1.0f; - - // The stacks to store parent loops during DFS - std::vector for_loop_stack_; - std::vector parallel_for_stack_; - std::vector vec_for_stack_; - std::vector unroll_for_stack_; - std::vector>> variable_definition_stack_; - - // GPU-related features - bool is_gpu_{false}; - int blockIdx_x_len_{1}; - int block_idx_y_len_{1}; - int block_idx_z_len_{1}; - int threadIdx_x_len_{1}; - int thread_idx_y_len_{1}; - int thread_idx_z_len_{1}; - int vthread_len_{1}; - int16_t cur_auto_unroll_max_step_{0}; - - // Store touch region information for all for loops. The format of this nested map: - // For a loop, for all its touched buffers, for all different accesses to the buffers, - // its (access type, number of touched elements, number of bytes of single element) - std::unordered_map>>> - for_touch_regions_; - - // The default cache line size in bytes - const int cache_line_size_ = 64; - - // Storage of buffer shape and dtype information. Needed because Load/Store - // nodes only do not contain this information. - BufferMap> buffer_shapes; - BufferMap buffer_dtypes; -}; - -// shifted log to incorporate the property that log2p(0) = 0 -inline float log2p(float x) { return x < 0 ? -std::log2(-x + 1) : std::log2(x + 1); } - -void GetPerStoreFeature(const PrimFunc& func, int cache_line_size, int max_n_bufs, - std::vector* ret, bool log_scale) { - PerStoreFeatureExtractor extractor(cache_line_size, func->buffer_map); - extractor(func->body); - - auto slog = log_scale ? log2p : [](float x) { return x; }; - - ret->push_back(extractor.buffer_features.size()); - - for (const auto& x : extractor.buffer_features) { - const FeatureSet& fea_set = x.second; - - /***** Group 1: Computation related features *****/ - ret->push_back(slog(fea_set.float_mad)); - ret->push_back(slog(fea_set.float_addsub)); - ret->push_back(slog(fea_set.float_mul)); - ret->push_back(slog(fea_set.float_divmod)); - ret->push_back(slog(fea_set.float_cmp)); - ret->push_back(slog(fea_set.float_math_func)); - ret->push_back(slog(fea_set.float_other_func)); - ret->push_back(slog(fea_set.int_mad)); - ret->push_back(slog(fea_set.int_addsub)); - ret->push_back(slog(fea_set.int_mul)); - ret->push_back(slog(fea_set.int_divmod)); - ret->push_back(slog(fea_set.int_cmp)); - ret->push_back(slog(fea_set.int_math_func)); - ret->push_back(slog(fea_set.int_other_func)); - ret->push_back(slog(fea_set.bool_op)); - ret->push_back(slog(fea_set.select_op)); - - ret->push_back(slog(fea_set.vec_num)); - ret->push_back(slog(fea_set.vec_prod)); - ret->push_back(slog(fea_set.vec_len)); - for (int i = 0; i <= static_cast(AnnotationPosType::kPosMixed); i++) { - ret->push_back(i == static_cast(fea_set.vec_type)); - } - - ret->push_back(slog(fea_set.unroll_num)); - ret->push_back(slog(fea_set.unroll_prod)); - ret->push_back(slog(fea_set.unroll_len)); - for (int i = 0; i <= static_cast(AnnotationPosType::kPosMixed); i++) { - ret->push_back(i == static_cast(fea_set.unroll_type)); - } - - ret->push_back(slog(fea_set.parallel_num)); - ret->push_back(slog(fea_set.parallel_prod)); - ret->push_back(slog(fea_set.parallel_len)); - for (int i = 0; i <= static_cast(AnnotationPosType::kPosMixed); i++) { - ret->push_back(i == static_cast(fea_set.parallel_type)); - } - - ret->push_back(fea_set.is_gpu); - ret->push_back(slog(fea_set.blockIdx_x_len)); - ret->push_back(slog(fea_set.blockIdx_y_len)); - ret->push_back(slog(fea_set.blockIdx_z_len)); - ret->push_back(slog(fea_set.threadIdx_x_len)); - ret->push_back(slog(fea_set.threadIdx_y_len)); - ret->push_back(slog(fea_set.threadIdx_z_len)); - ret->push_back(slog(fea_set.vthread_len)); - - /***** Group 2: Buffer access related features *****/ - // sort according to pair (lines, bytes) - std::vector> buf_order_key; - for (const auto& acc_fea : fea_set.access_feas) { - buf_order_key.emplace_back(acc_fea.lines, acc_fea.bytes); - } - std::vector buf_order(buf_order_key.size()); - std::iota(buf_order.begin(), buf_order.end(), 0); - - auto cmp = [&buf_order_key](int l, int r) { - return buf_order_key[l].first > buf_order_key[r].first || - (buf_order_key[l].first == buf_order_key[r].first && - buf_order_key[l].second > buf_order_key[r].second); - }; - std::sort(buf_order.begin(), buf_order.end(), cmp); - int n_bufs = std::min(max_n_bufs, static_cast(buf_order.size())); - buf_order.resize(n_bufs); - - for (int idx : buf_order) { - const auto& acc_fea = fea_set.access_feas[idx]; - for (int j = 0; j <= static_cast(BufferAccessType::kReadWrite); ++j) { - ret->push_back(j == static_cast(acc_fea.acc_type)); - } - ret->push_back(slog(acc_fea.bytes)); - ret->push_back(slog(acc_fea.unique_bytes)); - ret->push_back(slog(acc_fea.lines)); - ret->push_back(slog(acc_fea.unique_lines)); - for (int j = 0; j <= static_cast(ReuseType::kNoReuse); ++j) { - ret->push_back(j == static_cast(acc_fea.reuse_type)); - } - ret->push_back(slog(acc_fea.reuse_dis_iter)); - ret->push_back(slog(acc_fea.reuse_dis_bytes)); - ret->push_back(slog(acc_fea.reuse_ct)); - ret->push_back(slog(acc_fea.bytes_d_reuse_ct)); - ret->push_back(slog(acc_fea.unique_bytes_d_reuse_ct)); - ret->push_back(slog(acc_fea.lines_d_reuse_ct)); - ret->push_back(slog(acc_fea.unique_lines_d_reuse_ct)); - ret->push_back(slog(acc_fea.stride)); - } - // - fill padding - for (int i = 0; i < max_n_bufs - n_bufs; ++i) { - for (int j = 0; j <= static_cast(BufferAccessType::kReadWrite); ++j) { // 3 - ret->push_back(0.0f); - } - ret->push_back(0.0f); - ret->push_back(0.0f); - ret->push_back(0.0f); - ret->push_back(0.0f); - for (int j = 0; j <= static_cast(ReuseType::kNoReuse); ++j) { // 3 - ret->push_back(0.0f); - } - ret->push_back(0.0f); - ret->push_back(0.0f); - ret->push_back(0.0f); - ret->push_back(0.0f); - ret->push_back(0.0f); - ret->push_back(0.0f); - ret->push_back(0.0f); - ret->push_back(0.0f); - } - - /***** Group 3: Arithmetic intensity related features *****/ - for (size_t i = 0; i < ARITH_INTENSITY_CURVE_SAMPLE_N; ++i) { - ret->push_back(slog(fea_set.arith_intensity_curve[i])); - } - - /***** Group 4: Allocation related features *****/ - ret->push_back(slog(fea_set.alloc_size)); - ret->push_back(slog(fea_set.alloc_prod)); - ret->push_back(slog(fea_set.alloc_outer_prod)); - ret->push_back(slog(fea_set.alloc_inner_prod)); - - /***** Group 5: Outer scope related features *****/ - ret->push_back(slog(fea_set.outer_prod)); - ret->push_back(slog(fea_set.num_loops)); - ret->push_back(slog(fea_set.auto_unroll_max_step)); - } -} - -void GetPerStoreFeatureName(int max_n_bufs, std::vector* ret) { - /***** Group 1: Computation related features *****/ - ret->push_back(("float_mad")); - ret->push_back(("float_addsub")); - ret->push_back(("float_mul")); - ret->push_back(("float_divmod")); - ret->push_back(("float_cmp")); - ret->push_back(("float_mathfunc")); - ret->push_back(("float_otherfunc")); - ret->push_back(("int_mad")); - ret->push_back(("int_addsub")); - ret->push_back(("int_mul")); - ret->push_back(("int_divmod")); - ret->push_back(("int_cmp")); - ret->push_back(("int_mathfunc")); - ret->push_back(("int_otherfunc")); - ret->push_back(("bool_op")); - ret->push_back(("select_op")); - ret->push_back(("vec_num")); - ret->push_back(("vec_prod")); - ret->push_back(("vec_len")); - ret->push_back(("vec_type.kPosNone")); - ret->push_back(("vec_type.kPosInnerSpatial")); - ret->push_back(("vec_type.kPosMiddleSpatial")); - ret->push_back(("vec_type.kPosOuterSpatial")); - ret->push_back(("vec_type.kPosInnerReduce")); - ret->push_back(("vec_type.kPosMiddleReduce")); - ret->push_back(("vec_type.kPosOuterReduce")); - ret->push_back(("vec_type.kPosMixed")); - ret->push_back(("unroll_num")); - ret->push_back(("unroll_prod")); - ret->push_back(("unroll_len")); - ret->push_back(("unroll_type.kPosNone")); - ret->push_back(("unroll_type.kPosInnerSpatial")); - ret->push_back(("unroll_type.kPosMiddleSpatial")); - ret->push_back(("unroll_type.kPosOuterSpatial")); - ret->push_back(("unroll_type.kPosInnerReduce")); - ret->push_back(("unroll_type.kPosMiddleReduce")); - ret->push_back(("unroll_type.kPosOuterReduce")); - ret->push_back(("unroll_type.kPosMixed")); - ret->push_back(("parallel_num")); - ret->push_back(("parallel_prod")); - ret->push_back(("parallel_len")); - ret->push_back(("parallel_type.kPosNone")); - ret->push_back(("parallel_type.kPosInnerSpatial")); - ret->push_back(("parallel_type.kPosMiddleSpatial")); - ret->push_back(("parallel_type.kPosOuterSpatial")); - ret->push_back(("parallel_type.kPosInnerReduce")); - ret->push_back(("parallel_type.kPosMiddleReduce")); - ret->push_back(("parallel_type.kPosOuterReduce")); - ret->push_back(("parallel_type.kPosMixed")); - ret->push_back(("is_gpu")); - ret->push_back(("blockIdx_x_len")); - ret->push_back(("blockIdx_y_len")); - ret->push_back(("blockIdx_z_len")); - ret->push_back(("threadIdx_x_len")); - ret->push_back(("threadIdx_y_len")); - ret->push_back(("threadIdx_z_len")); - ret->push_back(("vthread_len")); - // section total: 57 - - /***** Group 2: Buffer access related features *****/ - for (size_t i = 0; i < static_cast(max_n_bufs); ++i) { - std::string prefix = "B" + std::to_string(i) + "."; - ret->push_back((prefix + "acc_type.kRead")); - ret->push_back((prefix + "acc_type.kWrite")); - ret->push_back((prefix + "acc_type.kReadWrite")); - ret->push_back((prefix + "bytes")); - ret->push_back((prefix + "unique_bytes")); - ret->push_back((prefix + "lines")); - ret->push_back((prefix + "unique_lines")); - ret->push_back((prefix + "reuse_type.kLoopMultipleRead")); - ret->push_back((prefix + "reuse_type.kSerialMultipleReadWrite")); - ret->push_back((prefix + "reuse_type.kNoReuse")); - ret->push_back((prefix + "reuse_dis_iter")); - ret->push_back((prefix + "reuse_dis_bytes")); - ret->push_back((prefix + "reuse_ct")); - ret->push_back((prefix + "bytes_d_reuse_ct")); - ret->push_back((prefix + "unique_bytes_d_reuse_ct")); - ret->push_back((prefix + "lines_d_reuse_ct")); - ret->push_back((prefix + "unique_lines_d_reuse_ct")); - ret->push_back((prefix + "stride")); - } - // section total : max_n_bufs * 18 - - /***** Group 3: Arithmetic intensity related features *****/ - for (size_t i = 0; i < ARITH_INTENSITY_CURVE_SAMPLE_N; ++i) { - ret->push_back(("arith_intensity_curve_" + std::to_string(i))); - } - // section total: ARITH_INTENSITY_CURVE_SAMPLE_N = 10 - - /***** Group 4: Allocation related features *****/ - ret->push_back(("alloc_size")); - ret->push_back(("alloc_prod")); - ret->push_back(("alloc_outer_prod")); - ret->push_back(("alloc_inner_prod")); - // section total : 4 - - /***** Group 5: Outer scope related features *****/ - ret->push_back(("outer_prod")); - ret->push_back(("num_loops")); - ret->push_back(("auto_unroll_max_step")); - // section total : 3 -} - -void GetPerStoreFeaturesWorkerFunc(const SearchTask& task, const State& state, int max_n_bufs, - std::vector* feature, std::atomic* error_ct) { - auto [sch, tensors] = task->compute_dag.ApplySteps(state->transform_steps); - - // When inlining, replace const matrices with const values. - // Produces wrong IR, but good enough for feature extraction, and - // can improve the speed of feature extraction/search. Must be - // called before ScheduleToModule to have an effect. - sch = sch.normalize_for_feature_extraction(); - - try { - const std::string& name = "main"; - auto pass_ctx = tvm::transform::PassContext::Current(); - - auto mod = ScheduleToModule(sch, Array{tensors.begin(), tensors.end()}, name, - std::unordered_map(), GlobalVarSupply()); - - bool disable_vectorize = - pass_ctx->GetConfig("tir.disable_vectorize", Bool(false)).value(); - bool instrument_bound_checkers = - pass_ctx->GetConfig("tir.instrument_bound_checkers", Bool(false)).value(); - - if (IsGPUTask(task)) { - auto pass_list = Array(); - // Phase 0 - pass_list.push_back(tir::transform::InjectPrefetch()); - pass_list.push_back(tir::transform::StorageFlatten(64, instrument_bound_checkers)); - // Phase 1 - pass_list.push_back(tir::transform::NarrowDataType(32)); - pass_list.push_back(tir::transform::Simplify()); - pass_list.push_back(tir::transform::VectorizeLoop(!disable_vectorize)); - pass_list.push_back(tir::transform::InjectVirtualThread()); - pass_list.push_back(tir::transform::StorageRewrite()); - pass_list.push_back(tir::transform::Simplify()); - tvm::Map gpu_params{ - {"max_shared_memory_per_block", task->hardware_params->max_shared_memory_per_block}, - {"max_local_memory_per_block", task->hardware_params->max_local_memory_per_block}, - {"max_threads_per_block", task->hardware_params->max_threads_per_block}, - {"max_vector_bytes", task->hardware_params->vector_unit_bytes}, - {"max_vthread", task->hardware_params->max_vthread_extent}, - }; - pass_list.push_back(tir::transform::VerifyGPUCode(gpu_params)); - const auto& optimize = tir::transform::Sequential(pass_list); - optimize(mod); - } - if (IsHexagonTask(task)) { - Target target = task->target; - const auto& optimize = tir::transform::Sequential({tir::transform::VerifyVTCMLimit(target)}); - optimize(mod); - } - const auto& optimize = - tir::transform::Sequential(Array{tir::transform::Simplify()}); - mod = optimize(std::move(mod)); - PrimFunc prim_func = Downcast(mod->Lookup(name)); - GetPerStoreFeature(prim_func, task->hardware_params->cache_line_bytes, max_n_bufs, feature); - } catch (Error& e) { - (*error_ct)++; - } -} - -void GetPerStoreFeaturesFromStates(const Array& states, const SearchTask& task, - int skip_first_n_feature_extraction, int max_n_bufs, - std::vector>* features) { - // extract features - features->assign(states.size(), std::vector()); - - std::atomic error_ct(0); - - support::parallel_for(skip_first_n_feature_extraction, states.size(), - [&task, &states, &max_n_bufs, &features, &error_ct](int i) { - GetPerStoreFeaturesWorkerFunc(task, states[i], max_n_bufs, - &(*features)[i], &error_ct); - }); -} - -void GetPerStoreFeaturesFromStates(const Array& states, const std::vector& tasks, - int skip_first_n_feature_extraction, int max_n_bufs, - std::vector>* features) { - // extract features - features->assign(states.size(), std::vector()); - - std::atomic error_ct(0); - - support::parallel_for(skip_first_n_feature_extraction, states.size(), - [&tasks, &states, &max_n_bufs, &features, &error_ct](int i) { - GetPerStoreFeaturesWorkerFunc(tasks[i], states[i], max_n_bufs, - &(*features)[i], &error_ct); - }); -} - -void GetPerStoreFeaturesFromFile(const std::string& filename, int max_lines, int max_n_bufs, - std::vector>* features, - std::vector* normalized_throughputs, - std::vector* task_ids) { - Array states; - std::vector tasks; - - normalized_throughputs->clear(); - task_ids->clear(); - - // (workload_key, target) -> (search_task, task_id) - std::unordered_map, std::pair> task_cache; - // task_id -> min_cost - std::vector min_costs; - - const auto* workload_key_to_tensors = - tvm::runtime::Registry::Get("auto_scheduler.workload_key_to_tensors"); - ICHECK(workload_key_to_tensors != nullptr); - - // read from file - RecordReader reader(filename); - auto cur_inp = make_object(); - auto cur_res = make_object(); - while (reader->ReadNext(cur_inp.get(), cur_res.get())) { - float cost = static_cast(FloatArrayMean(cur_res->costs)); - const std::string& workload_key = cur_inp->task->workload_key; - - SearchTask task; - size_t task_id; - std::pair key(workload_key, cur_inp->task->target->str()); - auto find_res = task_cache.find(key); - if (find_res == task_cache.end()) { - // rebuild task - Array tensors = (*workload_key_to_tensors)(workload_key); - Target target = cur_inp->task->target; - Target target_host = cur_inp->task->target_host; - CheckAndUpdateHostConsistency(&target, &target_host); - task = SearchTask(ComputeDAG(tensors), workload_key, target, target_host, - cur_inp->task->hardware_params, cur_inp->task->layout_rewrite_option, - cur_inp->task->task_input_names); - task_id = task_cache.size(); - - // compute min cost for each task - task_cache.insert(std::make_pair(key, std::make_pair(task, task_id))); - min_costs.push_back(cost); - } else { - std::tie(task, task_id) = find_res->second; - min_costs[task_id] = std::min(min_costs[task_id], cost); - } - - tasks.push_back(std::move(task)); - task_ids->push_back(task_id); - states.push_back(cur_inp->state); - normalized_throughputs->push_back(cost); - - if (max_lines > 0 && static_cast(states.size()) >= max_lines) { - break; - } - } - - for (size_t i = 0; i < normalized_throughputs->size(); ++i) { - (*normalized_throughputs)[i] = min_costs[(*task_ids)[i]] / (*normalized_throughputs)[i]; - } - - GetPerStoreFeaturesFromStates(states, tasks, 0, max_n_bufs, features); -} - -void GetPerStoreFeaturesFromMeasurePairs(const Array& inputs, - const Array& results, - int skip_first_n_feature_extraction, int max_n_bufs, - std::vector>* features, - std::vector* normalized_throughputs, - std::vector* task_ids) { - Array states; - std::vector tasks; - - normalized_throughputs->clear(); - task_ids->clear(); - - // (workload_key, target) -> (search_task, task_id) - std::unordered_map, std::pair> task_cache; - // task_id -> min_cost - std::vector min_costs; - - const auto* workload_key_to_tensors = - tvm::runtime::Registry::Get("auto_scheduler.workload_key_to_tensors"); - ICHECK(workload_key_to_tensors != nullptr); - - tasks.reserve(inputs.size()); - normalized_throughputs->reserve(inputs.size()); - task_ids->reserve(inputs.size()); - for (size_t i = 0; i < inputs.size(); ++i) { - float cost = static_cast(FloatArrayMean(results[i]->costs)); - const std::string& workload_key = inputs[i]->task->workload_key; - SearchTask task; - - size_t task_id; - std::pair key(workload_key, inputs[i]->task->target->str()); - auto find_res = task_cache.find(key); - if (find_res == task_cache.end()) { - if (inputs[i]->task->compute_dag.defined()) { // the measure input is complete - task = inputs[i]->task; - } else { - // The measure input is incomplete, rebuild task for incomplete measure pairs read from file - try { - Array tensors = (*workload_key_to_tensors)(workload_key); - Target target = inputs[i]->task->target; - Target target_host = inputs[i]->task->target_host; - CheckAndUpdateHostConsistency(&target, &target_host); - task = - SearchTask(ComputeDAG(tensors), workload_key, target, target_host, - inputs[i]->task->hardware_params, inputs[i]->task->layout_rewrite_option, - inputs[i]->task->task_input_names); - } catch (std::exception& e) { - // Cannot build ComputeDAG from workload key, the task may have not been registered in - // this search round - continue; - } - } - task_id = task_cache.size(); - - // compute min cost for each task - task_cache.insert(std::make_pair(key, std::make_pair(task, task_id))); - min_costs.push_back(cost); - } else { - std::tie(task, task_id) = find_res->second; - min_costs[task_id] = std::min(min_costs[task_id], cost); - } - - tasks.push_back(std::move(task)); - task_ids->push_back(task_id); - states.push_back(inputs[i]->state); - normalized_throughputs->push_back(cost); - } - - for (size_t i = 0; i < normalized_throughputs->size(); ++i) { - (*normalized_throughputs)[i] = min_costs[(*task_ids)[i]] / (*normalized_throughputs)[i]; - } - - GetPerStoreFeaturesFromStates(states, tasks, skip_first_n_feature_extraction, max_n_bufs, - features); -} - -/* - * \brief Serialize a two-dimensional variable-size feature vector with normalized throughputs - * and task ids to a one-dimensional flatten byte array. - * - * For faster data copy between c++ and python, the c++ part returns features in a single - * flatten array using a packed format. The python part then unpacks the flatten array. - * - * The packed format for n records is: - * { - * int n; - * int sizes[n+2]; // The sizes for the following arrays - * - * float features_0[size[0]]; // The features for record 0 - * float features_1[size[1]]; // The features for record 1 - * ... - * float features_i[size[i]]; // The features for record i - * ... // until i == n - 1 - * - * float throughputs[sizes[n]]; // The normalized throughputs for n records - * int task_ids[size[n+1]]; // The task ids for n records - * - * } - * To implement this format, we also store int as float, so we can store all numbers - * into a single float array. - */ -TVMByteArray SerializeFeatures(std::vector>&& features, - std::vector&& normalized_throughputs, - std::vector&& task_ids, std::vector* out_data) { - size_t total_bytes = 0; - std::vector size_vector; - - int n = features.size(); - - // serialize sizes - size_t size_vector_size = 1 + n + 2; - total_bytes += size_vector_size * sizeof(int); - - size_vector.reserve(size_vector_size); - size_vector.push_back(features.size()); - for (const auto& x : features) { - size_vector.push_back(static_cast(x.size())); - total_bytes += sizeof(float) * x.size(); - } - size_vector.push_back(static_cast(normalized_throughputs.size())); - total_bytes += sizeof(float) * normalized_throughputs.size(); - size_vector.push_back(static_cast(task_ids.size())); - total_bytes += sizeof(int) * task_ids.size(); - - ICHECK_EQ(size_vector.size(), size_vector_size); - - // allocate memory - out_data->reserve(total_bytes); - char* ptr = out_data->data(); - - // serialize size_vector - memmove(ptr, reinterpret_cast(size_vector.data()), size_vector.size() * sizeof(int)); - ptr += size_vector.size() * sizeof(int); - - // serialize features - for (auto& x : features) { - memmove(ptr, x.data(), sizeof(float) * x.size()); - ptr += sizeof(float) * x.size(); - x.clear(); - } - - // serialize normalized_throughputs - memmove(ptr, reinterpret_cast(normalized_throughputs.data()), - normalized_throughputs.size() * sizeof(int)); - ptr += normalized_throughputs.size() * sizeof(int); - - // serialize task_ids - memmove(ptr, reinterpret_cast(task_ids.data()), task_ids.size() * sizeof(int)); - ptr += task_ids.size() * sizeof(int); - - ICHECK_EQ(ptr - out_data->data(), total_bytes); - - return TVMByteArray{out_data->data(), total_bytes}; -} - -TVM_REGISTER_GLOBAL("auto_scheduler.GetPerStoreFeaturesFromFile") - .set_body([](TVMArgs args, TVMRetValue* ret) { - std::string filename = args[0]; - int max_lines = args[1]; - int max_n_bufs = args[2]; - - std::vector> features; - std::vector normalized_throughputs; - std::vector task_ids; - - GetPerStoreFeaturesFromFile(filename, max_lines, max_n_bufs, &features, - &normalized_throughputs, &task_ids); - - std::vector byte_data; - *ret = SerializeFeatures(std::move(features), std::move(normalized_throughputs), - std::move(task_ids), &byte_data); - }); - -TVM_REGISTER_GLOBAL("auto_scheduler.GetPerStoreFeaturesFromMeasurePairs") - .set_body([](TVMArgs args, TVMRetValue* ret) { - Array inputs = args[0]; - Array results = args[1]; - int skip_first_n_feature_extraction = args[2]; - int max_n_bufs = args[3]; - - std::vector> features; - std::vector normalized_throughputs; - std::vector task_ids; - - GetPerStoreFeaturesFromMeasurePairs(inputs, results, skip_first_n_feature_extraction, - max_n_bufs, &features, &normalized_throughputs, - &task_ids); - - std::vector byte_data; - *ret = SerializeFeatures(std::move(features), std::move(normalized_throughputs), - std::move(task_ids), &byte_data); - }); - -TVM_REGISTER_GLOBAL("auto_scheduler.GetPerStoreFeaturesFromStates") - .set_body([](TVMArgs args, TVMRetValue* ret) { - Array states = args[0]; - SearchTask task = args[1]; - int max_n_bufs = args[2]; - - std::vector> features; - std::vector normalized_throughputs; - std::vector task_ids; - - GetPerStoreFeaturesFromStates(states, task, 0, max_n_bufs, &features); - - std::vector byte_data; - *ret = SerializeFeatures(std::move(features), std::move(normalized_throughputs), - std::move(task_ids), &byte_data); - }); - -TVM_REGISTER_GLOBAL("auto_scheduler.GetPerStoreFeatureNames") - .set_body([](TVMArgs args, TVMRetValue* ret) { - int max_n_bufs = args[0]; - std::vector names; - - GetPerStoreFeatureName(max_n_bufs, &names); - - Array arr; - for (const auto& x : names) { - arr.push_back(x); - } - *ret = arr; - }); - -TVM_REGISTER_GLOBAL("auto_scheduler.FeaturesFromPrimFunc") - .set_body_typed([](const PrimFunc& func, int cache_line_size, int max_n_bufs, bool log_scale) { - std::vector vec; - GetPerStoreFeature(func, cache_line_size, max_n_bufs, &vec, log_scale); - int64_t num_feature_rows = vec[0]; // first element is number of rows - int64_t row_length = 0; - if (num_feature_rows != 0) { - row_length = (vec.size() - 1) / num_feature_rows; - } - auto ary = - runtime::NDArray::Empty({num_feature_rows, row_length}, {kDLFloat, 32, 1}, {kDLCPU, 0}); - // NDArray is row major by default - ary.CopyFromBytes(vec.data() + 1, sizeof(float) * num_feature_rows * row_length); - return ary; - }); - -} // namespace auto_scheduler -} // namespace tvm diff --git a/src/auto_scheduler/loop_state.cc b/src/auto_scheduler/loop_state.cc deleted file mode 100755 index 517f7ff91f55..000000000000 --- a/src/auto_scheduler/loop_state.cc +++ /dev/null @@ -1,575 +0,0 @@ -/* - * 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 auto_scheduler/loop_state.cc - * \brief An lightweight IR (intermediate representation) for loop structures. - * see auto_scheduler/loop_state.h for more explanation. - */ - -#include -#include -#include -#include -#include - -#include - -#include "utils.h" - -namespace tvm { -namespace auto_scheduler { - -TVM_REGISTER_OBJECT_TYPE(StepNode); -TVM_REGISTER_NODE_TYPE(StageNode); -TVM_REGISTER_NODE_TYPE(StateNode); -TVM_REGISTER_NODE_TYPE(IteratorNode); - -/********** Iterator **********/ -Iterator::Iterator(String name, Range range, IteratorKind iter_kind, IteratorAnnotation annotation, - const std::vector* orig_iters) { - auto node = make_object(); - node->name = std::move(name); - node->range = std::move(range); - node->iter_kind = iter_kind; - node->annotation = annotation; - if (orig_iters != nullptr) { - node->orig_iters = *orig_iters; - } - data_ = std::move(node); -} - -/********** Stage **********/ -Stage::Stage(te::Operation op) { - auto node = make_object(); - if (op->IsInstance()) { - node->op_type = StageKind::kCompute; - auto* pop = op.as(); - for (const auto& axis : pop->axis) { - node->iters.push_back(Iterator(CleanName(axis->var->name_hint), axis->dom, - IteratorKind::kSpatial, IteratorAnnotation::kNone)); - } - for (const auto& axis : pop->reduce_axis) { - node->iters.push_back(Iterator(CleanName(axis->var->name_hint), axis->dom, - IteratorKind::kReduction, IteratorAnnotation::kNone)); - } - } else if (op->IsInstance()) { - node->op_type = StageKind::kPlaceholder; - } else { - LOG(FATAL) << "Unsupported operator type" << op->_type_key; - } - - node->compute_at = ComputeAtKind::kRoot; - node->op = std::move(op); - node->attrs.auto_unroll_max_step = 0; - node->attrs.storage_offset = 0; - data_ = std::move(node); -} - -Stage::Stage(te::Operation op, StageKind op_type, const Array& iters, - ComputeAtKind compute_at, StageAttributes attrs) { - auto node = make_object(); - node->op = std::move(op); - node->op_type = op_type; - node->iters = iters; - node->compute_at = compute_at; - node->attrs = attrs; - data_ = std::move(node); -} - -/********** AttachMap **********/ -void AttachMap::SetComputeAtIter(int stage_id, int target_stage_id, int target_iter_id) { - AttachMapNode* pnode = CopyOnWrite(); - - // Delete the current entry of this stage - DeleteStageEntry(pnode, stage_id); - - // Store the new stage/iterator relations to map - IterKey iter_key(target_stage_id, target_iter_id); - pnode->stage_to_attach_iter[stage_id] = iter_key; - pnode->iter_to_attached_stages[iter_key].push_back(stage_id); -} - -void AttachMap::DeleteStage(int stage_id) { - AttachMapNode* pnode = CopyOnWrite(); - // Delete the original stage entry - DeleteStageEntry(pnode, stage_id); -} - -void AttachMap::UpdateIters(const std::vector& original_iters, - const std::vector& new_iters) { - ICHECK_EQ(original_iters.size(), new_iters.size()); - AttachMapNode* pnode = CopyOnWrite(); - std::unordered_map> new_iter_to_attached_stages; - for (size_t i = 0; i < original_iters.size(); ++i) { - auto entry = pnode->iter_to_attached_stages.find(original_iters[i]); - // We get > from this map - if (entry == pnode->iter_to_attached_stages.end()) { - // Skip if this iterator does not have any attach relations - continue; - } - - // Update the attaching target of an stage to the new iter in `stage_to_attach_iter` - for (const auto& s : entry->second) { - pnode->stage_to_attach_iter[s] = new_iters[i]; - } - - // Remove the original iterator relation from `iter_to_attached_stages` and add the new - // iterator to it - std::vector attached_stages = std::move(entry->second); - pnode->iter_to_attached_stages.erase(entry); - new_iter_to_attached_stages[new_iters[i]] = std::move(attached_stages); - } - - // Update new entries - for (auto& it : new_iter_to_attached_stages) { - pnode->iter_to_attached_stages[it.first] = std::move(it.second); - } -} - -void AttachMap::DeleteStageEntry(AttachMapNode* pnode, int stage_id) { - auto old_entry = pnode->stage_to_attach_iter.find(stage_id); - // We get from this map - if (old_entry != pnode->stage_to_attach_iter.end()) { - // Delete the stage in `iter_to_attached_stages`, if the corresponding iterator does not have - // any attached stage, delete this iterm too - auto entry2 = pnode->iter_to_attached_stages.find(old_entry->second); - // We get > from this map - FindAndDeleteItem(&entry2->second, stage_id); - if (entry2->second.size() == 0) { - pnode->iter_to_attached_stages.erase(entry2); - } - // Delete the stage in `stage_to_attach_iter` - pnode->stage_to_attach_iter.erase(old_entry); - } -} - -AttachMap AttachMap::ApplyStageIdOffset(int start_id, int offset) const { - AttachMap map = AttachMap(make_object()); - auto pmap = map.CopyOnWrite(); - for (const auto& x : operator->()->stage_to_attach_iter) { - auto key = x.first; - if (key >= start_id) { - key += offset; - } - auto value = x.second; - if (value.first >= start_id) { - value.first += offset; - } - pmap->stage_to_attach_iter.insert(std::make_pair(key, value)); - } - for (const auto& x : operator->()->iter_to_attached_stages) { - auto key = x.first; - if (key.first >= start_id) { - key.first += offset; - } - auto value = x.second; - for (auto& i : value) { - if (i >= start_id) { - i += offset; - } - } - pmap->iter_to_attached_stages.insert(std::make_pair(key, value)); - } - return map; -} - -/********** State **********/ -State::State(const Array& ops) { - auto node = make_object(); - for (const auto& op : ops) { - node->stages.push_back(Stage(op)); - } - node->attach_map = AttachMap(make_object()); - node->concrete = true; - data_ = std::move(node); -} - -/********** Schedule primitives apis for state **********/ -Iterator State::bind(int stage_id, const Iterator& it, IteratorAnnotation thread_type) { - const Stage& stage = operator->()->stages[stage_id]; - if (thread_type < IteratorAnnotation::kVThread || thread_type > IteratorAnnotation::kThreadZ) { - LOG(FATAL) << "thread_type error, valid: kVThread, kBlockX, kBlockY, " - << "kThreadX, kThreadY, kBlockZ, kThreadZ"; - } - AnnotationStep step = AnnotationStep(stage_id, GetIndex(stage->iters, it), thread_type); - CopyOnWrite()->transform_steps.push_back(step); - return step->ApplyToState(this); -} - -Iterator State::parallel(int stage_id, const Iterator& it) { - const Stage& stage = operator->()->stages[stage_id]; - AnnotationStep step = - AnnotationStep(stage_id, GetIndex(stage->iters, it), IteratorAnnotation::kParallel); - CopyOnWrite()->transform_steps.push_back(step); - return step->ApplyToState(this); -} - -Iterator State::unroll(int stage_id, const Iterator& it, int max_unroll) { - const Stage& stage = operator->()->stages[stage_id]; - - // Don't unroll if the extent is larger than max_unroll - if (max_unroll != -1 && it->range.defined()) { - if (auto imm = it->range->extent.as()) { - if (imm->value > max_unroll) { - return it; - } - } - } - - AnnotationStep step = - AnnotationStep(stage_id, GetIndex(stage->iters, it), IteratorAnnotation::kUnroll); - CopyOnWrite()->transform_steps.push_back(step); - return step->ApplyToState(this); -} - -Iterator State::vectorize(int stage_id, const Iterator& it) { - const Stage& stage = operator->()->stages[stage_id]; - AnnotationStep step = - AnnotationStep(stage_id, GetIndex(stage->iters, it), IteratorAnnotation::kVectorize); - CopyOnWrite()->transform_steps.push_back(step); - return step->ApplyToState(this); -} - -Iterator State::fuse(int stage_id, const Array& iters) { - const Stage& stage = operator->()->stages[stage_id]; - Array indices; - GetIndices(stage->iters, iters, &indices); - FuseStep step = FuseStep(stage_id, indices); - CopyOnWrite()->transform_steps.push_back(step); - return step->ApplyToState(this); -} - -void State::pragma(int stage_id, const Iterator& it, const String& pragma_type) { - const Stage& stage = operator->()->stages[stage_id]; - PragmaStep step = PragmaStep(stage_id, GetIndex(stage->iters, it), pragma_type); - CopyOnWrite()->transform_steps.push_back(step); - return step->ApplyToState(this); -} - -void State::reorder(int stage_id, const Array& order) { - const Stage& stage = operator->()->stages[stage_id]; - ICHECK_EQ(order.size(), stage->iters.size()) << "The order of all iterators " - << "should be specified"; - Array after_ids; - GetIndices(stage->iters, order, &after_ids); - ReorderStep step = ReorderStep(stage_id, after_ids); - CopyOnWrite()->transform_steps.push_back(step); - step->ApplyToState(this); -} - -Array State::split(int stage_id, const Iterator& it, - const Array>& lengths, bool inner_to_outer) { - const Stage& stage = operator->()->stages[stage_id]; - SplitStep step = - SplitStep(stage_id, GetIndex(stage->iters, it), - it->range.defined() ? it->range->extent : PrimExpr(), lengths, inner_to_outer); - CopyOnWrite()->transform_steps.push_back(step); - return step->ApplyToState(this); -} - -Array State::follow_split(int stage_id, const Iterator& it, int src_step_id, - int n_split) { - const Stage& stage = operator->()->stages[stage_id]; - FollowSplitStep step = - FollowSplitStep(stage_id, GetIndex(stage->iters, it), src_step_id, n_split); - CopyOnWrite()->transform_steps.push_back(step); - return step->ApplyToState(this); -} - -Array State::follow_fused_split(int stage_id, const Iterator& it, - const Array& src_step_ids, int level, - bool factor_or_nparts) { - const Stage& stage = operator->()->stages[stage_id]; - FollowFusedSplitStep step = FollowFusedSplitStep(stage_id, GetIndex(stage->iters, it), - src_step_ids, level, factor_or_nparts); - CopyOnWrite()->transform_steps.push_back(step); - return step->ApplyToState(this); -} - -void State::storage_align(int stage_id, const Iterator& it, int factor, int offset) { - const Stage& stage = operator->()->stages[stage_id]; - StorageAlignStep step = StorageAlignStep(stage_id, GetIndex(stage->iters, it), factor, offset); - CopyOnWrite()->transform_steps.push_back(step); - return step->ApplyToState(this); -} - -void State::compute_at(int stage_id, int target_stage_id, const Iterator& target_iter) { - const Stage& target_stage = operator->()->stages[target_stage_id]; - ComputeAtStep step = - ComputeAtStep(stage_id, target_stage_id, GetIndex(target_stage->iters, target_iter)); - CopyOnWrite()->transform_steps.push_back(step); - step->ApplyToState(this); -} - -void State::compute_inline(int stage_id) { - ComputeInlineStep step = ComputeInlineStep(stage_id); - CopyOnWrite()->transform_steps.push_back(step); - step->ApplyToState(this); -} - -void State::compute_root(int stage_id) { - ComputeRootStep step = ComputeRootStep(stage_id); - CopyOnWrite()->transform_steps.push_back(step); - step->ApplyToState(this); -} - -int State::cache_read(int stage_id, const String& scope_name, - const Array& reader_stage_ids, const ComputeDAG& dag) { - CacheReadStep step = CacheReadStep(stage_id, scope_name, reader_stage_ids); - CopyOnWrite()->transform_steps.push_back(step); - return step->ApplyToState(this, dag); -} - -int State::cache_write(int stage_id, const String& scope_name, const ComputeDAG& dag) { - CacheWriteStep step = CacheWriteStep(stage_id, scope_name); - CopyOnWrite()->transform_steps.push_back(step); - return step->ApplyToState(this, dag); -} - -int State::rfactor(int stage_id, const Iterator& it, int factor_iter_id, const ComputeDAG& dag) { - const Stage& stage = operator->()->stages[stage_id]; - RfactorStep step = RfactorStep(stage_id, GetIndex(stage->iters, it), factor_iter_id); - CopyOnWrite()->transform_steps.push_back(step); - return step->ApplyToState(this, dag); -} - -// Print stage to ostream -void PrintStage(std::ostream* os, int stage_id, const State& state, size_t base_indent, - bool delete_trivial_loop) { - const Stage& stage = state->stages[stage_id]; - - if (stage->attrs.auto_unroll_max_step != 0) { - for (size_t j = 0; j < base_indent; ++j) { - *os << " "; - } - *os << stage->op->name << " auto_unroll: " << stage->attrs.auto_unroll_max_step << "\n"; - } - if (stage->attrs.storage_offset != 0) { - for (size_t j = 0; j < base_indent; ++j) { - *os << " "; - } - *os << stage->op->name << " storage_offset: " << stage->attrs.storage_offset << "\n"; - } - - size_t indent = 0; - for (size_t i = 0; i < stage->iters.size(); ++i) { - const Iterator& iter = stage->iters[i]; - - if (!(delete_trivial_loop && iter->range.defined() && is_one(iter->range->extent))) { - for (size_t j = 0; j < base_indent + indent; ++j) { - *os << " "; - } - *os << IteratorAnnotationString[static_cast(iter->annotation)] << " "; - if (iter->range.defined()) { - *os << iter->name << " (" << iter->range->min << "," << iter->range->extent << ")"; - } else { - *os << iter->name << " (None)"; - } - *os << "\n"; - - indent += 2; - } - - if (state.defined()) { - IterKey iter_key(stage_id, i); - auto pair = state->attach_map->iter_to_attached_stages.find(iter_key); - if (pair != state->attach_map->iter_to_attached_stages.end()) { - // Print the attached stage - for (const auto& attach_stage_id : pair->second) { - PrintStage(os, attach_stage_id, state, base_indent + indent, delete_trivial_loop); - } - } - } - } - - for (size_t j = 0; j < base_indent + indent; ++j) { - *os << " "; - } - *os << stage->op->name << " = ...\n"; -} - -// Print state to ostream -void PrintState(std::ostream* os, const State& state, bool delete_trivial_loop) { - // Gather placeholders - Array placeholders; - for (const auto& stage : state->stages) { - if (stage->op_type == StageKind::kPlaceholder) { - placeholders.push_back(stage->op->name); - } - } - - *os << "Placeholder: "; - for (size_t i = 0; i < placeholders.size(); ++i) { - *os << placeholders[i]; - if (i != placeholders.size() - 1) { - *os << ", "; - } - } - *os << "\n"; - - // Print all stages - for (size_t i = 0; i < state->stages.size(); ++i) { - const Stage& stage = state->stages[i]; - if (stage->op_type == StageKind::kPlaceholder) { - continue; - } else if (stage->op_type == StageKind::kCompute) { - if (stage->compute_at == ComputeAtKind::kRoot) { - PrintStage(os, i, state, 0, delete_trivial_loop); - } - } else { - LOG(FATAL) << "Invalid op type"; - } - } -} - -String State::ToStr(bool delete_trivial_loop) const { - std::ostringstream os; - PrintState(&os, (*this), delete_trivial_loop); - return os.str(); -} - -TVM_STATIC_IR_FUNCTOR(ReprPrinter, vtable) - .set_dispatch([](const ObjectRef& ref, ReprPrinter* p) { - const auto& stage = tvm::Downcast(ref); - p->stream << stage->GetTypeKey() << "(" << stage.get() << ": " << stage->op->name << ")"; - }); - -TVM_STATIC_IR_FUNCTOR(ReprPrinter, vtable) - .set_dispatch([](const ObjectRef& ref, ReprPrinter* p) { - PrintState(&p->stream, tvm::Downcast(ref), true); - }); - -/********** State interface API for ffi **********/ -TVM_REGISTER_GLOBAL("auto_scheduler.StateBind") - .set_body_typed([](State state, int stage_id, const Iterator& it, int thread_type) { - const auto& res = state.bind(stage_id, it, IteratorAnnotation(thread_type)); - return Array{state, res}; - }); - -TVM_REGISTER_GLOBAL("auto_scheduler.StateParallel") - .set_body_typed([](State state, int stage_id, const Iterator& it) { - const auto& res = state.parallel(stage_id, it); - return Array{state, res}; - }); - -TVM_REGISTER_GLOBAL("auto_scheduler.StateUnroll") - .set_body_typed([](State state, int stage_id, const Iterator& it, int max_unroll) { - const auto& res = state.unroll(stage_id, it, max_unroll); - return Array{state, res}; - }); - -TVM_REGISTER_GLOBAL("auto_scheduler.StateVectorize") - .set_body_typed([](State state, int stage_id, const Iterator& it) { - const auto& res = state.vectorize(stage_id, it); - return Array{state, res}; - }); - -TVM_REGISTER_GLOBAL("auto_scheduler.StateFuse") - .set_body_typed([](State state, int stage_id, const Array& iters) { - const auto& res = state.fuse(stage_id, iters); - return Array{state, res}; - }); - -TVM_REGISTER_GLOBAL("auto_scheduler.StatePragma") - .set_body_typed([](State state, int stage_id, const Iterator& it, const String& pragma_type) { - state.pragma(stage_id, it, pragma_type); - return state; - }); - -TVM_REGISTER_GLOBAL("auto_scheduler.StateReorder") - .set_body_typed([](State state, int stage_id, const Array& order) { - state.reorder(stage_id, order); - return state; - }); - -TVM_REGISTER_GLOBAL("auto_scheduler.StateSplit") - .set_body_typed([](State state, int stage_id, const Iterator& it, - const Array>& lengths, bool inner_to_outer) { - const auto& res = state.split(stage_id, it, lengths, inner_to_outer); - return Array{state, res}; - }); - -TVM_REGISTER_GLOBAL("auto_scheduler.StateFollowSplit") - .set_body_typed([](State state, int stage_id, const Iterator& it, int src_step_id, - int n_split) { - const auto& res = state.follow_split(stage_id, it, src_step_id, n_split); - return Array{state, Array(res)}; - }); - -TVM_REGISTER_GLOBAL("auto_scheduler.StateFollowFusedSplit") - .set_body_typed([](State state, int stage_id, const Iterator& it, - const Array& src_step_ids, int level, bool factor_or_nparts) { - const auto& res = - state.follow_fused_split(stage_id, it, src_step_ids, level, factor_or_nparts); - return Array{state, Array(res)}; - }); - -TVM_REGISTER_GLOBAL("auto_scheduler.StateStorageAlign") - .set_body_typed([](State state, int stage_id, const Iterator& it, int factor, int offset) { - state.storage_align(stage_id, it, factor, offset); - return state; - }); - -TVM_REGISTER_GLOBAL("auto_scheduler.StateComputeAt") - .set_body_typed([](State state, int stage_id, int target_stage_id, - const Iterator& target_iter) { - state.compute_at(stage_id, target_stage_id, target_iter); - return state; - }); - -TVM_REGISTER_GLOBAL("auto_scheduler.StateComputeInline") - .set_body_typed([](State state, int stage_id) { - state.compute_inline(stage_id); - return state; - }); - -TVM_REGISTER_GLOBAL("auto_scheduler.StateComputeRoot") - .set_body_typed([](State state, int stage_id) { - state.compute_root(stage_id); - return state; - }); - -TVM_REGISTER_GLOBAL("auto_scheduler.StateCacheRead") - .set_body_typed([](State state, int stage_id, const String& scope_name, - const Array& reader_stage_ids, const ComputeDAG& dag) { - int res = state.cache_read(stage_id, scope_name, reader_stage_ids, dag); - return Array{state, Integer(res)}; - }); - -TVM_REGISTER_GLOBAL("auto_scheduler.StateCacheWrite") - .set_body_typed([](State state, int stage_id, const String& scope_name, - const ComputeDAG& task_dag) { - int res = state.cache_write(stage_id, scope_name, task_dag); - return Array{state, Integer(res)}; - }); - -TVM_REGISTER_GLOBAL("auto_scheduler.StateRfactor") - .set_body_typed([](State state, int stage_id, const Iterator& it, int factor_iter_id, - const ComputeDAG& dag) { - int res = state.rfactor(stage_id, it, factor_iter_id, dag); - return Array{state, Integer(res)}; - }); - -TVM_REGISTER_GLOBAL("auto_scheduler.StateEqual").set_body_typed([](State state1, State state2) { - return std::equal_to()(state1, state2); -}); - -} // namespace auto_scheduler -} // namespace tvm diff --git a/src/auto_scheduler/measure.cc b/src/auto_scheduler/measure.cc deleted file mode 100755 index abb77581e7ee..000000000000 --- a/src/auto_scheduler/measure.cc +++ /dev/null @@ -1,428 +0,0 @@ -/* - * 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 auto_scheduler/measure.cc - * \brief Distributed measurement infrastructure to measure the runtime costs of tensor programs. - */ - -#include -#include - -#include - -#include "search_policy/empty_policy.h" -#include "search_policy/sketch_policy.h" -#include "utils.h" - -namespace tvm { -namespace auto_scheduler { - -TVM_REGISTER_NODE_TYPE(MeasureInputNode); -TVM_REGISTER_NODE_TYPE(BuildResultNode); -TVM_REGISTER_NODE_TYPE(MeasureResultNode); -TVM_REGISTER_OBJECT_TYPE(MeasureCallbackNode); -TVM_REGISTER_OBJECT_TYPE(PythonBasedMeasureCallbackNode); -TVM_REGISTER_OBJECT_TYPE(ProgramRunnerNode); -TVM_REGISTER_OBJECT_TYPE(ProgramBuilderNode); -TVM_REGISTER_OBJECT_TYPE(ProgramMeasurerNode); -TVM_REGISTER_OBJECT_TYPE(LocalBuilderNode); -TVM_REGISTER_OBJECT_TYPE(LocalRunnerNode); -TVM_REGISTER_OBJECT_TYPE(RPCRunnerNode); - -static const char* ErrorNoToStr[] = { - "NoError", - "InstantiationError", - "CompileHostError", - "CompileDeviceError", - "RuntimeDeviceError", - "WrongAnswerError", - "BuildTimeoutError", - "RunTimeoutError", - "UnknownError", -}; - -/********** Measure input and result **********/ -MeasureInput::MeasureInput(SearchTask task, State state) { - auto node = make_object(); - node->task = std::move(task); - node->state = std::move(state); - data_ = std::move(node); -} - -MeasureInput MeasureInputNode::copy() const { - auto node = make_object(); - node->task = task; - node->state = state; - return MeasureInput(node); -} - -BuildResult::BuildResult(String filename, Array args, int error_no, String error_msg, - double time_cost) { - auto node = make_object(); - node->filename = std::move(filename); - node->args = std::move(args); - node->error_no = error_no; - node->error_msg = std::move(error_msg); - node->time_cost = time_cost; - data_ = std::move(node); -} - -MeasureResult::MeasureResult(Array costs, int error_no, String error_msg, double all_cost, - double timestamp) { - auto node = make_object(); - node->costs = std::move(costs); - node->error_no = error_no; - node->error_msg = std::move(error_msg); - node->all_cost = all_cost; - node->timestamp = timestamp; - data_ = std::move(node); -} - -MeasureResult MeasureResultNode::copy() const { - auto node = make_object(); - node->costs = costs; - node->error_no = error_no; - node->error_msg = error_msg; - node->all_cost = all_cost; - node->timestamp = timestamp; - return MeasureResult(node); -} - -/********** LocalBuilder **********/ -LocalBuilder::LocalBuilder(int timeout, int n_parallel, const String& build_func) { - auto node = make_object(); - node->timeout = timeout; - node->n_parallel = n_parallel; - node->build_func = build_func; - data_ = std::move(node); -} - -Array LocalBuilderNode::Build(const Array& inputs, int verbose) { - if (const auto* f = runtime::Registry::Get("auto_scheduler.local_builder.build")) { - Array results = (*f)(inputs, timeout, n_parallel, build_func, verbose); - return results; - } - LOG(FATAL) << "auto_scheduler.local_builder.build is not registered. " - << "This is a function registered in Python, " - << "make sure the TVM Python runtime has been loaded successfully."; - throw; -} - -/********** LocalRunner **********/ -LocalRunner::LocalRunner(int timeout, int number, int repeat, int min_repeat_ms, - double cooldown_interval, bool enable_cpu_cache_flush, int device) { - ObjectPtr node = make_object(); - node->timeout = timeout; - node->number = number; - node->repeat = repeat; - node->min_repeat_ms = min_repeat_ms; - node->cooldown_interval = cooldown_interval; - node->enable_cpu_cache_flush = enable_cpu_cache_flush; - node->device = device; - data_ = std::move(node); -} - -Array LocalRunnerNode::Run(const Array& inputs, - const Array& build_results, int verbose) { - if (const auto* f = runtime::Registry::Get("auto_scheduler.local_runner.run")) { - Array results = - (*f)(inputs, build_results, timeout, number, repeat, min_repeat_ms, cooldown_interval, - enable_cpu_cache_flush, verbose, device); - return results; - } - LOG(FATAL) << "auto_scheduler.local_runner.run is not registered. " - << "This is a function registered in Python, " - << "make sure the TVM Python runtime has been loaded successfully."; - throw; -} - -/********** RPCRunner **********/ -RPCRunner::RPCRunner(const String& key, const String& host, int port, int priority, int n_parallel, - int timeout, int number, int repeat, int min_repeat_ms, - double cooldown_interval, bool enable_cpu_cache_flush, int device) { - auto node = make_object(); - node->key = key; - node->host = host; - node->port = port; - node->priority = priority; - node->timeout = timeout; - node->n_parallel = n_parallel; - node->number = number; - node->repeat = repeat; - node->min_repeat_ms = min_repeat_ms; - node->cooldown_interval = cooldown_interval; - node->enable_cpu_cache_flush = enable_cpu_cache_flush; - node->device = device; - data_ = std::move(node); -} - -Array RPCRunnerNode::Run(const Array& inputs, - const Array& build_results, int verbose) { - if (const auto* f = runtime::Registry::Get("auto_scheduler.rpc_runner.run")) { - Array results = - (*f)(inputs, build_results, key, host, port, priority, n_parallel, timeout, number, repeat, - min_repeat_ms, cooldown_interval, enable_cpu_cache_flush, verbose, device); - return results; - } else { - LOG(FATAL) << "auto_scheduler.rpc_runner.run is not registered. " - << "This is a function registered in Python, " - << "make sure the TVM Python runtime has been loaded successfully."; - } - return Array(); -} - -/********** MeasureCallback **********/ -PythonBasedMeasureCallback::PythonBasedMeasureCallback(PackedFunc callback_func) { - auto node = make_object(); - node->callback_func = std::move(callback_func); - data_ = std::move(node); -} - -void PythonBasedMeasureCallbackNode::Callback(const SearchPolicy& policy, - const Array& inputs, - const Array& results) { - if (auto* sketch_policy = static_cast(policy.operator->())) { - callback_func(GetRef(sketch_policy), inputs, results); - } else if (auto* empty_policy = static_cast(policy.operator->())) { - callback_func(GetRef(empty_policy), inputs, results); - } else { - LOG(FATAL) << "Unrecognized search policy type. Expect SketchPolicy or EmptyPolicy"; - } -} - -/********** ProgramMeasurer **********/ -ProgramMeasurer::ProgramMeasurer(ProgramBuilder builder, ProgramRunner runner, - Optional> callbacks, int verbose, - int max_continuous_error) { - auto node = make_object(); - node->builder = std::move(builder); - node->runner = std::move(runner); - node->callbacks = std::move(callbacks); - node->verbose = verbose; - node->max_continuous_error = max_continuous_error < 0 - ? ProgramMeasurerNode::DEFAULT_MAX_CONTINUOUS_ERROR - : max_continuous_error; - data_ = std::move(node); -} - -void ProgramMeasurerNode::Reset() { - ct = error_ct = 0; - best_flops.clear(); - best_ct.clear(); - best_state.clear(); - has_valid.clear(); -} - -Array ProgramMeasurerNode::Measure(const SearchTask& task, - const SearchPolicy& policy, - const Array& inputs, - int batch_size) { - auto t_begin = std::chrono::high_resolution_clock::now(); - - Array results; - results.reserve(inputs.size()); - - if (batch_size == -1) { - // set default batch size - batch_size = builder->n_parallel * 2; - } - - int old_verbosity = verbose; - - StdCout(verbose) << "Get " << inputs.size() << " programs to measure:" << std::endl; - - for (size_t i = 0; i < inputs.size(); i += batch_size) { - Array input_batch(inputs.begin() + i, - inputs.begin() + std::min(i + batch_size, inputs.size())); - Array result_batch; - - // build and run - SilentMeasure(task, input_batch, &result_batch); - - // update current best state according to the new measure result - for (size_t j = 0; j < input_batch.size(); ++j) { - const String& workload_key = input_batch[j]->task->workload_key; - double flops; - - if (result_batch[j]->error_no == 0) { - flops = task->compute_dag->flop_ct / FloatArrayMean(result_batch[j]->costs); - error_ct = 0; - has_valid.insert(workload_key); - } else { - flops = 0.0; - error_ct++; - } - - if (flops > best_flops[workload_key]) { - best_flops[workload_key] = flops; - best_state[workload_key] = input_batch[j]->state; - best_ct[workload_key] = ct; - } - - ct++; - StdCout(verbose, 2) << std::fixed << std::setprecision(2) << Chars('=', 50) << "\n" - << "No: " << ct << "\tGFLOPS: " << flops / 1e9 << " / " - << best_flops[workload_key] / 1e9 << "\tresults: " << result_batch[j] - << "\n" - << Chars('=', 50) << "\n" - << input_batch[j]->state << "\n"; - } - - // Call callback functions - if (callbacks) { - for (const auto& callback : callbacks.value()) { - callback->Callback(policy, input_batch, result_batch); - } - } - - // Store result batch - for (auto& res : result_batch) { - results.push_back(res); - } - - if (error_ct > max_continuous_error) { - LOG(WARNING) << "Too many errors happened during tuning. Switching to debug mode." - << std::endl; - verbose = 2; - } else { - verbose = old_verbosity; - } - } - - PrintTimeElapsed(t_begin, "measurement", verbose); - - return results; -} - -void ProgramMeasurerNode::SilentMeasure(const SearchTask& task, const Array& inputs, - Array* results) { - results->clear(); - results->reserve(inputs.size()); - - // Call builder and runner - Array build_res_batch = builder->Build(inputs, verbose); - Array result_batch = runner->Run(inputs, build_res_batch, verbose); - - // Store result batch - for (auto& res : result_batch) { - results->push_back(res); - } -} - -/********** Printing functions **********/ -TVM_STATIC_IR_FUNCTOR(ReprPrinter, vtable) - .set_dispatch([](const ObjectRef& ref, ReprPrinter* p) { - p->stream << "MeasureInput()"; - }); - -TVM_STATIC_IR_FUNCTOR(ReprPrinter, vtable) - .set_dispatch([](const ObjectRef& ref, ReprPrinter* p) { - auto* node = static_cast(ref.get()); - if (node->error_no == static_cast(MeasureErrorNO::kNoError)) { - p->stream << "MeasureResult(cost:["; - auto old_config = p->stream.precision(4); - for (size_t i = 0; i < node->costs.size(); ++i) { - auto pf = node->costs[i].as(); - ICHECK(pf != nullptr); - p->stream << pf->value; - if (i != node->costs.size() - 1) { - p->stream << ","; - } - } - p->stream.precision(old_config); - p->stream << "], "; - p->stream << "error_no:" << 0 << ", " - << "all_cost:" << node->all_cost << ", " - << "Tstamp:" << node->timestamp << ")"; - } else { - p->stream << "MeasureResult(" - << "error_type:" << ErrorNoToStr[node->error_no] << ", " - << "error_msg:" << node->error_msg << ", " - << "all_cost:" << node->all_cost << ", " - << "Tstamp:" << node->timestamp << ")"; - } - }); - -TVM_STATIC_IR_FUNCTOR(ReprPrinter, vtable) - .set_dispatch([](const ObjectRef& ref, ReprPrinter* p) { - auto* node = static_cast(ref.get()); - p->stream << "BuildResult(" << node->filename << ", " << node->error_no << ", " - << node->time_cost << ")"; - }); - -/********** Measure interface API for ffi **********/ -TVM_REGISTER_GLOBAL("auto_scheduler.MeasureInput").set_body_typed([](SearchTask task, State state) { - return MeasureInput(task, state); -}); - -TVM_REGISTER_GLOBAL("auto_scheduler.BuildResult") - .set_body_typed([](String filename, Array args, int error_no, String error_msg, - double time_cost) { - return BuildResult(filename, args, error_no, error_msg, time_cost); - }); - -TVM_REGISTER_GLOBAL("auto_scheduler.MeasureResult") - .set_body_typed([](Array costs, int error_no, String error_msg, double all_cost, - double timestamp) { - return MeasureResult(costs, error_no, error_msg, all_cost, timestamp); - }); - -TVM_REGISTER_GLOBAL("auto_scheduler.PythonBasedMeasureCallback") - .set_body_typed([](PackedFunc callback_func) { - return PythonBasedMeasureCallback(callback_func); - }); - -TVM_REGISTER_GLOBAL("auto_scheduler.ProgramMeasurer") - .set_body_typed([](ProgramBuilder builder, ProgramRunner runner, - Array callbacks, int verbose, int max_continuous_error) { - return ProgramMeasurer(builder, runner, callbacks, verbose, max_continuous_error); - }); - -TVM_REGISTER_GLOBAL("auto_scheduler.ProgramBuilderBuild") - .set_body_typed([](const ProgramBuilder& builder, const Array& inputs, - int verbose) { return builder->Build(inputs, verbose); }); - -TVM_REGISTER_GLOBAL("auto_scheduler.ProgramRunnerRun") - .set_body_typed([](const ProgramRunner& runner, const Array& inputs, - const Array& build_results, - int verbose) { return runner->Run(inputs, build_results, verbose); }); - -TVM_REGISTER_GLOBAL("auto_scheduler.LocalBuilder") - .set_body_typed([](int timeout, int n_parallel, const String& build_func) { - return LocalBuilder(timeout, n_parallel, build_func); - }); - -TVM_REGISTER_GLOBAL("auto_scheduler.LocalRunner") - .set_body_typed([](int timeout, int number, int repeat, int min_repeat_ms, - double cooldown_interval, bool enable_cpu_cache_flush, int device) { - return LocalRunner(timeout, number, repeat, min_repeat_ms, cooldown_interval, - enable_cpu_cache_flush, device); - }); - -TVM_REGISTER_GLOBAL("auto_scheduler.RPCRunner") - .set_body_typed([](const String& key, const String& host, int port, int priority, - int n_parallel, int timeout, int number, int repeat, int min_repeat_ms, - double cooldown_interval, bool enable_cpu_cache_flush, int device) { - return RPCRunner(key, host, port, priority, n_parallel, timeout, number, repeat, - min_repeat_ms, cooldown_interval, enable_cpu_cache_flush, device); - }); - -} // namespace auto_scheduler -} // namespace tvm diff --git a/src/auto_scheduler/measure_record.cc b/src/auto_scheduler/measure_record.cc deleted file mode 100644 index af37443d91e2..000000000000 --- a/src/auto_scheduler/measure_record.cc +++ /dev/null @@ -1,486 +0,0 @@ -/* - * 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 auto_scheduler/measure_record.cc - * \brief Json serialization format for dumping and loading tuning records. - */ - -#include -#include -#include -#include -#include - -#include -#include -#include -#include -#include - -#include "utils.h" - -// Json serialization handler for MeasureInput, MeasureResult -// (and recursively for SearchTask, State, Step, ...) -namespace dmlc { -namespace json { - -template <> -struct Handler<::tvm::Array<::tvm::auto_scheduler::Stage>> { - inline static void Write(dmlc::JSONWriter* writer, - const ::tvm::Array<::tvm::auto_scheduler::Stage>& data) { - writer->BeginArray(false); - writer->EndArray(); - } - inline static void Read(dmlc::JSONReader* reader, - ::tvm::Array<::tvm::auto_scheduler::Stage>* data) { - bool s; - reader->BeginArray(); - s = reader->NextArrayItem(); - ICHECK(!s); - } -}; - -template <> -struct Handler<::tvm::Array<::tvm::auto_scheduler::Step>> { - inline static void Write(dmlc::JSONWriter* writer, - const ::tvm::Array<::tvm::auto_scheduler::Step>& data) { - writer->BeginArray(false); - for (const auto& step : data) { - writer->WriteArraySeperator(); - writer->BeginArray(false); - step->WriteToRecord(writer); - writer->EndArray(); - } - writer->EndArray(); - } - - inline static void Read(dmlc::JSONReader* reader, - ::tvm::Array<::tvm::auto_scheduler::Step>* data) { - bool s; - reader->BeginArray(); - data->clear(); - while (reader->NextArrayItem()) { - reader->BeginArray(); - data->push_back(::tvm::auto_scheduler::StepReadFromRecord(reader)); - s = reader->NextArrayItem(); - ICHECK(!s); - } - } -}; - -template <> -struct Handler<::tvm::auto_scheduler::StateNode> { - inline static void Write(dmlc::JSONWriter* writer, const ::tvm::auto_scheduler::StateNode& data) { - writer->BeginArray(false); - writer->WriteArrayItem(data.stages); - writer->WriteArrayItem(data.transform_steps); - writer->EndArray(); - } - inline static void Read(dmlc::JSONReader* reader, ::tvm::auto_scheduler::StateNode* data) { - bool s; - reader->BeginArray(); - s = reader->NextArrayItem(); - ICHECK(s); - reader->Read(&data->stages); - s = reader->NextArrayItem(); - ICHECK(s); - reader->Read(&data->transform_steps); - s = reader->NextArrayItem(); - ICHECK(!s); - } -}; - -template <> -struct Handler<::tvm::auto_scheduler::HardwareParamsNode> { - inline static void Write(dmlc::JSONWriter* writer, - const ::tvm::auto_scheduler::HardwareParamsNode& data) { - writer->BeginArray(false); - writer->WriteArrayItem(data.num_cores); - writer->WriteArrayItem(data.vector_unit_bytes); - writer->WriteArrayItem(data.cache_line_bytes); - writer->WriteArrayItem(data.max_shared_memory_per_block); - writer->WriteArrayItem(data.max_local_memory_per_block); - writer->WriteArrayItem(data.max_threads_per_block); - writer->WriteArrayItem(data.max_vthread_extent); - writer->WriteArrayItem(data.warp_size); - writer->EndArray(); - } - inline static void Read(dmlc::JSONReader* reader, - ::tvm::auto_scheduler::HardwareParamsNode* data) { - bool s; - reader->BeginArray(); - s = reader->NextArrayItem(); - CHECK(s); - reader->Read(&data->num_cores); - s = reader->NextArrayItem(); - CHECK(s); - reader->Read(&data->vector_unit_bytes); - s = reader->NextArrayItem(); - CHECK(s); - reader->Read(&data->cache_line_bytes); - s = reader->NextArrayItem(); - CHECK(s); - reader->Read(&data->max_shared_memory_per_block); - s = reader->NextArrayItem(); - CHECK(s); - reader->Read(&data->max_local_memory_per_block); - s = reader->NextArrayItem(); - CHECK(s); - reader->Read(&data->max_threads_per_block); - s = reader->NextArrayItem(); - CHECK(s); - reader->Read(&data->max_vthread_extent); - s = reader->NextArrayItem(); - CHECK(s); - reader->Read(&data->warp_size); - s = reader->NextArrayItem(); - CHECK(!s); - } -}; - -template <> -struct Handler<::tvm::auto_scheduler::SearchTaskNode> { - inline static void Write(dmlc::JSONWriter* writer, - const ::tvm::auto_scheduler::SearchTaskNode& data) { - writer->BeginArray(false); - writer->WriteArrayItem(std::string(data.workload_key)); - writer->WriteArrayItem(data.target->str()); - writer->WriteArrayItem(*data.hardware_params.get()); - ::tvm::Target target = data.target; - ::tvm::Target target_host = data.target_host; - ::tvm::CheckAndUpdateHostConsistency(&target, &target_host); - if (target_host.defined()) { - writer->WriteArrayItem(target_host->str()); - } else { - writer->WriteArrayItem(std::string("")); - } - writer->WriteArrayItem(static_cast(data.layout_rewrite_option)); - writer->WriteArraySeperator(); - writer->BeginArray(false); - for (const auto& i : data.task_input_names) { - writer->WriteArrayItem(std::string(i)); - } - writer->EndArray(); - writer->EndArray(); - } - inline static void Read(dmlc::JSONReader* reader, ::tvm::auto_scheduler::SearchTaskNode* data) { - bool s; - std::string str_value; - int int_value; - auto hardware_params_node = ::tvm::make_object<::tvm::auto_scheduler::HardwareParamsNode>(); - reader->BeginArray(); - s = reader->NextArrayItem(); - ICHECK(s); - reader->Read(&str_value); - data->workload_key = std::move(str_value); - s = reader->NextArrayItem(); - ICHECK(s); - reader->Read(&str_value); - data->target = ::tvm::Target(str_value); - s = reader->NextArrayItem(); - if (s) { - reader->Read(hardware_params_node.get()); - s = reader->NextArrayItem(); - data->hardware_params = ::tvm::auto_scheduler::HardwareParams(hardware_params_node); - if (s) { - reader->Read(&str_value); - if (!str_value.empty()) { - data->target_host = ::tvm::Target(str_value); - ::tvm::CheckAndUpdateHostConsistency(&data->target, &data->target_host); - } - s = reader->NextArrayItem(); - ICHECK(s); - reader->Read(&int_value); - data->layout_rewrite_option = ::tvm::auto_scheduler::LayoutRewriteOption(int_value); - s = reader->NextArrayItem(); - if (s) { - reader->BeginArray(); - s = reader->NextArrayItem(); - while (s) { - reader->Read(&str_value); - data->task_input_names.push_back(str_value); - s = reader->NextArrayItem(); - } - // Process the end of array - s = reader->NextArrayItem(); - } - ICHECK(!s); - } - } - } -}; - -template <> -struct Handler<::tvm::auto_scheduler::MeasureInputNode> { - inline static void Write(dmlc::JSONWriter* writer, - const ::tvm::auto_scheduler::MeasureInputNode& data) { - writer->BeginArray(false); - writer->WriteArrayItem(*data.task.operator->()); - writer->WriteArrayItem(*data.state.operator->()); - writer->EndArray(); - } - inline static void Read(dmlc::JSONReader* reader, ::tvm::auto_scheduler::MeasureInputNode* data) { - auto task_node = ::tvm::make_object<::tvm::auto_scheduler::SearchTaskNode>(); - auto state_node = ::tvm::make_object<::tvm::auto_scheduler::StateNode>(); - state_node->concrete = true; - - bool s; - reader->BeginArray(); - s = reader->NextArrayItem(); - ICHECK(s); - reader->Read(task_node.get()); - s = reader->NextArrayItem(); - ICHECK(s); - reader->Read(state_node.get()); - s = reader->NextArrayItem(); - ICHECK(!s); - - data->task = ::tvm::auto_scheduler::SearchTask(task_node); - data->state = ::tvm::auto_scheduler::State(state_node); - } -}; - -template <> -struct Handler<::tvm::auto_scheduler::MeasureResultNode> { - inline static void Write(dmlc::JSONWriter* writer, - const ::tvm::auto_scheduler::MeasureResultNode& data) { - writer->BeginArray(false); - writer->WriteArraySeperator(); - writer->BeginArray(false); - for (const auto& x : data.costs) { - auto pf = x.as<::tvm::tir::FloatImmNode>(); - ICHECK(pf != nullptr) << "Cost can only contain float values"; - writer->WriteArrayItem(pf->value); - } - writer->EndArray(); - writer->WriteArrayItem(data.error_no); - writer->WriteArrayItem(data.all_cost); - writer->WriteArrayItem(static_cast((data.timestamp))); - writer->EndArray(); - } - inline static void Read(dmlc::JSONReader* reader, - ::tvm::auto_scheduler::MeasureResultNode* data) { - std::vector double_list; - bool s; - reader->BeginArray(); - s = reader->NextArrayItem(); - ICHECK(s); - reader->Read(&double_list); - data->costs.clear(); - for (const auto& i : double_list) { - data->costs.push_back(::tvm::FloatImm(::tvm::DataType::Float(64), i)); - } - s = reader->NextArrayItem(); - ICHECK(s); - reader->Read(&data->error_no); - s = reader->NextArrayItem(); - ICHECK(s); - reader->Read(&data->all_cost); - s = reader->NextArrayItem(); - ICHECK(s); - reader->Read(&data->timestamp); - s = reader->NextArrayItem(); - ICHECK(!s); - } -}; - -} // namespace json -} // namespace dmlc - -namespace tvm { -namespace auto_scheduler { - -TVM_REGISTER_OBJECT_TYPE(RecordToFileNode); -TVM_REGISTER_OBJECT_TYPE(RecordReaderNode); - -RecordToFile::RecordToFile(String filename) { - auto node = make_object(); - node->filename = std::move(filename); - data_ = std::move(node); -} - -void WriteMeasureRecords(std::ostream* os, const Array& inputs, - const Array& results, const std::string log_version) { - dmlc::JSONWriter writer(os); - for (size_t i = 0; i < inputs.size(); ++i) { - writer.BeginObject(false); - writer.WriteObjectKeyValue("i", *inputs[i].operator->()); - writer.WriteObjectKeyValue("r", *results[i].operator->()); - writer.WriteObjectKeyValue("v", log_version); - writer.EndObject(); - *os << "\n"; - } -} - -void ReadMeasureRecord(const std::string& str, MeasureInputNode* inp, MeasureResultNode* res, - std::string* log_version) { - std::istringstream ss(str); - dmlc::JSONReader reader(&ss); - std::string key; - - reader.BeginObject(); - while (reader.NextObjectItem(&key)) { - if (key == "i") { - reader.Read(inp); - } else if (key == "r") { - reader.Read(res); - } else if (key == "v") { - reader.Read(log_version); - } else { - LOG(FATAL) << "Invalid key in json log: " << key; - } - } -} - -void RecordToFileNode::Callback(const SearchPolicy& policy, const Array& inputs, - const Array& results) { - std::ofstream ofs(filename, std::ofstream::app); - WriteMeasureRecords(&ofs, inputs, results); -} - -RecordReader::RecordReader(String filename) { - auto node = make_object(); - node->filename = filename; - node->infile.open(filename, std::ifstream::in); - data_ = std::move(node); -} - -RecordReaderNode::~RecordReaderNode() { infile.close(); } - -bool RecordReaderNode::ReadNext(MeasureInputNode* inp, MeasureResultNode* res) { - std::string log_version; - - while (std::getline(infile, cur_line_)) { - if (cur_line_[0] == '#' || cur_line_[0] == ' ') { - // skip comment lines begin with '#' or ' ' - continue; - } - ReadMeasureRecord(cur_line_, inp, res, &log_version); - return true; - } - - return false; -} - -std::pair, Array> RecordReaderNode::ReadLines(int max_size, - int skip_size) { - auto inp = make_object(); - auto res = make_object(); - Array inputs; - Array results; - - while (ReadNext(inp.get(), res.get())) { - if (skip_size > 0) { - skip_size--; - continue; - } - - inputs.push_back(inp->copy()); - results.push_back(res->copy()); - - if (max_size > 0 && static_cast(inputs.size()) >= max_size) { - break; - } - } - - return std::make_pair(inputs, results); -} - -TVM_REGISTER_GLOBAL("auto_scheduler.RecordToFile").set_body_typed([](const String& filename) { - return RecordToFile(filename); -}); - -TVM_REGISTER_GLOBAL("auto_scheduler.RecordReader").set_body_typed([](const String& filename) { - return RecordReader(filename); -}); - -TVM_REGISTER_GLOBAL("auto_scheduler.RecordReaderReadLines") - .set_body_typed([](RecordReader reader, int size, int skip_size) { - const auto& res = reader->ReadLines(size, skip_size); - return Array{res.first, res.second}; - }); - -TVM_REGISTER_GLOBAL("auto_scheduler.RecordReaderReadNext").set_body_typed([](RecordReader reader) { - auto inp = make_object(); - auto res = make_object(); - if (reader->ReadNext(inp.get(), res.get())) { - return Array{ObjectRef(inp), ObjectRef(res)}; - } else { - return Array(); - } -}); - -TVM_REGISTER_GLOBAL("auto_scheduler.ReadMeasureRecord").set_body_typed([](const std::string& str) { - auto inp = make_object(); - auto res = make_object(); - std::string log_version; - ReadMeasureRecord(str, inp.get(), res.get(), &log_version); - return Array{ObjectRef(inp), ObjectRef(res)}; -}); - -TVM_REGISTER_GLOBAL("auto_scheduler.WriteMeasureRecords") - .set_body_typed([](MeasureInput inp, MeasureResult res) { - auto inps = Array({inp}); - auto ress = Array({res}); - std::ostringstream ss; - WriteMeasureRecords(&ss, inps, ress); - return String(ss.str()); - }); - -TVM_REGISTER_GLOBAL("auto_scheduler.SaveRecords") - .set_body_typed([](String filename, Array in, Array res) { - std::ofstream ofs(filename, std::ofstream::app); - WriteMeasureRecords(&ofs, in, res); - }); - -TVM_REGISTER_GLOBAL("auto_scheduler.SerializeMeasureInput") - .set_body_typed([](const MeasureInput& input) { - std::ostringstream os; - dmlc::JSONWriter writer(&os); - writer.Write(*input.get()); - return os.str(); - }); - -TVM_REGISTER_GLOBAL("auto_scheduler.DeserializeMeasureInput").set_body_typed([](String json) { - std::istringstream ss(json); - dmlc::JSONReader reader(&ss); - auto inp = make_object(); - reader.Read(inp.get()); - return ObjectRef(inp); -}); - -TVM_REGISTER_GLOBAL("auto_scheduler.SerializeSearchTask") - .set_body_typed([](const SearchTask& search_task) { - std::ostringstream os; - dmlc::JSONWriter writer(&os); - writer.Write(*search_task.get()); - return os.str(); - }); - -TVM_REGISTER_GLOBAL("auto_scheduler.DeserializeSearchTask").set_body_typed([](String json) { - std::istringstream ss(json); - dmlc::JSONReader reader(&ss); - auto search_task = make_object(); - reader.Read(search_task.get()); - return ObjectRef(search_task); -}); - -} // namespace auto_scheduler -} // namespace tvm diff --git a/src/auto_scheduler/search_policy/empty_policy.cc b/src/auto_scheduler/search_policy/empty_policy.cc deleted file mode 100644 index 79f98793d848..000000000000 --- a/src/auto_scheduler/search_policy/empty_policy.cc +++ /dev/null @@ -1,125 +0,0 @@ -/* - * 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 auto_scheduler/search_policy/empty_policy.cc - * \brief A simple example of the search policy which always returns the initial naive schedule - * (state). - */ - -#include "empty_policy.h" - -#include -#include - -#include - -#include "utils.h" - -namespace tvm { -namespace auto_scheduler { - -TVM_REGISTER_NODE_TYPE(EmptyPolicyNode); - -EmptyPolicy::EmptyPolicy(SearchTask task, Optional> init_search_callbacks) { - auto node = make_object(); - node->search_task = task; - - // Run init_search_callbacks before the search process - // This Interface is usually used to set some init status - if (init_search_callbacks) { - node->RunCallbacks(init_search_callbacks.value()); - } - - data_ = std::move(node); -} - -State EmptyPolicyNode::Search(int num_measure_trials, int early_stopping, - int num_measures_per_round, ProgramMeasurer measurer) { - // Basic design principe: `SearchOneRound()` several times to get candidate states, - // measure them and return the best one - // Measure is disabled if num_measure_trials <= 1 - if (num_measure_trials <= 1) { - const auto& res = SearchOneRound(); - ICHECK_GT(res.size(), 0); - - return res[0]; - } else { - Array inputs; - Array results; - - measurer->Reset(); - int ct = 0; - // In each round, we call SearchOneRound to get several candidate states, - // then use ProgramMeasurer to measure their performance. - while (ct < num_measure_trials) { - const auto& res = SearchOneRound(); - ct += res.size(); - // Build MeasureInputs for measuring - inputs.clear(); - for (const auto& state : res) { - inputs.push_back(MeasureInput(search_task, state)); - } - // Perform measurement. - // ProgramMeasurer will record the state with best performance during measure process - results = measurer->Measure(search_task, GetRef(this), inputs); - } - - // Return a state with best measured performance - return measurer->best_state[search_task->workload_key]; - } -} - -std::pair, Array> EmptyPolicyNode::ContinueSearchOneRound( - int num_measure, ProgramMeasurer measurer) { - Array best_states; - Array inputs; - Array results; - - // Search one round to get promising states - PrintTitle("Search", verbose); - best_states = SearchOneRound(); - - // Measure these states - PrintTitle("Measure", verbose); - for (const auto& state : best_states) { - inputs.push_back(MeasureInput(search_task, state)); - } - results = measurer->Measure(search_task, GetRef(this), inputs); - - return std::make_pair(std::move(inputs), std::move(results)); -} - -// As an example policy, EmptyPolicy always returns a init state -Array EmptyPolicyNode::SearchOneRound() { - Array res; - - // Simply return the initial naive schedule (state). - res.push_back(search_task->compute_dag->init_state); - - return res; -} - -TVM_REGISTER_GLOBAL("auto_scheduler.EmptyPolicy") - .set_body_typed([](SearchTask task, Optional> init_search_callbacks) { - return EmptyPolicy(task, init_search_callbacks); - }); - -} // namespace auto_scheduler -} // namespace tvm diff --git a/src/auto_scheduler/search_policy/empty_policy.h b/src/auto_scheduler/search_policy/empty_policy.h deleted file mode 100644 index 2219ebce83f0..000000000000 --- a/src/auto_scheduler/search_policy/empty_policy.h +++ /dev/null @@ -1,77 +0,0 @@ -/* - * 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 auto_scheduler/search_policy/empty_policy.h - * \brief A simple example of the search policy which always returns the initial naive schedule - * (state). - */ - -#ifndef TVM_AUTO_SCHEDULER_SEARCH_POLICY_EMPTY_POLICY_H_ -#define TVM_AUTO_SCHEDULER_SEARCH_POLICY_EMPTY_POLICY_H_ - -#include -#include -#include - -#include - -namespace tvm { -namespace auto_scheduler { - -/*! - * \brief A simple example of the search policy which always returns the initial naive schedule - * (state). - * The key implementation for this structure is `Search()`, check `empty_policy.cc` for more - * details. - */ -class EmptyPolicyNode : public SearchPolicyNode { - public: - State Search(int num_measure_trials, int early_stopping, int num_measures_per_round, - ProgramMeasurer measurer) final; - - std::pair, Array> ContinueSearchOneRound( - int num_measure, ProgramMeasurer measurer) final; - - static constexpr const char* _type_key = "auto_scheduler.EmptyPolicy"; - TVM_DECLARE_FINAL_OBJECT_INFO(EmptyPolicyNode, SearchPolicyNode); - - private: - /*! - * \brief Use a sub function to generate several candidate states in each search round. - * \returns The generated states - */ - Array SearchOneRound(); -}; - -/*! - * \brief Managed reference to EmptyPolicyNode. - * \sa EmptyPolicyNode - */ -class EmptyPolicy : public SearchPolicy { - public: - explicit EmptyPolicy(SearchTask task, Optional> init_search_callbacks); - - TVM_DEFINE_MUTABLE_OBJECT_REF_METHODS(EmptyPolicy, SearchPolicy, EmptyPolicyNode); -}; - -} // namespace auto_scheduler -} // namespace tvm - -#endif // TVM_AUTO_SCHEDULER_SEARCH_POLICY_EMPTY_POLICY_H_ diff --git a/src/auto_scheduler/search_policy/search_policy.cc b/src/auto_scheduler/search_policy/search_policy.cc deleted file mode 100644 index 196bee8ff0e2..000000000000 --- a/src/auto_scheduler/search_policy/search_policy.cc +++ /dev/null @@ -1,121 +0,0 @@ -/* - * 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 auto_scheduler/search_policy/search_policy.cc - * \brief The base class of search policies. - */ - -#include -#include -#include - -#include "utils.h" - -namespace tvm { -namespace auto_scheduler { - -TVM_REGISTER_OBJECT_TYPE(SearchCallbackNode); -TVM_REGISTER_OBJECT_TYPE(SearchPolicyNode); -TVM_REGISTER_OBJECT_TYPE(PreloadMeasuredStatesNode); - -void SearchPolicyNode::PreloadMeasuredStates(const String& log_file) { - RecordReader reader = RecordReader(log_file); - const auto& res = reader->ReadLines(-1); - size_t log_size = res.first.size(); - ICHECK_EQ(log_size, res.second.size()); - if (log_size) { - Array measured_states; - std::vector measured_throughputs; - for (size_t i = 0; i < log_size; i++) { - const auto& inp = res.first[i]; - if (inp->task->workload_key == search_task->workload_key && - inp->task->target->kind->name.compare(search_task->target->kind->name) == 0) { - State state = search_task->compute_dag->init_state; - auto pstate = state.CopyOnWrite(); - pstate->transform_steps = inp->state->transform_steps; - for (const auto& step : pstate->transform_steps) { - StepApplyToState(step, &state, search_task->compute_dag); - } - measured_states.push_back(std::move(state)); - measured_throughputs.push_back( - res.second[i]->error_no == 0 ? (1.0 / FloatArrayMean(res.second[i]->costs)) : 0.0); - } - } - // We can assume the recorded states will all be valid after infer bound - measured_states = search_task->compute_dag.InferBound(measured_states); - for (size_t i = 0; i < measured_states.size(); i++) { - auto& state = measured_states[i]; - const auto& state_str = state.ToStr(); - if (!measured_states_set_.count(state_str)) { - measured_states_set_.insert(state_str); - if (measured_throughputs[i] != 0.0) { - measured_states_vector_.emplace_back(std::move(state)); - measured_states_throughputs_.emplace_back(measured_throughputs[i]); - } - } - } - - StdCout(verbose) << "SearchPolicy: Loaded " << measured_states_set_.size() - << " measurement records from " << log_file << " for " - << search_task->workload_key << std::endl; - } else { - StdCout(verbose) << "SearchPolicy: No measurement records found in " << log_file << " for " - << search_task->workload_key << std::endl; - } -} - -void SearchPolicyNode::RunCallbacks(const Array& callbacks) { - for (const auto& callback : callbacks) { - callback->Callback(this); - } -} - -PreloadMeasuredStates::PreloadMeasuredStates(String filename) { - auto node = make_object(); - node->filename = std::move(filename); - data_ = std::move(node); -} - -void PreloadMeasuredStatesNode::Callback(SearchPolicyNode* policy) { - policy->PreloadMeasuredStates(filename); -} - -TVM_REGISTER_GLOBAL("auto_scheduler.SearchPolicyRunCallbacks") - .set_body_typed([](SearchPolicy policy, Optional> callbacks) { - if (callbacks) { - policy->RunCallbacks(callbacks.value()); - } - }); - -TVM_REGISTER_GLOBAL("auto_scheduler.SearchPolicyContinueSearchOneRound") - .set_body_typed([](SearchPolicy policy, int num_measure, ProgramMeasurer measurer) { - auto [inputs, results] = policy->ContinueSearchOneRound(num_measure, measurer); - return Array{inputs, results}; - }); - -TVM_REGISTER_GLOBAL("auto_scheduler.SearchPolicySetVerbose") - .set_body_typed([](SearchPolicy policy, int verbose) { policy->verbose = verbose; }); - -TVM_REGISTER_GLOBAL("auto_scheduler.PreloadMeasuredStates").set_body_typed([](String filename) { - return PreloadMeasuredStates(filename); -}); - -} // namespace auto_scheduler -} // namespace tvm diff --git a/src/auto_scheduler/search_policy/sketch_policy.cc b/src/auto_scheduler/search_policy/sketch_policy.cc deleted file mode 100644 index 8b0faed5b5b6..000000000000 --- a/src/auto_scheduler/search_policy/sketch_policy.cc +++ /dev/null @@ -1,728 +0,0 @@ -/* - * 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 auto_scheduler/search_policy/sketch_search_policy.h - * \brief The search policy that searches in a hierarchical search space defined by sketches. - * The policy randomly samples programs from the space defined by sketches - * and use evolutionary search to fine-tune them. - */ - -#include "sketch_policy.h" - -#include -#include - -#include -#include -#include -#include -#include -#include -#include -#include -#include -#include -#include - -#include "sketch_policy_rules.h" - -namespace tvm { -namespace auto_scheduler { - -/********** Sketch generation rules **********/ -static RuleSkipStage rule_skip_stage; -static RuleAlwaysInline rule_always_inline; -static RuleMultiLevelTiling rule_multi_level_tiling; -static RuleMultiLevelTilingWithFusion rule_multi_level_tiling_with_fusion; -static RuleAddCacheRead rule_add_cache_read_stage; -static RuleAddCacheWrite rule_add_cache_write_stage; -static RuleAddRfactor rule_add_rfactor; -static RuleCrossThreadReduction rule_cross_thread_reduction; -static RuleSimplifyComputeWithConstTensor rule_simplify_compute_with_const_tensor; -static RuleSpecialComputeLocationGPU rule_special_compute_location_gpu; - -/********** Init population rules **********/ -static InitFillTileSize init_fill_tile_size; -static InitChangeComputeLocation init_change_compute_location; -static InitParallel init_parallel; -static InitUnroll init_unroll; -static InitVectorization init_vectorization; -static InitThreadBind init_thread_bind; - -/********** Sketch policy **********/ -TVM_REGISTER_NODE_TYPE(SketchPolicyNode); - -SketchPolicy::SketchPolicy(SearchTask task, CostModel program_cost_model, - Map params, int seed, int verbose, - Optional> init_search_callbacks) { - auto node = make_object(); - node->search_task = std::move(task); - node->program_cost_model = std::move(program_cost_model); - node->rand_gen = std::mt19937(seed); - node->params = std::move(params); - node->verbose = verbose; - node->sample_init_min_pop_ = - GetIntParam(node->params, SketchParamKey::SampleInitPopulation::min_population); - - if (init_search_callbacks) { - PrintTitle("Call init-search callbacks", verbose); - // Candidates: - // - auto_scheduler.PreloadMeasuredStates: Load already measured states to - // `measured_states_set_`, `measured_states_vector_` and `measured_states_throughputs_`. - // - auto_scheduler.PreloadCustomSketchRule: Add user custom sketch rules to `sketch_rules`, - // these rules will be processed prior to the default rules. - node->RunCallbacks(init_search_callbacks.value()); - } - - // NOTE: There are strong dependency among the rules below, - // so the order to push them into the vector should be considered carefully. - if (IsCPUTask(node->search_task)) { - // Sketch Generation Rules - node->sketch_rules.push_back(&rule_always_inline); - node->sketch_rules.push_back(&rule_simplify_compute_with_const_tensor); - node->sketch_rules.push_back(&rule_add_rfactor); - node->sketch_rules.push_back(&rule_add_cache_write_stage); - node->sketch_rules.push_back(&rule_multi_level_tiling_with_fusion); - node->sketch_rules.push_back(&rule_multi_level_tiling); - node->sketch_rules.push_back(&rule_skip_stage); - - // Initial Population Generation Rules - node->init_rules.push_back(&init_fill_tile_size); - node->init_rules.push_back(&init_change_compute_location); - node->init_rules.push_back(&init_parallel); - node->init_rules.push_back(&init_unroll); - node->init_rules.push_back(&init_vectorization); - - // Mutation Rules for Evolutionary Search - node->mutation_rules.push_back(std::make_shared(0.90)); - node->mutation_rules.push_back(std::make_shared(0.04)); - node->mutation_rules.push_back(std::make_shared(0.05)); - node->mutation_rules.push_back(std::make_shared(0.01)); - } else if (IsGPUTask(node->search_task)) { - // Sketch Generation Rules - if (node->search_task->target->GetAttr("device", "") == "mali") { - node->sketch_rules.push_back(&rule_always_inline); - node->sketch_rules.push_back(&rule_simplify_compute_with_const_tensor); - node->sketch_rules.push_back(&rule_add_rfactor); - node->sketch_rules.push_back(&rule_add_cache_write_stage); - node->sketch_rules.push_back(&rule_multi_level_tiling_with_fusion); - node->sketch_rules.push_back(&rule_multi_level_tiling); - node->sketch_rules.push_back(&rule_skip_stage); - } else { - node->sketch_rules.push_back(&rule_add_cache_read_stage); - node->sketch_rules.push_back(&rule_special_compute_location_gpu); - node->sketch_rules.push_back(&rule_always_inline); - node->sketch_rules.push_back(&rule_simplify_compute_with_const_tensor); - node->sketch_rules.push_back(&rule_cross_thread_reduction); - node->sketch_rules.push_back(&rule_add_cache_write_stage); - node->sketch_rules.push_back(&rule_multi_level_tiling_with_fusion); - node->sketch_rules.push_back(&rule_multi_level_tiling); - node->sketch_rules.push_back(&rule_skip_stage); - } - - // Initial Population Generation Rules - node->init_rules.push_back(&init_fill_tile_size); - node->init_rules.push_back(&init_thread_bind); - node->init_rules.push_back(&init_unroll); - - if (node->search_task->target->GetAttr("device", "") == "mali") { - node->init_rules.push_back(&init_vectorization); - } - - // Mutation Rules for Evolutionary Search - node->mutation_rules.push_back(std::make_shared(0.90)); - node->mutation_rules.push_back(std::make_shared(0.10)); - } else { - LOG(FATAL) << "No default sketch rules for target: " << node->search_task->target; - } - - data_ = std::move(node); -} - -State SketchPolicyNode::Search(int n_trials, int early_stopping, int num_measure_per_iter, - ProgramMeasurer measurer) { - num_measure_per_iter_ = num_measure_per_iter; - - if (n_trials <= 1) { - // No measurement is allowed - const Array& best_states = SearchOneRound(0); - ICHECK_GT(best_states.size(), 0); - return best_states[0]; - } else { - int num_random = - static_cast(GetDoubleParam(params, SketchParamKey::eps_greedy) * num_measure_per_iter); - early_stopping = early_stopping < 0 ? std::numeric_limits::max() >> 1 : early_stopping; - measurer->Reset(); - - int ct = 0; - int empty_retry_count = GetIntParam(params, SketchParamKey::empty_retry_count); - Array best_states, random_states; - Array inputs; - Array results; - while (ct < n_trials) { - if (!inputs.empty()) { - auto t_begin = std::chrono::high_resolution_clock::now(); - - // Retrain the cost model before the next search round - PrintTitle("Train cost model", verbose); - program_cost_model->Update(inputs, results); - - PrintTimeElapsed(t_begin, "training", verbose); - } - - // Search one round to get promising states - PrintTitle("Search", verbose); - best_states = SearchOneRound(num_random * 3, &random_states); - - // Infer bound. This is necessary for computing the correct ToStr() for redundancy check - best_states = search_task->compute_dag.InferBound(best_states); - random_states = search_task->compute_dag.InferBound(random_states); - - // Pick `num_measure_per_iter` states to measure, check hash to remove already measured state - // Also pick some random states to do eps-greedy - inputs = PickStatesWithEpsGreedy(best_states, random_states, n_trials - ct); - - // Currently it's hard to detect if all of the search space has been traversed - // Stop if no extra valid states found in several retries - if (inputs.empty()) { - if (empty_retry_count-- > 0) { - continue; - } else { - StdCout(verbose) << "It seems all candidates in the search space have been measured." - << std::endl; - break; - } - } else { - // Reset the retry count - empty_retry_count = GetIntParam(params, SketchParamKey::empty_retry_count); - } - - // Measure candidate states - PrintTitle("Measure", verbose); - results = measurer->Measure(search_task, GetRef(this), inputs); - ct += inputs.size(); - - // Check if reach the early stopping condition - if (ct - measurer->best_ct[search_task->workload_key] > early_stopping && - measurer->has_valid.count(search_task->workload_key)) { - StdCout(verbose) << "Stop early since no performance improvement in the last " - << early_stopping << " measurements trials.\n"; - break; - } - - // Update measured states throughputs. These states will join the EvolutionarySearch in later - // search rounds. - for (const auto& res : results) { - measured_states_throughputs_.push_back(1.0 / FloatArrayMean(res->costs)); - } - } - PrintTitle("Done", verbose); - - return measurer->best_state[search_task->workload_key]; - } -} - -std::pair, Array> SketchPolicyNode::ContinueSearchOneRound( - int num_measure, ProgramMeasurer measurer) { - num_measure_per_iter_ = num_measure; - - Array best_states, random_states; - Array inputs; - Array results; - int num_random = static_cast(GetDoubleParam(params, "eps_greedy") * num_measure); - - // Search one round to get promising states - PrintTitle("Search", verbose); - best_states = SearchOneRound(num_random * 3, &random_states); - - // Infer bound. This is necessary for computing the correct ToStr() for redundancy check - best_states = search_task->compute_dag.InferBound(best_states); - random_states = search_task->compute_dag.InferBound(random_states); - - // Pick `num_measure_per_iter` states to measure, check hash to remove already measured state - // Also pick some random states to do eps-greedy - inputs = PickStatesWithEpsGreedy(best_states, random_states, num_measure); - - // Measure candidate states - PrintTitle("Measure", verbose); - results = measurer->Measure(search_task, GetRef(this), inputs); - - // Update measured states throughputs. These states will join the EvolutionarySearch in later - // search rounds. - for (const auto& res : results) { - measured_states_throughputs_.push_back(1.0 / FloatArrayMean(res->costs)); - } - - auto t_begin = std::chrono::high_resolution_clock::now(); - - // Update the cost model - PrintTitle("Train cost model", verbose); - program_cost_model->Update(inputs, results); - - PrintTimeElapsed(t_begin, "training", verbose); - - return std::make_pair(std::move(inputs), std::move(results)); -} - -Array SketchPolicyNode::SearchOneRound(int num_random_states, Array* random_states) { - // Get parameters - int population = GetIntParam(params, SketchParamKey::EvolutionarySearch::population); - int num_use_measured = std::min( - static_cast(measured_states_vector_.size()), - static_cast( - GetDoubleParam(params, SketchParamKey::SampleInitPopulation::use_measured_ratio) * - population)); - - // 1. Generate sketches - if (sketch_cache_.empty()) { - sketch_cache_ = GenerateSketches(); - } - - // 2. Sample the init population - Array init_population = SampleInitPopulation(sketch_cache_); - - // 3. Perform evolutionary search. - // Also insert already measured good states to the initial population - std::vector indices = Argsort(measured_states_throughputs_); - for (int i = 0; i < num_use_measured; i++) { - init_population.push_back(measured_states_vector_[indices[i]]); - } - // Sample some random states for eps-greedy - if (num_random_states > 0 && random_states != nullptr) { - *random_states = RandomSampleStates(init_population, &rand_gen, num_random_states); - } - return EvolutionarySearch(init_population, num_measure_per_iter_ * 2); -} - -Array SketchPolicyNode::GenerateSketches() { - const State& init_state = search_task->compute_dag->init_state; - - // Two ping pong buffers to avoid copy - Array states_buf1{init_state}, states_buf2; - Array* pnow = &states_buf1; - Array* pnext = &states_buf2; - - // A map that maps state to its current working position (stage_id) - std::unordered_map cur_stage_id_map; - cur_stage_id_map[init_state] = static_cast(init_state->stages.size()) - 1; - - // Derivation rule based enumeration - Array out_states; - while (!pnow->empty()) { - pnext->clear(); - for (const State& state : *pnow) { - int stage_id = cur_stage_id_map[state]; - - // Reaches to the terminal stage - if (stage_id < 0) { - out_states.push_back(state); - continue; - } - - // Try all derivation rules - for (const auto& rule : sketch_rules) { - auto cond = rule->MeetCondition(*this, state, stage_id); - if (cond != SketchGenerationRule::ConditionKind::kSkip) { - for (const auto& pair : rule->Apply(*this, state, stage_id)) { - cur_stage_id_map[pair.first] = pair.second; - pnext->push_back(pair.first); - } - // Skip the rest rules - if (cond == SketchGenerationRule::ConditionKind::kApplyAndSkipRest) { - break; - } - } - } - } - std::swap(pnow, pnext); - } - - // Hack for rfactor: Replace the split factor for rfactor to the undefined Expr(), - // so later we can sample random value for the split factor. - // Why don't we use Expr() when doing the split for rfactor at the first time? - // Because during ApplySteps, a rfactor with undefined Expr() will crash TVM. - // So rfactor with undefined Expr() will conflict with cache_write, cache_read, rfactor - // in other stages - for (size_t i = 0; i < out_states.size(); ++i) { - auto state = out_states[i]; - auto pstate = state.CopyOnWrite(); - for (size_t step_id = 0; step_id < pstate->transform_steps.size(); ++step_id) { - if (pstate->transform_steps[step_id]->IsInstance()) { - ICHECK_GE(step_id, 1); - int split_step_id = static_cast(step_id - 1); - auto step = pstate->transform_steps[split_step_id].as(); - ICHECK(step != nullptr); - pstate->transform_steps.Set( - split_step_id, SplitStep(step->stage_id, step->iter_id, step->extent, {NullOpt}, - step->inner_to_outer)); - } - } - out_states.Set(i, std::move(state)); - } - - StdCout(verbose) << "Generate Sketches\t\t#s: " << out_states.size() << std::endl; - return out_states; -} - -Array SketchPolicyNode::SampleInitPopulation(const Array& sketches) { - // Use this population as the parallel degree to do sampling - int population = GetIntParam(params, SketchParamKey::EvolutionarySearch::population); - - auto tic_begin = std::chrono::high_resolution_clock::now(); - - int fail_ct = 0; - Array out_states; - std::vector rand_gens; - rand_gens.reserve(population); - for (int i = 0; i < population; i++) { - rand_gens.push_back(std::mt19937(rand_gen())); - } - - std::unordered_set explored_state_strs; - size_t iter = 1; - size_t unchange_cnt = 0; - while (static_cast(out_states.size()) < sample_init_min_pop_) { - std::vector temp_states(population); - - // Sample a batch of states randomly - support::parallel_for(0, population, [this, &temp_states, &sketches, &rand_gens](int index) { - // Randomly choose a sketch - State tmp_s = sketches[(rand_gens[index])() % sketches.size()]; - // Apply random annotation rules one by one - bool valid = true; - for (const auto& rule : init_rules) { - if (rule->Apply(this, &tmp_s, &rand_gens[index]) == - PopulationGenerationRule::ResultKind::kInvalid) { - valid = false; - break; - } - } - if (valid) { - temp_states[index] = std::move(tmp_s); - } - }); - - // Filter out the states that were failed to apply initial rules - Array cand_states; - for (auto tmp_s : temp_states) { - if (tmp_s.defined()) { - cand_states.push_back(std::move(tmp_s)); - } else { - fail_ct++; - } - } - - unchange_cnt++; - if (!cand_states.empty()) { - // Run the cost model to make filter out states that failed to extract features. - // This may happen due to illegal schedules or the schedules that uses too much - // memory on GPU. - std::vector pop_scores; - pop_scores.reserve(cand_states.size()); - cand_states = search_task->compute_dag.InferBound(cand_states); - PruneInvalidState(search_task, &cand_states); - program_cost_model->Predict(search_task, cand_states, &pop_scores); - - for (size_t i = 0; i < cand_states.size(); i++) { - const auto state_str = cand_states[i].ToStr(); - if (pop_scores[i] > -1e10 && explored_state_strs.count(state_str) == 0) { - explored_state_strs.insert(state_str); - out_states.push_back(std::move(cand_states[i])); - unchange_cnt = 0; // Reset the counter once we found a valid state - } else { - fail_ct++; - } - } - } - - if (iter % 5 == 0) { - double duration = std::chrono::duration_cast>( - std::chrono::high_resolution_clock::now() - tic_begin) - .count(); - StdCout(verbose) << "Sample Iter: " << iter << std::fixed << std::setprecision(4) - << "\t#Pop: " << out_states.size() << "\t#Target: " << sample_init_min_pop_ - << "\tfail_ct: " << fail_ct << "\tTime elapsed: " << std::fixed - << std::setprecision(2) << duration << std::endl; - } - - if (unchange_cnt == 5) { - // Reduce the target size to avoid too-long time in this phase if no valid state was found - // in the past iterations - if (sample_init_min_pop_ > 1) { - sample_init_min_pop_ /= 2; - StdCout(verbose) << "#Target has been reduced to " << sample_init_min_pop_ - << " due to too many failures or duplications" << std::endl; - } - unchange_cnt = 0; - } - iter++; - } - - double duration = std::chrono::duration_cast>( - std::chrono::high_resolution_clock::now() - tic_begin) - .count(); - StdCout(verbose) << "Sample Initial Population\t#s: " << out_states.size() - << "\tfail_ct: " << fail_ct << "\tTime elapsed: " << std::fixed - << std::setprecision(2) << duration << std::endl; - return out_states; -} - -Array SketchPolicyNode::EvolutionarySearch(const Array& init_population, - int out_size) { - Array best_states; - auto tic_begin = std::chrono::high_resolution_clock::now(); - - size_t population = GetIntParam(params, SketchParamKey::EvolutionarySearch::population); - double mutation_prob = GetDoubleParam(params, SketchParamKey::EvolutionarySearch::mutation_prob); - int num_iters = GetIntParam(params, SketchParamKey::EvolutionarySearch::num_iters); - - bool is_cost_model_reasonable = !program_cost_model->IsInstance(); - if (!is_cost_model_reasonable && num_iters > 2) { - num_iters = 2; - StdCout(verbose) << "GA iteration number has been adjusted to " << num_iters - << " due to random cost model" << std::endl; - } - - // Two ping pong buffers to avoid copy. - Array states_buf1{init_population}, states_buf2; - states_buf1.reserve(population); - states_buf2.reserve(population); - Array* pnow = &states_buf1; - Array* pnext = &states_buf2; - - // A heap to keep the best states during evolution - using StateHeapItem = std::pair; - auto cmp = [](const StateHeapItem& left, const StateHeapItem& right) { - return left.second > right.second; - }; - std::vector heap; - std::unordered_set in_heap(measured_states_set_); - heap.reserve(out_size); - - // auxiliary global variables - std::vector pop_scores; - std::vector pop_selection_probs; - float max_score = -1e-10f; - pop_scores.reserve(population); - pop_selection_probs.reserve(population); - std::uniform_real_distribution<> dis(0.0, 1.0); - - // mutation rules - int mutation_success_ct, mutation_fail_ct; - mutation_success_ct = mutation_fail_ct = 0; - std::vector rule_weights; - std::vector rule_selection_probs; - for (const auto& rule : mutation_rules) { - rule_weights.push_back(rule->weight); - } - ComputePrefixSumProb(rule_weights, &rule_selection_probs); - - // Genetic Algorithm - for (int k = 0; k < num_iters + 1; ++k) { - // Maintain the heap - *pnow = search_task->compute_dag.InferBound(*pnow); - PruneInvalidState(search_task, pnow); - program_cost_model->Predict(search_task, *pnow, &pop_scores); - - for (size_t i = 0; i < pnow->size(); ++i) { - const State& state = (*pnow)[i]; - std::string state_str = state.ToStr(); - - if (in_heap.count(state_str) == 0) { - if (static_cast(heap.size()) < out_size) { - heap.emplace_back((*pnow)[i], pop_scores[i]); - std::push_heap(heap.begin(), heap.end(), cmp); - in_heap.insert(state_str); - } else if (pop_scores[i] > heap.front().second) { - std::string old_state_str = heap.front().first.ToStr(); - in_heap.erase(old_state_str); - in_heap.insert(state_str); - - std::pop_heap(heap.begin(), heap.end(), cmp); - heap.back() = StateHeapItem(state, pop_scores[i]); - std::push_heap(heap.begin(), heap.end(), cmp); - } - if (pop_scores[i] > max_score) { - max_score = pop_scores[i]; - } - } - } - - // Print statistical information - if (k % 5 == 0 || k == num_iters) { - StdCout(verbose) << "GA Iter: " << k; - if (!heap.empty()) { - StdCout(verbose) << std::fixed << std::setprecision(4) << "\tMax score: " << max_score - << std::fixed << std::setprecision(4) - << "\tMin score: " << heap.front().second; - } else { - StdCout(verbose) << "\tMax score: N/A\tMin score: N/A"; - } - StdCout(verbose) << "\t#Pop: " << heap.size() << "\t#M+: " << mutation_success_ct / (k + 1) - << "\t#M-: " << mutation_fail_ct / (k + 1) << std::endl; - } - if (k == num_iters) { - break; - } - - // Compute selection probability - ComputePrefixSumProb(pop_scores, &pop_selection_probs); - - // TODO(merrymercy, comaniac): add crossover. - - // Do mutation - while (pnext->size() < population) { - State tmp_s = (*pnow)[RandomChoose(pop_selection_probs, &rand_gen)]; - - if (dis(rand_gen) < mutation_prob) { - const auto& rule = mutation_rules[RandomChoose(rule_selection_probs, &rand_gen)]; - if (rule->Apply(this, &tmp_s, &rand_gen) == PopulationGenerationRule::ResultKind::kValid) { - pnext->push_back(std::move(tmp_s)); - mutation_success_ct++; - } else { - mutation_fail_ct++; - } - } else { - pnext->push_back(std::move(tmp_s)); - } - } - - std::swap(pnext, pnow); - pnext->clear(); - } - - // Copy best states in the heap to out_states - std::sort(heap.begin(), heap.end(), cmp); - for (auto& item : heap) { - best_states.push_back(std::move(item.first)); - } - - double duration = std::chrono::duration_cast>( - std::chrono::high_resolution_clock::now() - tic_begin) - .count(); - StdCout(verbose) << "EvolutionarySearch\t\t#s: " << best_states.size() - << "\tTime elapsed: " << std::fixed << std::setprecision(2) << duration - << std::endl; - return best_states; -} - -Array SketchPolicyNode::PickStatesWithEpsGreedy(const Array& best_states, - const Array& random_states, - int remaining_n_trials) { - int num_random = - static_cast(GetDoubleParam(params, SketchParamKey::eps_greedy) * num_measure_per_iter_); - int num_good = num_measure_per_iter_ - num_random; - - Array inputs; - size_t offset_best = 0, offset_random = 0; - - while (static_cast(inputs.size()) < std::min(num_measure_per_iter_, remaining_n_trials)) { - State state; - - bool has_best = offset_best < best_states.size(); - bool has_random = offset_random < random_states.size(); - - if (static_cast(inputs.size()) < num_good) { - // prefer best states - if (has_best) { - state = best_states[offset_best++]; - } else if (has_random) { - state = random_states[offset_random++]; - } else { - break; - } - } else { - // prefer random states - if (has_random) { - state = random_states[offset_random++]; - } else if (has_best) { - state = best_states[offset_best++]; - } else { - break; - } - } - - // Check if it has already been measured - std::string state_str = state.ToStr(); - if (!measured_states_set_.count(state_str)) { - measured_states_set_.insert(std::move(state_str)); - measured_states_vector_.push_back(state); - inputs.push_back(MeasureInput(search_task, state)); - } - } - - return inputs; -} - -/********** PreloadCustomSketchRule **********/ -TVM_REGISTER_OBJECT_TYPE(PreloadCustomSketchRuleNode); - -PreloadCustomSketchRule::PreloadCustomSketchRule(PackedFunc meet_condition_func, - PackedFunc apply_func, String rule_name) { - auto node = make_object(); - node->meet_condition_func = std::move(meet_condition_func); - node->apply_func = std::move(apply_func); - node->rule_name = std::move(rule_name); - data_ = std::move(node); -} - -void PreloadCustomSketchRuleNode::Callback(SearchPolicyNode* policy) { - CHECK(policy->IsInstance()); - auto sketch_policy = dynamic_cast(policy); - sketch_policy->sketch_rules.push_back( - new RuleCustomSketch(meet_condition_func, apply_func, rule_name)); - StdCout(policy->verbose) << "Custom sketch rule \"" << rule_name << "\" added." << std::endl; -} - -TVM_REGISTER_GLOBAL("auto_scheduler.SketchPolicy") - .set_body_typed([](SearchTask task, CostModel program_cost_model, Map params, - int seed, int verbose, - Optional> init_search_callbacks) { - return SketchPolicy(task, program_cost_model, params, seed, verbose, init_search_callbacks); - }); - -TVM_REGISTER_GLOBAL("auto_scheduler.SketchPolicyGenerateSketches") - .set_body_typed([](SketchPolicy policy) { return policy->GenerateSketches(); }); - -TVM_REGISTER_GLOBAL("auto_scheduler.SketchPolicySampleInitialPopulation") - .set_body_typed([](SketchPolicy policy) { - const Array& sketches = policy->GenerateSketches(); - - Array init_population = policy->SampleInitPopulation(sketches); - return init_population; - }); - -TVM_REGISTER_GLOBAL("auto_scheduler.SketchPolicyEvolutionarySearch") - .set_body_typed([](SketchPolicy policy, Array init_population, int out_size) { - Array states = policy->EvolutionarySearch(init_population, out_size); - return states; - }); - -TVM_REGISTER_GLOBAL("auto_scheduler.PrintTitle").set_body_typed([](std::string title) { - PrintTitle(title, 1); -}); - -TVM_REGISTER_GLOBAL("auto_scheduler.PreloadCustomSketchRule") - .set_body_typed([](PackedFunc meet_condition_func, PackedFunc apply_func, String rule_name) { - return PreloadCustomSketchRule(meet_condition_func, apply_func, rule_name); - }); - -} // namespace auto_scheduler -} // namespace tvm diff --git a/src/auto_scheduler/search_policy/sketch_policy.h b/src/auto_scheduler/search_policy/sketch_policy.h deleted file mode 100644 index faf058b45b19..000000000000 --- a/src/auto_scheduler/search_policy/sketch_policy.h +++ /dev/null @@ -1,237 +0,0 @@ -/* - * 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 auto_scheduler/search_policy/sketch_policy.h - * \brief This search policy constructs a search space according to the compute declaration. - * It then randomly samples programs from the search space and uses evolutionary search with a - * learned cost model to fine tune the sampled programs. - * The final optimized programs are sent to actual hardware for measurement. - * The above process is repeated until the auto-scheduler runs out of time budget. - * - * Reference: - * L. Zheng, C. Jia, M. Sun, Z. Wu, C. Yu, et al. "Ansor : Generating High-Performance Tensor - * Programs for Deep Learning." (OSDI 2020). - */ - -#ifndef TVM_AUTO_SCHEDULER_SEARCH_POLICY_SKETCH_POLICY_H_ -#define TVM_AUTO_SCHEDULER_SEARCH_POLICY_SKETCH_POLICY_H_ - -#include -#include - -#include -#include -#include -#include -#include -#include - -#include "sketch_policy_rules.h" -#include "utils.h" - -namespace tvm { -namespace auto_scheduler { - -/*! \brief String keys used in parameter map of SketchPolicy. */ -struct SketchParamKey { - /*! \brief Always allocate this percentage of measurements to random sampled states. */ - static constexpr const char* eps_greedy = "eps_greedy"; - /*! \brief Retry several times if SearchOneRound gets no valid state. */ - static constexpr const char* empty_retry_count = "retry_search_one_round_on_empty"; - - struct SampleInitPopulation { - /*! \brief The minimal size of valid population in the initial sampling. */ - static constexpr const char* min_population = "sample_init_min_population"; - /*! \brief The maximum percentage of measured states in the initial sampling. */ - static constexpr const char* use_measured_ratio = "sample_init_use_measured_ratio"; - }; - - struct EvolutionarySearch { - /*! \brief The population size of evolutionary search. */ - static constexpr const char* population = "evolutionary_search_population"; - /*! \brief The number of iterations performed by generic algorithm.*/ - static constexpr const char* num_iters = "evolutionary_search_num_iters"; - /*! \brief The mutation probability.*/ - static constexpr const char* mutation_prob = "evolutionary_search_mutation_prob"; - }; - - struct MultiLevelTiling { - /*! \brief The structure of multi-level tiling for CPU. */ - static constexpr const char* cpu_structure = "cpu_multi_level_tiling_structure"; - /*! \brief The structure of multi-level tiling for GPU. */ - static constexpr const char* gpu_structure = "gpu_multi_level_tiling_structure"; - }; - - /*! \brief The max inner most split factor. */ - static constexpr const char* max_innermost_split_factor = "max_innermost_split_factor"; - /*! \brief The max vectorize size. */ - static constexpr const char* max_vectorize_size = "max_vectorize_size"; - /*! \brief Whether disable compute location changing. */ - static constexpr const char* disable_change_compute_location = "disable_change_compute_location"; -}; - -class SketchPolicy; - -/*! - * \brief The search policy that searches in a hierarchical search space defined by sketches. - * The policy randomly samples programs from the space defined by sketches - * and use evolutionary search to fine-tune them. - */ -class SketchPolicyNode : public SearchPolicyNode { - public: - /*! \brief The cost model to estimate the complete schedules. */ - CostModel program_cost_model; - /*! \brief The parameters map for this search policy. */ - Map params; - /*! \brief The rules to generate sketches. */ - std::vector sketch_rules; - /*! \brief The rules to generate initial population. */ - std::vector init_rules; - /*! \brief The rules to mutate states in the evolutionary search. */ - std::vector> mutation_rules; - /*! \brief Random generator. */ - std::mt19937 rand_gen; - /*! \brief Memorize split space for Split. */ - SplitFactorizationMemo split_memo; - - State Search(int num_measure_trials, int early_stopping, int num_measures_per_round, - ProgramMeasurer measurer) final; - - std::pair, Array> ContinueSearchOneRound( - int num_measure, ProgramMeasurer measurer) final; - - /*! - * \brief Generate sketches. - * \return The generated sketches(states). - */ - Array GenerateSketches(); - - /*! - * \brief Sample the init population. - * \param sketches The initial sketches for the sampled population - * \return The generated states (the initial population). - */ - Array SampleInitPopulation(const Array& sketches); - - /*! - * \brief Perform evolutionary search. - * \param init_populations The states generated from init population. - * \param out_size The number of expected output states. - * \return The generated states after evolutionary search. - */ - Array EvolutionarySearch(const Array& init_populations, int out_size); - - static constexpr const char* _type_key = "auto_scheduler.SketchPolicy"; - - TVM_DECLARE_FINAL_OBJECT_INFO(SketchPolicyNode, SearchPolicyNode); - - private: - /*! - * \brief Run one round of the search pipeline. - * \param num_random_states Number of states that are picked randomly, this is used for - * eps-greedy policy. - * \param random_states The picked random states, used as one of the output of this function. - * \return The best several states generated in this search round. - */ - Array SearchOneRound(int num_random_states, Array* random_states = nullptr); - - /*! - * \brief Pick states from best states and random states with eps-greedy policy. - * \param best_states States picked by cost model. - * \param random_states States picked randomly. - * \param remaining_n_trials The remaining number of states need to be generated. - * \return The generated states to be measured, wrapped in MeasureInput. - */ - Array PickStatesWithEpsGreedy(const Array& best_states, - const Array& random_states, - int remaining_n_trials); - - /*! \brief The number of states to measure per iteration. */ - int num_measure_per_iter_; - - /*! \brief The cached sketches */ - Array sketch_cache_; - - /*! \brief The minimul output population of SampleInitPopulation */ - int sample_init_min_pop_; - - friend class SketchPolicy; -}; - -/*! - * \brief Managed reference to SketchPolicyNode. - * \sa SketchPolicyNode - */ -class SketchPolicy : public SearchPolicy { - public: - /*! - * \brief The constructor. - * \param task The SearchTask for the computation declaration. - * \param program_cost_model The cost model for complete programs. - * \param params The parameters map for this search process. - * \param seed The random seed of this search process. - * \param verbose Verbose level. 0 for silent, 1 to output information during schedule - * search. - * \param init_search_callbacks SearchCallback to be called before schedule search. - */ - SketchPolicy(SearchTask task, CostModel program_cost_model, Map params, - int seed, int verbose, Optional> init_search_callbacks); - - TVM_DEFINE_MUTABLE_OBJECT_REF_METHODS(SketchPolicy, SearchPolicy, SketchPolicyNode); -}; - -/*! \brief Pre-search callback function to load custom rules for sketch generation */ -class PreloadCustomSketchRuleNode : public SearchCallbackNode { - public: - /*! \brief The condition check function of this rule. */ - PackedFunc meet_condition_func; - /*! \brief The apply function of this rule. */ - PackedFunc apply_func; - /*! \brief The name of this rule. */ - String rule_name; - - void Callback(SearchPolicyNode* policy) final; - - static constexpr const char* _type_key = "auto_scheduler.PreloadCustomSketchRule"; - TVM_DECLARE_FINAL_OBJECT_INFO(PreloadCustomSketchRuleNode, SearchCallbackNode); -}; - -/*! - * \brief Managed reference to PreloadCustomSketchRuleNode. - * \sa PreloadCustomSketchRuleNode - */ -class PreloadCustomSketchRule : public SearchCallback { - public: - /*! - * \brief The constructor. - * \param meet_condition_func The condition check function of this rule. - * \param apply_func The apply function of this rule. - * \param rule_name The name of this rule. - */ - PreloadCustomSketchRule(PackedFunc meet_condition_func, PackedFunc apply_func, String rule_name); - - TVM_DEFINE_MUTABLE_OBJECT_REF_METHODS(PreloadCustomSketchRule, SearchCallback, - PreloadCustomSketchRuleNode); -}; - -} // namespace auto_scheduler -} // namespace tvm - -#endif // TVM_AUTO_SCHEDULER_SEARCH_POLICY_SKETCH_POLICY_H_ diff --git a/src/auto_scheduler/search_policy/sketch_policy_rules.cc b/src/auto_scheduler/search_policy/sketch_policy_rules.cc deleted file mode 100644 index 0bf6da255d2a..000000000000 --- a/src/auto_scheduler/search_policy/sketch_policy_rules.cc +++ /dev/null @@ -1,1242 +0,0 @@ -/* - * 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 auto_scheduler/search_policy/sketch_policy_rules.cc - * \brief Rules for generating the sketches, sampling the initial population, and mutating the - * population in SketchPolicy. - */ - -#include "sketch_policy_rules.h" - -#include -#include -#include -#include - -#include "sketch_policy.h" - -namespace tvm { -namespace auto_scheduler { - -static std::vector auto_unroll_configs_cpu = {0, 16, 64, 512}; -static std::vector auto_unroll_configs_gpu = {0, 16, 64, 512, 1024}; - -/********** Sketch Generation Rule **********/ -/********** RuleSkipStage **********/ - -SketchGenerationRule::ConditionKind RuleSkipStage::MeetCondition(const SketchPolicyNode& policy, - const State& state, - int stage_id) const { - // This rule should be the last rule, always return true to decrease the stage index count - return ConditionKind::kApply; -} - -std::vector> RuleSkipStage::Apply(const SketchPolicyNode& policy, - const State& state, int stage_id) const { - return {std::make_pair(state, stage_id - 1)}; -} - -/********** RuleAlwaysInline **********/ -inline bool ShouldAlwaysBeInlined(const SketchPolicyNode& policy, const State& state, - int stage_id) { - const SearchTask& task = policy.search_task; - const Stage& stage = state->stages[stage_id]; - - // Check the inline limitation of TE - if (stage->op_type == StageKind::kPlaceholder || IsOutputOp(task, state, stage_id) || - HasReduceIter(stage)) { - return false; - } - - if (IsGPUTask(task)) { // Greedily inline all inlinable ops on gpu - return true; - } else { - // Only always-inline strict-inlinable ops on cpu. - // The computation location of other ops will be tuned by InitChangeComputeLocation - // and MutateComputeLocation. - return IsStrictlyInlineable(task, state, stage_id); - } -} - -SketchGenerationRule::ConditionKind RuleAlwaysInline::MeetCondition(const SketchPolicyNode& policy, - const State& state, - int stage_id) const { - return ShouldAlwaysBeInlined(policy, state, stage_id) ? ConditionKind::kApplyAndSkipRest - : ConditionKind::kSkip; -} - -std::vector> RuleAlwaysInline::Apply(const SketchPolicyNode& policy, - const State& state, int stage_id) const { - State tmp_s = state; - tmp_s.compute_inline(stage_id); - return {std::make_pair(std::move(tmp_s), stage_id - 1)}; -} - -/********** RuleMultiLevelTiling **********/ - -SketchGenerationRule::ConditionKind RuleMultiLevelTiling::MeetCondition( - const SketchPolicyNode& policy, const State& state, int stage_id) const { - return NeedsMultilevelTiling(policy.search_task, state, stage_id) - ? ConditionKind::kApplyAndSkipRest - : ConditionKind::kSkip; -} - -std::vector> RuleMultiLevelTiling::Apply(const SketchPolicyNode& policy, - const State& state, - int stage_id) const { - const std::string& multi_level_tiling_structure = - IsGPUTask(policy.search_task) - ? GetStringParam(policy.params, SketchParamKey::MultiLevelTiling::gpu_structure) - : GetStringParam(policy.params, SketchParamKey::MultiLevelTiling::cpu_structure); - State tmp_s = DoMultiLevelTiling(state, stage_id, multi_level_tiling_structure); - return {std::make_pair(std::move(tmp_s), stage_id - 1)}; -} - -/********** RuleMultiLevelTilingWithFusion **********/ - -SketchGenerationRule::ConditionKind RuleMultiLevelTilingWithFusion::MeetCondition( - const SketchPolicyNode& policy, const State& state, int stage_id) const { - if (NeedsMultilevelTiling(policy.search_task, state, stage_id) && - HasSingleElementwiseMatchedConsumer(policy.search_task, state, stage_id)) { - // Always do fusion for stage with cache_write or is in GPU policy - return HasCacheWriteStage(state, stage_id) || IsGPUTask(policy.search_task) - ? ConditionKind::kApplyAndSkipRest - : ConditionKind::kApply; - } - return ConditionKind::kSkip; -} - -std::vector> RuleMultiLevelTilingWithFusion::Apply( - const SketchPolicyNode& policy, const State& state, int stage_id) const { - int target_stage_id; - ICHECK( - HasSingleElementwiseMatchedConsumer(policy.search_task, state, stage_id, &target_stage_id)); - const std::string& multi_level_tiling_structure = - IsGPUTask(policy.search_task) - ? GetStringParam(policy.params, SketchParamKey::MultiLevelTiling::gpu_structure) - : GetStringParam(policy.params, SketchParamKey::MultiLevelTiling::cpu_structure); - std::vector spatial_split_step_ids; - State base_state = - DoMultiLevelTiling(state, stage_id, multi_level_tiling_structure, &spatial_split_step_ids); - - std::vector> ret; - std::vector follow_tiling_levels = - IsGPUTask(policy.search_task) ? std::vector{3} : std::vector{1, 2}; - for (int level : follow_tiling_levels) { - if (tolower(multi_level_tiling_structure[level - 1]) != 's') { - continue; - } - State tmp_s = base_state; - tmp_s = FollowTiling(tmp_s, target_stage_id, spatial_split_step_ids, level); - const Iterator& target_iter = - tmp_s->stages[target_stage_id]->iters[level * spatial_split_step_ids.size() - 1]; - tmp_s.compute_at(stage_id, target_stage_id, target_iter); - ret.emplace_back(std::move(tmp_s), stage_id - 1); - } - - return ret; -} - -/********** RuleAddCacheRead **********/ - -SketchGenerationRule::ConditionKind RuleAddCacheRead::MeetCondition(const SketchPolicyNode& policy, - const State& state, - int stage_id) const { - const SearchTask& task = policy.search_task; - - // Don't cache_read a stage if it has multiple consumers - const std::set& consumers = GetConsumers(task, state, stage_id); - - if (consumers.size() == 0) return ConditionKind::kSkip; - // Don't cache_read a stage if its consumer does not need multi-level tiling - int target_stage_id = *consumers.begin(); - if (!NeedsMultilevelTiling(task, state, target_stage_id)) { - return ConditionKind::kSkip; - } - - // Don't cache_read a stage if its consumer does cross-thread reduction - if (HasCrossThreadReduction(state, target_stage_id)) { - return ConditionKind::kSkip; - } - - // Only direct producers can be cache read - const std::set& producers = GetDirectProducers(task, state, target_stage_id); - if (producers.find(stage_id) == producers.end()) { - return ConditionKind::kSkip; - } - - return ConditionKind::kApplyAndSkipRest; -} - -std::vector> RuleAddCacheRead::Apply(const SketchPolicyNode& policy, - const State& state, int stage_id) const { - const SearchTask& task = policy.search_task; - const std::set& consumers = GetConsumers(task, state, stage_id); - State tmp_s = state; - - int target_stage_id_offset = 0; - for (int orig_target_stage_id : consumers) { - int target_stage_id = orig_target_stage_id + target_stage_id_offset; - - // Cache read add shared memory - int added_stage_id = tmp_s.cache_read(stage_id, "shared", {target_stage_id}, task->compute_dag); - target_stage_id_offset++; - target_stage_id++; - - const auto& share_read_pos = - GetLastReduceIteratorInOutermostReduceTile(tmp_s->stages[target_stage_id]); - tmp_s.compute_at(added_stage_id, target_stage_id, share_read_pos); - } - - return {std::make_pair(tmp_s, stage_id)}; -} - -/********** RuleAddCacheWrite **********/ - -SketchGenerationRule::ConditionKind RuleAddCacheWrite::MeetCondition(const SketchPolicyNode& policy, - const State& state, - int stage_id) const { - // Add cache write if a stage needs multi-level tiling, but does not have a element-wise - // matched consumer - if (NeedsMultilevelTiling(policy.search_task, state, stage_id) && - !HasSingleElementwiseMatchedConsumer(policy.search_task, state, stage_id)) { - // An apply and skip rule will be handled in RuleMultiLevelTilingWithFusion - return IsGPUTask(policy.search_task) ? ConditionKind::kApplyAndSkipRest : ConditionKind::kApply; - } - - return ConditionKind::kSkip; -} - -std::vector> RuleAddCacheWrite::Apply(const SketchPolicyNode& policy, - const State& state, - int stage_id) const { - State tmp_s = state; - tmp_s.cache_write(stage_id, "local", policy.search_task->compute_dag); - return {std::make_pair(std::move(tmp_s), stage_id)}; -} - -/********** RuleAddRfactor **********/ - -SketchGenerationRule::ConditionKind RuleAddRfactor::MeetCondition(const SketchPolicyNode& policy, - const State& state, - int stage_id) const { - return (NeedsRfactor(policy.search_task, state, stage_id) && !HasCacheWriteStage(state, stage_id)) - ? ConditionKind::kApply - : ConditionKind::kSkip; -} - -std::vector> RuleAddRfactor::Apply(const SketchPolicyNode& policy, - const State& state, int stage_id) const { - // Fuse all reduction iters - Array space_iters, reduce_iters; - Iterator fused_reduce_iter; - State base_state = - FuseAllReductionIterators(state, stage_id, &fused_reduce_iter, &space_iters, &reduce_iters); - - // TODO(merrymercy): We can do more analysis here to generate less and more efficient sketches. - // In some cases, we only need rfactor for more parallel - // In some cases, we only need rfactor for vectorization. - // Now we will generate two versions and let the search figure out the bette one. - - // Split reduction iters - const auto& split_res = base_state.split(stage_id, fused_reduce_iter, {Integer(1)}); - int factor_axis_id = static_cast(space_iters.size()); - std::vector> ret; - for (const auto& split_iter : split_res) { - State tmp_s = base_state; - int rstage_id = - tmp_s.rfactor(stage_id, split_iter, factor_axis_id, policy.search_task->compute_dag); - - // reorder the space iterator to innermost for vectorization - if (split_iter == split_res[1]) { - Array new_order; - for (size_t i = 0; i < tmp_s->stages[rstage_id]->iters.size(); ++i) { - if (i != space_iters.size()) { - new_order.push_back(tmp_s->stages[rstage_id]->iters[i]); - } - } - new_order.push_back(tmp_s->stages[rstage_id]->iters[space_iters.size()]); - tmp_s.reorder(rstage_id, new_order); - } - - ret.emplace_back(std::move(tmp_s), rstage_id - 1); - } - - return ret; -} - -/********** RuleSimplifyComputeWithConstTensor **********/ - -SketchGenerationRule::ConditionKind RuleSimplifyComputeWithConstTensor::MeetCondition( - const SketchPolicyNode& policy, const State& state, int stage_id) const { - return state->stages[stage_id]->op->attrs.count(SearchPolicyKey::simplify_const_tensor_indices) - ? ConditionKind::kApplyAndSkipRest - : ConditionKind::kSkip; -} - -std::vector> RuleSimplifyComputeWithConstTensor::Apply( - const SketchPolicyNode& policy, const State& state, int stage_id) const { - std::set const_tensor_indices = GetIterNameSetParam( - state->stages[stage_id]->op->attrs, SearchPolicyKey::simplify_const_tensor_indices); - - State tmp_s = state; - Array> tiled_outer_iters; - Array unrolled_inner_iters; - - // Currently set to 2 - size_t tile_level = 2; - - for (const auto& iter : state->stages[stage_id]->iters) { - if (const_tensor_indices.count(iter->name)) { - // unroll indices of const tensors - unrolled_inner_iters.push_back(tmp_s.unroll(stage_id, iter)); - } else { - // tile other space indices - ICHECK(iter->iter_kind == IteratorKind::kSpatial); - tiled_outer_iters.push_back( - tmp_s.split(stage_id, iter, Array>(tile_level - 1, NullOpt))); - } - } - - // reorder them - Array new_order; - for (size_t i = 0; i < tile_level; ++i) { - for (size_t j = 0; j < tiled_outer_iters.size(); ++j) { - new_order.push_back(tiled_outer_iters[j][i]); - } - } - new_order.insert(new_order.end(), unrolled_inner_iters.begin(), unrolled_inner_iters.end()); - tmp_s.reorder(stage_id, new_order); - - return {std::make_pair(tmp_s, stage_id - 1)}; -} - -/********** RuleCrossThreadReduction **********/ - -SketchGenerationRule::ConditionKind RuleCrossThreadReduction::MeetCondition( - const SketchPolicyNode& policy, const State& state, int stage_id) const { - ICHECK(IsGPUTask(policy.search_task)); - - // If it is an intermediate state created by RuleAddCacheWrite, - // we just skip it. - if (HasCacheWriteStage(state, stage_id)) { - return ConditionKind::kSkip; - } - - const auto& op = state->stages[stage_id]->op; - if (op->IsInstance()) { - // Compute the product of lengths of all space iters and all reduce iters - auto [cum_space_len, cum_reduce_len] = - GetCumulativeSpaceAndReductionLength(state->stages[stage_id]); - - if (NeedsMultilevelTiling(policy.search_task, state, stage_id)) { - // Avoid rfactor if we have enough parallelism on space iters - if (cum_space_len > policy.search_task->hardware_params->max_threads_per_block) { - return ConditionKind::kSkip; - } - - return cum_space_len < cum_reduce_len ? ConditionKind::kApply : ConditionKind::kSkip; - } else if (cum_reduce_len > 1) { - // Try rfactor for other reduction operators - return cum_reduce_len > policy.search_task->hardware_params->warp_size ? ConditionKind::kApply - : ConditionKind::kSkip; - } - } - - return ConditionKind::kSkip; -} - -std::vector> RuleCrossThreadReduction::Apply(const SketchPolicyNode& policy, - const State& state, - int stage_id) const { - const SearchTask& task = policy.search_task; - State tmp_s = state; - - // fuse all reduction iters - Array space_iters, reduce_iters; - Iterator fused_reduce_iter; - tmp_s = - FuseAllReductionIterators(tmp_s, stage_id, &fused_reduce_iter, &space_iters, &reduce_iters); - - // Check the opportunity for kernel fusion - bool fusible = false; - int target_stage_id = GetSingleConsumerId(policy.search_task, tmp_s, stage_id); - int num_common_outer = -1; - if (target_stage_id >= 0) { - num_common_outer = - GetNumCommonOuterIterator(policy.search_task, tmp_s, stage_id, target_stage_id); - if (num_common_outer > 0 && - !NeedsMultilevelTiling(policy.search_task, state, target_stage_id)) { - fusible = true; - } - } - - if (fusible) { - const Stage& target_stage = state->stages[target_stage_id]; - std::vector split_step_ids; - - GetSplitStepIds(tmp_s, target_stage_id, &split_step_ids); - - if (split_step_ids.size() == 0) { - // If the target stage does not have split step, - // it must be a simple stage without reduce iters. - // We then should do a split for it. - ICHECK(!HasReduceIter(target_stage)); - const auto& split_res = tmp_s.split(target_stage_id, target_stage->iters.back(), - {Integer(task->hardware_params->warp_size)}); - tmp_s.bind(target_stage_id, split_res[1], IteratorAnnotation::kThreadX); - split_step_ids.push_back(tmp_s->transform_steps.size() - 2); - } - - ICHECK_EQ(split_step_ids.size(), 1); - - const Iterator& target_iter = tmp_s->stages[target_stage_id]->iters[num_common_outer - 1]; - const auto& split_res = tmp_s.follow_split(stage_id, fused_reduce_iter, split_step_ids[0], 1); - tmp_s.bind(stage_id, split_res[1], IteratorAnnotation::kThreadX); - tmp_s.compute_at(stage_id, target_stage_id, target_iter); - } else { - const auto& split_res = - tmp_s.split(stage_id, fused_reduce_iter, {Integer(task->hardware_params->warp_size)}); - tmp_s.bind(stage_id, split_res[1], IteratorAnnotation::kThreadX); - } - - return {std::make_pair(std::move(tmp_s), stage_id - 1)}; -} - -/********** RuleSpecialComputeLocationGPU **********/ - -SketchGenerationRule::ConditionKind RuleSpecialComputeLocationGPU::MeetCondition( - const SketchPolicyNode& policy, const State& state, int stage_id) const { - if (GetProducers(policy.search_task, state, stage_id).empty()) { - return ConditionKind::kSkip; - } - - if (!ShouldAlwaysBeInlined(policy, state, stage_id)) { - return ConditionKind::kSkip; - } - - const std::set& consumers = GetConsumers(policy.search_task, state, stage_id); - if (consumers.size() == 1 && state->stages[*consumers.begin()]->op->attrs.count( - SearchPolicyKey::simplify_const_tensor_indices)) { - return ConditionKind::kApplyAndSkipRest; - } - - return ConditionKind::kSkip; -} - -std::vector> RuleSpecialComputeLocationGPU::Apply( - const SketchPolicyNode& policy, const State& state, int stage_id) const { - State tmp_s = state; - const std::set& consumers = GetConsumers(policy.search_task, state, stage_id); - ICHECK_EQ(consumers.size(), 1); - - // Get the last outer space iterator that is not unrolled. - const Stage& target_stage = state->stages[*consumers.begin()]; - for (size_t i = 0; i < target_stage->iters.size(); ++i) { - if (target_stage->iters[i]->annotation == IteratorAnnotation::kUnroll) { - ICHECK_GT(i, 0); - - tmp_s.compute_at(stage_id, *consumers.begin(), target_stage->iters[i - 1]); - break; - } - } - - return {std::make_pair(std::move(tmp_s), stage_id - 1)}; -} - -/********** RuleCustomSketch **********/ - -SketchGenerationRule::ConditionKind RuleCustomSketch::MeetCondition(const SketchPolicyNode& policy, - const State& state, - int stage_id) const { - auto ret = meet_condition_func_(tvm::runtime::GetRef(&policy), state, stage_id); - if (ret.type_code() == 0) { - return ConditionKind(static_cast(ret)); - } else { - LOG(WARNING) << "Wrong rule condition value. Apply the rule and skip the rest"; - return ConditionKind::kApplyAndSkipRest; - } -} - -std::vector> RuleCustomSketch::Apply(const SketchPolicyNode& policy, - const State& state, int stage_id) const { - Array> apply_ret = - apply_func_(tvm::runtime::GetRef(&policy), state, stage_id); - std::vector> ret; - for (const auto& item : apply_ret) { - CHECK_EQ(item.size(), 2); - auto next = item[1].as(); - ICHECK(next); - ret.emplace_back(Downcast(item[0]), next->value); - } - return ret; -} - -/********** Init Population **********/ - -PopulationGenerationRule::ResultKind InitFillTileSize::Apply(SketchPolicyNode* policy, State* state, - std::mt19937* rand_gen) const { - SplitFactorizationMemo split_memo; - int max_innermost_split_factor = - GetIntParam(policy->params, SketchParamKey::max_innermost_split_factor); - - StateNode* pstate = state->CopyOnWrite(); - // Scan the transformation history and randomly fill tiles size for all SplitStep - for (size_t step_id = 0; step_id < (*state)->transform_steps.size(); ++step_id) { - if (auto ps = (*state)->transform_steps[step_id].as()) { - bool all_defined = true; - for (const auto& len : ps->lengths) { - if (!len) { - all_defined = false; - break; - } - } - if (all_defined) { - continue; - } - - ICHECK(ps->extent); - int extent = GetIntImm(ps->extent.value()); - const auto& candidate_lens = split_memo.GetFactorizationSchemes(extent, ps->lengths.size(), - max_innermost_split_factor); - ICHECK(!candidate_lens.empty()); - const auto& candidate_lengths = candidate_lens[(*rand_gen)() % candidate_lens.size()]; - - pstate->transform_steps.Set( - step_id, - SplitStep(ps->stage_id, ps->iter_id, ps->extent, - Array>(candidate_lengths.begin(), candidate_lengths.end()), - ps->inner_to_outer)); - } - } - pstate->concrete = true; - - return ResultKind::kValid; -} - -PopulationGenerationRule::ResultKind InitChangeComputeLocation::Apply( - SketchPolicyNode* policy, State* state, std::mt19937* rand_gen) const { - if (GetIntParam(policy->params, SketchParamKey::disable_change_compute_location)) { - return ResultKind::kValid; - } - - for (int stage_id = static_cast((*state)->stages.size()) - 1; stage_id >= 0; stage_id--) { - const Stage& stage = (*state)->stages[stage_id]; - // Skip the inlined stages and placeholders - if (stage->op_type == StageKind::kPlaceholder || stage->compute_at == ComputeAtKind::kInlined) { - continue; - } - // Skip the tiled stages - if (IsTiled(stage) || NeedsMultilevelTiling(policy->search_task, *state, stage_id)) { - continue; - } - - std::vector> candidates = - GetComputeLocationCandidates(policy->search_task, *state, stage_id); - - int choice = (*rand_gen)() % (candidates.size() + 2); - - if (choice == 0) { - if (!HasReduceIter(stage)) { - const auto& stage_to_attach_iter = (*state)->attach_map->stage_to_attach_iter; - if (stage_to_attach_iter.find(stage_id) != stage_to_attach_iter.end()) { - state->compute_inline(stage_id); - } - } - } else if (choice == 1) { - state->compute_root(stage_id); - } else { - choice = choice - 2; - const Stage& stage = (*state)->stages[candidates[choice].first]; - state->compute_at(stage_id, candidates[choice].first, - stage->iters[candidates[choice].second]); - } - } - - try { - *state = policy->search_task->compute_dag.InferBound(*state); - } catch (std::exception& e) { - return ResultKind::kInvalid; - } - return ResultKind::kValid; -} - -PopulationGenerationRule::ResultKind InitParallel::Apply(SketchPolicyNode* policy, State* state, - std::mt19937* rand_gen) const { - std::function - annotate_parallel; - annotate_parallel = [&annotate_parallel](const SketchPolicyNode& policy, State* state, - int stage_id, int iter_offset) { - const Stage& stage = (*state)->stages[stage_id]; - - Array to_fuse; - int64_t parallel_degree = 1; - - // Try to fuse and parallel the outermost n iterators - // Stop if we meet reduce iterator or we have enough parallel degree - size_t iter_id = iter_offset; - for (; iter_id < stage->iters.size(); ++iter_id) { - const Iterator& it = stage->iters[iter_id]; - if (it->iter_kind == IteratorKind::kReduction || - it->annotation != IteratorAnnotation::kNone) { - break; - } - to_fuse.push_back(it); - parallel_degree *= GetExtent(it); - - if (parallel_degree > policy.search_task->hardware_params->num_cores * 16) { - break; - } - - if ((*state)->attach_map->iter_to_attached_stages.count(std::make_pair(stage_id, iter_id))) { - break; - } - } - - if (parallel_degree == 1) { - auto res = - (*state)->attach_map->iter_to_attached_stages.find(std::make_pair(stage_id, iter_id)); - if (res != (*state)->attach_map->iter_to_attached_stages.end()) { - for (int attached_stage_id : res->second) { - annotate_parallel(policy, state, attached_stage_id, 0); - } - annotate_parallel(policy, state, stage_id, iter_id + 1); - } - } - - if (!to_fuse.empty()) { - if (to_fuse.size() == 1) { - state->parallel(stage_id, to_fuse[0]); - } else { - Iterator fused_iter = state->fuse(stage_id, to_fuse); - state->parallel(stage_id, fused_iter); - } - } - }; - - for (size_t stage_id = 0; stage_id < (*state)->stages.size(); ++stage_id) { - const Stage& stage = (*state)->stages[stage_id]; - if (stage->compute_at != ComputeAtKind::kRoot || stage->op_type == StageKind::kPlaceholder) { - continue; - } - - annotate_parallel(*policy, state, stage_id, 0); - } - - return ResultKind::kValid; -} - -PopulationGenerationRule::ResultKind InitUnroll::Apply(SketchPolicyNode* policy, State* state, - std::mt19937* rand_gen) const { - std::vector& auto_unroll_configs = - IsGPUTask(policy->search_task) ? auto_unroll_configs_gpu : auto_unroll_configs_cpu; - for (size_t stage_id = 0; stage_id < (*state)->stages.size(); ++stage_id) { - const Stage& stage = (*state)->stages[stage_id]; - // Skip the inlined stage and placeholder stage - if (stage->compute_at == ComputeAtKind::kInlined || stage->op_type == StageKind::kPlaceholder) { - continue; - } - - // Handle always_unroll_inner attr - if (stage->op->attrs.count(SearchPolicyKey::always_unroll_inner)) { - const auto& to_unroll_name_set = - GetIterNameSetParam(stage->op->attrs, SearchPolicyKey::always_unroll_inner); - - // Unroll the space iterators and reduce iterators listed in the attrs in the innermost - // tile - std::set visited_names; - for (int n = static_cast(stage->iters.size()) - 1; n >= 0; n--) { - const Iterator& it = stage->iters[n]; - - // If we meet two iterators that come from a same original iterator, - // then we are out of the innermost tile - size_t size_before = visited_names.size(); - ExtractOriginalIterators(it->name, &visited_names); - if (size_before == visited_names.size()) { - break; - } - - std::set name; - ExtractOriginalIterators(it->name, &name); - if (name.size() == 1 && to_unroll_name_set.count(*name.begin())) { - if (it->annotation == IteratorAnnotation::kNone) { - state->unroll(stage_id, it); - } - } - } - } - - if (HasReduceIter(stage)) { - // Use auto unroll for multi level tiled stage - int value = auto_unroll_configs[(*rand_gen)() % auto_unroll_configs.size()]; - state->pragma(stage_id, (*state)->stages[stage_id]->iters[0], - std::string("auto_unroll_max_step") + "$" + std::to_string(value)); - } - } - - return ResultKind::kValid; -} - -PopulationGenerationRule::ResultKind InitVectorization::Apply(SketchPolicyNode* policy, - State* state, - std::mt19937* rand_gen) const { - for (size_t stage_id = 0; stage_id < (*state)->stages.size(); ++stage_id) { - const Stage& stage = (*state)->stages[stage_id]; - // Skip the inlined stage and placeholder stage - if (stage->compute_at == ComputeAtKind::kInlined || stage->op_type == StageKind::kPlaceholder) { - continue; - } - - // Try to fuse and vectorize the space iterators in the inner most tile - int64_t cum_length_prod = 1; - - int num_fusible = 0; - while (num_fusible < static_cast(stage->iters.size())) { - int iter_id = static_cast(stage->iters.size()) - 1 - num_fusible; - // Stop if this iterator has been a compute at attach point - if ((*state)->attach_map->iter_to_attached_stages.count(std::make_pair(stage_id, iter_id))) { - break; - } - - const Iterator& it = stage->iters[iter_id]; - // Stop if we meet a reduce iterator or annotated iterator - if (it->iter_kind == IteratorKind::kReduction || - it->annotation != IteratorAnnotation::kNone) { - break; - } - - // Stop if the memory access is not continuous (vectorizable) - // Note: The check is too hard, so we use heuristic here - if (IsTiled(stage) && num_fusible != 0) { - // If the stage is tiled, then the memory access must not be continuous - // for the innermost two iterators - break; - } - - cum_length_prod *= GetExtent(it); - if (cum_length_prod > GetIntParam(policy->params, SketchParamKey::max_vectorize_size)) { - break; - } - - num_fusible++; - } - - if (num_fusible > 1) { - // Select a random range to fuse - num_fusible = 1 + (*rand_gen)() % (num_fusible - 1); - } - - if (num_fusible == 1) { - state->vectorize(stage_id, stage->iters.back()); - } else if (num_fusible > 1) { - Array to_fuse(stage->iters.end() + (-num_fusible), stage->iters.end()); - state->vectorize(stage_id, state->fuse(stage_id, to_fuse)); - } - } - - return ResultKind::kValid; -} - -PopulationGenerationRule::ResultKind InitThreadBind::Apply(SketchPolicyNode* policy, State* state, - std::mt19937* rand_gen) const { - // Collect all stages that are roots of stages that perform multi-level tiling. - std::set multi_level_tiling_root_set; - for (size_t stage_id = 0; stage_id < (*state)->stages.size(); ++stage_id) { - if (NeedsMultilevelTiling(policy->search_task, *state, stage_id)) { - const Stage& stage = (*state)->stages[stage_id]; - if (stage->compute_at == ComputeAtKind::kInlined) { - continue; - } else if (stage->compute_at != ComputeAtKind::kIter) { - // This stage is not multi-level tiled, - // so it must be produced by RuleCrossThreadReduction. - ICHECK(HasCrossThreadReduction(*state, stage_id)); - } else { - const auto res = (*state)->attach_map->stage_to_attach_iter.find(stage_id); - ICHECK(res != (*state)->attach_map->stage_to_attach_iter.end()); - multi_level_tiling_root_set.insert(res->second.first); - } - } - } - - *state = policy->search_task->compute_dag.InferBound(*state); - - for (int stage_id = (*state)->stages.size() - 1; stage_id >= 0; --stage_id) { - const Stage& stage = (*state)->stages[stage_id]; - - if (stage->compute_at == ComputeAtKind::kInlined || stage->op_type == StageKind::kPlaceholder) { - continue; - } - - // Deal with the cross-thread reduction generated by RuleCrossThreadReduction - if (HasCrossThreadReduction(*state, stage_id)) { - if (stage->compute_at != ComputeAtKind::kRoot) { - continue; - } - - Iterator fused_it; - *state = std::move(FuseAllOuterSpaceIterators(*state, stage_id, &fused_it)); - state->bind(stage_id, fused_it, IteratorAnnotation::kBlockX); - continue; - } - - // Skip if this stage has already been annotaed with threadIdx.x - if (HasAnnotatedIter(stage, IteratorAnnotation::kThreadX)) { - continue; - } - - if (stage->compute_at == ComputeAtKind::kRoot) { - // This stage has not been tiled, but in GPU schedule, we must tile the root stage - // to do thread binding - if (!multi_level_tiling_root_set.count(stage_id)) { - Iterator fused_it; - *state = FuseAllOuterSpaceIterators(*state, stage_id, &fused_it); - - if (GetExtent(fused_it) <= policy->search_task->hardware_params->warp_size) { - state->bind(stage_id, fused_it, IteratorAnnotation::kThreadX); - } else { - // Set threadIdx.x = default_warp_size by default. - // The later EvolutionarySearch will try more possibility - const auto& split_its = state->split( - stage_id, fused_it, {Integer(policy->search_task->hardware_params->warp_size)}); - state->bind(stage_id, split_its[0], IteratorAnnotation::kBlockX); - state->bind(stage_id, split_its[1], IteratorAnnotation::kThreadX); - } - continue; - } - - // Otherwise, this is a tiled root stage, we assume it should be tiled with 3 space level - // in the outer iterators. - // The remaining part deals with the thread binding for multi-level tiled stages - auto pop = stage->op.as(); - std::vector to_fuse; - int total_space_extent = 1; - for (const auto& i : pop->root_iter_vars()) { - ICHECK(i->dom.defined()); - const auto& pint = i->dom->extent.as(); - ICHECK(pint); - total_space_extent *= pint->value; - } - - bool check_min_thread_extent = true; - // If the total space extent is too small, disable the check of minimal thread extent - if (total_space_extent <= policy->search_task->hardware_params->warp_size * 2) { - check_min_thread_extent = false; - } - - // Fuse the outermost space tile as blockIdx - for (size_t i = 0; i < pop->axis.size(); i++) { - const auto& it = (*state)->stages[stage_id]->iters[i]; - // There may be some iterators that are marked with no split, stop if reaches next - // tiling level - if (!StrEndsWith(it->name, ".0")) { - break; - } - to_fuse.push_back(it); - } - const auto& blockidx_it = state->fuse(stage_id, to_fuse); - state->bind(stage_id, blockidx_it, IteratorAnnotation::kBlockX); - - // Fuse the second outermost space tile as vthread - to_fuse.clear(); - for (size_t i = 1; i < pop->axis.size() + 1; i++) { - const auto& it = (*state)->stages[stage_id]->iters[i]; - // There may be some iterators that are marked with no split, stop if reaches next - // tiling level - if (!StrEndsWith(it->name, ".1")) { - break; - } - to_fuse.push_back((*state)->stages[stage_id]->iters[i]); - } - const auto& vthread_it = state->fuse(stage_id, to_fuse); - if (GetExtent(vthread_it) > policy->search_task->hardware_params->max_vthread_extent) { - return ResultKind::kInvalid; - } - state->bind(stage_id, vthread_it, IteratorAnnotation::kVThread); - - // Fuse the third outermost space tile as threadIdx - to_fuse.clear(); - for (size_t i = 2; i < pop->axis.size() + 2; i++) { - const auto& it = (*state)->stages[stage_id]->iters[i]; - // There may be some iterators that are marked with no split, stop if reaches next - // tiling level - if (!StrEndsWith(it->name, ".2")) { - break; - } - to_fuse.push_back((*state)->stages[stage_id]->iters[i]); - } - const auto& threadidx_it = state->fuse(stage_id, to_fuse); - if (check_min_thread_extent && - GetExtent(threadidx_it) < policy->search_task->hardware_params->warp_size) { - return ResultKind::kInvalid; - } - state->bind(stage_id, threadidx_it, IteratorAnnotation::kThreadX); - } else if (stage->compute_at == ComputeAtKind::kIter && - StrEndsWith(stage->op->name, ".shared")) { - // Do cooperative fetching for the cache read stage. - // Get spatial_split_step_ids from the root stage - const auto& it = (*state)->attach_map->stage_to_attach_iter.find(stage_id); - ICHECK(it != (*state)->attach_map->stage_to_attach_iter.end()); - Array spatial_split_step_ids = GetSpatialSplitStepIds(*state, it->second.first); - - // Fuse all iterators to do cooperative fetching - Iterator fused = state->fuse(stage_id, (*state)->stages[stage_id]->iters); - // Split out an extra iterator for vectorization - // The later EvolutionarySearch will try more possibility - const auto& iters0 = state->split(stage_id, fused, {Integer(1)}); - state->vectorize(stage_id, iters0[1]); - // Follow split to keep a same thread extent with the root stage - const auto& iters1 = - state->follow_fused_split(stage_id, iters0[0], spatial_split_step_ids, 1, true); - state->bind(stage_id, iters1[1], IteratorAnnotation::kThreadX); - } - } - return ResultKind::kValid; -} - -PopulationGenerationRule::ResultKind MutateTileSize::Apply(SketchPolicyNode* policy, State* state, - std::mt19937* rand_gen) const { - int max_innermost_split_factor = - GetIntParam(policy->params, SketchParamKey::max_innermost_split_factor); - - // Extract all SplitStep - std::vector split_step_ids; - for (size_t i = 0; i < (*state)->transform_steps.size(); ++i) { - if (auto ps = (*state)->transform_steps[i].as()) { - if (!ps->extent.defined() || !ps->extent.value()->IsInstance()) { - continue; - } - auto innermost_factor = ps->lengths.back().value_or(max_innermost_split_factor + 1); - if (GetIntImm(innermost_factor) <= max_innermost_split_factor) { - split_step_ids.push_back(i); - } - } - } - if (split_step_ids.empty()) { - // No tile size could be mutated. - return ResultKind::kInvalid; - } - - // Select a SplitStep with extent larger than one to mutate. - int retry_ct = 0; - int64_t extent = 1; - int step_id; - const SplitStepNode* ps; - - do { - step_id = split_step_ids[(*rand_gen)() % split_step_ids.size()]; - ps = (*state)->transform_steps[step_id].as(); - ICHECK(ps != nullptr); - extent = GetIntImm(ps->extent.value()); - retry_ct += 1; - } while (retry_ct < static_cast(split_step_ids.size()) << 2 && (extent == 1 || extent == 0)); - - if (extent <= 1) { - // Cannot find a step with extent larger than one. - return ResultKind::kInvalid; - } - - // Fetch the current tile sizes. - std::vector lengths(ps->lengths.size() + 1, 1); - for (int i = 0; i < static_cast(ps->lengths.size()); ++i) { - lengths[i + 1] = GetIntImm(ps->lengths[i].value()); - } - lengths[0] = extent / ElementProduct(lengths); - - // Random permute the tile size order. - std::vector random_perm; - RandomPermutation(lengths.size(), &random_perm, rand_gen); - - // Try to divide a factor from one tile size and multiple it to another. - for (size_t i = 0; i < random_perm.size(); ++i) { - size_t src_idx = random_perm[i]; - int length = lengths[src_idx]; - if (length <= 1) { - continue; - } - - // Divide one factor from lengths[src_idx] and multiply it to lengths[dst_idx] - size_t dst_idx = random_perm[(i + 1) % random_perm.size()]; - const std::vector& factors = policy->split_memo.GetFactors(length); - ICHECK_GE(factors.size(), 1); - - int divide_factor; - if (dst_idx == lengths.size() - 1) { - // Maintain the restriction of hardware_params.max_innermost_split_factor. - int max_factor_index = static_cast(factors.size()) - 1; - for (; max_factor_index >= 1; max_factor_index--) { - if (factors[max_factor_index] * lengths[dst_idx] <= max_innermost_split_factor) { - break; - } - } - if (max_factor_index == 0) { - // Failed on this dst_idx, try next one. - continue; - } - divide_factor = factors[1 + (*rand_gen)() % (max_factor_index)]; - } else { - divide_factor = factors[1 + (*rand_gen)() % (factors.size() - 1)]; - } - - // Divide one factor from lengths[src_idx] and multiply it to lengths[dst_idx]. - Array new_lengths; - for (size_t j = 1; j < lengths.size(); ++j) { - if (j == src_idx) { - new_lengths.push_back(Integer(lengths[j] / divide_factor)); - } else if (j == dst_idx) { - new_lengths.push_back(Integer(lengths[j] * divide_factor)); - } else { - new_lengths.push_back(Integer(lengths[j])); - } - } - - ICHECK_LE(GetIntImm(new_lengths.back()), max_innermost_split_factor); - - StateNode* pstate = state->CopyOnWrite(); - pstate->transform_steps.Set( - step_id, SplitStep(ps->stage_id, ps->iter_id, ps->extent, - Array>(new_lengths.begin(), new_lengths.end()), - ps->inner_to_outer)); - return ResultKind::kValid; - } - return ResultKind::kInvalid; -} - -PopulationGenerationRule::ResultKind MutateAutoUnroll::Apply(SketchPolicyNode* policy, State* state, - std::mt19937* rand_gen) const { - // Extract all auto_unroll_max_step pragma steps. - std::vector pragma_steps; - for (size_t i = 0; i < (*state)->transform_steps.size(); ++i) { - if (auto ps = (*state)->transform_steps[i].as()) { - if (StrStartsWith(ps->pragma_type, "auto_unroll_max_step")) { - pragma_steps.push_back(i); - } - } - } - if (pragma_steps.empty()) { - return ResultKind::kInvalid; - } - - std::vector& auto_unroll_configs = - IsGPUTask(policy->search_task) ? auto_unroll_configs_gpu : auto_unroll_configs_cpu; - - // Randomly pick up an auto unroll pragma step - auto step_id = pragma_steps[(*rand_gen)() % pragma_steps.size()]; - auto ps = (*state)->transform_steps[step_id].as(); - ICHECK(ps); - - // Mutate its value to a random candidates - int val = auto_unroll_configs[(*rand_gen)() % auto_unroll_configs.size()]; - StateNode* pstate = state->CopyOnWrite(); - pstate->transform_steps.Set( - step_id, PragmaStep(ps->stage_id, ps->iter_id, - std::string("auto_unroll_max_step") + "$" + std::to_string(val))); - Stage new_stage = pstate->stages[ps->stage_id]; - new_stage.CopyOnWrite()->attrs.auto_unroll_max_step = val; - pstate->stages.Set(ps->stage_id, new_stage); - return ResultKind::kValid; -} - -PopulationGenerationRule::ResultKind MutateComputeLocation::Apply(SketchPolicyNode* policy, - State* state, - std::mt19937* rand_gen) const { - if (GetIntParam(policy->params, SketchParamKey::disable_change_compute_location)) { - return ResultKind::kInvalid; - } - - // Extract all compute_at steps. - std::vector compute_at_steps; - for (size_t s = 0; s < (*state)->transform_steps.size(); ++s) { - if (auto ps = (*state)->transform_steps[s].as()) { - int stage_inc = GetTargetStageIDInState(*state, s) - ps->stage_id; - - if (IsTiled((*state)->stages[ps->stage_id + stage_inc])) { - continue; - } - - if (NeedsMultilevelTiling(policy->search_task, *state, ps->stage_id + stage_inc)) { - continue; - } - compute_at_steps.push_back(s); - } - } - if (compute_at_steps.empty()) { - return ResultKind::kInvalid; - } - - // Randomly pick one step - size_t step_id = compute_at_steps[(*rand_gen)() % compute_at_steps.size()]; - auto ps = (*state)->transform_steps[step_id].as(); - int stage_inc = GetTargetStageIDInState(*state, step_id) - ps->stage_id; - ICHECK(ps != nullptr); - - // Randomly pick a new computation location - std::vector> candidates = - GetComputeLocationCandidates(policy->search_task, *state, ps->stage_id + stage_inc); - if (candidates.empty()) { - return ResultKind::kInvalid; - } - int choice = (*rand_gen)() % (candidates.size()); - int new_compute_at_stage_id = candidates[choice].first; - int new_compute_at_iter_id = candidates[choice].second; - - // Replay a new state. - State tmp_s = policy->search_task->compute_dag->init_state; - for (size_t s = 0; s < (*state)->transform_steps.size(); ++s) { - if (s == step_id) { - tmp_s.CopyOnWrite()->transform_steps.push_back( - ComputeAtStep(ps->stage_id, new_compute_at_stage_id - stage_inc, new_compute_at_iter_id)); - } else { - tmp_s.CopyOnWrite()->transform_steps.push_back((*state)->transform_steps[s]); - } - try { - StepApplyToState(tmp_s->transform_steps.back(), &tmp_s, policy->search_task->compute_dag); - } catch (Error& e) { - return ResultKind::kInvalid; - } - } - - *state = tmp_s; - return ResultKind::kValid; -} - -PopulationGenerationRule::ResultKind MutateParallel::Apply(SketchPolicyNode* policy, State* state, - std::mt19937* rand_gen) const { - // This mutation rule only focuses on a case that parallel was added to - // the outermost loop and the loop is generated by fusing other loops. - // In short, we mutate the fusion step before the parallel step. - - // Extract all parallel steps. - std::vector parallel_steps; - for (size_t s = 0; s < (*state)->transform_steps.size(); ++s) { - auto ps = (*state)->transform_steps[s].as(); - if (!ps || ps->annotation != IteratorAnnotation::kParallel) { - continue; - } - - // Skip non-outermost loop or the parallel step without fusion beforehand. - if (ps->iter_id != 0 || s == 0 || !(*state)->transform_steps[s - 1].as()) { - continue; - } - auto fuse_step = (*state)->transform_steps[s - 1].as(); - if (fuse_step->fused_ids[0] != 0) { - continue; - } - - parallel_steps.push_back(s); - } - if (parallel_steps.empty()) { - return ResultKind::kInvalid; - } - - // Randomly pick one parallel step. - size_t step_id = parallel_steps[(*rand_gen)() % parallel_steps.size()]; - - // Replay a new state until the picked fuse step. - State tmp_s = policy->search_task->compute_dag->init_state; - for (size_t s = 0; s < step_id - 1; ++s) { - const auto& step = (*state)->transform_steps[s]; - tmp_s.CopyOnWrite()->transform_steps.push_back(step); - StepApplyToState(step, &tmp_s, policy->search_task->compute_dag); - } - - // Compute all possible fusion granularities - auto fuse_step = (*state)->transform_steps[step_id - 1].as(); - int stage_id = fuse_step->stage_id; - const Stage& stage = tmp_s->stages[stage_id]; - size_t max_fusable_iter_id; - for (max_fusable_iter_id = 0; max_fusable_iter_id < stage->iters.size(); ++max_fusable_iter_id) { - const Iterator& it = stage->iters[max_fusable_iter_id]; - if (it->iter_kind == IteratorKind::kReduction || it->annotation != IteratorAnnotation::kNone) { - break; - } - - if (tmp_s->attach_map->iter_to_attached_stages.count( - std::make_pair(stage_id, max_fusable_iter_id))) { - break; - } - } - - if (max_fusable_iter_id == 0) { - return ResultKind::kInvalid; - } - - // Randomly pick one granularity - int fuse_to_iter_id = (*rand_gen)() % max_fusable_iter_id + 1; - Array fused_ids; - for (int i = 0; i < fuse_to_iter_id; ++i) { - fused_ids.push_back(i); - } - int iter_offset = fuse_step->fused_ids.back()->value - fused_ids.back()->value; - if (iter_offset == 0) { - return ResultKind::kInvalid; - } - - // Replay the mutated fused and annotation step. - auto new_fuse_step = FuseStep(stage_id, fused_ids); - tmp_s.CopyOnWrite()->transform_steps.push_back(new_fuse_step); - StepApplyToState(new_fuse_step, &tmp_s, policy->search_task->compute_dag); - tmp_s.CopyOnWrite()->transform_steps.push_back((*state)->transform_steps[step_id]); - StepApplyToState((*state)->transform_steps[step_id], &tmp_s, policy->search_task->compute_dag); - - // Replay the rest steps. - for (size_t s = step_id + 1; s < (*state)->transform_steps.size(); ++s) { - auto step = (*state)->transform_steps[s]; - if (step->stage_id == stage_id) { - // Since we changed the loop structure, iter ID in later steps to the same stage - // has to be adjusted. - if (auto ps = step.as()) { - if (ps->iter_id == 0) { - step = AnnotationStep(ps->stage_id, 0, ps->annotation); - } else { - ICHECK_LE(ps->iter_id + iter_offset, tmp_s->stages[stage_id]->iters.size()); - step = AnnotationStep(ps->stage_id, ps->iter_id + iter_offset, ps->annotation); - } - } else if (auto ps = step.as()) { - if (ps->iter_id == 0) { - step = PragmaStep(ps->stage_id, 0, ps->pragma_type); - } else { - ICHECK_LE(ps->iter_id + iter_offset, tmp_s->stages[stage_id]->iters.size()); - step = PragmaStep(ps->stage_id, ps->iter_id + iter_offset, ps->pragma_type); - } - } else { - return ResultKind::kInvalid; - } - } - if (IsStageNumberChangingStep(step)) { - // For these steps, we have to update stage_id because these steps will make stage_id - // out-dated. But here we just simply give up this mutation for simplicity. - // This is not an issue because this will never happend in normal cases where all these steps - // are before parallel steps. - return ResultKind::kInvalid; - } - tmp_s.CopyOnWrite()->transform_steps.push_back(step); - try { - StepApplyToState(tmp_s->transform_steps.back(), &tmp_s, policy->search_task->compute_dag); - } catch (Error& e) { - return ResultKind::kInvalid; - } - } - - *state = tmp_s; - return ResultKind::kValid; -} - -} // namespace auto_scheduler -} // namespace tvm diff --git a/src/auto_scheduler/search_policy/sketch_policy_rules.h b/src/auto_scheduler/search_policy/sketch_policy_rules.h deleted file mode 100644 index fc1916b8c67d..000000000000 --- a/src/auto_scheduler/search_policy/sketch_policy_rules.h +++ /dev/null @@ -1,245 +0,0 @@ -/* - * 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 auto_scheduler/search_policy/sketch_policy_rules.h - * \brief Rules for generating the sketches, sampling the initial population, and mutating the - * population in SketchPolicy. - */ - -#ifndef TVM_AUTO_SCHEDULER_SEARCH_POLICY_SKETCH_POLICY_RULES_H_ -#define TVM_AUTO_SCHEDULER_SEARCH_POLICY_SKETCH_POLICY_RULES_H_ - -#include -#include - -#include -#include -#include - -#include "utils.h" - -namespace tvm { -namespace auto_scheduler { - -class SketchPolicyNode; - -/********** Sketch Generation Rule **********/ - -/*! \brief The base class for derivation rules used in the sketch generation. */ -class SketchGenerationRule { - public: - /*! \brief Result enumeration of the condition function. */ - enum class ConditionKind : int { - /*! \brief Skip this rule and continue to try the next rules. */ - kSkip = 0, - /*! \brief Apply this rule and continue to try the next rules. */ - kApply = 1, - /*! \brief Apply this rule and skip the rest rules. */ - kApplyAndSkipRest = 2 - }; - - /*! - * \brief Condition check function of this rule. - * \param policy The SketchPolicyNode of this rule, some information may be used during - * the condition checking. - * \param state The original state to be checked. - * \param stage_id The index of the stage to process this condition check. - * \return The condition check result of this rule. - */ - virtual ConditionKind MeetCondition(const SketchPolicyNode& policy, const State& state, - int stage_id) const = 0; - - /*! - * \brief Apply function of this rule. - * \param policy The SketchPolicyNode of this rule, some information may be used during - * the rule applying. - * \param state The original state to apply this rule. - * \param stage_id The index of the next stage to apply this rule. - * \return The state after applying this rule, and index of the next stage. - */ - virtual std::vector> Apply(const SketchPolicyNode& policy, - const State& state, int stage_id) const = 0; - - /*! - * \brief Get the name of this rule. - * \return A string of the rule name. - */ - virtual std::string GetRuleName() const = 0; -}; - -#define DEFINE_SKETCH_GENERATION_RULE(rule_name) \ - class rule_name : public SketchGenerationRule { \ - public: \ - ConditionKind MeetCondition(const SketchPolicyNode& policy, const State& state, \ - int stage_id) const final; \ - std::vector> Apply(const SketchPolicyNode& policy, const State& state, \ - int stage_id) const final; \ - std::string GetRuleName() const final { return #rule_name; } \ - }; - -/*! \brief The rule that simply skips the current stage. It returns an unchanged state and move to - * the next stage. */ -DEFINE_SKETCH_GENERATION_RULE(RuleSkipStage); - -/*! \brief The rule that inlines simple elementwise ops. - * \note This rule only inlines the strictly inlineable stages. Stages marked as not strictly - * inlineable will have a chance to try different compute at location in InitPopulation later. - */ -DEFINE_SKETCH_GENERATION_RULE(RuleAlwaysInline); - -/*! \brief The rule that performs multi-level tiling. */ -DEFINE_SKETCH_GENERATION_RULE(RuleMultiLevelTiling); - -/*! \brief The rule that performs multi-level tiling and fuses later consumers. */ -DEFINE_SKETCH_GENERATION_RULE(RuleMultiLevelTilingWithFusion); - -/*! \brief The rule that adds a cache read stage. Mainly used for GPU cooperative fetching, - * Currently only support 1 to 1 match cache read. */ -DEFINE_SKETCH_GENERATION_RULE(RuleAddCacheRead); - -/*! \brief The rule that adds a cache write stage. */ -DEFINE_SKETCH_GENERATION_RULE(RuleAddCacheWrite); - -/*! \brief The rule that adds rfactor stage. */ -DEFINE_SKETCH_GENERATION_RULE(RuleAddRfactor); - -/*! \brief The rule that deals with compute ops that perform "fake reduction" with const tensors. - * This kind of op comes from winograd transformation. */ -DEFINE_SKETCH_GENERATION_RULE(RuleSimplifyComputeWithConstTensor); - -/*! \brief The rule that use cross thread reduction for GPU. */ -DEFINE_SKETCH_GENERATION_RULE(RuleCrossThreadReduction); - -/*! \brief Handle special cases in Winograd transformation for GPU. We need to change the compute - * location of the producers of compute ops that perform "fake reduction" with const tensors. */ -DEFINE_SKETCH_GENERATION_RULE(RuleSpecialComputeLocationGPU); - -/*! \brief The rule that allows users to generate custom sketches. */ -class RuleCustomSketch : public SketchGenerationRule { - public: - RuleCustomSketch(PackedFunc meet_condition_func, PackedFunc apply_func, - String rule_name = "CustomSketchRule") - : meet_condition_func_(std::move(meet_condition_func)), - apply_func_(std::move(apply_func)), - rule_name_(std::move(rule_name)) {} - - ConditionKind MeetCondition(const SketchPolicyNode& policy, const State& state, - int stage_id) const final; - - std::vector> Apply(const SketchPolicyNode& policy, const State& state, - int stage_id) const final; - - std::string GetRuleName() const final { return rule_name_; } - - private: - PackedFunc meet_condition_func_; - PackedFunc apply_func_; - String rule_name_; -}; - -/********** Init Population **********/ - -/*! \brief The base class for rules used to annotate the sketches to get the initial population. */ -class PopulationGenerationRule { - public: - /*! \brief Result enumeration of the apply function. */ - enum class ResultKind : int { kValid = 0, kInvalid = 1 }; - - /*! - * \brief Apply function of this rule. - * \param policy The SketchPolicyNode of this rule, some member may get changed during the - * rule applying. (e.g. random number generator) - * \param state The state to apply this rule, update inplace. - * \return The result of this rule, indicate if there's any valid state generated. - */ - virtual ResultKind Apply(SketchPolicyNode* policy, State* state, - std::mt19937* rand_gen) const = 0; - - /*! \brief The deconstructor */ - virtual ~PopulationGenerationRule() = default; -}; - -// A helper to define population initialization rules -#define DEFINE_INIT_POPULATION_RULE(rule_name) \ - class rule_name : public PopulationGenerationRule { \ - public: \ - ResultKind Apply(SketchPolicyNode* policy, State* state, std::mt19937* rand_gen) const final; \ - }; - -/*! \brief The rule that fills the incomplete SplitSteps. */ -DEFINE_INIT_POPULATION_RULE(InitFillTileSize); - -/*! \brief The rule that randomly changes the computation location for some stages that do not - * need tiling and are not strictly inlineable(e.g. data padding). */ -DEFINE_INIT_POPULATION_RULE(InitChangeComputeLocation); - -/*! \brief The rule that annotates parallel for CPU. */ -DEFINE_INIT_POPULATION_RULE(InitParallel); - -/*! \brief The rule that annotates unroll. */ -DEFINE_INIT_POPULATION_RULE(InitUnroll); - -/*! \brief The rule that annotates vectorization. */ -DEFINE_INIT_POPULATION_RULE(InitVectorization); - -/*! \brief The rule that annotates thread binding for GPU. */ -DEFINE_INIT_POPULATION_RULE(InitThreadBind); - -/********** Mutation **********/ - -/*! \brief The base class for mutation rules used in the evolutionary search. */ -class PopulationMutationRule : public PopulationGenerationRule { - public: - /* \brief The constructor - * \param selection_weight the probabiliy of applying this rule is - * proportional to this weight - */ - explicit PopulationMutationRule(double selection_weight) : weight(selection_weight) {} - - /* \brief The weight of this rule */ - double weight; -}; - -// A helper to define mutation rules used in the evolutionary search -#define DEFINE_MUTATE_POPULATION_RULE(rule_name) \ - class rule_name : public PopulationMutationRule { \ - public: \ - explicit rule_name(double weight) : PopulationMutationRule(weight) {} \ - ResultKind Apply(SketchPolicyNode* policy, State* state, std::mt19937* rand_gen) const final; \ - }; - -/*! \brief The rule that mutates tile size by randomly dividing a tile size by a factor - and multipling it to another tile size. */ -DEFINE_MUTATE_POPULATION_RULE(MutateTileSize); - -/*! \brief The rule that mutates the number of fused outer iterators annotated by parallel. */ -DEFINE_MUTATE_POPULATION_RULE(MutateParallel); - -/*! \brief The rule that randomly changes the computation location for some stages that do not - * need tiling and are not strictly inlineable(e.g. data padding). */ -DEFINE_MUTATE_POPULATION_RULE(MutateComputeLocation); - -/*! \brief The rule that mutates the value of a randomly selected auto unroll pragma step. */ -DEFINE_MUTATE_POPULATION_RULE(MutateAutoUnroll); - -} // namespace auto_scheduler -} // namespace tvm - -#endif // TVM_AUTO_SCHEDULER_SEARCH_POLICY_SKETCH_POLICY_RULES_H_ diff --git a/src/auto_scheduler/search_policy/utils.cc b/src/auto_scheduler/search_policy/utils.cc deleted file mode 100644 index ac1cf2dd82c9..000000000000 --- a/src/auto_scheduler/search_policy/utils.cc +++ /dev/null @@ -1,485 +0,0 @@ -/* - * 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 auto_scheduler/search_policy/utils.cc - * \brief Common utilities - */ - -#include "utils.h" - -#include - -namespace tvm { -namespace auto_scheduler { - -Array GetSpatialSplitStepIds(const State& s, int stage_id) { - const auto& stage = s->stages[stage_id]; - const auto& pop = s->stages[stage_id]->op.as(); - ICHECK(pop != nullptr); - const std::set& no_split_at_inner_name_set = - stage->op->attrs.count(SearchPolicyKey::no_split_at_inner) - ? GetIterNameSetParam(stage->op->attrs, SearchPolicyKey::no_split_at_inner) - : std::set(); - size_t reduce_count = 0; - for (const auto axis : pop->reduce_axis) { - if (!no_split_at_inner_name_set.count(axis->var->name_hint)) { - reduce_count++; - } - } - - Array spatial_split_step_ids; - for (int i = s->transform_steps.size() - 1; i >= 0; --i) { - if (IsStageNumberChangingStep(s->transform_steps[i])) { - if (stage_id > s->transform_steps[i]->stage_id) { - stage_id--; - } - } else if (auto ps = s->transform_steps[i].as()) { - if (stage_id == ps->stage_id) { - // Assume SplitStep on reduction axes are always after SplitStep on spatial axes. - if (reduce_count) { - reduce_count--; - } else { - spatial_split_step_ids.push_back(i); - } - } - } - } - - return spatial_split_step_ids; -} - -std::vector> GetComputeLocationCandidates(const SearchTask& task, - const State& state, int stage_id) { - int target_stage_id = GetSingleConsumerId(task, state, stage_id); - if (target_stage_id < 0) { - return {}; - } - const Stage& target_stage = state->stages[target_stage_id]; - - std::vector> candidates; - bool target_compute_at_other = target_stage->compute_at == ComputeAtKind::kIter; - bool target_is_tiled = IsTiled(target_stage); - - bool visited_reduce = false; - // Enumerate compute_at location at target_stage - // TODO(merrymercy): More analysis here to make smarter choices - for (size_t i = 0; i < target_stage->iters.size(); ++i) { - const Iterator& target_iter = target_stage->iters[i]; - if (target_iter->iter_kind == IteratorKind::kReduction) { - visited_reduce = true; - if (!target_is_tiled) { // Do not go into reduce iter - break; - } - } else if (target_iter->iter_kind == IteratorKind::kSpatial) { - if (visited_reduce) { // Do not go into inner tile - break; - } - } - - if (target_iter->annotation == IteratorAnnotation::kUnroll) { - // Do not go into the unroll region of const tensor indices - break; - } - - if (GetExtent(target_iter) == 1) { - // Skip iterators with length of 1 - continue; - } - if (target_compute_at_other && target_iter->iter_kind == IteratorKind::kSpatial && - StrEndsWith(target_iter->name, ".0")) { - // Skip the first level iterators if target stage compute_at another stage - // In this case, the lengths of first level iterators are always one - continue; - } - candidates.emplace_back(target_stage_id, i); - - if (state->attach_map->iter_to_attached_stages.count(std::make_pair(target_stage_id, i))) { - break; - } - } - - // if the target_stage is already compute_at another stage X, try also compute_at X - // We call stage X as `target_target_stage` - if (target_compute_at_other) { - int target_target_stage_id; - target_target_stage_id = state->attach_map->stage_to_attach_iter.at(target_stage_id).first; - const Stage& target_target_stage = state->stages[target_target_stage_id]; - - for (size_t i = 0; i < target_target_stage->iters.size(); ++i) { - const Iterator& target_target_iter = target_target_stage->iters[i]; - if (target_target_iter->iter_kind == IteratorKind::kReduction || - state->attach_map->iter_to_attached_stages.count( - std::make_pair(target_target_stage_id, i))) { - break; - } - - if (target_target_iter->annotation == IteratorAnnotation::kUnroll) { - // Do not go into the unroll region of const tensor indices - break; - } - - if (GetExtent(target_target_iter) == 1) { // skip iterators with length of 1 - continue; - } - - candidates.emplace_back(target_target_stage_id, i); - } - } - - return candidates; -} - -State DoMultiLevelTiling(const State& state, int stage_id, const std::string& format, - std::vector* spatial_split_step_ids) { - // Temporal object to be used if the input pointer is nullptr - std::vector temp_split_step_ids; - if (spatial_split_step_ids == nullptr) { - spatial_split_step_ids = &temp_split_step_ids; - } - spatial_split_step_ids->clear(); - - std::vector> space_levels; - std::vector> reduce_levels; - std::vector space_outer, space_inner, reduce_outer, reduce_inner; - - size_t n_space = - std::count(format.begin(), format.end(), 's') + std::count(format.begin(), format.end(), 'S'); - size_t n_reduce = - std::count(format.begin(), format.end(), 'r') + std::count(format.begin(), format.end(), 'R'); - if (n_space + n_reduce != format.size()) { - LOG(FATAL) << "Invalid multi-level tiling format: " << format; - } - space_levels.resize(n_space); - reduce_levels.resize(n_reduce); - - State tmp_s = state; - const Stage& stage = state->stages[stage_id]; - const std::set& no_split_at_inner_name_set = - stage->op->attrs.count(SearchPolicyKey::no_split_at_inner) - ? GetIterNameSetParam(stage->op->attrs, SearchPolicyKey::no_split_at_inner) - : std::set(); - - auto sr_levels = [&](int size, const Iterator& iter, std::vector>& levels) { - ICHECK_GE(size, 1); - if (size == 1) { - levels[0].push_back(iter); - } else { - Array split_res = - tmp_s.split(stage_id, iter, Array>(size - 1, NullOpt)); - for (int i = 0; i < size; i++) { - levels[i].push_back(split_res[i]); - } - if (iter->iter_kind == IteratorKind::kSpatial) { - spatial_split_step_ids->push_back(tmp_s->transform_steps.size() - 1); - } - } - }; - - for (const auto& iter : state->stages[stage_id]->iters) { - if (!no_split_at_inner_name_set.count(iter->name)) { - if (iter->iter_kind == IteratorKind::kSpatial) { - sr_levels(n_space, iter, space_levels); - } else if (iter->iter_kind == IteratorKind::kReduction) { - sr_levels(n_reduce, iter, reduce_levels); - } else { - LOG(FATAL) << "Invalid iter type: " << int(iter->iter_kind); - } - } else { - if (iter->iter_kind == IteratorKind::kSpatial) { - space_inner.push_back(iter); - } else if (iter->iter_kind == IteratorKind::kReduction) { - reduce_inner.push_back(iter); - } else { - LOG(FATAL) << "Invalid iter type: " << int(iter->iter_kind); - } - } - } - - auto fill_levels = [&](std::vector& levels_iter, std::vector& fill) { - if (!fill.empty()) { - levels_iter.insert(levels_iter.begin(), std::make_move_iterator(fill.begin()), - std::make_move_iterator(fill.end())); - } - }; - if (!space_levels.empty()) { - fill_levels(space_levels.front(), space_outer); - fill_levels(space_levels.back(), space_inner); - } - if (!reduce_levels.empty()) { - fill_levels(reduce_levels.front(), reduce_outer); - fill_levels(reduce_levels.back(), reduce_inner); - } - - Array order; - int space_ct = 0, reduce_ct = 0; - for (const auto c : format) { - if (c == 's' || c == 'S') { - order.insert(order.end(), std::make_move_iterator(space_levels[space_ct].begin()), - std::make_move_iterator(space_levels[space_ct].end())); - space_ct++; - } else if (c == 'r' || c == 'R') { - order.insert(order.end(), std::make_move_iterator(reduce_levels[reduce_ct].begin()), - std::make_move_iterator(reduce_levels[reduce_ct].end())); - reduce_ct++; - } else { - LOG(FATAL) << "Invalid multi level tiling format: " << format; - } - } - - tmp_s.reorder(stage_id, order); - return tmp_s; -} - -State FollowTiling(const State& state, int stage_id, const std::vector& split_step_ids, - int n_split) { - if (n_split < 1 || n_split > 3) { - LOG(FATAL) << "Invalid split parts, currently only support 1, 2 and 3"; - } - // Apply up to three-level tiling structure: space_L0, space_L1, space_L2 - std::vector space_0, space_1, space_2, space_3, tmp_order; - Array split_res; - - auto pop = state->stages[stage_id]->op.as(); - ICHECK(pop != nullptr); - const Stage& stage = state->stages[stage_id]; - const std::set& no_split_at_inner_name_set = - stage->op->attrs.count(SearchPolicyKey::no_split_at_inner) - ? GetIterNameSetParam(stage->op->attrs, SearchPolicyKey::no_split_at_inner) - : std::set(); - int no_split_at_inner_name_in_stage_cnt = 0; - for (const auto& iter : state->stages[stage_id]->iters) { - no_split_at_inner_name_in_stage_cnt += no_split_at_inner_name_set.count(iter->name); - } - - ICHECK_EQ(state->stages[stage_id]->iters.size() - no_split_at_inner_name_in_stage_cnt, - split_step_ids.size()); - - State tmp_s = state; - int ct = 0; - for (const auto& iter : state->stages[stage_id]->iters) { - if (iter->iter_kind == IteratorKind::kSpatial) { - // For spatial iterator, split it into multi iterators - if (!no_split_at_inner_name_set.count(iter->name)) { - IteratorAnnotation ann_type = iter->annotation; - split_res = tmp_s.follow_split(stage_id, iter, split_step_ids[ct], n_split); - // Restore annotation. Move unroll and vectorize to inner, move parallel - // to outer - switch (ann_type) { - case IteratorAnnotation::kUnroll: - split_res.Set(n_split, tmp_s.unroll(stage_id, split_res[n_split])); - break; - case IteratorAnnotation::kVectorize: - split_res.Set(n_split, tmp_s.vectorize(stage_id, split_res[n_split])); - break; - case IteratorAnnotation::kParallel: - split_res.Set(0, tmp_s.parallel(stage_id, split_res[0])); - break; - default: - break; - } - - space_0.push_back(split_res[0]); - space_1.push_back(split_res[1]); - if (n_split >= 2) { - space_2.push_back(split_res[2]); - if (n_split == 3) { - space_3.push_back(split_res[3]); - } - } - ct++; - } else { - if (no_split_at_inner_name_set.count(iter->name)) { - if (n_split == 1) { - space_1.push_back(iter); - } else if (n_split == 2) { - space_2.push_back(iter); - } else { - ICHECK_EQ(n_split, 3); - space_3.push_back(iter); - } - } - } - } else { - LOG(FATAL) << "Invalid iter type: " << int(iter->iter_kind); - } - } - - if (n_split == 3) { - ConcatenateMove(&tmp_order, &space_0, &space_1, &space_2, &space_3); - } else if (n_split == 2) { - ConcatenateMove(&tmp_order, &space_0, &space_1, &space_2); - } else { - ConcatenateMove(&tmp_order, &space_0, &space_1); - } - tmp_s.reorder(stage_id, tmp_order); - return tmp_s; -} - -// Return whether a state has nested parallel, which is invalid on CPUs -bool HasNestedParallel(const State& state) { - std::function count_parallel_ct; - - count_parallel_ct = [&state, &count_parallel_ct](int stage_id, size_t* parallel_ct) { - const Stage& stage = state->stages[stage_id]; - - if (stage->compute_at == ComputeAtKind::kInlined) { - return; - } - - for (size_t i = 0; i < stage->iters.size(); ++i) { - if (stage->iters[i]->annotation == IteratorAnnotation::kParallel) { - (*parallel_ct)++; - } - - IterKey iter_key(stage_id, i); - auto pair = state->attach_map->iter_to_attached_stages.find(iter_key); - if (pair != state->attach_map->iter_to_attached_stages.end()) { - for (const auto& attach_stage_id : pair->second) { - count_parallel_ct(attach_stage_id, parallel_ct); - } - } - } - }; - - for (size_t stage_id = 0; stage_id < state->stages.size(); ++stage_id) { - size_t parallel_ct = 0; - - if (state->stages[stage_id]->compute_at == ComputeAtKind::kRoot) { - count_parallel_ct(stage_id, ¶llel_ct); - if (parallel_ct >= 2) { - return true; - } - } - } - - return false; -} - -void PruneInvalidState(const SearchTask& task, Array* states) { - size_t pt = 0; - for (size_t i = 0; i < states->size(); ++i) { - if (!(*states)[i].defined()) { - continue; - } - if (!IsGPUTask(task) && HasNestedParallel((*states)[i])) { - continue; - } - - if (i != pt) { - states->Set(pt, (*states)[i]); - } - pt++; - } - - if (pt == 0) { - LOG(FATAL) << "Internal error: All states are invalid."; - } else { - states->resize(pt); - } -} - -/********** SplitFactorizationMemo **********/ -const Array>& SplitFactorizationMemo::GetFactorizationSchemes( - int extent, int n_lengths, int max_innermost_factor) { - QueryKey key = std::make_tuple(extent, n_lengths, max_innermost_factor); - const auto& it = memory_.find(key); - if (it != memory_.end()) { - return it->second; - } - - tmp_stack_ = Array(n_lengths, Integer()); - results_ = &memory_[key]; - n_lengths_ = n_lengths; - - DfsEnumerate(0, extent, max_innermost_factor); - - return *results_; -} - -void SplitFactorizationMemo::DfsEnumerate(int now, int remaining_length, int max_innermost_factor) { - if (now == n_lengths_) { - if (tmp_stack_.back().as()->value <= max_innermost_factor) { - results_->push_back(tmp_stack_); - } - } else { - for (const auto& f : GetFactors(remaining_length)) { - tmp_stack_.Set(now, Integer(f)); - DfsEnumerate(now + 1, remaining_length / f, max_innermost_factor); - } - } -} - -const std::vector& SplitFactorizationMemo::GetFactors(int n) { - auto it = factor_memory_.find(n); - if (it != factor_memory_.end()) { - return it->second; - } - - std::vector& res = factor_memory_[n]; - int step = n % 2 == 0 ? 1 : 2; - for (size_t i = 1; i < static_cast(std::sqrt(n)) + 1; i += step) { - if (n % i == 0) { - res.push_back(i); - if (n / i != i) { - res.push_back(n / i); - } - } - } - std::sort(res.begin(), res.end()); - return res; -} - -/********** Utils interface API for ffi **********/ - -TVM_REGISTER_GLOBAL("auto_scheduler.SearchPolicyUtilsGetConsumers") - .set_body_typed([](const SearchTask& task, const State& state, int stage_id) { - const std::set& consumers = GetConsumers(task, state, stage_id); - tvm::Map ret; - for (const auto& i : consumers) { - ret.Set(Integer(i), Integer(i)); - } - return ret; - }); - -TVM_REGISTER_GLOBAL("auto_scheduler.SearchPolicyUtilsIsElementwiseMatch") - .set_body_typed([](const SearchTask& task, const State& state, int stage_id, - int target_stage_id) { - return ElementwiseMatch(task, state, stage_id, target_stage_id); - }); - -TVM_REGISTER_GLOBAL("auto_scheduler.SearchPolicyUtilsIsTiled") - .set_body_typed([](const Stage& stage) { return IsTiled(stage); }); - -TVM_REGISTER_GLOBAL("auto_scheduler.SearchPolicyUtilsHasCacheReadStage") - .set_body_typed([](const State& s, int stage_id) { return HasCacheReadStage(s, stage_id); }); - -TVM_REGISTER_GLOBAL("auto_scheduler.SearchPolicyUtilsHasCacheWriteStage") - .set_body_typed([](const State& s, int stage_id) { return HasCacheWriteStage(s, stage_id); }); - -TVM_REGISTER_GLOBAL("auto_scheduler.SearchPolicyUtilsHasRfactorStage") - .set_body_typed([](const State& s, int stage_id) { return HasRfactorStage(s, stage_id); }); - -TVM_REGISTER_GLOBAL("auto_scheduler.SearchPolicyUtilsHasCrossThreadReduction") - .set_body_typed([](const State& s, int stage_id) { - return HasCrossThreadReduction(s, stage_id); - }); - -} // namespace auto_scheduler -} // namespace tvm diff --git a/src/auto_scheduler/search_policy/utils.h b/src/auto_scheduler/search_policy/utils.h deleted file mode 100644 index cc6b0ab23756..000000000000 --- a/src/auto_scheduler/search_policy/utils.h +++ /dev/null @@ -1,716 +0,0 @@ -/* - * 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 auto_scheduler/search_policy/utils.h - * \brief Common utilities for search policies. - */ - -#ifndef TVM_AUTO_SCHEDULER_SEARCH_POLICY_UTILS_H_ -#define TVM_AUTO_SCHEDULER_SEARCH_POLICY_UTILS_H_ - -#include -#include -#include -#include -#include - -#include -#include -#include -#include -#include -#include -#include -#include -#include - -#include "../utils.h" - -namespace tvm { -namespace auto_scheduler { - -/*! \brief Return whether the search task is targeting a CPU. */ -inline bool IsCPUTask(const SearchTask& task) { - return (task)->target->GetTargetDeviceType() == kDLCPU; -} - -/*! \brief Return whether the search task is targeting a GPU. */ -inline bool IsGPUTask(const SearchTask& task) { - int device_type = (task)->target->GetTargetDeviceType(); - return device_type == kDLCUDA || device_type == kDLOpenCL || device_type == kDLVulkan || - device_type == kDLMetal || device_type == kDLROCM || device_type == kOpenGL; -} - -/*! \brief Return whether the search task is targeting a Hexagon. */ -inline bool IsHexagonTask(const SearchTask& task) { - return (task)->target->GetTargetDeviceType() == kDLHexagon; -} - -/*! \brief Return whether the search task is targeting a CUDA GPU. */ -inline bool IsCUDATask(const SearchTask& task) { - return (task)->target->GetTargetDeviceType() == kDLCUDA; -} - -/*! \brief Return whether the search task is targeting a OpenCL GPU. */ -inline bool IsOpenCLTask(const SearchTask& task) { - return (task)->target->GetTargetDeviceType() == kDLOpenCL; -} - -/*! \brief Argsort. Order: largest to smallest */ -template -inline std::vector Argsort(const std::vector& scores) { - std::vector index; - index.reserve(scores.size()); - for (size_t i = 0; i < scores.size(); ++i) { - index.push_back(i); - } - auto cmp = [&scores](int l, int r) { return scores[l] > scores[r]; }; - std::sort(index.begin(), index.end(), cmp); - return index; -} - -/*! \brief Convert operation to stage id. */ -inline int OperationToStage(const te::Operation& op, const State& state) { - for (size_t i = 0; i < state->stages.size(); ++i) { - if (op == state->stages[i]->op) { - return i; - } - } - LOG(FATAL) << "Cannot find op: " << op; -} - -/********** Get Parameters **********/ - -/*! \brief Get an integer from a tvm str Map. */ -inline int GetIntParam(const Map& attr_dict, const std::string& key) { - ICHECK_GT(attr_dict.count(key), 0) << "Cannot find key: \"" << key << "\" in " << attr_dict; - auto pint = attr_dict[key].as(); - ICHECK(pint != nullptr); - return pint->value; -} - -/*! \brief Get a double from a tvm str Map. */ -inline double GetDoubleParam(const Map& attr_dict, const std::string& key) { - ICHECK_GT(attr_dict.count(key), 0) << "Cannot find key: \"" << key << "\" in " << attr_dict; - auto pdouble = attr_dict[key].as(); - ICHECK(pdouble != nullptr); - return pdouble->value; -} - -/*! \brief Get a string from a tvm str Map. */ -inline std::string GetStringParam(const Map& attr_dict, const std::string& key) { - ICHECK_GT(attr_dict.count(key), 0) << "Cannot find key: \"" << key << "\" in " << attr_dict; - const auto& target = attr_dict[key]; - if (auto pstr = target.as()) { - return pstr->value; - } else if (auto pstr = target.as()) { - return pstr->data; - } else { - LOG(FATAL) << "Could not convert object " << target << " of type " << target->GetTypeKey() - << " to string"; - } -} - -/*! \brief Get a iterator name set from a tvm str Map. */ -inline std::set GetIterNameSetParam(const Map& attr_dict, - const std::string& key) { - std::set ret; - ICHECK_GT(attr_dict.count(key), 0) << "Cannot find key: \"" << key << "\" in " << attr_dict; - auto names = attr_dict[key].as(); - ICHECK(names != nullptr); - for (const auto& name : *names) { - ret.insert(name.as()->data); - } - return ret; -} - -/********** Checks with ComputeDAG **********/ - -/*! \brief Return whether an op is strictly-inlineable. */ -inline bool IsStrictlyInlineable(const SearchTask& task, const State& state, int stage_id) { - if (state->current_compute_dag) { - return state->current_compute_dag.as()->access_analyzer.IsStrictlyInlineable( - state->stages[stage_id]->op); - } else { - return task->compute_dag->access_analyzer.IsStrictlyInlineable(state->stages[stage_id]->op); - } -} - -/*! \brief Return whether an op is an output op. */ -inline bool IsOutputOp(const SearchTask& task, const State& state, int stage_id) { - if (state->current_compute_dag) { - return state->current_compute_dag.as()->access_analyzer.IsOutput( - state->stages[stage_id]->op); - } else { - return task->compute_dag->access_analyzer.IsOutput(state->stages[stage_id]->op); - } -} - -/*! \brief Return whether an op needs multi level tiling. */ -inline bool NeedsMultilevelTiling(const SearchTask& task, const State& state, int stage_id) { - if (state->current_compute_dag) { - return state->current_compute_dag.as()->access_analyzer.NeedsMultiLevelTiling( - state->stages[stage_id]->op); - } else { - return task->compute_dag->access_analyzer.NeedsMultiLevelTiling(state->stages[stage_id]->op); - } -} - -/*! \brief Get all consumers for a stage. This function propagates the relation for inlined ops. */ -inline std::set GetConsumers(const SearchTask& task, const State& state, int stage_id) { - std::unordered_set consumers; - std::set ret; - - if (state->current_compute_dag) { - consumers = state->current_compute_dag.as()->access_analyzer.GetConsumers( - state, state->stages[stage_id]->op); - } else { - consumers = task->compute_dag->access_analyzer.GetConsumers(state, state->stages[stage_id]->op); - } - - for (const auto& op : consumers) { - ret.insert(OperationToStage(op, state)); - } - return ret; -} - -/*! \brief Check if a stage has single consumer or all of its consumers share a common root, return - * the target consumer root or -1. */ -inline int GetSingleConsumerId(const SearchTask& task, const State& state, int stage_id) { - const std::set& consumers = GetConsumers(task, state, stage_id); - if (consumers.empty()) { - return -1; - } - - if (consumers.size() == 1) { - return *consumers.begin(); - } else { - // Check all consumers share a common root - int common_root_id = -1; - bool mismatch = false; - for (const auto& consumer_stage_id : consumers) { - int root_id = -1; - if (state->stages[consumer_stage_id]->compute_at == ComputeAtKind::kRoot) { - root_id = consumer_stage_id; - } else if (state->stages[consumer_stage_id]->compute_at == ComputeAtKind::kIter) { - root_id = state->attach_map->stage_to_attach_iter.at(consumer_stage_id).first; - } else { - LOG(FATAL) << "Invalid case"; - } - - if (common_root_id == -1) { - common_root_id = root_id; - } else { - if (common_root_id != root_id) { - mismatch = true; - break; - } - } - } - - return mismatch ? -1 : common_root_id; - } -} - -/*! \brief Get all producers for a stage. This function propagates the relation for inlined ops. */ -inline std::set GetProducers(const SearchTask& task, const State& state, int stage_id) { - std::unordered_set producers; - std::set ret; - - if (state->current_compute_dag) { - producers = state->current_compute_dag.as()->access_analyzer.GetProducers( - state, state->stages[stage_id]->op); - } else { - producers = task->compute_dag->access_analyzer.GetProducers(state, state->stages[stage_id]->op); - } - - for (const auto& op : producers) { - ret.insert(OperationToStage(op, state)); - } - return ret; -} - -/*! \brief Get all producers for a stage. This function DOES NOT propagates the relation for - * inlined ops. */ -inline std::set GetDirectProducers(const SearchTask& task, const State& state, int stage_id) { - std::unordered_set producers; - std::set ret; - - if (state->current_compute_dag) { - producers = state->current_compute_dag.as()->access_analyzer.GetDirectProducers( - state->stages[stage_id]->op); - } else { - producers = task->compute_dag->access_analyzer.GetDirectProducers(state->stages[stage_id]->op); - } - - for (const auto& op : producers) { - ret.insert(OperationToStage(op, state)); - } - return ret; -} - -/*! \brief Get the number of common outer iterators. This function propagates the relation for - * chains with multiple ops. */ -inline int GetNumCommonOuterIterator(const SearchTask& task, const State& state, int stage_id, - int target_stage_id) { - if (state->current_compute_dag) { - return state->current_compute_dag.as() - ->access_analyzer.GetNumCommonOuterIterator(state->stages[stage_id]->op, - state->stages[target_stage_id]->op); - } else { - return task->compute_dag->access_analyzer.GetNumCommonOuterIterator( - state->stages[stage_id]->op, state->stages[target_stage_id]->op); - } -} - -/*! \brief Return whether two ops are elementwise-matched. */ -inline bool ElementwiseMatch(const SearchTask& task, const State& state, int stage_id, - int target_stage_id) { - const auto& op = state->stages[stage_id]->op; - const auto& target_op = state->stages[target_stage_id]->op; - if (state->current_compute_dag) { - return state->current_compute_dag.as()->access_analyzer.ElementWiseMatch( - op, target_op); - } else { - return task->compute_dag->access_analyzer.ElementWiseMatch(op, target_op); - } -} - -/********** Get informations from Stage/Iterator **********/ - -/*! \brief Return the extent of an iterator. */ -inline int64_t GetExtent(const Iterator& it) { - if (it->range.defined()) { - if (auto pint = it->range->extent.as()) { - return pint->value; - } - } - return -1; -} - -/*! \brief Compute the product of lengths of all space iters and all reduce iters, respectively. */ -inline std::pair GetCumulativeSpaceAndReductionLength(const Stage& stage) { - int64_t cum_space_len = 1, cum_reduce_len = 1; - for (const auto& iter : stage->iters) { - if (iter->iter_kind == IteratorKind::kSpatial) { - cum_space_len *= GetExtent(iter); - } else if (iter->iter_kind == IteratorKind::kReduction) { - cum_reduce_len *= GetExtent(iter); - } - } - return std::make_pair(cum_space_len, cum_reduce_len); -} - -/*! \brief Return whether this stage needs rfactor. */ -inline bool NeedsRfactor(const SearchTask& task, const State& state, int stage_id) { - const auto& op = state->stages[stage_id]->op; - if (op->IsInstance()) { - // Compute the product of lengths of all space iters and all reduce iters - int cum_space_len, cum_reduce_len; - std::tie(cum_space_len, cum_reduce_len) = - GetCumulativeSpaceAndReductionLength(state->stages[stage_id]); - - if (NeedsMultilevelTiling(task, state, stage_id)) { - // Do not use rfactor if we have enough parallelism on space iters - if (cum_space_len > cum_reduce_len || cum_space_len > task->hardware_params->num_cores * 16) { - return false; - } else { - return true; - } - } else if (cum_reduce_len > 1) { - // Always try rfactor for reduction ops - return cum_reduce_len > task->hardware_params->num_cores; - } - } - - return false; -} - -/*! \brief Return whether the stage has reduce iterators. */ -inline bool HasReduceIter(const Stage& stage) { - for (const auto& iter : stage->iters) { - if (iter->iter_kind != IteratorKind::kSpatial) { - return true; - } - } - return false; -} - -/*! \brief Return whether the stage has specific annotated iterators. */ -inline bool HasAnnotatedIter(const Stage& stage, IteratorAnnotation type) { - for (const auto& iter : stage->iters) { - if (iter->annotation == type) { - return true; - } - } - return false; -} - -/*! \brief Return whether the stage has only one consumer and they are elementwise-matched. */ -inline bool HasSingleElementwiseMatchedConsumer(const SearchTask& task, const State& state, - int stage_id, int* target_stage_id = nullptr) { - // Temporal object to be used if the input pointer is nullptr - int temp_target_stage_id; - if (target_stage_id == nullptr) { - target_stage_id = &temp_target_stage_id; - } - const std::set& consumers = GetConsumers(task, state, stage_id); - if (consumers.size() == 1) { - *target_stage_id = *consumers.begin(); - if (ElementwiseMatch(task, state, stage_id, *target_stage_id) && - (!(HasReduceIter(state->stages[stage_id]) && - HasReduceIter(state->stages[*target_stage_id]))) && - (!StrEndsWith(state->stages[*target_stage_id]->op->name, ".shared"))) { - return true; - } - } - return false; -} - -/*! \brief Return whether the step changes the number of stages */ -inline bool IsStageNumberChangingStep(const Step& step) { - return step->IsInstance() || step->IsInstance() || - step->IsInstance(); -} - -/*! \brief Return whether the state does cache_read for stage_id. */ -inline bool HasCacheReadStage(const State& s, int stage_id) { - for (int i = static_cast(s->transform_steps.size()) - 1; i >= 0; --i) { - if (auto ps = s->transform_steps[i].as()) { - if (stage_id == ps->stage_id) { - return true; - } - } - - if (IsStageNumberChangingStep(s->transform_steps[i])) { - if (stage_id > s->transform_steps[i]->stage_id) { - stage_id--; - } - } - } - return false; -} - -/*! \brief Return whether the state does cache_write for stage_id. */ -inline bool HasCacheWriteStage(const State& s, int stage_id) { - for (int i = static_cast(s->transform_steps.size()) - 1; i >= 0; --i) { - if (auto ps = s->transform_steps[i].as()) { - if (stage_id == ps->stage_id) { - return true; - } - } - - if (IsStageNumberChangingStep(s->transform_steps[i])) { - if (stage_id > s->transform_steps[i]->stage_id) { - stage_id--; - } - } - } - return false; -} - -/*! \brief Return whether the state does rfactor for stage_id. */ -inline bool HasRfactorStage(const State& s, int stage_id) { - for (int i = static_cast(s->transform_steps.size()) - 1; i >= 0; --i) { - if (auto ps = s->transform_steps[i].as()) { - if (stage_id == ps->stage_id) { - return true; - } - } - - if (IsStageNumberChangingStep(s->transform_steps[i])) { - if (stage_id > s->transform_steps[i]->stage_id) { - stage_id--; - } - } - } - return false; -} - -/*! \brief Return whether the stage does cross thread reduction. */ -inline bool HasCrossThreadReduction(const State& state, int stage_id) { - std::function check_stage = [](const Stage& in_stage) { - for (const auto& iter : in_stage->iters) { - if (iter->annotation == IteratorAnnotation::kThreadX && - iter->iter_kind == IteratorKind::kReduction) { - return true; - } - } - return false; - }; - - // Check the stage itself - if (check_stage(state->stages[stage_id])) { - return true; - } - - // Check the attached stages - for (size_t iter_id = 0; iter_id < state->stages[stage_id]->iters.size(); iter_id++) { - const auto& res = - state->attach_map->iter_to_attached_stages.find(std::make_pair(stage_id, iter_id)); - if (res != state->attach_map->iter_to_attached_stages.end()) { - for (int attached_stage_id : res->second) { - if (check_stage(state->stages[attached_stage_id])) { - return true; - } - } - } - } - - return false; -} - -/*! \brief Return whether the stage has been tiled already. */ -inline bool IsTiled(const Stage& stage) { - auto op = stage->op.as(); - ICHECK(op != nullptr); - return stage->iters.size() != op->axis.size() + op->reduce_axis.size(); -} - -/*! \brief Extract primitive iterators from a nested fused or splitted iterator's name. */ -inline void ExtractOriginalIterators(const std::string& name, std::set* rets) { - size_t last_pos = 0; - for (size_t i = 0; i < name.size(); ++i) { - if (name[i] == '@' || name[i] == '.') { // '@' for fuse and '.' for split - if (!isdigit(name[last_pos]) && name[last_pos] != '@' && name[last_pos] != '.') { - rets->insert(name.substr(last_pos, i - last_pos)); - } - last_pos = i + 1; - } - } - - if (last_pos < name.size() && !isdigit(name[last_pos]) && name[last_pos] != '@' && - name[last_pos] != '.') { - rets->insert(name.substr(last_pos, name.size() - last_pos)); - } -} - -/*! \brief Get the last reduce iterator in the outermost reduce tile. */ -inline Iterator GetLastReduceIteratorInOutermostReduceTile(const Stage& stage) { - auto pop = stage->op.as(); - ICHECK(pop != nullptr); - std::set original_names; - - const std::set& no_split_at_inner_name_set = - stage->op->attrs.count(SearchPolicyKey::no_split_at_inner) - ? GetIterNameSetParam(stage->op->attrs, SearchPolicyKey::no_split_at_inner) - : std::set(); - size_t reduce_axis_size = 0; - for (const auto axis : pop->reduce_axis) { - if (!no_split_at_inner_name_set.count(axis->var->name_hint)) { - reduce_axis_size++; - } - } - if (reduce_axis_size) { - for (const auto& iter : stage->iters) { - if (iter->iter_kind == IteratorKind::kReduction) { - ExtractOriginalIterators(iter->name, &original_names); - if (original_names.size() == reduce_axis_size) { - return iter; - } - } - } - } else { - // Return the first reduce iterator - for (const auto& iter : stage->iters) { - if (iter->iter_kind == IteratorKind::kReduction) { - return iter; - } - } - } - - LOG(FATAL) << "Cannot find the iterator."; -} - -/*! \brief Get the target stage id of a history step in the new state. - * We need this because the stage_id in the history may be stale due to later steps */ -inline int GetTargetStageIDInState(const State& s, int step_id) { - int stage_inc = 0; - - for (size_t i = step_id + 1; i < s->transform_steps.size(); ++i) { - if (IsStageNumberChangingStep(s->transform_steps[i])) { - if (s->transform_steps[i]->stage_id <= s->transform_steps[step_id]->stage_id + stage_inc) - stage_inc++; - } - } - return s->transform_steps[step_id]->stage_id + stage_inc; -} - -/*! \brief Get all split steps for one stage. */ -inline void GetSplitStepIds(const State& s, int stage_id, std::vector* split_step_ids) { - for (int i = static_cast(s->transform_steps.size()) - 1; i >= 0; --i) { - if (auto ps = s->transform_steps[i].as()) { - if (stage_id == ps->stage_id) { - split_step_ids->push_back(i); - } - } - - if (IsStageNumberChangingStep(s->transform_steps[i])) { - if (stage_id > s->transform_steps[i]->stage_id) { - stage_id--; - } - } - } -} - -/*! \brief Fuse all reduction iterators. */ -inline State FuseAllReductionIterators(const State& state, int stage_id, Iterator* fused_iter, - Array* space_iters, - Array* reduce_iters) { - space_iters->clear(); - reduce_iters->clear(); - - for (const auto& iter : state->stages[stage_id]->iters) { - if (iter->iter_kind == IteratorKind::kSpatial) { - space_iters->push_back(iter); - } else if (iter->iter_kind == IteratorKind::kReduction) { - reduce_iters->push_back(iter); - } - } - - ICHECK(!reduce_iters->empty()); - State tmp_s = state; - if (reduce_iters->size() > 1) { - *fused_iter = tmp_s.fuse(stage_id, *reduce_iters); - } else { - *fused_iter = (*reduce_iters)[0]; - } - return tmp_s; -} - -/*! \brief Fuse all outer level space iterators. */ -inline State FuseAllOuterSpaceIterators(const State& state, int stage_id, Iterator* fused_iter) { - std::vector to_fuse; - for (size_t iter_id = 0; iter_id < state->stages[stage_id]->iters.size(); ++iter_id) { - const auto& it = state->stages[stage_id]->iters[iter_id]; - // Stop at reduce iterator or annotated iterator - if (it->iter_kind == IteratorKind::kReduction || it->annotation != IteratorAnnotation::kNone) { - break; - } - // Stop at compute_at attach point - if (state->attach_map->iter_to_attached_stages.count(std::make_pair(stage_id, iter_id - 1))) { - break; - } - to_fuse.push_back(it); - } - - State tmp_s = state; - if (to_fuse.size() == 1) { - *fused_iter = to_fuse[0]; - } else { - *fused_iter = tmp_s.fuse(stage_id, to_fuse); - } - return tmp_s; -} - -/*! \brief Random sample states. */ -inline Array RandomSampleStates(const Array& in_states, std::mt19937* random_gen, - size_t out_size) { - Array out_states; - for (size_t i = 0; i < out_size; i++) { - out_states.push_back(in_states[(*random_gen)() % in_states.size()]); - } - return out_states; -} - -/*! \brief Compute prefix-sum probability based on the given weights */ -inline void ComputePrefixSumProb(const std::vector& weights, - std::vector* prefix_sum_probs) { - // Compute selection probabilities. - float sum = 0.0; - prefix_sum_probs->resize(weights.size()); - for (size_t i = 0; i < weights.size(); ++i) { - sum += std::max(weights[i], 0.0f); - (*prefix_sum_probs)[i] = sum; - } - for (size_t i = 0; i < weights.size(); ++i) { - (*prefix_sum_probs)[i] /= sum; - } -} - -/*! \brief Random choose an index according to a prefix sum probability. */ -inline int RandomChoose(const std::vector& prefix_sum_probs, std::mt19937* random_gen) { - std::uniform_real_distribution<> dis(0.0, 1.0); - double x = dis(*random_gen); - - ICHECK(!prefix_sum_probs.empty()); - - return std::lower_bound(prefix_sum_probs.begin(), prefix_sum_probs.end(), x) - - prefix_sum_probs.begin(); -} - -/*! \brief Print a title */ -inline void PrintTitle(const std::string& title, int verbose) { - StdCout(verbose) << Chars('-', 70) << "\n" - << Chars('-', 30) << " [ " << title << " ]\n" - << Chars('-', 70) << std::endl; -} - -/*! - * \brief Enumerate all possible factorization schemes for splitting an axes. - * \note This class will memorize the results for reuse. - */ -class SplitFactorizationMemo { - public: - using QueryKey = std::tuple; - - const Array>& GetFactorizationSchemes(int extent, int n_lengths, - int max_innermost_factor); - const std::vector& GetFactors(int n); - - private: - void DfsEnumerate(int now, int remaining_length, int max_innermost_factor); - - std::unordered_map>> memory_; - - int n_lengths_; - Array tmp_stack_; - Array>* results_; - std::unordered_map> factor_memory_; -}; - -/*! \brief Get the indexes of SplitStep that processes on spatial iterator. */ -Array GetSpatialSplitStepIds(const State& s, int stage_id); - -/*! \brief Get the possible compute locations for a stage. */ -std::vector> GetComputeLocationCandidates(const SearchTask& task, - const State& state, int stage_id); - -// Apply multi-level tiling structure according to a string format, -// where "S" stands a space level, "R" stands for a reduction level. -// For example, if the format is "SSRSRS", then we will -// use tiling structure: space_L0, space_L1, reduce_L0, space_L2, reduce_L1, space_L3 -// For example, if apply "SSRSRS" to matrix multiplication, -// we have space iterators i and j, reduce iterator k. -// Then the tiling structure is : i0, j0, i1, j1, k0, i2, j2, k1, i3, j3 -State DoMultiLevelTiling(const State& state, int stage_id, const std::string& format, - std::vector* spatial_split_step_ids = nullptr); - -// Apply tiling structure: space, space, space, ..., with tile sizes from other SplitStep -State FollowTiling(const State& state, int stage_id, const std::vector& split_step_ids, - int n_split); - -// Prune invalid states and return the results in-place. -void PruneInvalidState(const SearchTask& task, Array* states); - -} // namespace auto_scheduler -} // namespace tvm - -#endif // TVM_AUTO_SCHEDULER_SEARCH_POLICY_UTILS_H_ diff --git a/src/auto_scheduler/search_task.cc b/src/auto_scheduler/search_task.cc deleted file mode 100755 index ca59c1f3b077..000000000000 --- a/src/auto_scheduler/search_task.cc +++ /dev/null @@ -1,217 +0,0 @@ -/* - * 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 auto_scheduler/search_task.cc - * \brief Meta information and hardware parameters for a search task. - */ - -#include -#include -#include -#include -#include - -#include - -namespace tvm { -namespace auto_scheduler { - -TVM_REGISTER_NODE_TYPE(HardwareParamsNode); -TVM_REGISTER_NODE_TYPE(SearchTaskNode); - -HardwareParams::HardwareParams(int num_cores, int vector_unit_bytes, int cache_line_bytes, - int max_shared_memory_per_block, int max_local_memory_per_block, - int max_threads_per_block, int max_vthread_extent, int warp_size) { - auto node = make_object(); - node->num_cores = num_cores; - node->vector_unit_bytes = vector_unit_bytes; - node->cache_line_bytes = cache_line_bytes; - node->max_shared_memory_per_block = max_shared_memory_per_block; - node->max_local_memory_per_block = max_local_memory_per_block; - node->max_threads_per_block = max_threads_per_block; - node->max_vthread_extent = max_vthread_extent; - node->warp_size = warp_size; - data_ = std::move(node); -} - -HardwareParams HardwareParamsNode::GetDefaultHardwareParams(const Target& target, - const Target& target_host) { - // There is no use of target_host so no updates here in the function. - const auto device_type = target->GetTargetDeviceType(); - if (device_type == kDLCPU) { - return HardwareParams(tvm::runtime::threading::MaxConcurrency(), 64, 64, 0, 0, 0, 0, 0); - } else if (device_type == kDLCUDA || device_type == kDLROCM) { - auto dev = Device{static_cast(device_type), 0}; - auto device_name = device_type == kDLCUDA ? "device_api.cuda" : "device_api.rocm"; - auto func = tvm::runtime::Registry::Get(device_name); - ICHECK(func != nullptr) << "Cannot find CUDA device_api in registry"; - auto device_api = static_cast(((*func)()).operator void*()); - - tvm::runtime::TVMRetValue ret; - device_api->GetAttr(dev, tvm::runtime::DeviceAttrKind::kMaxSharedMemoryPerBlock, &ret); - int max_shared_memory_per_block = ret; - - // There is no explicit local memory limition in CUDA runtime, - // so we can use INT32_MAX to disalbe the check on local_memory. - int max_local_memory_per_block = INT32_MAX; - - device_api->GetAttr(dev, tvm::runtime::DeviceAttrKind::kMaxThreadsPerBlock, &ret); - int max_threads_per_block = ret; - - device_api->GetAttr(dev, tvm::runtime::DeviceAttrKind::kWarpSize, &ret); - int warp_size = ret; - - int max_vthread_extent = warp_size / 4; - return HardwareParams(-1, 16, 64, max_shared_memory_per_block, max_local_memory_per_block, - max_threads_per_block, max_vthread_extent, warp_size); - } else if (device_type == kDLMetal) { - // Reference: https://developer.apple.com/metal/Metal-Feature-Set-Tables.pdf - // This setting looks working for Metal GPUs later than A10 - int max_shared_memory_per_block = 32 * 1024; - int max_local_memory_per_block = INT32_MAX; // skip the check on local memory - int max_threads_per_block = 1024; - int warp_size = 8; - int max_vthread_extent = warp_size / 4; - return HardwareParams(-1, 16, 64, max_shared_memory_per_block, max_local_memory_per_block, - max_threads_per_block, max_vthread_extent, warp_size); - } else if (target->GetTargetDeviceType() == kDLOpenCL) { - if (target->GetAttr("device", "") == "mali") { - // We cannot use device API to get hardware attributes like CUDA, - // because like Mali target is normally on the remote machine. - int max_shared_memory_per_block = 32768; - int max_local_memory_per_block = INT32_MAX; // skip the check on local memory - int max_threads_per_block = 256; - int warp_size = 1; - int max_vthread_extent = 1; - return HardwareParams(-1, 16, 64, max_shared_memory_per_block, max_local_memory_per_block, - max_threads_per_block, max_vthread_extent, warp_size); - } else if (target->GetAttr("device", "") == "adreno") { - int max_shared_memory_per_block = 32768; - int max_local_memory_per_block = 32768; - int max_threads_per_block = 256; - int warp_size = 1; - int max_vthread_extent = 1; - return HardwareParams(-1, 16, 64, max_shared_memory_per_block, max_local_memory_per_block, - max_threads_per_block, max_vthread_extent, warp_size); - } else { - // add other opencl target - auto dev = Device{static_cast(device_type), 0}; - auto device_name = "device_api.opencl"; - auto func = tvm::runtime::Registry::Get(device_name); - ICHECK(func != nullptr) << "Cannot find OpenCL device_api in registry"; - auto device_api = static_cast(((*func)()).operator void*()); - - tvm::runtime::TVMRetValue ret; - device_api->GetAttr(dev, tvm::runtime::DeviceAttrKind::kMaxSharedMemoryPerBlock, &ret); - int max_shared_memory_per_block = ret; - - int max_local_memory_per_block = INT32_MAX; - - device_api->GetAttr(dev, tvm::runtime::DeviceAttrKind::kMaxThreadsPerBlock, &ret); - int max_threads_per_block = ret; - - device_api->GetAttr(dev, tvm::runtime::DeviceAttrKind::kWarpSize, &ret); - int warp_size = ret; - - if (warp_size == 1) { - LOG(WARNING) - << "Warp size 1 is not recommended for OpenCL devices. Tuning might crash or stuck"; - } - - int max_vthread_extent = std::max(1, warp_size / 4); - return HardwareParams(-1, 16, 64, max_shared_memory_per_block, max_local_memory_per_block, - max_threads_per_block, max_vthread_extent, warp_size); - } - } else if (device_type == kDLVulkan) { - auto dev = Device{static_cast(device_type), 0}; - auto device_name = "device_api.vulkan"; - auto func = tvm::runtime::Registry::Get(device_name); - ICHECK(func != nullptr) << "Cannot find Vulkan device_api in registry"; - auto device_api = static_cast(((*func)()).operator void*()); - - tvm::runtime::TVMRetValue ret; - device_api->GetAttr(dev, tvm::runtime::DeviceAttrKind::kMaxSharedMemoryPerBlock, &ret); - int max_shared_memory_per_block = ret; - - int max_local_memory_per_block = INT32_MAX; - - device_api->GetAttr(dev, tvm::runtime::DeviceAttrKind::kMaxThreadsPerBlock, &ret); - int max_threads_per_block = ret; - - device_api->GetAttr(dev, tvm::runtime::DeviceAttrKind::kWarpSize, &ret); - int warp_size = ret; - - int max_vthread_extent = std::max(1, warp_size / 4); - - return HardwareParams(-1, 16, 64, max_shared_memory_per_block, max_local_memory_per_block, - max_threads_per_block, max_vthread_extent, warp_size); - } else { - LOG(FATAL) << "No default hardware parameters for target: " << target; - } - return HardwareParams(); -} - -SearchTask::SearchTask(ComputeDAG compute_dag, String workload_key, Target target, - Target target_host, Optional hardware_params, - LayoutRewriteOption layout_rewrite_option, Array task_input_names, - String desc) { - CheckAndUpdateHostConsistency(&target, &target_host); - auto node = make_object(); - node->compute_dag = std::move(compute_dag); - node->workload_key = std::move(workload_key); - node->desc = std::move(desc); - node->target = std::move(target); - node->target_host = std::move(target_host); - if (hardware_params) { - node->hardware_params = hardware_params.value(); - } else { - node->hardware_params = - HardwareParamsNode::GetDefaultHardwareParams(node->target, node->target_host); - } - node->layout_rewrite_option = layout_rewrite_option; - node->task_input_names = std::move(task_input_names); - data_ = std::move(node); -} - -TVM_REGISTER_GLOBAL("auto_scheduler.HardwareParams") - .set_body_typed([](int num_cores, int vector_unit_bytes, int cache_line_bytes, - int max_shared_memory_per_block, int max_local_memory_per_block, - int max_threads_per_block, int max_vthread_extent, int warp_size) { - return HardwareParams(num_cores, vector_unit_bytes, cache_line_bytes, - max_shared_memory_per_block, max_local_memory_per_block, - max_threads_per_block, max_vthread_extent, warp_size); - }); - -TVM_REGISTER_GLOBAL("auto_scheduler.GetDefaultHardwareParams") - .set_body_typed([](Target target, Target target_host) { - return HardwareParamsNode::GetDefaultHardwareParams(target, target_host); - }); - -TVM_REGISTER_GLOBAL("auto_scheduler.SearchTask") - .set_body_typed([](ComputeDAG compute_dag, String workload_key, Target target, - Target target_host, Optional hardware_params, - int layout_rewrite_option, Array task_input_names, String desc) { - CheckAndUpdateHostConsistency(&target, &target_host); - return SearchTask(compute_dag, workload_key, target, target_host, hardware_params, - LayoutRewriteOption(layout_rewrite_option), task_input_names, desc); - }); - -} // namespace auto_scheduler -} // namespace tvm diff --git a/src/auto_scheduler/transform_step.cc b/src/auto_scheduler/transform_step.cc deleted file mode 100644 index 73acb7c0e70a..000000000000 --- a/src/auto_scheduler/transform_step.cc +++ /dev/null @@ -1,1879 +0,0 @@ -/* - * 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 auto_scheduler/transform_step.cc - * \brief Transformation steps. These steps are used to manipulate the LoopState. - * They are similar to the schedule primitives in te::Stage. - */ - -#include -#include -#include -#include -#include -#include - -#include -#include -#include - -#include "utils.h" - -namespace dmlc { -namespace json { - -template <> -struct Handler<::tvm::Array<::tvm::Integer>> { - inline static void Write(dmlc::JSONWriter* writer, const ::tvm::Array<::tvm::Integer>& array) { - writer->BeginArray(false); - for (const auto& i : array) { - ICHECK(i.defined()); - writer->WriteArrayItem(i->value); - } - writer->EndArray(); - } - inline static void Read(dmlc::JSONReader* reader, ::tvm::Array<::tvm::Integer>* array) { - array->clear(); - reader->BeginArray(); - while (reader->NextArrayItem()) { - int value; - Handler::Read(reader, &value); - array->push_back(value); - } - } -}; - -template <> -struct Handler<::tvm::Array<::tvm::Optional<::tvm::Integer>>> { - inline static void Write(dmlc::JSONWriter* writer, - const ::tvm::Array<::tvm::Optional<::tvm::Integer>>& array) { - writer->BeginArray(false); - for (const auto& i : array) { - ICHECK(i); - writer->WriteArrayItem(i.value()->value); - } - writer->EndArray(); - } - inline static void Read(dmlc::JSONReader* reader, - ::tvm::Array<::tvm::Optional<::tvm::Integer>>* array) { - array->clear(); - reader->BeginArray(); - while (reader->NextArrayItem()) { - int value; - Handler::Read(reader, &value); - array->push_back(::tvm::Integer(value)); - } - } -}; - -} // namespace json -} // namespace dmlc - -namespace tvm { -namespace auto_scheduler { - -// Update the te::stage to tir::IterVar axis mapping -void UpdateStageToAxesMap(const te::Stage& stage, StageToAxesMap* stage_to_axes) { - if (auto pop = stage->op.as()) { - Array axes; - for (const auto& axis : pop->axis) { - axes.push_back(axis); - } - for (const auto& axis : pop->reduce_axis) { - axes.push_back(axis); - } - stage_to_axes->Set(stage, std::move(axes)); - } else if (stage->op->IsInstance()) { - {} // do nothing on Placeholder - } else { - LOG(FATAL) << "Invalid op " << stage->op; - } -} - -const char* IteratorAnnotationString[] = { - "for", // kNone = 0 - "unroll", // kUnroll = 1 - "vectorize", // kVectorize = 2 - "parallel", // kParallel = 3 - "vthread", // kVThread = 4 - "blockIdx.x", // kBlockX = 5 - "threadIdx.x", // kThreadX = 6 - "blockIdx.y", // kBlockY = 7 - "threadIdx.y", // kThreadY = 8 - "blockIdx.z", // kBlockZ = 9 - "threadIdx.z", // kThreadZ = 10 - "tensorize" // kTensorized = 11 -}; - -StepNode* Step::CopyOnWrite() { - CHECK(data_ != nullptr); - if (!data_.unique()) { - if (const auto& ps = as()) { - auto n = make_object(*ps); - ObjectPtr(std::move(n)).swap(data_); - } else if (const auto& ps = as()) { - auto n = make_object(*ps); - ObjectPtr(std::move(n)).swap(data_); - } else if (const auto& ps = as()) { - auto n = make_object(*ps); - ObjectPtr(std::move(n)).swap(data_); - } else if (const auto& ps = as()) { - auto n = make_object(*ps); - ObjectPtr(std::move(n)).swap(data_); - } else if (const auto& ps = as()) { - auto n = make_object(*ps); - ObjectPtr(std::move(n)).swap(data_); - } else if (const auto& ps = as()) { - auto n = make_object(*ps); - ObjectPtr(std::move(n)).swap(data_); - } else if (const auto& ps = as()) { - auto n = make_object(*ps); - ObjectPtr(std::move(n)).swap(data_); - } else if (const auto& ps = as()) { - auto n = make_object(*ps); - ObjectPtr(std::move(n)).swap(data_); - } else if (const auto& ps = as()) { - auto n = make_object(*ps); - ObjectPtr(std::move(n)).swap(data_); - } else if (const auto& ps = as()) { - auto n = make_object(*ps); - ObjectPtr(std::move(n)).swap(data_); - } else if (const auto& ps = as()) { - auto n = make_object(*ps); - ObjectPtr(std::move(n)).swap(data_); - } else if (const auto& ps = as()) { - auto n = make_object(*ps); - ObjectPtr(std::move(n)).swap(data_); - } else if (const auto& ps = as()) { - auto n = make_object(*ps); - ObjectPtr(std::move(n)).swap(data_); - } else if (const auto& ps = as()) { - auto n = make_object(*ps); - ObjectPtr(std::move(n)).swap(data_); - } else { - LOG(FATAL) << "Invalid step: " << (*this); - } - } - return static_cast(data_.get()); -} - -Step StepReadFromRecord(dmlc::JSONReader* reader) { - std::string name; - bool s; - s = reader->NextArrayItem(); - ICHECK(s); - reader->Read(&name); - if (name == AnnotationStepNode::record_prefix_str) { - return AnnotationStep(reader); - } else if (name == FuseStepNode::record_prefix_str) { - return FuseStep(reader); - } else if (name == PragmaStepNode::record_prefix_str) { - return PragmaStep(reader); - } else if (name == ReorderStepNode::record_prefix_str) { - return ReorderStep(reader); - } else if (name == SplitStepNode::record_prefix_str) { - return SplitStep(reader); - } else if (name == FollowSplitStepNode::record_prefix_str) { - return FollowSplitStep(reader); - } else if (name == FollowFusedSplitStepNode::record_prefix_str) { - return FollowFusedSplitStep(reader); - } else if (name == StorageAlignStepNode::record_prefix_str) { - return StorageAlignStep(reader); - } else if (name == ComputeAtStepNode::record_prefix_str) { - return ComputeAtStep(reader); - } else if (name == ComputeInlineStepNode::record_prefix_str) { - return ComputeInlineStep(reader); - } else if (name == ComputeRootStepNode::record_prefix_str) { - return ComputeRootStep(reader); - } else if (name == CacheReadStepNode::record_prefix_str) { - return CacheReadStep(reader); - } else if (name == CacheWriteStepNode::record_prefix_str) { - return CacheWriteStep(reader); - } else if (name == RfactorStepNode::record_prefix_str) { - return RfactorStep(reader); - } else { - LOG(FATAL) << "Invalid step format: " << name; - } - return Step(); -} - -void StepApplyToState(const Step& step, State* state, const ComputeDAG& dag) { - // We need this runtime dispatcher because different steps have different function signatures - if (auto ps = step.as()) { - ps->ApplyToState(state); - } else if (auto ps = step.as()) { - ps->ApplyToState(state); - } else if (auto ps = step.as()) { - ps->ApplyToState(state); - } else if (auto ps = step.as()) { - ps->ApplyToState(state); - } else if (auto ps = step.as()) { - ps->ApplyToState(state); - } else if (auto ps = step.as()) { - ps->ApplyToState(state); - } else if (auto ps = step.as()) { - ps->ApplyToState(state); - } else if (auto ps = step.as()) { - ps->ApplyToState(state); - } else if (auto ps = step.as()) { - ps->ApplyToState(state); - } else if (auto ps = step.as()) { - ps->ApplyToState(state); - } else if (auto ps = step.as()) { - ps->ApplyToState(state); - } else if (auto ps = step.as()) { - ps->ApplyToState(state, dag); - } else if (auto ps = step.as()) { - ps->ApplyToState(state, dag); - } else if (auto ps = step.as()) { - ps->ApplyToState(state, dag); - } else { - LOG(FATAL) << "Invalid step: " << step; - } -} - -void StepApplyToSchedule(const Step& step, Array* stages, StageToAxesMap* stage_to_axes, - te::Schedule* schedule, const Array& transform_steps) { - if (auto ps = step.as()) { - ps->ApplyToSchedule(stages, stage_to_axes); - } else if (auto ps = step.as()) { - ps->ApplyToSchedule(stages, stage_to_axes); - } else if (auto ps = step.as()) { - ps->ApplyToSchedule(stages, stage_to_axes); - } else if (auto ps = step.as()) { - ps->ApplyToSchedule(stages, stage_to_axes); - } else if (auto ps = step.as()) { - ps->ApplyToSchedule(stages, stage_to_axes); - } else if (auto ps = step.as()) { - ps->ApplyToSchedule(stages, stage_to_axes, transform_steps); - } else if (auto ps = step.as()) { - ps->ApplyToSchedule(stages, stage_to_axes, transform_steps); - } else if (auto ps = step.as()) { - ps->ApplyToSchedule(stages, stage_to_axes); - } else if (auto ps = step.as()) { - ps->ApplyToSchedule(stages, stage_to_axes); - } else if (auto ps = step.as()) { - ps->ApplyToSchedule(stages, stage_to_axes); - } else if (auto ps = step.as()) { - ps->ApplyToSchedule(stages, stage_to_axes); - } else if (auto ps = step.as()) { - ps->ApplyToSchedule(stages, stage_to_axes, schedule); - } else if (auto ps = step.as()) { - ps->ApplyToSchedule(stages, stage_to_axes, schedule); - } else if (auto ps = step.as()) { - ps->ApplyToSchedule(stages, stage_to_axes, schedule); - } else { - LOG(FATAL) << "Invalid Step: " << step; - } -} - -String StepPrintAsPythonAPI(const Step& step, Array* stages, - StageToAxesMap* stage_to_axes, te::Schedule* schedule, - const Array& transform_steps) { - if (auto ps = step.as()) { - return ps->PrintAsPythonAPI(stages, stage_to_axes); - } else if (auto ps = step.as()) { - return ps->PrintAsPythonAPI(stages, stage_to_axes); - } else if (auto ps = step.as()) { - return ps->PrintAsPythonAPI(stages, stage_to_axes); - } else if (auto ps = step.as()) { - return ps->PrintAsPythonAPI(stages, stage_to_axes); - } else if (auto ps = step.as()) { - return ps->PrintAsPythonAPI(stages, stage_to_axes); - } else if (auto ps = step.as()) { - return ps->PrintAsPythonAPI(stages, stage_to_axes, transform_steps); - } else if (auto ps = step.as()) { - return ps->PrintAsPythonAPI(stages, stage_to_axes, transform_steps); - } else if (auto ps = step.as()) { - return ps->PrintAsPythonAPI(stages, stage_to_axes); - } else if (auto ps = step.as()) { - return ps->PrintAsPythonAPI(stages, stage_to_axes); - } else if (auto ps = step.as()) { - return ps->PrintAsPythonAPI(stages, stage_to_axes); - } else if (auto ps = step.as()) { - return ps->PrintAsPythonAPI(stages, stage_to_axes); - } else if (auto ps = step.as()) { - return ps->PrintAsPythonAPI(stages, stage_to_axes, schedule); - } else if (auto ps = step.as()) { - return ps->PrintAsPythonAPI(stages, stage_to_axes, schedule); - } else if (auto ps = step.as()) { - return ps->PrintAsPythonAPI(stages, stage_to_axes, schedule); - } else { - LOG(FATAL) << "Invalid Step: " << step; - } - return ""; -} - -/********** Steps working on single stage **********/ - -/********** Annotation **********/ -AnnotationStep::AnnotationStep(int stage_id, int iter_id, IteratorAnnotation ann) { - auto node = make_object(); - node->stage_id = stage_id; - node->iter_id = iter_id; - node->annotation = ann; - data_ = std::move(node); -} - -AnnotationStep::AnnotationStep(dmlc::JSONReader* reader) { - auto node = make_object(); - bool s; - s = reader->NextArrayItem(); - ICHECK(s); - reader->Read(&node->stage_id); - s = reader->NextArrayItem(); - ICHECK(s); - reader->Read(&node->iter_id); - s = reader->NextArrayItem(); - ICHECK(s); - int int_val; - reader->Read(&int_val); - node->annotation = IteratorAnnotation(int_val); - data_ = std::move(node); -} - -void AnnotationStepNode::WriteToRecord(dmlc::JSONWriter* writer) const { - writer->WriteArraySeperator(); - writer->WriteString(record_prefix_str); - writer->WriteArrayItem(stage_id); - writer->WriteArrayItem(iter_id); - writer->WriteArrayItem(static_cast(annotation)); -} - -Iterator AnnotationStepNode::ApplyToState(State* state) const { - const Stage& stage = (*state)->stages[stage_id]; - Iterator it = stage->iters[iter_id]; - - ICHECK(it->annotation == IteratorAnnotation::kNone); - Iterator new_it = Iterator(it->name, it->range, it->iter_kind, annotation, &it->orig_iters); - Stage new_stage = stage; - new_stage.CopyOnWrite()->iters.Set(iter_id, new_it); - state->CopyOnWrite()->stages.Set(stage_id, std::move(new_stage)); - return new_it; -} - -void AnnotationStepNode::ApplyToSchedule(Array* stages, - StageToAxesMap* stage_to_axes) const { - te::Stage stage = (*stages)[stage_id]; - const Array& axes = (*stage_to_axes)[stage]; - - switch (annotation) { - case IteratorAnnotation::kUnroll: - stage.unroll(axes[iter_id]); - break; - case IteratorAnnotation::kVectorize: - stage.vectorize(axes[iter_id]); - break; - case IteratorAnnotation::kParallel: - stage.parallel(axes[iter_id]); - break; - case IteratorAnnotation::kVThread: - case IteratorAnnotation::kBlockX: - case IteratorAnnotation::kBlockY: - case IteratorAnnotation::kBlockZ: - case IteratorAnnotation::kThreadX: - case IteratorAnnotation::kThreadY: - case IteratorAnnotation::kThreadZ: - stage.bind(axes[iter_id], - te::thread_axis(Range(), IteratorAnnotationString[static_cast(annotation)])); - break; - case IteratorAnnotation::kNone: - break; - default: - LOG(FATAL) << "Invalid Annotation " << static_cast(annotation); - break; - } - - stages->Set(stage_id, std::move(stage)); -} - -String AnnotationStepNode::PrintAsPythonAPI(Array* stages, - StageToAxesMap* stage_to_axes) const { - std::stringstream ss; - const auto& stage = (*stages)[stage_id]; - const auto& iter = (*stage_to_axes)[stage][iter_id]; - const auto& op_name = CleanName(stage->op->name); - - ss << "s[" << op_name << "]."; - switch (annotation) { - case IteratorAnnotation::kUnroll: - ss << "unroll("; - break; - case IteratorAnnotation::kVectorize: - ss << "vectorize("; - break; - case IteratorAnnotation::kParallel: - ss << "parallel("; - break; - case IteratorAnnotation::kVThread: - case IteratorAnnotation::kBlockX: - case IteratorAnnotation::kBlockY: - case IteratorAnnotation::kBlockZ: - case IteratorAnnotation::kThreadX: - case IteratorAnnotation::kThreadY: - case IteratorAnnotation::kThreadZ: - ss << "bind("; - break; - case IteratorAnnotation::kNone: - break; - default: - LOG(FATAL) << "Invalid annotation " << static_cast(annotation); - break; - } - ss << CleanName(iter->var->name_hint, op_name); - switch (annotation) { - case IteratorAnnotation::kVThread: - case IteratorAnnotation::kBlockX: - case IteratorAnnotation::kBlockY: - case IteratorAnnotation::kBlockZ: - case IteratorAnnotation::kThreadX: - case IteratorAnnotation::kThreadY: - case IteratorAnnotation::kThreadZ: - ss << ", te.thread_axis(\"" << IteratorAnnotationString[static_cast(annotation)] - << "\")"; - break; - default: - break; - } - ss << ")\n"; - - ApplyToSchedule(stages, stage_to_axes); - return ss.str(); -} - -/********** Fuse **********/ -FuseStep::FuseStep(int stage_id, const Array& fused_ids) { - auto node = make_object(); - node->stage_id = stage_id; - for (const auto& x : fused_ids) { - ICHECK(x->IsInstance()); - } - node->fused_ids = fused_ids; - data_ = std::move(node); -} - -FuseStep::FuseStep(dmlc::JSONReader* reader) { - auto node = make_object(); - bool s; - s = reader->NextArrayItem(); - ICHECK(s); - reader->Read(&node->stage_id); - s = reader->NextArrayItem(); - ICHECK(s); - reader->Read(&node->fused_ids); - data_ = std::move(node); -} - -void FuseStepNode::WriteToRecord(dmlc::JSONWriter* writer) const { - writer->WriteArraySeperator(); - writer->WriteString(record_prefix_str); - writer->WriteArrayItem(stage_id); - writer->WriteArrayItem(fused_ids); -} - -Iterator FuseStepNode::ApplyToState(State* state) const { - const Stage& stage = (*state)->stages[stage_id]; - size_t old_iter_size = static_cast(stage->iters.size()); - - String new_name; - PrimExpr new_extent = 1; - IteratorKind new_iter_kind = IteratorKind::kSpecial; - std::vector orig_iters; - - for (size_t i = 0; i < fused_ids.size(); ++i) { - if (i > 0) { - ICHECK_EQ(fused_ids[i]->value, fused_ids[i - 1]->value + 1); - } - if (i != fused_ids.size() - 1) { - const auto& iter_to_attached_stage = (*state)->attach_map->iter_to_attached_stages; - if (iter_to_attached_stage.find(std::make_pair(stage_id, fused_ids[i].IntValue())) != - iter_to_attached_stage.end()) { - LOG(FATAL) << "Invalid Fuse. Trying to fuse iterators that have been attached by some " - << "stages. State before fusion:\n" - << (*state); - } - } - - const Iterator& it = stage->iters[fused_ids[i].IntValue()]; - orig_iters.push_back(it); - new_name = new_name + it->name + "@"; - - if (it->range.defined() && new_extent.defined()) { - new_extent = new_extent * it->range->extent; - } else { - new_extent = PrimExpr(); - } - - if (i == 0) { - new_iter_kind = it->iter_kind; - } else { - if (new_iter_kind != it->iter_kind) { - new_iter_kind = IteratorKind::kMixed; - } - } - } - - Range range; - if (new_extent.defined()) { - range = Range::FromMinExtent(0, new_extent); - } - Iterator new_it = - Iterator(new_name, range, new_iter_kind, IteratorAnnotation::kNone, &orig_iters); - Array new_iters; - - if (fused_ids.empty()) { - new_iters.push_back(new_it); - } else { - new_iters.insert(new_iters.end(), stage->iters.begin(), - stage->iters.begin() + fused_ids.front().IntValue()); - new_iters.push_back(new_it); - new_iters.insert(new_iters.end(), stage->iters.begin() + fused_ids.back().IntValue() + 1, - stage->iters.end()); - } - - StateNode* pstate = state->CopyOnWrite(); - pstate->stages.Set(stage_id, - Stage(stage->op, stage->op_type, new_iters, stage->compute_at, stage->attrs)); - - if (fused_ids.empty()) { - return new_it; - } - - // Two vectors are used to represent the iterator relation before and after fuse - // The original iterators in AttachMap will be updated with the new iterators - std::vector from_iters; - std::vector to_iters; - const size_t begin_id = fused_ids.front().IntValue(), end_id = fused_ids.back().IntValue(); - for (size_t i = 0; i < old_iter_size; ++i) { - if (i <= begin_id) { - continue; - } else if (i > end_id) { - // move forward - from_iters.emplace_back(stage_id, i); - to_iters.emplace_back(stage_id, i - end_id + begin_id); - } else { - // move to the fused id - from_iters.emplace_back(stage_id, i); - to_iters.emplace_back(stage_id, begin_id); - } - } - pstate->attach_map.UpdateIters(from_iters, to_iters); - - return new_it; -} - -IterVar FuseStepNode::ApplyToSchedule(Array* stages, - StageToAxesMap* stage_to_axes) const { - auto stage = (*stages)[stage_id]; - const Array& axes = stage_to_axes->at(stage); - - Array to_fuse; - for (const auto& i : fused_ids) { - to_fuse.push_back(axes[i.IntValue()]); - } - IterVar fused_axis; - stage.fuse(to_fuse, &fused_axis); - - Array new_axes; - if (fused_ids.empty()) { - new_axes.push_back(fused_axis); - } else { - new_axes.insert(new_axes.end(), axes.begin(), axes.begin() + fused_ids.front().IntValue()); - new_axes.push_back(fused_axis); - new_axes.insert(new_axes.end(), axes.begin() + fused_ids.back().IntValue() + 1, axes.end()); - } - - stage_to_axes->Set(stage, std::move(new_axes)); - stages->Set(stage_id, std::move(stage)); - return fused_axis; -} - -String FuseStepNode::PrintAsPythonAPI(Array* stages, - StageToAxesMap* stage_to_axes) const { - const auto& stage = (*stages)[stage_id]; - const auto& op_name = CleanName(stage->op->name); - std::stringstream to_fuse; - - for (size_t i = 0; i < fused_ids.size(); ++i) { - to_fuse << CleanName(stage_to_axes->at(stage)[fused_ids[i].IntValue()]->var->name_hint, - op_name); - if (i != fused_ids.size() - 1) { - to_fuse << ", "; - } - } - - std::stringstream ss; - const auto& fused = ApplyToSchedule(stages, stage_to_axes); - - ss << CleanName(fused->var->name_hint, op_name) << " = s[" << op_name << "].fuse(" - << to_fuse.str() << ")\n"; - - return ss.str(); -} - -/********** Pragma **********/ -PragmaStep::PragmaStep(int stage_id, int iter_id, String pragma_type) { - auto node = make_object(); - node->stage_id = stage_id; - node->iter_id = iter_id; - node->pragma_type = std::move(pragma_type); - data_ = std::move(node); -} - -PragmaStep::PragmaStep(dmlc::JSONReader* reader) { - auto node = make_object(); - bool s; - s = reader->NextArrayItem(); - ICHECK(s); - reader->Read(&node->stage_id); - s = reader->NextArrayItem(); - ICHECK(s); - reader->Read(&node->iter_id); - s = reader->NextArrayItem(); - ICHECK(s); - std::string string_value; - reader->Read(&string_value); - node->pragma_type = std::move(string_value); - data_ = std::move(node); -} - -void PragmaStepNode::WriteToRecord(dmlc::JSONWriter* writer) const { - writer->WriteArraySeperator(); - writer->WriteString(record_prefix_str); - writer->WriteArrayItem(stage_id); - writer->WriteArrayItem(iter_id); - writer->WriteArraySeperator(); - writer->WriteString(pragma_type); -} - -void PragmaStepNode::ApplyToState(State* state) const { - if (pragma_type == "debug_skip_region") { - StateNode* pstate = state->CopyOnWrite(); - pstate->attach_map.DeleteStage(stage_id); - } else if (StrStartsWith(pragma_type, "auto_unroll_max_step")) { - StateNode* pstate = state->CopyOnWrite(); - Stage stage = pstate->stages[stage_id]; - size_t pos = 0; - for (; pos < pragma_type.size(); ++pos) { - if ((*(pragma_type.c_str() + pos)) == '$') { - break; - } - } - ICHECK_LT(pos, pragma_type.size()) << "max step value not found."; - stage.CopyOnWrite()->attrs.auto_unroll_max_step = atoi(pragma_type.c_str() + pos + 1); - pstate->stages.Set(stage_id, std::move(stage)); - } else { - LOG(FATAL) << "Unsupported pragma: " << pragma_type; - } -} - -void PragmaStepNode::ApplyToSchedule(Array* stages, - StageToAxesMap* stage_to_axes) const { - te::Stage stage = (*stages)[stage_id]; - const Array& axes = (*stage_to_axes)[stage]; - if (StrStartsWith(pragma_type, "auto_unroll_max_step")) { - size_t pos = 0; - for (; pos < pragma_type.size(); ++pos) { - if ((*(pragma_type.c_str() + pos)) == '$') { - break; - } - } - ICHECK_LT(pos, pragma_type.size()) << "max step value not found."; - int value = atoi(pragma_type.c_str() + pos + 1); - if (iter_id < static_cast(axes.size())) { - stage.pragma(axes[iter_id], "auto_unroll_max_step", value); - stage.pragma(axes[iter_id], "unroll_explicit", true); - } - } else { - ICHECK_LT(iter_id, axes.size()); - stage.pragma(axes[iter_id], pragma_type); - } - stages->Set(stage_id, std::move(stage)); -} - -String PragmaStepNode::PrintAsPythonAPI(Array* stages, - StageToAxesMap* stage_to_axes) const { - std::stringstream ss; - const auto& stage = (*stages)[stage_id]; - const auto& op_name = CleanName(stage->op->name); - - if (StrStartsWith(pragma_type, "auto_unroll_max_step")) { - size_t pos = 0; - for (; pos < pragma_type.size(); ++pos) { - if ((*(pragma_type.c_str() + pos)) == '$') { - break; - } - } - ICHECK_LT(pos, pragma_type.size()) << "max step value not found."; - int value = atoi(pragma_type.c_str() + pos + 1); - ss << "s[" << op_name << "].pragma(" - << CleanName((*stage_to_axes)[stage][iter_id]->var->name_hint, op_name) - << ", \"auto_unroll_max_step\", " << value << ")\n"; - ss << "s[" << op_name << "].pragma(" - << CleanName((*stage_to_axes)[stage][iter_id]->var->name_hint, op_name) - << ", \"unroll_explicit\", True)\n"; - } else { - ss << "s[" << op_name << "].pragma(" - << CleanName((*stage_to_axes)[stage][iter_id]->var->name_hint, op_name) << ", \"" - << pragma_type << "\")\n"; - } - - ApplyToSchedule(stages, stage_to_axes); - return ss.str(); -} - -/********** Reorder **********/ -ReorderStep::ReorderStep(int stage_id, const Array& after_ids) { - auto node = make_object(); - node->stage_id = stage_id; - for (const auto& x : after_ids) { - ICHECK(x->IsInstance()); - } - node->after_ids = after_ids; - data_ = std::move(node); -} - -ReorderStep::ReorderStep(dmlc::JSONReader* reader) { - auto node = make_object(); - bool s; - s = reader->NextArrayItem(); - ICHECK(s); - reader->Read(&node->stage_id); - s = reader->NextArrayItem(); - ICHECK(s); - reader->Read(&node->after_ids); - data_ = std::move(node); -} - -void ReorderStepNode::WriteToRecord(dmlc::JSONWriter* writer) const { - writer->WriteArraySeperator(); - writer->WriteString(record_prefix_str); - writer->WriteArrayItem(stage_id); - writer->WriteArrayItem(after_ids); -} - -void ReorderStepNode::ApplyToState(State* state) const { - const Stage& stage = (*state)->stages[stage_id]; - Array iters; - for (auto x : after_ids) { - iters.push_back(stage->iters[x.IntValue()]); - } - state->CopyOnWrite()->stages.Set( - stage_id, Stage(stage->op, stage->op_type, iters, stage->compute_at, stage->attrs)); -} - -void ReorderStepNode::ApplyToSchedule(Array* stages, - StageToAxesMap* stage_to_axes) const { - auto stage = (*stages)[stage_id]; - const Array& axes = stage_to_axes->at(stage); - ICHECK_EQ(after_ids.size(), axes.size()); - - Array new_axes; - new_axes.reserve(axes.size()); - for (auto i : after_ids) { - new_axes.push_back(axes[i.IntValue()]); - } - stage.reorder(new_axes); - - stage_to_axes->Set(stage, std::move(new_axes)); - stages->Set(stage_id, std::move(stage)); -} - -String ReorderStepNode::PrintAsPythonAPI(Array* stages, - StageToAxesMap* stage_to_axes) const { - const auto& stage = (*stages)[stage_id]; - const auto& op_name = CleanName(stage->op->name); - std::stringstream ss; - - ss << "s[" << op_name << "].reorder("; - for (size_t i = 0; i < after_ids.size(); ++i) { - ss << CleanName((*stage_to_axes)[stage][after_ids[i].IntValue()]->var->name_hint, op_name); - if (i != after_ids.size() - 1) { - ss << ", "; - } - } - ss << ")\n"; - - ApplyToSchedule(stages, stage_to_axes); - return ss.str(); -} - -/********** Split **********/ -// common part for SplitStep, FollowSplitStep, and FollowFusedSplitStep -Array ApplySplitToState(State* state, int stage_id, int iter_id, - const Array>& lengths, bool inner_to_outer) { - const Stage& stage = (*state)->stages[stage_id]; - const Iterator& it = stage->iters[iter_id]; - size_t old_iter_size = stage->iters.size(); - bool concrete = true; - - Optional tosplit_min, tosplit_extent; - if (it->range.defined()) { - tosplit_min = it->range->min; - tosplit_extent = it->range->extent; - } else { - tosplit_min = NullOpt; - tosplit_extent = NullOpt; - } - - Array outs; - for (size_t i = 0; i < lengths.size(); ++i) { - Optional l; - String name; - if (inner_to_outer) { - l = lengths[lengths.size() - i - 1]; - name = it->name + "." + std::to_string(lengths.size() - i); - } else { - l = lengths[i]; - name = it->name + "." + std::to_string(i); - } - Iterator res; - if (l && tosplit_min && tosplit_extent) { - res = Iterator(name, Range::FromMinExtent(tosplit_min.value(), l.value()), it->iter_kind, - IteratorAnnotation::kNone); - tosplit_min = Integer(0); - tosplit_extent = indexdiv(tosplit_extent.value() + l.value() - 1, l.value()); - } else { - res = Iterator(name, Range(), it->iter_kind, IteratorAnnotation::kNone); - tosplit_min = NullOpt; - tosplit_extent = NullOpt; - if (!l.defined()) { - concrete = false; - } - } - outs.push_back(std::move(res)); - } - - Range range; - if (tosplit_min && tosplit_extent) { - range = Range::FromMinExtent(tosplit_min.value(), tosplit_extent.value()); - } - if (inner_to_outer) { - outs.push_back(Iterator(it->name + ".0", range, it->iter_kind, IteratorAnnotation::kNone)); - // Reverse the Iterator array - Array temp(outs.rbegin(), outs.rend()); - outs = std::move(temp); - } else { - outs.push_back(Iterator(it->name + "." + std::to_string(lengths.size()), range, it->iter_kind, - IteratorAnnotation::kNone)); - } - - Array new_iters; - new_iters.insert(new_iters.end(), stage->iters.begin(), stage->iters.begin() + iter_id); - new_iters.insert(new_iters.end(), outs.begin(), outs.end()); - new_iters.insert(new_iters.end(), stage->iters.begin() + iter_id + 1, stage->iters.end()); - - StateNode* pstate = state->CopyOnWrite(); - pstate->stages.Set(stage_id, - Stage(stage->op, stage->op_type, new_iters, stage->compute_at, stage->attrs)); - pstate->concrete &= concrete; - - // Two vectors are used to represent the iterator relation before and after split - // The original iterators in AttachMap will be updated with the new iterators - std::vector from_iters; - std::vector to_iters; - for (size_t i = iter_id; i < old_iter_size; ++i) { - from_iters.emplace_back(stage_id, i); - to_iters.emplace_back(stage_id, i + lengths.size()); - } - pstate->attach_map.UpdateIters(from_iters, to_iters); - - return outs; -} - -Array ApplySplitToSchedule(Array* stages, StageToAxesMap* stage_to_axes, - int stage_id, int iter_id, - const Array>& lengths, bool inner_to_outer) { - auto stage = (*stages)[stage_id]; - const Array& axes = stage_to_axes->at(stage); - - Array outs; - if (inner_to_outer) { - IterVar outer = axes[iter_id], inner; - for (int i = static_cast(lengths.size()) - 1; i >= 0; i--) { - IterVar to_split = outer; - stage.split(to_split, lengths[i].value(), &outer, &inner); - outs.push_back(inner); - } - outs.push_back(outer); - } else { - IterVar outer, inner = axes[iter_id]; - for (size_t i = 0; i < lengths.size(); i++) { - IterVar to_split = inner; - stage.split_by_nparts(to_split, lengths[i].value(), &outer, &inner); - outs.push_back(outer); - } - outs.push_back(inner); - } - - Array new_axes; - new_axes.insert(new_axes.end(), axes.begin(), axes.begin() + iter_id); - if (inner_to_outer) { - for (auto x = outs.rbegin(); x != outs.rend(); ++x) { - new_axes.push_back((*x)); - } - } else { - for (const auto& x : outs) { - new_axes.push_back(x); - } - } - new_axes.insert(new_axes.end(), axes.begin() + iter_id + 1, axes.end()); - - stage_to_axes->Set(stage, std::move(new_axes)); - stages->Set(stage_id, std::move(stage)); - return outs; -} - -String PrintSplitAsPythonAPI(Array* stages, StageToAxesMap* stage_to_axes, int stage_id, - int iter_id, const Array>& lengths, - bool inner_to_outer) { - const auto& stage = (*stages)[stage_id]; - auto to_split = stage_to_axes->at(stage)[iter_id]; - const auto& func_name = CleanName(stage->op->name); - const auto& outs = - ApplySplitToSchedule(stages, stage_to_axes, stage_id, iter_id, lengths, inner_to_outer); - ICHECK_EQ(outs.size(), lengths.size() + 1); - - std::stringstream ss; - int size = static_cast(lengths.size()); - if (inner_to_outer) { - for (int i = size - 1; i >= 0; i--) { - ss << CleanName(outs[size - i]->var->name_hint, func_name) << ", " - << CleanName(outs[size - i - 1]->var->name_hint, func_name) << " = s[" << func_name - << "].split(" << CleanName(to_split->var->name_hint, func_name) - << ", factor=" << lengths[i] << ")\n"; - to_split = outs[size - i]; - } - } else { - for (int i = 0; i < size; i++) { - ss << CleanName(outs[i]->var->name_hint, func_name) << ", " - << CleanName(outs[i + 1]->var->name_hint, func_name) << " = s[" << func_name << "].split(" - << CleanName(to_split->var->name_hint, func_name) << ", nparts=" << lengths[i] << ")\n"; - to_split = outs[i + 1]; - } - } - - return ss.str(); -} - -SplitStep::SplitStep(int stage_id, int iter_id, Optional extent, - const Array>& lengths, bool inner_to_outer) { - auto node = make_object(); - node->stage_id = stage_id; - // Extent can be a irreducible expression in some special cases - if (extent && extent.value()->IsInstance()) { - node->extent = tvm::Downcast(extent.value()); - } - node->iter_id = iter_id; - node->lengths = lengths; - node->inner_to_outer = inner_to_outer; - data_ = std::move(node); -} - -SplitStep::SplitStep(dmlc::JSONReader* reader) { - auto node = make_object(); - bool s; - s = reader->NextArrayItem(); - ICHECK(s); - reader->Read(&node->stage_id); - s = reader->NextArrayItem(); - ICHECK(s); - reader->Read(&node->iter_id); - int int_val; - s = reader->NextArrayItem(); - ICHECK(s); - reader->Read(&int_val); - if (int_val) { - node->extent = Integer(int_val); - } - s = reader->NextArrayItem(); - ICHECK(s); - reader->Read(&node->lengths); - s = reader->NextArrayItem(); - ICHECK(s); - reader->Read(&node->inner_to_outer); - data_ = std::move(node); -} - -void SplitStepNode::WriteToRecord(dmlc::JSONWriter* writer) const { - writer->WriteArraySeperator(); - writer->WriteString(record_prefix_str); - writer->WriteArrayItem(stage_id); - writer->WriteArrayItem(iter_id); - writer->WriteArrayItem(extent ? GetIntImm(extent.value()) : 0); - writer->WriteArrayItem(lengths); - writer->WriteArrayItem(static_cast(inner_to_outer)); -} - -Array SplitStepNode::ApplyToState(State* state) const { - return ApplySplitToState(state, stage_id, iter_id, lengths, inner_to_outer); -} - -Array SplitStepNode::ApplyToSchedule(Array* stages, - StageToAxesMap* stage_to_axes) const { - return ApplySplitToSchedule(stages, stage_to_axes, stage_id, iter_id, lengths, inner_to_outer); -} - -String SplitStepNode::PrintAsPythonAPI(Array* stages, - StageToAxesMap* stage_to_axes) const { - return PrintSplitAsPythonAPI(stages, stage_to_axes, stage_id, iter_id, lengths, inner_to_outer); -} - -/********** Follow Split **********/ -FollowSplitStep::FollowSplitStep(int stage_id, int iter_id, int src_step_id, int n_split) { - auto node = make_object(); - node->stage_id = stage_id; - node->iter_id = iter_id; - node->src_step_id = src_step_id; - node->n_split = n_split; - data_ = std::move(node); -} - -void FollowSplitStepNode::WriteToRecord(dmlc::JSONWriter* writer) const { - writer->WriteArraySeperator(); - writer->WriteString(record_prefix_str); - writer->WriteArrayItem(stage_id); - writer->WriteArrayItem(iter_id); - writer->WriteArrayItem(src_step_id); - writer->WriteArrayItem(n_split); -} - -Array> FollowSplitStepNode::ExtractSplitLengths( - const Array& transform_steps) const { - // Make sure src_step_id is within the range of transform_steps. - ICHECK_LT(src_step_id, transform_steps.size()); - auto ps = transform_steps[src_step_id].as(); - ICHECK(ps != nullptr); - - // Make sure the size of ps->lengths is not smaller than n_split-1. - // Note that the number of actual splitting factors of src_step is ps->lengths.size()+1. - ICHECK_LE(n_split, ps->lengths.size() + 1); - ICHECK(ps != nullptr); - - Array> lengths; - lengths.reserve(n_split); - int j = 0; - // Get the first (n_split-1) split factors of followed src_step. - for (; j < n_split - 1; ++j) { - lengths.push_back(ps->lengths[j]); - } - - // Get the last split factor of src_step for splitting level if n_split is smaller than - // ps->lengths.size()+1. - PrimExpr last_factor = 1; - for (; j < static_cast(ps->lengths.size()); ++j) { - if (ps->lengths[j]) { - last_factor *= ps->lengths[j].value(); - } else { - last_factor = PrimExpr(); - break; - } - } - if (last_factor.defined()) { - lengths.push_back(Downcast(last_factor)); - } else { - lengths.push_back(NullOpt); - } - - return lengths; -} - -FollowSplitStep::FollowSplitStep(dmlc::JSONReader* reader) { - auto node = make_object(); - bool s; - s = reader->NextArrayItem(); - ICHECK(s); - reader->Read(&node->stage_id); - s = reader->NextArrayItem(); - ICHECK(s); - reader->Read(&node->iter_id); - s = reader->NextArrayItem(); - ICHECK(s); - reader->Read(&node->src_step_id); - s = reader->NextArrayItem(); - ICHECK(s); - reader->Read(&node->n_split); - data_ = std::move(node); -} - -Array FollowSplitStepNode::ApplyToState(State* state) const { - return ApplySplitToState(state, stage_id, iter_id, ExtractSplitLengths((*state)->transform_steps), - true); -} - -Array FollowSplitStepNode::ApplyToSchedule(Array* stages, - StageToAxesMap* stage_to_axes, - const Array& transform_steps) const { - return ApplySplitToSchedule(stages, stage_to_axes, stage_id, iter_id, - ExtractSplitLengths(transform_steps), true); -} - -String FollowSplitStepNode::PrintAsPythonAPI(Array* stages, - StageToAxesMap* stage_to_axes, - const Array& transform_steps) const { - return PrintSplitAsPythonAPI(stages, stage_to_axes, stage_id, iter_id, - ExtractSplitLengths(transform_steps), true); -} - -/********** Follow Fused Split **********/ -FollowFusedSplitStep::FollowFusedSplitStep(int stage_id, int iter_id, - const Array& src_step_ids, int level, - bool factor_or_nparts) { - auto node = make_object(); - node->stage_id = stage_id; - node->iter_id = iter_id; - node->src_step_ids = src_step_ids; - node->level = level; - node->factor_or_nparts = factor_or_nparts; - data_ = std::move(node); -} - -FollowFusedSplitStep::FollowFusedSplitStep(dmlc::JSONReader* reader) { - auto node = make_object(); - bool s; - s = reader->NextArrayItem(); - ICHECK(s); - reader->Read(&node->stage_id); - s = reader->NextArrayItem(); - ICHECK(s); - reader->Read(&node->iter_id); - s = reader->NextArrayItem(); - ICHECK(s); - reader->Read(&node->src_step_ids); - s = reader->NextArrayItem(); - ICHECK(s); - reader->Read(&node->level); - s = reader->NextArrayItem(); - ICHECK(s); - reader->Read(&node->factor_or_nparts); - data_ = std::move(node); -} - -void FollowFusedSplitStepNode::WriteToRecord(dmlc::JSONWriter* writer) const { - writer->WriteArraySeperator(); - writer->WriteString(record_prefix_str); - writer->WriteArrayItem(stage_id); - writer->WriteArrayItem(iter_id); - writer->WriteArrayItem(src_step_ids); - writer->WriteArrayItem(level); - writer->WriteArrayItem(static_cast(factor_or_nparts)); -} - -Optional FollowFusedSplitStepNode::ExtractSplitLength( - const Array& transform_steps) const { - PrimExpr ret(1); - - for (auto src_step_id : src_step_ids) { - // Make sure the src_step_id is within the range of transform_steps. - ICHECK_LT(src_step_id.IntValue(), transform_steps.size()); - auto ps = transform_steps[src_step_id.IntValue()].as(); - ICHECK(ps != nullptr); - // Multiple the splitting factor on corresponding splitting level of src_steps. - if (ps->lengths[level] && ret.defined()) { - ret *= ps->lengths[level].value(); - } else { - return NullOpt; - } - } - return Downcast(ret); -} - -Array FollowFusedSplitStepNode::ApplyToState(State* state) const { - return ApplySplitToState(state, stage_id, iter_id, - {ExtractSplitLength((*state)->transform_steps)}, factor_or_nparts); -} - -Array FollowFusedSplitStepNode::ApplyToSchedule(Array* stages, - StageToAxesMap* stage_to_axes, - const Array& transform_steps) const { - return ApplySplitToSchedule(stages, stage_to_axes, stage_id, iter_id, - {ExtractSplitLength(transform_steps)}, factor_or_nparts); -} - -String FollowFusedSplitStepNode::PrintAsPythonAPI(Array* stages, - StageToAxesMap* stage_to_axes, - const Array& transform_steps) const { - return PrintSplitAsPythonAPI(stages, stage_to_axes, stage_id, iter_id, - {ExtractSplitLength(transform_steps)}, factor_or_nparts); -} - -/********** Storage Align **********/ -StorageAlignStep::StorageAlignStep(int stage_id, int iter_id, int factor, int offset) { - auto node = make_object(); - node->stage_id = stage_id; - node->iter_id = iter_id; - node->factor = factor; - node->offset = offset; - data_ = std::move(node); -} - -StorageAlignStep::StorageAlignStep(dmlc::JSONReader* reader) { - auto node = make_object(); - bool s; - s = reader->NextArrayItem(); - ICHECK(s); - reader->Read(&node->stage_id); - s = reader->NextArrayItem(); - ICHECK(s); - reader->Read(&node->iter_id); - s = reader->NextArrayItem(); - ICHECK(s); - reader->Read(&node->factor); - s = reader->NextArrayItem(); - ICHECK(s); - reader->Read(&node->offset); - data_ = std::move(node); -} - -void StorageAlignStepNode::WriteToRecord(dmlc::JSONWriter* writer) const { - writer->WriteArraySeperator(); - writer->WriteString(record_prefix_str); - writer->WriteArrayItem(stage_id); - writer->WriteArrayItem(iter_id); - writer->WriteArrayItem(factor); - writer->WriteArrayItem(offset); -} - -void StorageAlignStepNode::ApplyToState(State* state) const { - StateNode* pstate = state->CopyOnWrite(); - Stage stage = pstate->stages[stage_id]; - stage.CopyOnWrite()->attrs.storage_offset = offset; - pstate->stages.Set(stage_id, std::move(stage)); -} - -void StorageAlignStepNode::ApplyToSchedule(Array* stages, - StageToAxesMap* stage_to_axes) const { - te::Stage stage = (*stages)[stage_id]; - const Array& axes = (*stage_to_axes)[stage]; - stage.storage_align(axes[iter_id], factor, offset); - stages->Set(stage_id, std::move(stage)); -} - -String StorageAlignStepNode::PrintAsPythonAPI(Array* stages, - StageToAxesMap* stage_to_axes) const { - std::stringstream ss; - const auto& stage = (*stages)[stage_id]; - const auto& op_name = CleanName(stage->op->name); - ss << "s[" << op_name << "].storage_align(" - << CleanName((*stage_to_axes)[stage][iter_id]->var->name_hint, op_name) << ", " << factor - << ", " << offset << ")\n"; - - ApplyToSchedule(stages, stage_to_axes); - return ss.str(); -} - -/********** Steps working on multiple stages **********/ - -/********** Compute At **********/ -ComputeAtStep::ComputeAtStep(int stage_id, int target_stage_id, int target_iter_id) { - auto node = make_object(); - node->stage_id = stage_id; - node->target_stage_id = target_stage_id; - node->target_iter_id = target_iter_id; - data_ = std::move(node); -} - -ComputeAtStep::ComputeAtStep(dmlc::JSONReader* reader) { - auto node = make_object(); - bool s; - s = reader->NextArrayItem(); - ICHECK(s); - reader->Read(&node->stage_id); - s = reader->NextArrayItem(); - ICHECK(s); - reader->Read(&node->target_stage_id); - s = reader->NextArrayItem(); - ICHECK(s); - reader->Read(&node->target_iter_id); - data_ = std::move(node); -} - -void ComputeAtStepNode::WriteToRecord(dmlc::JSONWriter* writer) const { - writer->WriteArraySeperator(); - writer->WriteString(record_prefix_str); - writer->WriteArrayItem(stage_id); - writer->WriteArrayItem(target_stage_id); - writer->WriteArrayItem(target_iter_id); -} -void ComputeAtStepNode::ApplyToState(State* state) const { - const Stage& stage = (*state)->stages[stage_id]; - - // Remove the bound information of each iterator since they may not be accurate after - // compute at - Array new_iters; - for (const Iterator& it : stage->iters) { - new_iters.push_back( - Iterator(it->name, Range(), it->iter_kind, it->annotation, &it->orig_iters)); - } - - StateNode* pstate = state->CopyOnWrite(); - pstate->stages.Set(stage_id, Stage(stage->op, stage->op_type, std::move(new_iters), - ComputeAtKind::kIter, stage->attrs)); - // Update attach map - pstate->attach_map.SetComputeAtIter(stage_id, target_stage_id, target_iter_id); -} - -void ComputeAtStepNode::ApplyToSchedule(Array* stages, - StageToAxesMap* stage_to_axes) const { - te::Stage stage = (*stages)[stage_id]; - const auto& target_stage = (*stages)[target_stage_id]; - const auto& target_axis = (*stage_to_axes)[target_stage][target_iter_id]; - stage.compute_at(target_stage, target_axis); - - stages->Set(stage_id, std::move(stage)); -} - -String ComputeAtStepNode::PrintAsPythonAPI(Array* stages, - StageToAxesMap* stage_to_axes) const { - std::stringstream ss; - const auto& stage = (*stages)[stage_id]; - const auto& target_stage = (*stages)[target_stage_id]; - const auto& op_name = CleanName(stage->op->name); - const auto& target_op_name = CleanName(target_stage->op->name); - ss << "s[" << op_name << "].compute_at(s[" << target_op_name << "], " - << CleanName((*stage_to_axes)[target_stage][target_iter_id]->var->name_hint, target_op_name) - << ")\n"; - ApplyToSchedule(stages, stage_to_axes); - return ss.str(); -} - -/********** Compute Inline **********/ -ComputeInlineStep::ComputeInlineStep(int stage_id) { - auto node = make_object(); - node->stage_id = stage_id; - data_ = std::move(node); -} - -ComputeInlineStep::ComputeInlineStep(dmlc::JSONReader* reader) { - auto node = make_object(); - bool s; - s = reader->NextArrayItem(); - ICHECK(s); - reader->Read(&node->stage_id); - data_ = std::move(node); -} - -void ComputeInlineStepNode::WriteToRecord(dmlc::JSONWriter* writer) const { - writer->WriteArraySeperator(); - writer->WriteString(record_prefix_str); - writer->WriteArrayItem(stage_id); -} - -void ComputeInlineStepNode::ApplyToState(State* state) const { - const Stage& stage = (*state)->stages[stage_id]; - - // Check the validity of compute_inline - for (size_t i = 0; i < stage->iters.size(); ++i) { - ICHECK_EQ((*state)->attach_map->iter_to_attached_stages.count(std::make_pair(stage_id, i)), 0) - << "Invalid compute_inline: There are some other stages that are attached to the " - << "target stage"; - } - - StateNode* pstate = state->CopyOnWrite(); - auto new_stage = pstate->stages[stage_id]; - new_stage.CopyOnWrite()->compute_at = ComputeAtKind::kInlined; - pstate->stages.Set(stage_id, std::move(new_stage)); - // Update attach map - pstate->attach_map.DeleteStage(stage_id); -} - -void ComputeInlineStepNode::ApplyToSchedule(Array* stages, - StageToAxesMap* stage_to_axes) const { - auto stage = (*stages)[stage_id]; - stage.compute_inline(); - stages->Set(stage_id, std::move(stage)); -} - -String ComputeInlineStepNode::PrintAsPythonAPI(Array* stages, - StageToAxesMap* stage_to_axes) const { - std::stringstream ss; - const auto& stage = (*stages)[stage_id]; - ss << "s[" << CleanName(stage->op->name) << "].compute_inline()\n"; - ApplyToSchedule(stages, stage_to_axes); - return ss.str(); -} - -/********** Compute Root **********/ -ComputeRootStep::ComputeRootStep(int stage_id) { - auto node = make_object(); - node->stage_id = stage_id; - data_ = std::move(node); -} - -ComputeRootStep::ComputeRootStep(dmlc::JSONReader* reader) { - auto node = make_object(); - bool s; - s = reader->NextArrayItem(); - ICHECK(s); - reader->Read(&node->stage_id); - data_ = std::move(node); -} - -void ComputeRootStepNode::WriteToRecord(dmlc::JSONWriter* writer) const { - writer->WriteArraySeperator(); - writer->WriteString(record_prefix_str); - writer->WriteArrayItem(stage_id); -} - -void ComputeRootStepNode::ApplyToState(State* state) const { - const Stage& stage = (*state)->stages[stage_id]; - - // Remove the bound information of each iterator since they may not be accurate after - // compute root - Array new_iters; - for (const Iterator& it : stage->iters) { - new_iters.push_back( - Iterator(it->name, Range(), it->iter_kind, it->annotation, &it->orig_iters)); - } - - StateNode* pstate = state->CopyOnWrite(); - pstate->stages.Set(stage_id, Stage(stage->op, stage->op_type, std::move(new_iters), - ComputeAtKind::kRoot, stage->attrs)); - // Update attach map - pstate->attach_map.DeleteStage(stage_id); -} - -void ComputeRootStepNode::ApplyToSchedule(Array* stages, - StageToAxesMap* stage_to_axes) const { - auto stage = (*stages)[stage_id]; - stage.compute_root(); - stages->Set(stage_id, std::move(stage)); -} - -String ComputeRootStepNode::PrintAsPythonAPI(Array* stages, - StageToAxesMap* stage_to_axes) const { - std::stringstream ss; - const auto& stage = (*stages)[stage_id]; - ss << "s[" << CleanName(stage->op->name) << "].compute_root()\n"; - ApplyToSchedule(stages, stage_to_axes); - return ss.str(); -} - -/********** Steps adding new stages **********/ - -/*! - * \brief Common part for steps that add new stages(e.g. CacheReadStep, CacheWriteStep, - * RfactorStep). This will return all steps that can change the number of stages in a ComputeDAG, - * and stop by the current step. - */ -Array GetFormerStageModifiableSteps(Step current_step, const Array& transform_steps) { - Array ret_steps; - for (size_t i = 0; i < transform_steps.size(); ++i) { - const Step& step = transform_steps[i]; - if (step->IsInstance() || step->IsInstance()) { - ret_steps.push_back(step); - } else if (step->IsInstance()) { - // add FuseStepNode required by rfactor - if (i >= 2 && transform_steps[i - 2]->IsInstance()) { - const Step& fuse_step = transform_steps[i - 2]; - if (fuse_step->stage_id == step->stage_id) { - ret_steps.push_back(fuse_step); - } - } - // add SplitStepNode required by rfactor - ICHECK_GE(i, 1); - ICHECK(transform_steps[i - 1]->IsInstance()); - const Step& split_step = transform_steps[i - 1]; - ICHECK_EQ(split_step->stage_id, step->stage_id); - ret_steps.push_back(split_step); - // add RfactorStepNode - ret_steps.push_back(step); - } - // A state may have multiple stage modifiable steps, stop by the current step to avoid - // replaying excess steps - if (step.same_as(current_step)) { - break; - } - } - return ret_steps; -} - -/********** Cache Read **********/ -CacheReadStep::CacheReadStep(int stage_id, String scope_name, - const Array& reader_stage_ids) { - auto node = make_object(); - node->stage_id = stage_id; - node->scope_name = std::move(scope_name); - node->reader_stage_ids = reader_stage_ids; - data_ = std::move(node); -} - -CacheReadStep::CacheReadStep(dmlc::JSONReader* reader) { - auto node = make_object(); - bool s; - s = reader->NextArrayItem(); - ICHECK(s); - reader->Read(&node->stage_id); - s = reader->NextArrayItem(); - ICHECK(s); - std::string string_value; - reader->Read(&string_value); - node->scope_name = std::move(string_value); - s = reader->NextArrayItem(); - ICHECK(s); - reader->Read(&node->reader_stage_ids); - data_ = std::move(node); -} - -void CacheReadStepNode::WriteToRecord(dmlc::JSONWriter* writer) const { - writer->WriteArraySeperator(); - writer->WriteString(record_prefix_str); - writer->WriteArrayItem(stage_id); - writer->WriteArraySeperator(); - writer->WriteString(scope_name); - writer->WriteArrayItem(reader_stage_ids); -} - -int CacheReadStepNode::ApplyToState(State* state, const ComputeDAG& dag) const { - StateNode* pstate = state->CopyOnWrite(); - const ComputeDAG& current_compute_dag = dag.ReplayAndGetDAG( - GetFormerStageModifiableSteps(GetRef(this), (*state)->transform_steps)); - - // target_stage -> target_stage + target_store - // Update the op of the target stage, insert a new cache read stage behind, update the op of - // later stages, then update the stage_id mapping in AttachMap - int added_stage_id = stage_id + 1; - Stage tmp_stage = pstate->stages[stage_id]; - tmp_stage.CopyOnWrite()->op = current_compute_dag->ops[stage_id]; - pstate->stages.Set(stage_id, std::move(tmp_stage)); - pstate->stages.insert(pstate->stages.begin() + added_stage_id, - Stage(current_compute_dag->ops[added_stage_id])); - for (size_t i = added_stage_id + 1; i < pstate->stages.size(); ++i) { - tmp_stage = pstate->stages[i]; - tmp_stage.CopyOnWrite()->op = current_compute_dag->ops[i]; - pstate->stages.Set(i, std::move(tmp_stage)); - } - pstate->attach_map = pstate->attach_map.ApplyStageIdOffset(added_stage_id); - pstate->current_compute_dag = std::move(current_compute_dag); - - return added_stage_id; -} - -te::Tensor CacheReadStepNode::ApplyToSchedule(Array* stages, - StageToAxesMap* stage_to_axes, - te::Schedule* schedule) const { - const te::Stage& stage = (*stages)[stage_id]; - Array readers; - for (const auto& i : reader_stage_ids) { - readers.push_back((*stages)[i.IntValue()]->origin_op); - } - auto out = schedule->cache_read(stage->origin_op.output(0), scope_name, readers); - - const auto& new_stage = (*schedule)[out->op]; - UpdateStageToAxesMap(new_stage, stage_to_axes); - stages->insert(stages->begin() + stage_id + 1, new_stage); - - return out; -} - -String CacheReadStepNode::PrintAsPythonAPI(Array* stages, StageToAxesMap* stage_to_axes, - te::Schedule* schedule) const { - std::stringstream ss; - // Since the original stage will be changed after schedule apply, keep a copy here - // These information will be used to print Python API string later - auto stage = (*stages)[stage_id]; - Array reader_stages; - for (size_t i = 0; i < reader_stage_ids.size(); ++i) { - reader_stages.push_back((*stages)[reader_stage_ids[i].IntValue()]); - } - auto out = ApplyToSchedule(stages, stage_to_axes, schedule); - - const auto& op_name = CleanName(out->op->name); - ss << op_name << " = " - << "s.cache_read(" << CleanName(stage->op->name) << ", \"" << scope_name << "\", [" - << CleanName(reader_stages[0]->op->name); - for (size_t i = 1; i < reader_stage_ids.size(); ++i) { - ss << ", " << CleanName(reader_stages[i]->op->name); - } - ss << "])\n"; - - // Print the iterators of the new added stage - const auto& iters = out->op->root_iter_vars(); - for (size_t i = 0; i < iters.size(); ++i) { - ss << CleanName(iters[i]->var->name_hint, op_name); - if (i != iters.size() - 1) { - ss << ", "; - } - } - ss << " = " - << "tuple(" << op_name << ".op.axis)\n"; - - return ss.str(); -} - -/********** Cache Write **********/ -CacheWriteStep::CacheWriteStep(int stage_id, String scope_name) { - auto node = make_object(); - node->stage_id = stage_id; - node->scope_name = std::move(scope_name); - data_ = std::move(node); -} - -CacheWriteStep::CacheWriteStep(dmlc::JSONReader* reader) { - auto node = make_object(); - bool s; - s = reader->NextArrayItem(); - ICHECK(s); - reader->Read(&node->stage_id); - s = reader->NextArrayItem(); - ICHECK(s); - std::string string_value; - reader->Read(&string_value); - node->scope_name = std::move(string_value); - data_ = std::move(node); -} - -void CacheWriteStepNode::WriteToRecord(dmlc::JSONWriter* writer) const { - writer->WriteArraySeperator(); - writer->WriteString(record_prefix_str); - writer->WriteArrayItem(stage_id); - writer->WriteArraySeperator(); - writer->WriteString(scope_name); -} - -int CacheWriteStepNode::ApplyToState(State* state, const ComputeDAG& dag) const { - StateNode* pstate = state->CopyOnWrite(); - int last_dag_op_size = pstate->current_compute_dag - ? pstate->current_compute_dag.value().as()->ops.size() - : dag->ops.size(); - const ComputeDAG& current_compute_dag = dag.ReplayAndGetDAG( - GetFormerStageModifiableSteps(GetRef(this), (*state)->transform_steps)); - int added_ops = current_compute_dag->ops.size() - last_dag_op_size; - // TODO(jcf94): Update this check to equal after fixing the cache write bug in TVM - ICHECK_GE(added_ops, 1); - - // target_stage -> cache_write_stage + target_stage - // Assume no step has been applied to the target stage before cache write. - // Insert a new cache write stage ahead, update the op of the target stage and later stages, then - // update the stage_id mapping in AttachMap - pstate->stages.insert(pstate->stages.begin() + stage_id, - Stage(current_compute_dag->ops[stage_id])); - pstate->stages.Set(stage_id + 1, Stage(current_compute_dag->ops[stage_id + 1])); - int next_stage_id = stage_id + 2; - // TODO(jc94): Fix the cache write bug in TVM and remove added_op == 2 support. - // TVM's cache_write has a bug with multi outputs. See - // `tests/python/auto_scheduler/test_auto_scheduler_loop_state.py::test_cache_read_write` test - // for more details - if (added_ops == 2) { - pstate->stages.insert(pstate->stages.begin() + next_stage_id, - Stage(current_compute_dag->ops[next_stage_id])); - next_stage_id++; - } else if (added_ops > 2) { - LOG(ERROR) << "Unexpected behavior of CacheWrite."; - } - for (size_t i = next_stage_id; i < current_compute_dag->ops.size(); ++i) { - Stage tmp_stage = pstate->stages[i]; - tmp_stage.CopyOnWrite()->op = current_compute_dag->ops[i]; - pstate->stages.Set(i, std::move(tmp_stage)); - } - pstate->attach_map = pstate->attach_map.ApplyStageIdOffset(stage_id, added_ops); - pstate->current_compute_dag = std::move(current_compute_dag); - - return stage_id; -} - -Array CacheWriteStepNode::ApplyToSchedule(Array* stages, - StageToAxesMap* stage_to_axes, - te::Schedule* schedule) const { - const te::Stage& stage = (*stages)[stage_id]; - Array tensor_array; - // If the target stage has multi outputs, TVM requires to cache_write - // all of them or schedule.cache_write will raise an error - for (auto i = 0; i < stage->op->num_outputs(); ++i) { - tensor_array.push_back(stage->origin_op.output(i)); - } - auto outs = schedule->cache_write(tensor_array, scope_name); - - UpdateStageToAxesMap(stage, stage_to_axes); - // Even if there is multi outputs, TVM schedule only generate one - // new stage - const auto& new_stage = (*schedule)[outs[0]->op]; - UpdateStageToAxesMap(new_stage, stage_to_axes); - stages->insert(stages->begin() + stage_id, new_stage); - - return outs; -} - -String CacheWriteStepNode::PrintAsPythonAPI(Array* stages, StageToAxesMap* stage_to_axes, - te::Schedule* schedule) const { - std::stringstream ss; - // Since the original stage will be changed after schedule apply, keep a copy here - // These information will be used to print Python API string later - te::Stage stage = (*stages)[stage_id]; - auto outs = ApplyToSchedule(stages, stage_to_axes, schedule); - - for (size_t i = 0; i < outs.size(); ++i) { - ss << CleanName(outs[i]->op->name) << ", "; - } - ss << "= " - << "s.cache_write([" << CleanName(stage->op.output(0)->op->name); - for (auto i = 1; i < stage->op->num_outputs(); ++i) { - ss << ", " << CleanName(stage->op.output(i)->op->name); - } - ss << "], \"" << scope_name << "\")\n"; - - // Print the iterators of the new added stage - for (const auto& out : outs) { - const auto& iters = out->op->root_iter_vars(); - const auto& op_name = CleanName(out->op->name); - for (size_t i = 0; i < iters.size(); ++i) { - ss << CleanName(iters[i]->var->name_hint, op_name); - if (i != iters.size() - 1) { - ss << ", "; - } - } - ss << " = " - << "tuple(" << op_name << ".op.axis)" - << " + " - << "tuple(" << op_name << ".op.reduce_axis)\n"; - } - - return ss.str(); -} - -/********** Rfactor **********/ -RfactorStep::RfactorStep(int stage_id, int iter_id, int factor_iter_id) { - auto node = make_object(); - node->stage_id = stage_id; - node->iter_id = iter_id; - node->factor_iter_id = factor_iter_id; - data_ = std::move(node); -} - -RfactorStep::RfactorStep(dmlc::JSONReader* reader) { - auto node = make_object(); - bool s; - s = reader->NextArrayItem(); - ICHECK(s); - reader->Read(&node->stage_id); - s = reader->NextArrayItem(); - ICHECK(s); - reader->Read(&node->iter_id); - s = reader->NextArrayItem(); - ICHECK(s); - reader->Read(&node->factor_iter_id); - data_ = std::move(node); -} - -void RfactorStepNode::WriteToRecord(dmlc::JSONWriter* writer) const { - writer->WriteArraySeperator(); - writer->WriteString(record_prefix_str); - writer->WriteArrayItem(stage_id); - writer->WriteArrayItem(iter_id); - writer->WriteArrayItem(factor_iter_id); -} - -int RfactorStepNode::ApplyToState(State* state, const ComputeDAG& dag) const { - StateNode* pstate = state->CopyOnWrite(); - const auto& compute_at_type = pstate->stages[stage_id]->compute_at; - const ComputeDAG& current_compute_dag = dag.ReplayAndGetDAG( - GetFormerStageModifiableSteps(GetRef(this), (*state)->transform_steps)); - - // target_stage -> rfactor_compute + target_stage - // Insert a new compute stage, update the target stage and later stage, then update the stage_id - // mapping in AttachMap - pstate->stages.insert(pstate->stages.begin() + stage_id, - Stage(current_compute_dag->ops[stage_id])); - // Maintain the compute_at type of the target stage - Stage target_stage = Stage(current_compute_dag->ops[stage_id + 1]); - target_stage.CopyOnWrite()->compute_at = compute_at_type; - pstate->stages.Set(stage_id + 1, std::move(target_stage)); - for (size_t i = stage_id + 2; i < pstate->stages.size(); ++i) { - Stage stage = pstate->stages[i]; - stage.CopyOnWrite()->op = current_compute_dag->ops[i]; - pstate->stages.Set(i, std::move(stage)); - } - pstate->attach_map = pstate->attach_map.ApplyStageIdOffset(stage_id); - pstate->current_compute_dag = std::move(current_compute_dag); - - return stage_id; -} - -Array RfactorStepNode::ApplyToSchedule(Array* stages, - StageToAxesMap* stage_to_axes, - te::Schedule* schedule) const { - const auto& stage = (*stages)[stage_id]; - const Array& axes = (*stage_to_axes)[stage]; - - const te::Tensor& tensor = stage->origin_op.output(0); - const IterVar& axis = axes[iter_id]; - auto outs = schedule->rfactor(tensor, axis, factor_iter_id); - - UpdateStageToAxesMap(stage, stage_to_axes); - const auto& new_stage = (*schedule)[outs[0]->op]; - UpdateStageToAxesMap(new_stage, stage_to_axes); - stages->insert(stages->begin() + stage_id, new_stage); - - return outs; -} - -String RfactorStepNode::PrintAsPythonAPI(Array* stages, StageToAxesMap* stage_to_axes, - te::Schedule* schedule) const { - std::stringstream ss; - const auto& stage = (*stages)[stage_id]; - - const auto& tensor_name = CleanName(stage->origin_op.output(0)->op->name); - const auto& axis_name = CleanName((*stage_to_axes)[stage][iter_id]->var->name_hint); - - const auto& outs = ApplyToSchedule(stages, stage_to_axes, schedule); - - for (size_t i = 0; i < outs.size(); ++i) { - ss << CleanName(outs[i]->op->name); - if (i != outs.size() - 1) { - ss << ", "; - } - } - ss << " = " - << "s.rfactor(" << tensor_name << ", " << axis_name << ", " << factor_iter_id << ")\n"; - - for (const auto& out : outs) { - const auto& iters = out->op->root_iter_vars(); - const auto& op_name = CleanName(out->op->name); - for (size_t i = 0; i < iters.size(); ++i) { - ss << CleanName(iters[i]->var->name_hint, op_name); - if (i != iters.size() - 1) { - ss << ", "; - } - } - ss << " = " - << "tuple(" << op_name << ".op.axis)" - << " + " - << "tuple(" << op_name << ".op.reduce_axis)\n"; - } - - const auto& output = (*stages)[stage_id + 1]->op.output(0); - const auto& iters = output->op->root_iter_vars(); - const auto& op_name = CleanName(output->op->name); - for (size_t i = 0; i < iters.size(); ++i) { - ss << CleanName(iters[i]->var->name_hint, op_name); - if (i != iters.size() - 1) { - ss << ", "; - } - } - ss << " = " - << "tuple(s[" << op_name << "].op.axis)" - << " + " - << "tuple(s[" << op_name << "].op.reduce_axis)\n"; - - return ss.str(); -} - -} // namespace auto_scheduler -} // namespace tvm diff --git a/src/auto_scheduler/utils.cc b/src/auto_scheduler/utils.cc deleted file mode 100755 index 68f503836cfb..000000000000 --- a/src/auto_scheduler/utils.cc +++ /dev/null @@ -1,36 +0,0 @@ -/* - * 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 auto_scheduler/utils.cc - * \brief Common utilities. - */ - -#include "utils.h" - -namespace tvm { -namespace auto_scheduler { - -NullStream& NullStream::Global() { - static NullStream stream; - return stream; -} - -} // namespace auto_scheduler -} // namespace tvm diff --git a/src/auto_scheduler/utils.h b/src/auto_scheduler/utils.h deleted file mode 100755 index f55cad00e4cc..000000000000 --- a/src/auto_scheduler/utils.h +++ /dev/null @@ -1,303 +0,0 @@ -/* - * 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 auto_scheduler/utils.h - * \brief Common utilities. - */ - -#ifndef TVM_AUTO_SCHEDULER_UTILS_H_ -#define TVM_AUTO_SCHEDULER_UTILS_H_ - -#include -#include - -#include -#include -#include -#include -#include -#include -#include -#include -#include -#include -#include -#include - -namespace std { - -/*! \brief Hash function for std::pair */ -template -struct hash> { - std::size_t operator()(const std::pair& k) const { - return ::dmlc::HashCombine(std::hash()(k.first), std::hash()(k.second)); - } -}; - -/*! \brief Hash function for std::tuple */ -template -struct hash> { - std::size_t operator()(const std::tuple& k) const { - return ::dmlc::HashCombine( - ::dmlc::HashCombine(std::hash()(std::get<0>(k)), std::hash()(std::get<1>(k))), - std::hash()(std::get<2>(k))); - } -}; - -} // namespace std - -namespace tvm { -namespace auto_scheduler { - -/********** Utilities for Array, std::vector, std::string **********/ -/*! \brief Get the first appearance index of elements in an Array */ -template -inline void GetIndices(const Array& array, const Array& to_locate, Array* indices) { - for (const auto& v : to_locate) { - auto it = std::find(array.begin(), array.end(), v); - if (it != array.end()) { - indices->push_back(it - array.begin()); - } else { - LOG(FATAL) << "Cannot find the item"; - } - } -} - -/*! \brief Get the first appearance index of an element in an Array */ -template -inline int GetIndex(const Array& array, const T& to_locate) { - for (size_t i = 0; i < array.size(); ++i) { - if (array[i] == to_locate) { - return i; - } - } - LOG(FATAL) << "Cannot find the item"; -} - -/*! \brief Delete the item in a std::vector if it exists. */ -template -inline void FindAndDeleteItem(std::vector* array, const T& to_delete) { - auto iter = std::find(array->begin(), array->end(), to_delete); - if (iter != array->end()) { - array->erase(iter); - } -} - -/*! \brief Compute the product of all elements in a vector */ -inline int64_t ElementProduct(const std::vector& array) { - int64_t ret = 1; - for (auto x : array) { - ret *= x; - } - return ret; -} - -/*! \brief Move elements from multiple vectors to one vector */ -template -std::vector& ConcatenateMove(std::vector* out, std::vector* in) { - out->insert(out->end(), std::make_move_iterator(in->begin()), std::make_move_iterator(in->end())); - return *out; -} - -/*! \brief Move elements from multiple vectors to one vector */ -template -std::vector& ConcatenateMove(std::vector* out, std::vector* first, Args... args) { - ConcatenateMove(out, first); - ConcatenateMove(out, args...); - return *out; -} - -/*! \brief Get a random permutation of integers [0, n-1] */ -template -void RandomPermutation(int n, std::vector* out, G* gen) { - out->assign(n, 0); - std::iota(out->begin(), out->end(), 0); - std::shuffle(out->begin(), out->end(), *gen); -} - -/*! \brief Replace a sub-string to another sub-string in a string */ -inline void StrReplace(std::string* base, const std::string& from, const std::string& to) { - auto pos = base->find(from); - while (pos != std::string::npos) { - base->replace(pos, from.size(), to); - pos = base->find(from, pos + to.size()); - } -} - -/*! \brief Return whether two int arrays are elementwise-equal */ -inline bool IntArrayEqual(const Array& arr1, const Array& arr2) { - if (arr1.size() != arr2.size()) { - return false; - } - - for (size_t i = 0; i < arr1.size(); ++i) { - auto int1 = arr1[i].as(); - auto int2 = arr2[i].as(); - ICHECK(int1 != nullptr); - ICHECK(int2 != nullptr); - if (int1->value != int2->value) { - return false; - } - } - return true; -} - -/********** Utilities for TVM Containers / ByteArray **********/ -/*! \brief Compute mean of a FloatImm array */ -inline double FloatArrayMean(const Array& float_array) { - double sum = 0; - if (float_array.empty()) { - return 0.0; - } - - for (const auto& x : float_array) { - auto floatimm = x.as(); - ICHECK(floatimm != nullptr); - sum += floatimm->value; - } - return sum / float_array.size(); -} - -/*! \brief Return whether a string starts with another substring */ -inline bool StrStartsWith(const String& a, const String& b) { - if (b.size() > a.size()) return false; - return std::equal(a.c_str(), a.c_str() + b.size(), b.c_str()); -} - -/*! \brief Return whether a string ends with another substring */ -inline bool StrEndsWith(const String& a, const String& b) { - if (b.size() > a.size()) return false; - return std::equal(a.c_str() + a.size() - b.size(), a.c_str() + a.size(), b.c_str()); -} - -/********** Other Utilities **********/ -/*! \brief Get an int value from an Expr */ -inline int64_t GetIntImm(const PrimExpr& expr) { - auto pint = expr.as(); - if (pint == nullptr) { - return 1; - } - return pint->value; -} - -/*! \brief Compute the product of the lengths of axes */ -inline int64_t AxisLengthProd(const Array& axes) { - int64_t ret = 1.0; - for (const auto& x : axes) { - if (const IntImmNode* imm = x->dom->extent.as()) { - ret *= imm->value; - } else { - return -1.0; - } - } - return ret; -} - -/*! - * \brief Clean the name of an iterator or an op to make it valid in python code. - * \param str The original name. - * \param prefix The name prefix to differentiate the same name (e.g., the same iterator names). - * \return The cleaned name. - */ -inline std::string CleanName(const std::string& str, const std::string& prefix = "") { - std::string ret = str; - StrReplace(&ret, ".", "_"); - StrReplace(&ret, "@", "_"); - StrReplace(&ret, "outer", "o"); - StrReplace(&ret, "inner", "i"); - if (prefix != "") { - return prefix + "_" + ret; - } - return ret; -} - -/*! \brief An empty output stream */ -class NullStream : public std::ostream { - public: - NullStream() : std::ostream(nullptr) {} - NullStream(const NullStream&) : std::ostream(nullptr) {} - static NullStream& Global(); -}; - -template -NullStream& operator<<(NullStream& os, const T& value) { - return os; -} - -/*! \brief Get std cout with verbose control */ -inline std::ostream& StdCout(int verbose, int setting = 1) { - return verbose >= setting ? std::cout : NullStream::Global(); -} - -/*! \brief Print multiple chars */ -inline std::string Chars(const char& str, int times) { - std::stringstream ret; - for (int i = 0; i < times; ++i) { - ret << str; - } - return ret.str(); -} - -/*! \brief Print the time elapsed */ -inline void PrintTimeElapsed(std::chrono::time_point t_begin, - const std::string& info, int verbose) { - double duration = std::chrono::duration_cast>( - std::chrono::high_resolution_clock::now() - t_begin) - .count(); - StdCout(verbose) << "Time elapsed for " << info << ": " << std::fixed << std::setprecision(2) - << duration << " s" << std::endl; -} - -/*! - * \brief Parse shape and axis names from layout string - */ -inline void ParseKernelLayout(const String& layout, Array* shape, - std::vector* axes) { - int32_t factor = 0; - std::string axis = ""; - for (char c : std::string(layout)) { - if (c >= 'A' && c <= 'z') { - axis += c; - if (factor != 0) { - shape->push_back(factor); - factor = 0; - } - } else if (c >= '0' && c <= '9') { - factor = factor * 10 + c - '0'; - if (!axis.empty()) { - axes->push_back(axis); - axis = ""; - } - } else { - LOG(FATAL) << "Invalid layout " << layout; - } - } - if (!axis.empty()) { - axes->push_back(axis); - } -} - -/*! \brief Get the base name before '_' of an axis */ -inline std::string AxisBaseName(const std::string& str) { return str.substr(0, str.rfind("_")); } - -} // namespace auto_scheduler -} // namespace tvm - -#endif // TVM_AUTO_SCHEDULER_UTILS_H_ diff --git a/src/autotvm/feature_visitor.cc b/src/autotvm/feature_visitor.cc deleted file mode 100644 index a7ae9fc56830..000000000000 --- a/src/autotvm/feature_visitor.cc +++ /dev/null @@ -1,115 +0,0 @@ -/* - * 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 feature_visitor.cc - * \brief Base class for feature extractor. - * These features are used for machine learning cost model - */ - -#include "feature_visitor.h" - -namespace tvm { -namespace autotvm { - -// for loop -void FeatureVisitor::VisitStmt_(const ForNode* op) { - const auto* extent = op->extent.as(); - int64_t loop_extent = -1; - if (extent != nullptr) loop_extent = extent->value; - AnnotationType ann = kSerial; - switch (op->kind) { - case ForKind ::kParallel: - ann = kParallel; - break; - case ForKind::kUnrolled: - ann = kUnrolled; - break; - case ForKind::kVectorized: - ann = kVectorized; - break; - case ForKind::kSerial: - ann = kSerial; - break; - case ForKind::kThreadBinding: - LOG(FATAL) << "Loop ThreadBinding is reserved for future used and " - << "not yet supported in TIR"; - break; - } - - if (EnterItervar_(op->loop_var, loop_extent, ann)) { - StmtExprVisitor::VisitStmt_(op); - ExitItervar_(); - } -} - -// parallel axis, virtual thread -void FeatureVisitor::VisitStmt_(const AttrStmtNode* op) { - if (op->attr_key == tir::attr::thread_extent || op->attr_key == tir::attr::virtual_thread) { - Var var = op->node.as()->var; - const auto* extent = op->value.as(); - ICHECK(extent); - - std::string name = var.get()->name_hint; - AnnotationType ann = kParallel; - if (op->attr_key == tir::attr::thread_extent) { - if (name == "blockIdx.x") - ann = kBlockX; - else if (name == "blockIdx.y") - ann = kBlockY; - else if (name == "blockIdx.z") - ann = kBlockZ; - else if (name == "threadIdx.x") - ann = kThreadX; - else if (name == "threadIdx.y") - ann = kThreadY; - else if (name == "threadIdx.z") - ann = kThreadZ; - else - LOG(FATAL) << "invalid thread itervar " + name; - } else { - ann = kVirtualThread; - } - - if (EnterItervar_(var, extent->value, ann)) { - StmtExprVisitor::VisitStmt_(op); - ExitItervar_(); - } - } else { - StmtExprVisitor::VisitStmt_(op); - } -} - -// memory access -void FeatureVisitor::VisitExpr_(const BufferLoadNode* op) { - ICHECK_EQ(op->indices.size(), 1) << "FeatureVisitor can only be used on flattened buffers"; - EnterMem_(op->buffer->data, op->indices[0]); - StmtExprVisitor::VisitExpr_(op); - ExitMem_(); -} - -void FeatureVisitor::VisitStmt_(const BufferStoreNode* op) { - ICHECK_EQ(op->indices.size(), 1) << "FeatureVisitor can only be used on flattened buffers"; - EnterMem_(op->buffer->data, op->indices[0]); - StmtExprVisitor::VisitStmt_(op); - ExitMem_(); -} - -} // namespace autotvm -} // namespace tvm diff --git a/src/autotvm/feature_visitor.h b/src/autotvm/feature_visitor.h deleted file mode 100644 index 3d34882c77db..000000000000 --- a/src/autotvm/feature_visitor.h +++ /dev/null @@ -1,99 +0,0 @@ -/* - * 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 feature_visitor.h - * \brief Base class for feature extractor. - * These features are used for machine learning cost model - */ - -#ifndef TVM_AUTOTVM_FEATURE_VISITOR_H_ -#define TVM_AUTOTVM_FEATURE_VISITOR_H_ - -#include -#include -#include - -#include - -namespace tvm { -namespace autotvm { - -using namespace tvm::tir; - -/*! - * \brief Type of for loop, used as one-hot encoding in features - */ -enum AnnotationType { - kBlockX, - kBlockY, - kBlockZ, - kThreadX, - kThreadY, - kThreadZ, - kUnrolled, - kVectorized, - kParallel, - kSerial, - kVirtualThread, - kNum, -}; - -/*! - * \brief A base class for feature extractor, used for processing - * for loop and memory access in the IR - */ -class FeatureVisitor : public StmtExprVisitor { - public: - // for loop - void VisitStmt_(const ForNode* op) final; - void VisitStmt_(const AttrStmtNode* op) final; - - // memory access - void VisitExpr_(const BufferLoadNode* op) final; - void VisitStmt_(const BufferStoreNode* op) final; - - using StmtExprVisitor::VisitExpr_; - using StmtExprVisitor::VisitStmt_; - - protected: - /*! - * \brief Enter a for loop node - * \param var The expression to be printed. - * \param length The output stream - * \param ann_type The type for the for loop - * \return skip Whether skip this node - */ - virtual bool EnterItervar_(tir::Var var, int64_t length, AnnotationType ann_type) = 0; - /*! \brief Exit a for loop subtree */ - virtual void ExitItervar_() = 0; - /*! - * \brief Enter a memory access node - * \param buffer_var The buffer to access. - * \param index Index expression - */ - virtual void EnterMem_(tir::Var buffer_var, tvm::PrimExpr index) = 0; - /*! \brief Exit a memory access node */ - virtual void ExitMem_() = 0; -}; - -} // namespace autotvm -} // namespace tvm - -#endif // TVM_AUTOTVM_FEATURE_VISITOR_H_ diff --git a/src/autotvm/touch_extractor.cc b/src/autotvm/touch_extractor.cc deleted file mode 100644 index dd3cf88f7bf6..000000000000 --- a/src/autotvm/touch_extractor.cc +++ /dev/null @@ -1,524 +0,0 @@ -/* - * 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 touch_extractor.cc - * \brief Extract feature of touch pattern of axes in lowered IR - */ - -#include "touch_extractor.h" - -#include -#include -#include -#include - -namespace tvm { -namespace autotvm { - -int ParallelLevel(AnnotationType ann) { - switch (ann) { - case kBlockX: - case kBlockY: - case kBlockZ: - return 2; - case kThreadX: - case kThreadY: - case kThreadZ: - case kParallel: - return 1; - default: - return 0; - } -} - -// get touch pattern from index expression -class IndexParser : public ExprVisitor { - public: - void Parse(PrimExpr expr) { - pattern_map.clear(); - this->VisitExpr(expr); - } - - void VisitExpr_(const VarNode* op) final { - // TODO(lmzheng): handle more index types (multiple occurrence) - if (pattern_map.count(op) == 0) { - pattern_map[op] = TouchPattern(); - pattern_map[op].stride = next_stride_; - next_stride_ = 1; - } - } - - void VisitExpr_(const MulNode* op) final { - if (op->a.as()) { - if (const auto stride = op->b.as()) { - next_stride_ = stride->value; - } - } - ExprVisitor::VisitExpr_(op); - } - - std::unordered_map pattern_map; - - private: - int64_t next_stride_ = 1; -}; - -// extract iter vars and their touch pattern from ir -bool TouchExtractor::EnterItervar_(Var var, int64_t length, AnnotationType ann_type) { - // do not insert duplicated occurrences of virtual thread - if (ann_type == kVirtualThread && itervar_map.count(var) != 0) { - skip_stack_size_.push_back(itervar_stack_.size()); - return true; - } else { - itervar_stack_.push_back(var); - topdown_product_ *= length; - - if (itervar_map.count(var) != 0) { - // find two duplicated axes - // these happens when we create tvm.thread_axis("threadIdx.x") once and - // bind it twice. Here we treat them as two axes - // so we create a snapshot for the old one and freeze it - Var old = Var(var.get()->name_hint); - itervar_map.insert({old, itervar_map[var]}); - itervar_map.erase(var); - } - - itervar_map.insert( - {var, ItervarFeature(var, length, static_cast(itervar_stack_.size()), ann_type, - topdown_product_, static_cast(itervar_counter_++))}); - } - - return true; -} - -void TouchExtractor::ExitItervar_() { - if (!skip_stack_size_.empty() && skip_stack_size_.back() == itervar_stack_.size()) { - skip_stack_size_.pop_back(); - return; - } - Var var = itervar_stack_.back(); - - // update count and reuse ratio for upper iter vars (includes self) - for (auto kv : itervar_map[var].touch_feature) { - if (kv.second.stride != 0) { // multiply count - for (auto stack_var : itervar_stack_) { - auto touch_pattern = itervar_map[stack_var].touch_feature.find(kv.first); - ICHECK(touch_pattern != itervar_map[stack_var].touch_feature.end()); - touch_pattern->second.count *= itervar_map[var].length; - } - } else { // multiply reuse ratio - for (auto stack_var : itervar_stack_) { - auto touch_pattern = itervar_map[stack_var].touch_feature.find(kv.first); - ICHECK(touch_pattern != itervar_map[stack_var].touch_feature.end()); - touch_pattern->second.reuse *= itervar_map[var].length; - } - } - } - itervar_stack_.pop_back(); - - int64_t length = itervar_map[var].length; - if (length != 0) topdown_product_ /= length; - int64_t bottomup_product = -1; - for (auto kv : itervar_map[var].touch_feature) { - bottomup_product = std::max(bottomup_product, kv.second.count * kv.second.reuse); - } - - itervar_map[var].bottomup_product = bottomup_product; - - // push base to upper parallel axis - int para_level = ParallelLevel(itervar_map[var].ann); - // if is the separate line of parallel level, push the base to upper parallel level - if (!itervar_stack_.empty() && - ParallelLevel(itervar_map[itervar_stack_.back()].ann) == para_level + 1) { - for (auto kv : itervar_map[var].touch_feature) { - for (auto stack_var : itervar_stack_) { - if (ParallelLevel(itervar_map[stack_var].ann) == para_level + 1) { - auto touch_pattern = itervar_map[stack_var].touch_feature.find(kv.first); - ICHECK(touch_pattern != itervar_map[stack_var].touch_feature.end()); - touch_pattern->second.thread_reuse = -kv.second.reuse; - touch_pattern->second.thread_count = -kv.second.count; - // NOTE: use minus as a flag to denote it is a base, - // indicating it is not the final value - } - } - } - } - - for (auto kv : itervar_map[var].touch_feature) { - if (kv.second.thread_count < 0) { - itervar_map[var].touch_feature[kv.first].thread_count = - kv.second.count / (-kv.second.thread_count); - itervar_map[var].touch_feature[kv.first].thread_reuse = - kv.second.reuse / (-kv.second.thread_reuse); - } - } -} - -void TouchExtractor::EnterMem_(Var buffer_var, PrimExpr index) { - std::string name = buffer_var.get()->name_hint; - TouchedBuffer buf = name + "_" + std::to_string(buffer_counter_[name]++); - - // extract touch pattern from index - IndexParser parser; - parser.Parse(index); - - // push up mem access info - for (auto var : itervar_stack_) { - auto x = parser.pattern_map.find(var.get()); - if (x != parser.pattern_map.end()) { - itervar_map[var].touch_feature[buf] = x->second; - } else { - itervar_map[var].touch_feature[buf] = TouchPattern(); - } - } -} - -void TouchExtractor::ExitMem_() {} - -/*! - * \brief Get axis-based feature for all axes - * \param stmt The statement to be extracted - * \param bool Whether take log for numerical feature - * \param ret_feature The buffer where the return value is stored - * - * \note The format of return value is - * (( - * ('_itervar_', var), - * ('_attr_', length, nest_level, topdown, bottomup, one_hot_annotation), - * ('_arith_', add_ct, mul_ct, div_ct), - * ('data_vec_0', stride, mod, count, reuse, thread_count, thread_reuse), - * ('conv_0', stride, mod, count, reuse, thread_count, thread_reuse), - * ), - * ( - * ('_itervar_', var2), - * ('_attr_', length, nest_level, one_hot_annotation), - * ('_arith_', add_ct, mul_ct, div_ct), - * ('kernel_vec_0', stride, mod, count, reuse, thread_count, thread_reuse), - * ('conv_1', stride, mod, count, reuse, thread_count, thread_reuse), - * )) - * - * Itervars are sorted according to their first occurrence position in IR. - * Buffers touched by an itervar are sorted by their unique names. - * - * \note If you want to flatten these features as the input of your model, - * You can use the faster one GetItervarFeatureFlatten below. - */ -void GetItervarFeature(Stmt stmt, bool take_log, Array>>* ret_feature) { - // extract - TouchExtractor touch_analyzer; - touch_analyzer.Analyze(stmt); - - // sort according to order - std::vector vars; - for (auto kv : touch_analyzer.itervar_map) { - vars.push_back(kv.first); - } - std::sort(vars.begin(), vars.end(), [&](const Var& lhs, const Var& rhs) -> bool { - return touch_analyzer.itervar_map[lhs].order < touch_analyzer.itervar_map[rhs].order; - }); - - // whether take log for numerical feature - std::function trans; - if (take_log) { - trans = [](int64_t x) { - if (x < 0) return -std::log(-x + 1) / std::log(2); - x = x + 1; - return std::log(x) / std::log(2); - }; - } else { - trans = [](int64_t x) { return x; }; - } - - // serialize for front end - for (auto var : vars) { - Array> feature_row; - ItervarFeature& fea = touch_analyzer.itervar_map[var]; - feature_row.push_back(Array{tvm::tir::StringImm("_itervar_"), var}); - - Array attr{ - tvm::tir::StringImm("_attr_"), - FloatImm(DataType::Float(32), trans(fea.length)), - IntImm(DataType::Int(32), fea.nest_level), - FloatImm(DataType::Float(32), trans(fea.topdown_product)), - FloatImm(DataType::Float(32), trans(fea.bottomup_product)), - }; - // one hot annotation - for (int i = 0; i < kNum; i++) { - attr.push_back(i == fea.ann); - } - feature_row.push_back(attr); - - // arithmetic - feature_row.push_back(Array{ - tvm::tir::StringImm("_arith_"), - FloatImm(DataType::Float(32), trans(fea.add_ct)), - FloatImm(DataType::Float(32), trans(fea.mul_ct)), - FloatImm(DataType::Float(32), trans(fea.div_ct)), - }); - - // touch map - std::vector bufs; - for (auto kv : fea.touch_feature) { - bufs.push_back(kv.first); - } - std::sort(bufs.begin(), bufs.end()); - for (auto k : bufs) { - TouchPattern& v = fea.touch_feature[k]; - feature_row.push_back(Array{ - tvm::tir::StringImm(k), - FloatImm(DataType::Float(32), trans(v.stride)), - FloatImm(DataType::Float(32), trans(v.mod)), - FloatImm(DataType::Float(32), trans(v.count)), - FloatImm(DataType::Float(32), trans(v.reuse)), - FloatImm(DataType::Float(32), trans(v.thread_count)), - FloatImm(DataType::Float(32), trans(v.thread_reuse)), - }); - } - - ret_feature->push_back(feature_row); - } -} - -/*! - * \brief Get axis-based feature for all axes and flatten them into a one-dimensional vector. - * \param stmt The statement to be extracted - * \param bool Whether take log for numerical feature - * \param ret_feature The buffer where the return value is stored - * - * \note See GetItervarFeature for more details about the return value. - * This is an optimized version of GetItervarFeature + Flatten. This runs much faster. - */ -void GetItervarFeatureFlatten(Stmt stmt, bool take_log, std::vector* ret_feature) { - // extract touch feature - TouchExtractor touch_analyzer; - touch_analyzer.Analyze(stmt); - - // sort according to order - std::vector vars; - for (auto kv : touch_analyzer.itervar_map) { - vars.push_back(kv.first); - } - std::sort(vars.begin(), vars.end(), [&](const Var& lhs, const Var& rhs) -> bool { - return touch_analyzer.itervar_map[lhs].order < touch_analyzer.itervar_map[rhs].order; - }); - - // whether take log for numerical feature - std::function trans; - if (take_log) { - trans = [](int64_t x) { - if (x < 0) return -std::log(-x + 1) / std::log(2); - x = x + 1; - return std::log(x) / std::log(2); - }; - } else { - trans = [](int64_t x) { return x; }; - } - - // serialize for front end - for (auto var : vars) { - ItervarFeature& fea = touch_analyzer.itervar_map[var]; - - ret_feature->push_back(trans(fea.length)); - ret_feature->push_back(fea.nest_level); - ret_feature->push_back(trans(fea.topdown_product)); - ret_feature->push_back(trans(fea.bottomup_product)); - - // one hot annotation - for (int i = 0; i < kNum; i++) { - ret_feature->push_back(i == fea.ann); - } - - // arithmetic - ret_feature->push_back(trans(fea.add_ct)); - ret_feature->push_back(trans(fea.mul_ct)); - ret_feature->push_back(trans(fea.div_ct)); - - // touch map - std::vector bufs; - for (auto kv : fea.touch_feature) { - bufs.push_back(kv.first); - } - std::sort(bufs.begin(), bufs.end()); - for (auto k : bufs) { - TouchPattern& v = fea.touch_feature[k]; - ret_feature->push_back(trans(v.stride)); - ret_feature->push_back(trans(v.mod)); - ret_feature->push_back(trans(v.count)); - ret_feature->push_back(trans(v.reuse)); - ret_feature->push_back(trans(v.thread_count)); - ret_feature->push_back(trans(v.thread_reuse)); - } - } -} - -/*! - * \brief Get curve sample feature (relation feature) and flatten them into a one-dimensional - * vector. \param stmt The statement to be extracted \param sample_n The number of points used for - * sampling a curve (along one dimension) \param ret_feature The buffer where the return value is - * stored - */ -void GetCurveSampleFeatureFlatten(Stmt stmt, int sample_n, std::vector* ret_feature) { - // extract touch feature - TouchExtractor touch_ext; - touch_ext.Analyze(stmt); - - // sort according to order - std::vector vars; - for (auto kv : touch_ext.itervar_map) { - vars.push_back(kv.first); - } - std::sort(vars.begin(), vars.end(), [&](const Var& lhs, const Var& rhs) -> bool { - return touch_ext.itervar_map[lhs].order < touch_ext.itervar_map[rhs].order; - }); - - int max_depth = 0; - std::map> reuse_curve; - std::map> count_curve; - std::map> topdown_curve; - std::map> bottomup_curve; - std::set innermost_buffers; - std::set added; - - // find maximum depth of loop nest - for (auto var : vars) { - ItervarFeature& fea = touch_ext.itervar_map[var]; - max_depth = std::max(max_depth, fea.nest_level); - } - - // mark inner most buffer - for (auto iter = vars.rbegin(); iter != vars.rend(); iter++) { - auto var = *iter; - ItervarFeature& fea = touch_ext.itervar_map[var]; - if (fea.nest_level == max_depth) { - for (auto kv : fea.touch_feature) { - // delete buffer no (e.g. 'A_0' -> 'A', 'A_1' -> 'A') - std::string raw_name = kv.first.substr(0, kv.first.rfind("_")); - - // delete memory scope (e.g. 'A.local' -> 'A', 'A.shared' -> 'A') - size_t pos = raw_name.find("."); - if (pos < kv.first.size()) raw_name = raw_name.substr(0, pos); - - // If there are multiple innermost buffers that are derived from a same raw buffer - // We only record the last occurrence (note the `iter` is in reverse order) - // e.g. `A.local`, `A.shared` are derived from `A`, if they all occurred at the inner most - // level, we will only record the last occurrence, - if (added.find(raw_name) == added.end()) { - innermost_buffers.insert(kv.first); - added.insert(raw_name); - } - } - } - } - - // pad the first point (zero) for all curves - for (auto buf : innermost_buffers) { - reuse_curve[buf].push_back(0); - count_curve[buf].push_back(0); - topdown_curve[buf].push_back(0); - bottomup_curve[buf].push_back(0); - } - - // extract curves - for (auto var : vars) { - ItervarFeature& fea = touch_ext.itervar_map[var]; - for (auto kv : fea.touch_feature) { - if (innermost_buffers.find(kv.first) != innermost_buffers.end()) { - reuse_curve[kv.first].emplace_back(std::log(kv.second.reuse) / std::log(2)); - count_curve[kv.first].emplace_back(std::log(kv.second.count) / std::log(2)); - topdown_curve[kv.first].emplace_back(std::log(fea.topdown_product) / std::log(2)); - bottomup_curve[kv.first].emplace_back(std::log(fea.bottomup_product) / std::log(2)); - } - } - } - - // sample relation in the curve - auto sample_curve = [&](const std::vector& x, const std::vector& y, - double weight) { - for (int i = 0; i < sample_n; i++) { - double xx = i * weight; - for (int j = static_cast(x.size()) - 1; j >= 0; j--) { - if (xx > x[j] - 1e-6) { - ret_feature->emplace_back(y[j]); - ret_feature->emplace_back(xx - x[j]); - break; - } - } - } - }; - - // serialize to frontend - for (auto k : innermost_buffers) { - std::vector& count = count_curve[k]; - std::vector& reuse = reuse_curve[k]; - std::vector& top_down = topdown_curve[k]; - - std::sort(count.begin(), count.end()); - std::sort(reuse.begin(), reuse.end()); - std::sort(top_down.begin(), top_down.end()); - - sample_curve(count, reuse, 1); - sample_curve(reuse, count, 1); - sample_curve(count, top_down, 1); - sample_curve(top_down, count, 1); - } -} - -// register API for front end -TVM_REGISTER_GLOBAL("autotvm.feature.GetItervarFeature") - .set_body([](TVMArgs args, TVMRetValue* ret) { - Stmt stmt = args[0]; - bool take_log = args[1]; - Array>> ret_feature; - - GetItervarFeature(stmt, take_log, &ret_feature); - - *ret = ret_feature; - }); - -TVM_REGISTER_GLOBAL("autotvm.feature.GetItervarFeatureFlatten") - .set_body([](TVMArgs args, TVMRetValue* ret) { - Stmt stmt = args[0]; - bool take_log = args[1]; - std::vector ret_feature; - - GetItervarFeatureFlatten(stmt, take_log, &ret_feature); - - TVMByteArray arr; - arr.size = sizeof(float) * ret_feature.size(); - arr.data = reinterpret_cast(ret_feature.data()); - *ret = arr; - }); - -TVM_REGISTER_GLOBAL("autotvm.feature.GetCurveSampleFeatureFlatten") - .set_body([](TVMArgs args, TVMRetValue* ret) { - Stmt stmt = args[0]; - int sample_n = args[1]; - std::vector ret_feature; - - GetCurveSampleFeatureFlatten(stmt, sample_n, &ret_feature); - - TVMByteArray arr; - arr.size = sizeof(float) * ret_feature.size(); - arr.data = reinterpret_cast(ret_feature.data()); - *ret = arr; - }); - -} // namespace autotvm -} // namespace tvm diff --git a/src/autotvm/touch_extractor.h b/src/autotvm/touch_extractor.h deleted file mode 100644 index 83260e1e0633..000000000000 --- a/src/autotvm/touch_extractor.h +++ /dev/null @@ -1,144 +0,0 @@ -/* - * 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 touch_extractor.h - * \brief Extract feature of touch pattern of axes in lowered IR - */ - -#ifndef TVM_AUTOTVM_TOUCH_EXTRACTOR_H_ -#define TVM_AUTOTVM_TOUCH_EXTRACTOR_H_ - -#include -#include -#include - -#include -#include -#include -#include -#include -#include - -#include "feature_visitor.h" - -namespace tvm { -namespace autotvm { - -using TouchedBuffer = std::string; - -// touch pattern buf[(stride * var) % mod) + other] -struct TouchPattern { - int64_t stride{0}; - int64_t mod{-1}; // -1 for +inf - - int64_t count{1}; - int64_t reuse{1}; - int64_t thread_count{0}; // count when move thread axis into innermost - int64_t thread_reuse{0}; // reuse ratio move thread axis into innermost -}; - -// all the feature of an iter var -struct ItervarFeature { - ItervarFeature(Var var, int64_t extent, int nest, AnnotationType ann_type, int64_t topdown, - int counter) - : length(extent), nest_level(nest), ann(ann_type), topdown_product(topdown), order(counter) {} - ItervarFeature() {} - - // Axis Attributes - int64_t length; - int nest_level; - AnnotationType ann; // one-hot axis type - int64_t topdown_product; // accumulative product of axis length, in top-down order - int64_t bottomup_product; // accumulative product of axis length, in bottom-up order - // bottomup_product = reuse * count for any touched buffer - - int order; // used for soring axis - - // Arithmetic feature - int add_ct{0}; - int mul_ct{0}; - int div_ct{0}; - - // Memory Touch Feature - std::unordered_map touch_feature; -}; - -// extract iter vars and their touch pattern from ir -class TouchExtractor : public FeatureVisitor { - public: - void Analyze(const Stmt& stmt) { operator()(stmt); } - - // arithmetic stats - void VisitExpr_(const AddNode* op) final { - if (op->dtype.is_float() || op->dtype.is_bfloat16()) { - itervar_map[itervar_stack_.back()].add_ct++; - } - FeatureVisitor::VisitExpr_(op); - } - - void VisitExpr_(const SubNode* op) final { - if (op->dtype.is_float() || op->dtype.is_bfloat16()) { - itervar_map[itervar_stack_.back()].add_ct++; - } - FeatureVisitor::VisitExpr_(op); - } - - void VisitExpr_(const MulNode* op) final { - if (op->dtype.is_float() || op->dtype.is_bfloat16()) { - itervar_map[itervar_stack_.back()].mul_ct++; - } - FeatureVisitor::VisitExpr_(op); - } - - void VisitExpr_(const DivNode* op) final { - if (op->dtype.is_float() || op->dtype.is_bfloat16()) { - itervar_map[itervar_stack_.back()].div_ct++; - } - FeatureVisitor::VisitExpr_(op); - } - - void VisitExpr_(const ModNode* op) final { - if (op->dtype.is_float() || op->dtype.is_bfloat16()) { - itervar_map[itervar_stack_.back()].div_ct++; - } - FeatureVisitor::VisitExpr_(op); - } - - std::unordered_map itervar_map; - - private: - bool EnterItervar_(Var var, int64_t length, AnnotationType ann_type); - void ExitItervar_(); - void EnterMem_(Var buffer_var, PrimExpr index); - void ExitMem_(); - - int64_t topdown_product_{1}; - std::map buffer_counter_; - size_t itervar_counter_{0}; - std::deque itervar_stack_; // use deque instead of stack for indexing - std::deque skip_stack_size_; - - using FeatureVisitor::VisitExpr_; -}; - -} // namespace autotvm -} // namespace tvm - -#endif // TVM_AUTOTVM_TOUCH_EXTRACTOR_H_ diff --git a/src/contrib/hybrid/codegen_hybrid.cc b/src/contrib/hybrid/codegen_hybrid.cc deleted file mode 100644 index d3f50c1c2459..000000000000 --- a/src/contrib/hybrid/codegen_hybrid.cc +++ /dev/null @@ -1,513 +0,0 @@ -/* - * 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 codegen_hybrid.cc - */ -#include "codegen_hybrid.h" - -#include -#include - -#include -#include - -namespace tvm { -namespace contrib { - -using runtime::TVMArgs; -using runtime::TVMRetValue; - -using namespace tir; - -std::string dot_to_underscore(std::string s) { - for (auto& ch : s) - if (ch == '.') ch = '_'; - return s; -} - -std::string CodeGenHybrid::Finish() { return stream.str(); } - -void CodeGenHybrid::PrintType(DataType t, std::ostream& os) { - if (t.is_float()) { - os << "float"; - ICHECK(t.bits() == 16 || t.bits() == 32 || t.bits() == 64); - } else if (t.is_int()) { - os << "int"; - ICHECK(t.bits() == 8 || t.bits() == 16 || t.bits() == 32 || t.bits() == 64); - } else if (t.is_bfloat16()) { - os << "bfloat"; - ICHECK(t.bits() == 16); - } else { - ICHECK(t.is_uint()) << "Unsupported type " << t; - os << "uint"; - ICHECK(t.bits() == 8 || t.bits() == 16 || t.bits() == 32 || t.bits() == 64); - } - os << t.bits(); -} - -void CodeGenHybrid::VisitExpr_(const IntImmNode* op, std::ostream& os) { // NOLINT(*) - os << op->value; -} - -void CodeGenHybrid::VisitExpr_(const FloatImmNode* op, std::ostream& os) { // NOLINT(*) - PrintType(op->dtype, os); - os << "(" << std::setprecision(20) << op->value << ")"; -} -void CodeGenHybrid::VisitExpr_(const StringImmNode* op, std::ostream& os) { // NOLINT(*) - os << "'" << op->value << "'"; -} - -template -inline void PrintBinaryExpr(const T* op, const char* opstr, - std::ostream& os, // NOLINT(*) - CodeGenHybrid* p) { - ICHECK(op->dtype.lanes() == 1) << "vec bin op not implemented"; - if (isalpha(opstr[0])) { - os << opstr << '('; - p->PrintExpr(op->a, os); - os << ", "; - p->PrintExpr(op->b, os); - os << ')'; - } else { - os << '('; - p->PrintExpr(op->a, os); - if (!strcmp(opstr, "&&")) opstr = "and"; - if (!strcmp(opstr, "||")) opstr = "or"; - os << ' ' << opstr << ' '; - p->PrintExpr(op->b, os); - os << ')'; - } -} - -inline void PrintBinaryIntrinsitc(const CallNode* op, const char* opstr, - std::ostream& os, // NOLINT(*) - CodeGenHybrid* p) { - ICHECK(op->dtype.lanes() == 1) << "vec bin intrin not implemented"; - ICHECK_EQ(op->args.size(), 2U); - os << '('; - p->PrintExpr(op->args[0], os); - os << opstr; - p->PrintExpr(op->args[1], os); - os << ')'; -} - -void CodeGenHybrid::VisitExpr_(const CastNode* op, std::ostream& os) { // NOLINT(*) - if (op->dtype == op->value.dtype()) { - PrintExpr(op->value, stream); - } else { - PrintType(op->dtype, os); - os << "("; - PrintExpr(op->value, os); - os << ")"; - } -} - -void CodeGenHybrid::VisitExpr_(const VarNode* op, std::ostream& os) { // NOLINT(*) - os << GetVarID(op); -} -void CodeGenHybrid::VisitExpr_(const AddNode* op, std::ostream& os) { // NOLINT(*) - PrintBinaryExpr(op, "+", os, this); -} -void CodeGenHybrid::VisitExpr_(const SubNode* op, std::ostream& os) { // NOLINT(*) - PrintBinaryExpr(op, "-", os, this); -} -void CodeGenHybrid::VisitExpr_(const MulNode* op, std::ostream& os) { // NOLINT(*) - PrintBinaryExpr(op, "*", os, this); -} - -void CodeGenHybrid::VisitExpr_(const DivNode* op, std::ostream& os) { // NOLINT(*) - if (op->dtype.is_int()) - PrintBinaryExpr(op, "//", os, this); - else - PrintBinaryExpr(op, "/", os, this); -} - -void CodeGenHybrid::VisitExpr_(const FloorDivNode* op, std::ostream& os) { // NOLINT(*) - if (op->dtype.is_int()) - PrintBinaryExpr(op, "//", os, this); - else - PrintBinaryExpr(op, "/", os, this); -} - -void CodeGenHybrid::VisitExpr_(const ModNode* op, std::ostream& os) { // NOLINT(*) - PrintBinaryExpr(op, "%", os, this); -} - -void CodeGenHybrid::VisitExpr_(const FloorModNode* op, std::ostream& os) { // NOLINT(*) - PrintBinaryExpr(op, "%", os, this); -} -void CodeGenHybrid::VisitExpr_(const MinNode* op, std::ostream& os) { // NOLINT(*) - PrintBinaryExpr(op, "min", os, this); -} -void CodeGenHybrid::VisitExpr_(const MaxNode* op, std::ostream& os) { // NOLINT(*) - PrintBinaryExpr(op, "max", os, this); -} -void CodeGenHybrid::VisitExpr_(const EQNode* op, std::ostream& os) { // NOLINT(*) - PrintBinaryExpr(op, "==", os, this); -} -void CodeGenHybrid::VisitExpr_(const NENode* op, std::ostream& os) { // NOLINT(*) - PrintBinaryExpr(op, "!=", os, this); -} -void CodeGenHybrid::VisitExpr_(const LTNode* op, std::ostream& os) { // NOLINT(*) - PrintBinaryExpr(op, "<", os, this); -} -void CodeGenHybrid::VisitExpr_(const LENode* op, std::ostream& os) { // NOLINT(*) - PrintBinaryExpr(op, "<=", os, this); -} -void CodeGenHybrid::VisitExpr_(const GTNode* op, std::ostream& os) { // NOLINT(*) - PrintBinaryExpr(op, ">", os, this); -} -void CodeGenHybrid::VisitExpr_(const GENode* op, std::ostream& os) { // NOLINT(*) - PrintBinaryExpr(op, ">=", os, this); -} -void CodeGenHybrid::VisitExpr_(const AndNode* op, std::ostream& os) { // NOLINT(*) - PrintBinaryExpr(op, "&&", os, this); -} -void CodeGenHybrid::VisitExpr_(const OrNode* op, std::ostream& os) { // NOLINT(*) - PrintBinaryExpr(op, "||", os, this); -} -void CodeGenHybrid::VisitExpr_(const NotNode* op, std::ostream& os) { // NOLINT(*) - os << "not "; - PrintExpr(op->a, os); -} - -void CodeGenHybrid::VisitExpr_(const ProducerLoadNode* op, std::ostream& os) { // NOLINT(*) - auto tensor = Downcast(op->producer); - - os << GetTensorID(tensor); - os << "["; - for (size_t i = 0; i < op->indices.size(); ++i) { - if (i) os << ", "; - std::stringstream idx; - PrintExpr(op->indices[i], idx); - os << idx.str(); - } - os << "]"; -} -void CodeGenHybrid::VisitExpr_(const CallNode* op, std::ostream& os) { // NOLINT(*) - if (op->op.same_as(builtin::bitwise_and())) { - PrintBinaryIntrinsitc(op, "&", os, this); - } else if (op->op.same_as(builtin::bitwise_xor())) { - PrintBinaryIntrinsitc(op, "^", os, this); - } else if (op->op.same_as(builtin::bitwise_or())) { - PrintBinaryIntrinsitc(op, "|", os, this); - } else if (op->op.same_as(builtin::shift_left())) { - PrintBinaryIntrinsitc(op, "<<", os, this); - } else if (op->op.same_as(builtin::shift_right())) { - PrintBinaryIntrinsitc(op, ">>", os, this); - } else if (op->op.same_as(builtin::bitwise_not())) { - ICHECK_EQ(op->args.size(), 1U); - os << "(~"; - PrintExpr(op->args[0], os); - os << ')'; - } else if (op->op.same_as(builtin::if_then_else())) { - PrintExpr(op->args[1], os); - os << " if "; - PrintExpr(op->args[0], os); - os << " else "; - PrintExpr(op->args[2], os); - } else if (op->op.same_as(builtin::call_pure_extern()) || - op->op.same_as(builtin::call_extern())) { - StringImm fname = Downcast(op->args[0]); - os << fname << "("; - for (size_t i = 1; i < op->args.size(); i++) { - PrintExpr(op->args[i], os); - if (i < op->args.size() - 1) { - os << ", "; - } - } - os << ")"; - } else { - auto* ptr_op = op->op.as(); - ICHECK(ptr_op != nullptr); - std::string name = ptr_op->name; - ICHECK_EQ(name.compare(0, 4, "tir."), 0); - os << name.substr(4) << "("; - for (size_t i = 0; i < op->args.size(); i++) { - PrintExpr(op->args[i], os); - if (i < op->args.size() - 1) { - os << ", "; - } - } - os << ")"; - } -} - -void CodeGenHybrid::VisitExpr_(const BufferLoadNode* op, std::ostream& os) { // NOLINT(*) - LOG(FATAL) << "Phase 0 has no BufferLoad(s)!"; -} - -void CodeGenHybrid::VisitStmt_(const BufferStoreNode* op) { - LOG(FATAL) << "Phase 0 has no BufferStore(s)!"; -} - -void CodeGenHybrid::VisitExpr_(const LetNode* op, std::ostream& os) { // NOLINT(*) - LOG(FATAL) << "Phase 0 has no Let(s)!"; -} - -void CodeGenHybrid::VisitStmt_(const AllocateNode* op) { - LOG(FATAL) << "Phase 0 has no Allocate(s)!"; -} - -void CodeGenHybrid::VisitExpr_(const RampNode* op, std::ostream& os) { // NOLINT(*) - LOG(FATAL) << "Ramp to be supported yet"; -} - -void CodeGenHybrid::VisitExpr_(const BroadcastNode* op, std::ostream& os) { // NOLINT(*) - LOG(FATAL) << "Broadcast: not supported "; -} - -void CodeGenHybrid::VisitExpr_(const SelectNode* op, std::ostream& os) { // NOLINT(*) - PrintExpr(op->true_value, os); - os << " if "; - PrintExpr(op->condition, os); - os << " else "; - PrintExpr(op->false_value, os); - os << "\n"; -} - -void CodeGenHybrid::VisitStmt_(const LetStmtNode* op) { - std::string value = PrintExpr(op->value); - stream << GetVarID(op->var.get()) << " = " << value << ";\n"; - PrintStmt(op->body); -} - -void CodeGenHybrid::VisitStmt_(const AttrStmtNode* op) { - if (op->attr_key == tir::attr::thread_extent) { - auto iter_var = op->node.as(); - ICHECK(iter_var); - binds_[iter_var->var.get()] = dot_to_underscore(iter_var->var->name_hint); - PrintIndent(); - stream << "for " << binds_[iter_var->var.get()] << " in bind('" << iter_var->var->name_hint - << "', "; - PrintExpr(op->value, stream); - stream << "):\n"; - indent_ += tab_; - PrintStmt(op->body); - indent_ -= tab_; - } else { - // For now we ignore the unsupported AttrStmt - PrintStmt(op->body); - } -} - -void CodeGenHybrid::VisitStmt_(const ProducerRealizeNode* op) { - auto tensor = Downcast(op->producer); - if (!op->storage_scope.empty()) { - PrintIndent(); - stream << GetTensorID(tensor) << " = allocate(("; - for (size_t i = 0; i < op->bounds.size(); ++i) { - if (i) stream << ", "; - stream << PrintExpr(op->bounds[i]->extent); - } - if (op->bounds.size() == 1) stream << ", "; - stream << "), '"; - PrintType(tensor->dtype, stream); - stream << "', '"; - stream << op->storage_scope << "')\n"; - } - PrintStmt(op->body); -} - -void CodeGenHybrid::VisitStmt_(const AssertStmtNode* op) { - PrintIndent(); - stream << "assert "; - PrintExpr(op->condition, stream); - stream << ", "; - PrintExpr(op->message, stream); - stream << "\n"; - PrintStmt(op->body); -} - -void CodeGenHybrid::VisitStmt_(const ProducerStoreNode* op) { - auto tensor = Downcast(op->producer); - PrintIndent(); - stream << GetTensorID(tensor); - stream << "["; - for (size_t i = 0; i < op->indices.size(); ++i) { - if (i) stream << ", "; - PrintExpr(op->indices[i], stream); - } - stream << "] = "; - PrintExpr(op->value, stream); - stream << "\n"; -} - -void CodeGenHybrid::VisitStmt_(const ForNode* op) { - std::string extent = PrintExpr(op->extent); - PrintIndent(); - std::string vid = GetVarID(op->loop_var.get()); - stream << "for " << vid << " in " - << "range(" << extent << "):\n"; - indent_ += tab_; - PrintStmt(op->body); - indent_ -= tab_; -} - -bool is_noop(const Stmt& stmt) { - if (!stmt.defined()) return true; - if (auto eval = stmt.as()) return is_const_int(eval->value); - return false; -} - -void CodeGenHybrid::VisitStmt_(const IfThenElseNode* op) { - std::string cond = PrintExpr(op->condition); - PrintIndent(); - stream << "if " << cond << ":\n"; - indent_ += tab_; - PrintStmt(op->then_case); - indent_ -= tab_; - - if (op->else_case && !is_noop(op->else_case.value())) { - PrintIndent(); - stream << "else:\n"; - indent_ += tab_; - PrintStmt(op->else_case.value()); - indent_ -= tab_; - } -} - -void CodeGenHybrid::VisitStmt_(const SeqStmtNode* op) { - for (Stmt stmt : op->seq) { - PrintStmt(stmt); - } -} - -void CodeGenHybrid::VisitStmt_(const EvaluateNode* op) { - if (is_const_int(op->value)) return; - std::string str = PrintExpr(op->value); - if (!str.empty()) stream << str << "\n"; -} - -void CodeGenHybrid::PrintIndent() { stream << std::string(indent_, ' '); } - -std::string CodeGenHybrid::GetVarID(const VarNode* v) { - if (binds_.count(v)) return binds_[v]; - auto key = std::make_pair(static_cast(v), 0); - if (id_map_.count(key)) { - return id_map_[key]; - } - return id_map_[key] = ids_allocated->FreshName(v->name_hint); -} - -std::string CodeGenHybrid::GetTensorID(const Tensor& tensor) { - auto key = std::make_pair(tensor->op.get(), tensor->value_index); - if (id_map_.count(key)) { - return id_map_[key]; - } - std::string name_hint = tensor->op->name; - if (tensor->op->num_outputs() > 1) { - name_hint += "_v" + std::to_string(tensor->value_index); - } - return id_map_[key] = ids_allocated->FreshName(name_hint); -} - -void CodeGenHybrid::ReserveKeywords() { - ids_allocated->ReserveName("def"); - ids_allocated->ReserveName("for"); - ids_allocated->ReserveName("in"); - ids_allocated->ReserveName("range"); - ids_allocated->ReserveName("True"); - ids_allocated->ReserveName("False"); - ids_allocated->ReserveName("unroll"); - ids_allocated->ReserveName("const_range"); - ids_allocated->ReserveName("parallel"); - ids_allocated->ReserveName("vectorize"); - ids_allocated->ReserveName("bind"); - ids_allocated->ReserveName("threadIdx.x"); - ids_allocated->ReserveName("threadIdx.y"); - ids_allocated->ReserveName("threadIdx.z"); - ids_allocated->ReserveName("blockIdx.x"); - ids_allocated->ReserveName("blockIdx.y"); - ids_allocated->ReserveName("blockIdx.z"); - ids_allocated->ReserveName("vthread"); - ids_allocated->ReserveName("allocate"); - ids_allocated->ReserveName("output_tensor"); - ids_allocated->ReserveName("sqrt"); - ids_allocated->ReserveName("log"); - ids_allocated->ReserveName("tanh"); - ids_allocated->ReserveName("power"); - ids_allocated->ReserveName("exp"); - ids_allocated->ReserveName("sigmoid"); - ids_allocated->ReserveName("popcount"); - ids_allocated->ReserveName("likely"); - ids_allocated->ReserveName("int8"); - ids_allocated->ReserveName("int16"); - ids_allocated->ReserveName("int32"); - ids_allocated->ReserveName("int64"); - ids_allocated->ReserveName("uint8"); - ids_allocated->ReserveName("uint16"); - ids_allocated->ReserveName("uint32"); - ids_allocated->ReserveName("uint64"); - ids_allocated->ReserveName("float16"); - ids_allocated->ReserveName("float32"); - ids_allocated->ReserveName("float64"); - ids_allocated->ReserveName("ceil_div"); - ids_allocated->ReserveName("max_num_threads"); -} - -void CodeGenHybrid::DumpStmt(const Stmt& stmt, const Array& inputs, - const Array& outputs, const std::string& name) { - ReserveKeywords(); - ids_allocated->ReserveName(name); - - stream << "def " << name << "("; - for (size_t i = 0; i < inputs.size(); ++i) { - if (i) stream << ", "; - if (auto tensor = inputs[i].as()) { - stream << GetTensorID(tensor.value()); - } else { - auto var = inputs[i].as(); - ICHECK(var) << "Input should either be a tensor or a variable!"; - stream << GetVarID(var); - } - } - stream << "):\n"; - indent_ += tab_; - for (size_t i = 0; i < outputs.size(); ++i) { - PrintIndent(); - stream << GetTensorID(outputs[i]) << " = output_tensor(("; - for (size_t j = 0; j < outputs[i]->shape.size(); ++j) { - if (j) stream << ", "; - PrintExpr(outputs[i]->shape[j], stream); - } - if (outputs[i]->shape.size() == 1) stream << ", "; - stream << "), '" << outputs[i]->dtype << "')\n"; - } - PrintStmt(stmt); - PrintIndent(); - stream << "return "; - for (size_t i = 0; i < outputs.size(); ++i) { - if (i) stream << ", "; - stream << GetTensorID(outputs[i]); - } - stream << "\n"; -} - -TVM_REGISTER_GLOBAL("hybrid._Dump").set_body([](TVMArgs args, TVMRetValue* rv) { - CodeGenHybrid codegen; - if (args.size() == 4) - codegen.DumpStmt(args[0], args[1], args[2], args[3]); - else - codegen.DumpStmt(args[0], args[1], args[2]); - *rv = codegen.Finish(); -}); -} // namespace contrib -} // namespace tvm diff --git a/src/contrib/hybrid/codegen_hybrid.h b/src/contrib/hybrid/codegen_hybrid.h deleted file mode 100644 index 58be2cf112e0..000000000000 --- a/src/contrib/hybrid/codegen_hybrid.h +++ /dev/null @@ -1,171 +0,0 @@ -/* - * 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 codegen_hybrid.h - * \brief Common utilities to generated C style code. - */ -#ifndef TVM_CONTRIB_HYBRID_CODEGEN_HYBRID_H_ -#define TVM_CONTRIB_HYBRID_CODEGEN_HYBRID_H_ - -#include -#include -#include -#include -#include -#include - -#include -#include -#include -#include -#include - -namespace tvm { -namespace contrib { - -using namespace te; -using namespace tir; -/*! - * \brief A base class to generate Hybrid Script. - * - * **NOTE** CodeGenHybrid does not aim at generating Python scripts consumed by Python2/3. - * For runtime support, please refer the decorator in ``tvm/python/hybrid/api.py``. - */ -class CodeGenHybrid : public ExprFunctor, - public StmtFunctor { - public: - /*! - * \brief Dump the given function body to hybrid script. - * \param stmt The function body to be dumped to hybrid script. - * \param inputs Input tensors of this schedule. - * \param outputs Output tensors of this schedule. - * \param name The name of the function. - */ - void DumpStmt(const Stmt& stmt, const Array& inputs, const Array& outputs, - const std::string& name = "hybrid_func"); - /*! - * \brief Finalize the compilation and return the code. - * \return The code. - */ - std::string Finish(); - /*! \brief Reserve keywords in avoid of name conflict. */ - void ReserveKeywords(); - /*! - * \brief Print the Stmt n to CodeGenHybrid->stream - * \param n The statement to be printed. - */ - void PrintStmt(const Stmt& n) { this->VisitStmt(n); } - /*! - * \brief Print the expression n(or its ssa id if in ssa mode) into os - * \param n The expression to be printed. - * \param os The output stream - */ - void PrintExpr(const PrimExpr& n, std::ostream& os) { this->VisitExpr(n, os); } - /*! - * \brief Same as PrintExpr, but simply returns result string - * \param n The expression to be printed. - */ - std::string PrintExpr(const PrimExpr& n) { - std::ostringstream os; - PrintExpr(n, os); - return os.str(); - } - // expression - void VisitExpr_(const VarNode* op, std::ostream& os) override; // NOLINT(*) - void VisitExpr_(const BufferLoadNode* op, std::ostream& os) override; // NOLINT(*) - void VisitExpr_(const LetNode* op, std::ostream& os) override; // NOLINT(*) - void VisitExpr_(const CallNode* op, std::ostream& os) override; // NOLINT(*) - void VisitExpr_(const ProducerLoadNode* op, std::ostream& os) override; // NOLINT(*) - void VisitExpr_(const AddNode* op, std::ostream& os) override; // NOLINT(*) - void VisitExpr_(const SubNode* op, std::ostream& os) override; // NOLINT(*) - void VisitExpr_(const MulNode* op, std::ostream& os) override; // NOLINT(*) - void VisitExpr_(const DivNode* op, std::ostream& os) override; // NOLINT(*) - void VisitExpr_(const ModNode* op, std::ostream& os) override; // NOLINT(*) - void VisitExpr_(const FloorDivNode* op, std::ostream& os) override; // NOLINT(*) - void VisitExpr_(const FloorModNode* op, std::ostream& os) override; // NOLINT(*) - void VisitExpr_(const MinNode* op, std::ostream& os) override; // NOLINT(*) - void VisitExpr_(const MaxNode* op, std::ostream& os) override; // NOLINT(*) - void VisitExpr_(const EQNode* op, std::ostream& os) override; // NOLINT(*) - void VisitExpr_(const NENode* op, std::ostream& os) override; // NOLINT(*) - void VisitExpr_(const LTNode* op, std::ostream& os) override; // NOLINT(*) - void VisitExpr_(const LENode* op, std::ostream& os) override; // NOLINT(*) - void VisitExpr_(const GTNode* op, std::ostream& os) override; // NOLINT(*) - void VisitExpr_(const GENode* op, std::ostream& os) override; // NOLINT(*) - void VisitExpr_(const AndNode* op, std::ostream& os) override; // NOLINT(*) - void VisitExpr_(const OrNode* op, std::ostream& os) override; // NOLINT(*) - void VisitExpr_(const CastNode* op, std::ostream& os) override; // NOLINT(*) - void VisitExpr_(const NotNode* op, std::ostream& os) override; // NOLINT(*) - void VisitExpr_(const SelectNode* op, std::ostream& os) override; // NOLINT(*) - void VisitExpr_(const RampNode* op, std::ostream& os) override; // NOLINT(*) - void VisitExpr_(const BroadcastNode* op, std::ostream& os) override; // NOLINT(*) - void VisitExpr_(const IntImmNode* op, std::ostream& os) override; // NOLINT(*) - void VisitExpr_(const FloatImmNode* op, std::ostream& os) override; // NOLINT(*) - void VisitExpr_(const StringImmNode* op, std::ostream& os) override; // NOLINT(*) - // statment - void VisitStmt_(const LetStmtNode* op) override; - void VisitStmt_(const BufferStoreNode* op) override; - void VisitStmt_(const ProducerStoreNode* op) override; - void VisitStmt_(const ForNode* op) override; - void VisitStmt_(const IfThenElseNode* op) override; - void VisitStmt_(const AllocateNode* op) override; - void VisitStmt_(const ProducerRealizeNode* op) override; - void VisitStmt_(const AttrStmtNode* op) override; - void VisitStmt_(const AssertStmtNode* op) override; - void VisitStmt_(const EvaluateNode* op) override; - void VisitStmt_(const SeqStmtNode* op) override; - /*! - * \brief Print Type represetnation of type t. - * \param t The type representation. - * \param os The stream to print the ctype into - */ - virtual void PrintType(DataType t, std::ostream& os); // NOLINT(*) - - private: - /*! \brief The current indent of the code dump. */ - int indent_{0}; - /*! \brief The tab size of code indent. */ - const int tab_{4}; - /*! \brief Print the current indent spaces. */ - inline void PrintIndent(); - /*! \brief NameSupply for allocated ids. */ - NameSupply ids_allocated; - /*! - * \brief Keys are either (tensors, value_index) or (variables, 0). - * Values are the corresponding IDs.*/ - std::map, std::string> id_map_; - /*! \brief Variables (keys) binded to the threads (values). */ - std::map binds_; - /*! \brief The output code string builder. */ - std::stringstream stream; - /*! - * \brief Get or allocate the ID for the given variable. - * \param v The given variable. - */ - std::string GetVarID(const VarNode* v); - /*! - * \brief Get or allocate the ID for the given tensor. - * \param tensor The tensor to allocate a name. - */ - std::string GetTensorID(const Tensor& tensor); -}; - -} // namespace contrib -} // namespace tvm -#endif // TVM_CONTRIB_HYBRID_CODEGEN_HYBRID_H_ diff --git a/src/contrib/tf_op/tvm_dso_op_kernels.cc b/src/contrib/tf_op/tvm_dso_op_kernels.cc deleted file mode 100644 index 78c10e4822c8..000000000000 --- a/src/contrib/tf_op/tvm_dso_op_kernels.cc +++ /dev/null @@ -1,327 +0,0 @@ -/* - * 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. - */ - -#ifdef TF_TVMDSOOP_ENABLE_GPU -#include -#endif -#include -#include -#include -#include -#include - -#include "tensorflow/core/framework/op_kernel.h" - -typedef Eigen::ThreadPoolDevice CPUDevice; -typedef Eigen::GpuDevice GPUDevice; -typedef tensorflow::gtl::InlinedVector ShapeContainer; - -using tensorflow::OpKernel; -using tensorflow::OpKernelConstruction; -using tensorflow::OpKernelContext; - -using tvm::runtime::TVMArgs; -using tvm::runtime::TVMArgsSetter; -using tvm::runtime::TVMRetValue; - -// Op utility trait for diffrent device type template -template -class TVMDSOOpTrait; - -// Buffer information used for actual computation. -// Each buffer is associated with one TensorFlow 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: - tensorflow::Tensor inline_tensor; - tensorflow::Tensor* tensor; - - size_t size; - size_t offset; - - int device_type; - - char* origin_buf; - char* buf; - - void CopyToOrigin() { - if (buf == origin_buf) { - return; - } - if (device_type == kDLCPU) { - memcpy(origin_buf, buf + offset, size); -#ifdef TF_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 TF_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"; - } - } -}; - -tensorflow::Status GetDLPackDtype(const tensorflow::Tensor& tf_tensor, DLDataType* res) { - auto dtype = tf_tensor.dtype(); - - if (dtype == tensorflow::DT_HALF) { - *res = {kDLFloat, 16, 1}; - } else if (dtype == tensorflow::DT_FLOAT) { - *res = {kDLFloat, 32, 1}; - } else if (dtype == tensorflow::DT_DOUBLE) { - *res = {kDLFloat, 64, 1}; - } else if (dtype == tensorflow::DT_INT8) { - *res = {kDLInt, 8, 1}; - } else if (dtype == tensorflow::DT_INT16) { - *res = {kDLInt, 16, 1}; - } else if (dtype == tensorflow::DT_INT32) { - *res = {kDLInt, 32, 1}; - } else if (dtype == tensorflow::DT_INT64) { - *res = {kDLInt, 64, 1}; - } else if (dtype == tensorflow::DT_UINT8) { - *res = {kDLUInt, 8, 1}; - } else if (dtype == tensorflow::DT_UINT16) { - *res = {kDLUInt, 16, 1}; - } else if (dtype == tensorflow::DT_UINT32) { - *res = {kDLUInt, 32, 1}; - } else if (dtype == tensorflow::DT_UINT64) { - *res = {kDLUInt, 64, 1}; - } else { - return tensorflow::Status(tensorflow::error::INTERNAL, "Fail to get dlpack datatype"); - } - return tensorflow::Status::OK(); -} - -// Ensure buffer used for actual computation take 64byte alignment -void EnsureAlignment(OpKernelContext* ctx, const tensorflow::Tensor& tensor, TensorAsBuf* out) { - char* buf = const_cast(tensor.tensor_data().data()); - out->origin_buf = buf; - out->size = tensor.TotalBytes(); - - int alignment = 64; - char* aligned = reinterpret_cast(((uint64_t)buf + alignment - 1) & (~(alignment - 1))); - if (buf == aligned) { - out->tensor = const_cast(&tensor); - out->buf = buf; - out->offset = 0; - } else { - tensorflow::TensorShape buf_shape; - tensorflow::int64 dims[1] = {(tensorflow::int64)(tensor.TotalBytes() + alignment)}; - tensorflow::TensorShapeUtils::MakeShape(dims, 1, &buf_shape); - - out->tensor = &out->inline_tensor; - ctx->allocate_temp(tensor.dtype(), buf_shape, out->tensor); - - buf = const_cast(out->tensor->tensor_data().data()); - char* buf_aligned = reinterpret_cast(((uint64_t)buf + alignment) & (~(alignment - 1))); - out->buf = buf; - out->offset = buf_aligned - buf; - } -} - -// Create DLPack tensor from TensorFlow tensor -tensorflow::Status MakeDLTensor(const TensorAsBuf& src, const DLDevice& dev, int64_t* tf_shape, - DLTensor* out) { - DLDataType dlpack_type; - const tensorflow::Tensor& tensor = *src.tensor; - - auto status = GetDLPackDtype(tensor, &dlpack_type); - if (!status.ok()) { - return status; - } - out->device = dev; - out->ndim = tensor.shape().dims(); - out->shape = tf_shape; - out->strides = nullptr; - out->byte_offset = 0; - out->dtype = dlpack_type; - out->data = src.buf + src.offset; - return tensorflow::Status::OK(); -} - -template <> -class TVMDSOOpTrait { - public: - static const int device_type = kDLCPU; - - static int device_id(OpKernelContext* context) { return 0; } - - static void make_shape_from_tensor(const tensorflow::Tensor& shape_tensor, - tensorflow::TensorShape* output_shape) { - tensorflow::int64 num_dims = shape_tensor.NumElements(); - const tensorflow::int64* dims = shape_tensor.flat().data(); - tensorflow::TensorShapeUtils::MakeShape(dims, num_dims, output_shape); - } -}; - -#ifdef TF_TVMDSOOP_ENABLE_GPU -template <> -class TVMDSOOpTrait { - public: - static const int device_type = kDLCUDA; - - static int device_id(OpKernelContext* context) { - auto device_base = context->device(); - auto gpu_device_info = device_base->tensorflow_gpu_device_info(); - return gpu_device_info->gpu_id; - } - - static void make_shape_from_tensor(const tensorflow::Tensor& shape_tensor, - tensorflow::TensorShape* output_shape) { - tensorflow::int64 num_dims = shape_tensor.NumElements(); - const tensorflow::int64* flat = shape_tensor.flat().data(); - tensorflow::int64* dims = new tensorflow::int64[num_dims]; - cudaMemcpy(dims, flat, sizeof(tensorflow::int64) * num_dims, cudaMemcpyDeviceToHost); - tensorflow::TensorShapeUtils::MakeShape(dims, num_dims, output_shape); - delete[] dims; - } -}; -#endif - -template -class TVMDSOOp : public OpKernel { - private: - tvm::runtime::PackedFunc tvm_func; - std::string lib_path; - std::string func_name; - - tensorflow::DataType output_dtype; - - bool has_static_output_shape; - std::vector static_output_shape; - - void initAttributes(OpKernelConstruction* context) { - context->GetAttr("lib_path", &lib_path); - context->GetAttr("func_name", &func_name); - context->GetAttr("output_dtype", &output_dtype); - - context->GetAttr("has_static_output_shape", &has_static_output_shape); - context->GetAttr("static_output_shape", &static_output_shape); - } - - public: - explicit TVMDSOOp(OpKernelConstruction* context) : OpKernel(context) { - // Get attr - initAttributes(context); - - // Load TVM function from dynamic library - tvm::runtime::Module mod_dylib = tvm::runtime::Module::LoadFromFile(lib_path); - tvm_func = mod_dylib.GetFunction(func_name); - ICHECK(tvm_func != nullptr); - } - - void Compute(tensorflow::OpKernelContext* context) override { - // the last input is output shape spec - const int num_inputs = context->num_inputs() - 1; - const int num_total_args = num_inputs + 1; - std::vector args(num_total_args); - std::vector buf_info(num_inputs); - std::vector shapes(num_inputs); - - tensorflow::Status status; - int device_id = TVMDSOOpTrait::device_id(context); - int device_type = TVMDSOOpTrait::device_type; - - DLDevice dl_dev = {DLDeviceType(device_type), device_id}; - - // Get output shape - tensorflow::TensorShape output_shape; - auto& output_shape_tensor = context->input(num_inputs); - if (has_static_output_shape) { - // use static output shape - const tensorflow::int64* dims = static_output_shape.data(); - tensorflow::TensorShapeUtils::MakeShape(dims, static_output_shape.size(), &output_shape); - } else if (output_shape_tensor.dims() == 1) { - // use shape tensor values as output shape - TVMDSOOpTrait::make_shape_from_tensor(output_shape_tensor, &output_shape); - } else { - // use input tensor shape by default - output_shape = context->input(0).shape(); - } - - for (int i = 0; i < num_inputs; ++i) { - // Grab the input tensor - auto& input_tensor = context->input(i); - - // Create shape container, should keep ref during execution - shapes[i] = input_tensor.shape().dim_sizes(); - auto shape_ptr = reinterpret_cast(shapes[i].data()); - - TensorAsBuf& input = buf_info[i]; - input.device_type = device_type; - - EnsureAlignment(context, input_tensor, &input); - input.CopyFromOrigin(); - - status = MakeDLTensor(input, dl_dev, shape_ptr, &args[i]); - OP_REQUIRES_OK(context, status); - } - - // Allocate output tensor - tensorflow::Tensor* output_tensor; - OP_REQUIRES_OK(context, context->allocate_output(0, output_shape, &output_tensor)); - // shape dimension buf should keel alive on stack - auto output_shape_dim_buf = output_tensor->shape().dim_sizes(); - auto output_shape_ptr = reinterpret_cast(output_shape_dim_buf.data()); - - TensorAsBuf output; - output.device_type = device_type; - EnsureAlignment(context, *output_tensor, &output); - - status = MakeDLTensor(output, dl_dev, output_shape_ptr, &args[num_inputs]); - OP_REQUIRES_OK(context, status); - - // Prepare PackedFunc arguments - std::vector tvm_values(num_total_args); - std::vector tvm_type_codes(num_total_args); - TVMArgsSetter setter(tvm_values.data(), tvm_type_codes.data()); - for (int k = 0; k < num_total_args; ++k) { - setter(k, &args[k]); - } - TVMRetValue rv; - tvm_func.CallPacked(TVMArgs(tvm_values.data(), tvm_type_codes.data(), num_total_args), &rv); - - output.CopyToOrigin(); - } -}; - -#ifdef TF_TVMDSOOP_ENABLE_GPU -REGISTER_KERNEL_BUILDER(Name("TvmDsoOp").Device(tensorflow::DEVICE_CPU), TVMDSOOp); -REGISTER_KERNEL_BUILDER(Name("TvmDsoOp").Device(tensorflow::DEVICE_GPU), TVMDSOOp); -#else -REGISTER_KERNEL_BUILDER(Name("TvmDsoOp").Device(tensorflow::DEVICE_CPU), TVMDSOOp); -#endif diff --git a/src/contrib/tf_op/tvm_dso_ops.cc b/src/contrib/tf_op/tvm_dso_ops.cc deleted file mode 100644 index 794494298d71..000000000000 --- a/src/contrib/tf_op/tvm_dso_ops.cc +++ /dev/null @@ -1,35 +0,0 @@ -/* - * 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 "tensorflow/core/framework/op.h" - -REGISTER_OP("TvmDsoOp") - .Input("input_args: ListT") - .Attr( - "ListT: list({float16, float32, float64, int8, int16, int32, int64, uint8, uint16," - "uint32, uint64})") - .Input("dynamic_output_shape: int64") - .Output("output: output_dtype") - .Attr("lib_path: string") - .Attr("func_name: string") - .Attr( - "output_dtype: {float16, float32, float64, int8, int16, int32, int64, uint8, uint16," - "uint32, uint64} = DT_FLOAT") - .Attr("static_output_shape: list(int) >= 0 = []") - .Attr("has_static_output_shape: bool"); diff --git a/src/contrib/torch/pt_call_tvm/tvm_class.cc b/src/contrib/torch/pt_call_tvm/tvm_class.cc deleted file mode 100644 index f5ae95a5a73d..000000000000 --- a/src/contrib/torch/pt_call_tvm/tvm_class.cc +++ /dev/null @@ -1,686 +0,0 @@ -/* - * 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::memory::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/tvm_module_wrapper/RuntimeModuleWrapperTVM.cc b/src/contrib/torch/tvm_module_wrapper/RuntimeModuleWrapperTVM.cc deleted file mode 100644 index 3e1c7e7c0edf..000000000000 --- a/src/contrib/torch/tvm_module_wrapper/RuntimeModuleWrapperTVM.cc +++ /dev/null @@ -1,307 +0,0 @@ -/* - * 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 "../../../runtime/graph_executor/graph_executor_factory.h" -#include "../../../support/base64.h" -#include "runtime_bridge.h" - -namespace tvm { -namespace contrib { - -/* - * TVM's FFI for passing module from python to C++ - */ -struct ThreadLocalStore { - tvm::runtime::Module mod; - static ThreadLocalStore* ThreadLocal() { - thread_local ThreadLocalStore tls; - return &tls; - } -}; - -TVM_REGISTER_GLOBAL("tvmtorch.save_runtime_mod").set_body_typed([](tvm::runtime::Module mod) { - ThreadLocalStore::ThreadLocal()->mod = mod; -}); - -/* - * Convert NDArray to DLPack extend tensor. It should be zero-cost. - * @param src Pointer to NDArray - * @return DLPack extended tensor - */ -DLPackTensorExt CreateDLpackTensorExt(tvm::runtime::NDArray* src) { - auto is_bool = src->DataType().is_bool(); - DLManagedTensor* tensor; - if (is_bool) { - // If we change DLDataType{kDLInt, 8, 1} to DataType::Bool() - // we will get `RuntimeError: Unsupported kUInt bits 1` - auto tmp = src->CreateView(src->Shape(), DLDataType{kDLInt, 8, 1}); - tensor = tmp.ToDLPack(); - } else { - tensor = src->ToDLPack(); - } - DLPackTensorExt ret{tensor, is_bool}; - return ret; -} - -/* - * Create an NDArray with boolean type. (One memory copy) - * @param src DLpack extended tensor - * @return a new NDArray - */ -tvm::runtime::NDArray CreateBoolNDarray(DLPackTensorExt* src) { - auto& tensor = src->dl_managed_tensor->dl_tensor; - std::vector shape; - for (int64_t i = 0; i < tensor.ndim; i++) { - shape.push_back(tensor.shape[i]); - } - auto ret = tvm::runtime::NDArray::Empty(shape, DataType::Bool(), tensor.device); - ret.CopyFrom(&src->dl_managed_tensor->dl_tensor); - return std::move(ret); -} - -bool IsZeroCopy(DLPackTensorExt* src) { - auto& dl_tensor = src->dl_managed_tensor->dl_tensor; - return tvm::runtime::NDArray::AbilityOfZeroCopyForDLTensor(&dl_tensor, dl_tensor.device); -} - -/* - * Create an NDArray from DLpack extended tensor. - * @param src DLpack extended tensor - * @return a new NDArray - */ -tvm::runtime::NDArray NDarrayFromDLpack(DLPackTensorExt* src) { - using tvm::runtime::NDArray; - - NDArray array; - auto& dl_tensor = src->dl_managed_tensor->dl_tensor; - if (src->is_bool) { - // one memory copy - // the code is similar to NewFromDLTensor except for the type - array = CreateBoolNDarray(src); - } else if (IsZeroCopy(src)) { - array = NDArray::FromExternalDLTensor(src->dl_managed_tensor->dl_tensor); - } else { - // one memory copy - array = NDArray::NewFromDLTensor(&dl_tensor, dl_tensor.device); - } - return array; -} - -} // namespace contrib -} // namespace tvm - -extern "C" { - -struct TVMContribTorchRuntimeModule { - tvm::runtime::Module mod; - - explicit TVMContribTorchRuntimeModule(tvm::runtime::Module& mod) : mod(mod) {} -}; - -bool tvm_contrib_torch_tensor_ability_of_zero_copy(DLPackTensorExt* src) { - return (!src->is_bool) && (tvm::contrib::IsZeroCopy(src)); -} - -TVMContribTorchRuntimeModule* tvm_contrib_torch_get_last_saved_runtime_module() { - return new TVMContribTorchRuntimeModule(tvm::contrib::ThreadLocalStore::ThreadLocal()->mod); -} - -void tvm_contrib_torch_operator_module_forward(TVMContribTorchRuntimeModule* runtime_module, - DLPackTensorExt* inputs, size_t input_size) { - tvm::runtime::PackedFunc run = runtime_module->mod.GetFunction("__tvm_main__"); - - std::vector tvm_values(input_size); - std::vector tvm_type_codes(input_size); - tvm::runtime::TVMArgsSetter setter(tvm_values.data(), tvm_type_codes.data()); - - std::vector input_cache(input_size); - - for (size_t k = 0; k < input_size; ++k) { - auto datum = tvm::contrib::NDarrayFromDLpack(&inputs[k]); // could have one memory copy - input_cache[k] = datum; // we keep the datum in a vector for future use, otherwise the datum - // will be freed after the loop - setter(k, datum); - } - - run.CallPacked(tvm::runtime::TVMArgs(tvm_values.data(), tvm_type_codes.data(), input_size), - nullptr); - - for (size_t k = 0; k < input_size; ++k) { - if (!tvm_contrib_torch_tensor_ability_of_zero_copy(&inputs[k])) - input_cache[k].CopyTo(&inputs[k].dl_managed_tensor->dl_tensor); - } -} - -TVMContribTorchRuntimeModule* tvm_contrib_torch_create_graph_runtime_module( - TVMContribTorchRuntimeModule* graph_executor_factory, DLManagedTensor* input_example) { - tvm::runtime::PackedFunc built_module = graph_executor_factory->mod.GetFunction("default"); - tvm::Device device_info = input_example->dl_tensor.device; - tvm::runtime::Module runtime_module = built_module(device_info); - return new TVMContribTorchRuntimeModule(runtime_module); -} - -size_t tvm_contrib_torch_graph_executor_module_forward(TVMContribTorchRuntimeModule* runtime_module, - DLPackTensorExt* inputs, size_t input_size, - DLPackTensorExt** outputs) { - tvm::runtime::PackedFunc run = runtime_module->mod.GetFunction("run"); - tvm::runtime::PackedFunc set_input = runtime_module->mod.GetFunction("set_input"); - tvm::runtime::PackedFunc get_output = runtime_module->mod.GetFunction("get_output"); - tvm::runtime::PackedFunc get_num_outputs = runtime_module->mod.GetFunction("get_num_outputs"); - - for (size_t k = 0; k < input_size; ++k) { - set_input(k, &inputs[k].dl_managed_tensor->dl_tensor); - } - - run(); - - int64_t output_length = get_num_outputs(); - - DLPackTensorExt* outputs_ptr = new DLPackTensorExt[output_length]; - *outputs = outputs_ptr; - - for (int64_t k = 0; k < output_length; ++k) { - tvm::runtime::NDArray results = get_output(k); - outputs_ptr[k] = tvm::contrib::CreateDLpackTensorExt(&results); - } - - return output_length; -} - -inline size_t b64strlen(const std::string b64str) { - ICHECK(b64str.size() % 4 == 0) << "invalid base64 encoding"; - size_t length = b64str.size() / 4 * 3; - if (b64str[b64str.size() - 2] == '=') { - length -= 2; - } else if (b64str[b64str.size() - 1] == '=') { - length -= 1; - } - return length; -} - -inline void b64decode(const std::string b64str, uint8_t* ret) { - size_t index = 0; - const auto length = b64str.size(); - for (size_t i = 0; i < length; i += 4) { - int8_t ch0 = tvm::support::base64::DecodeTable[(int32_t)b64str[i]]; - int8_t ch1 = tvm::support::base64::DecodeTable[(int32_t)b64str[i + 1]]; - int8_t ch2 = tvm::support::base64::DecodeTable[(int32_t)b64str[i + 2]]; - int8_t ch3 = tvm::support::base64::DecodeTable[(int32_t)b64str[i + 3]]; - uint8_t st1 = (ch0 << 2) + (ch1 >> 4); - ret[index++] = st1; - if (b64str[i + 2] != '=') { - uint8_t st2 = ((ch1 & 0b1111) << 4) + (ch2 >> 2); - ret[index++] = st2; - if (b64str[i + 3] != '=') { - uint8_t st3 = ((ch2 & 0b11) << 6) + ch3; - ret[index++] = st3; - } - } - } - ICHECK(b64strlen(b64str) == index) << "base64 decoding fails"; -} - -/*! - * \brief Export TVM runtime module to base64 stream including its submodules. - * Note that this targets modules that are binary serializable and DSOExportable. - * \param module The runtime module to export - * \return std::string The content of exported file - */ -std::string ExportModuleToBase64(tvm::runtime::Module module) { - static const tvm::runtime::PackedFunc* f_to_str = - tvm::runtime::Registry::Get("export_runtime_module"); - ICHECK(f_to_str) << "IndexError: Cannot find the packed function " - "`export_runtime_module` in the global registry"; - return (*f_to_str)(module); -} - -struct Deleter { // deleter - explicit Deleter(std::string file_name) { this->file_name = file_name; } - void operator()(FILE* p) const { - fclose(p); - ICHECK(remove(file_name.c_str()) == 0) - << "remove temporary file (" << file_name << ") unsuccessfully"; - } - std::string file_name; -}; - -/*! - * \brief Import TVM runtime module from base64 stream - * Note that this targets modules that are binary serializable and DSOExportable. - * \param base64str base64 stream, which are generated by `ExportModuleToBase64`. - * \return runtime::Module runtime module constructed from the given stream - */ -tvm::runtime::Module ImportModuleFromBase64(std::string base64str) { - auto length = b64strlen(base64str); - - std::vector bytes(length); // bytes stream - b64decode(base64str, bytes.data()); - - auto now = std::chrono::system_clock::now(); - auto in_time_t = std::chrono::system_clock::to_time_t(now); - std::stringstream datetime; - datetime << std::put_time(std::localtime(&in_time_t), "%Y-%m-%d-%X"); - const std::string file_name = "tmp-module-" + datetime.str() + ".so"; - LOG(INFO) << file_name; - std::unique_ptr pFile(fopen(file_name.c_str(), "wb"), Deleter(file_name)); - fwrite(bytes.data(), sizeof(uint8_t), length, pFile.get()); - fflush(pFile.get()); - - std::string load_f_name = "runtime.module.loadfile_so"; - const tvm::runtime::PackedFunc* f = tvm::runtime::Registry::Get(load_f_name); - ICHECK(f != nullptr) << "Loader for `.so` files is not registered," - << " resolved to (" << load_f_name << ") in the global registry." - << "Ensure that you have loaded the correct runtime code, and" - << "that you are on the correct hardware architecture."; - tvm::runtime::Module ret = (*f)(file_name, ""); - return ret; -} - -char* tvm_contrib_torch_encode(TVMContribTorchRuntimeModule* runtime_module) { - std::string std = ExportModuleToBase64(runtime_module->mod); - char* ret = new char[std.length() + 1]; - snprintf(ret, std.length() + 1, "%s", std.c_str()); - return ret; -} - -TVMContribTorchRuntimeModule* tvm_contrib_torch_decode(const char* state) { - tvm::runtime::Module ret = ImportModuleFromBase64(state); - return new TVMContribTorchRuntimeModule(ret); -} - -void tvm_contrib_torch_free_runtime_module(TVMContribTorchRuntimeModule* module_ptr) { - delete module_ptr; -} - -void tvm_contrib_torch_free_dlpack_tensor_ext_array(DLPackTensorExt* dlpack_ptr) { - delete[] dlpack_ptr; -} - -void tvm_contrib_torch_free_encoding(char* encoding) { delete[] encoding; } -} diff --git a/src/contrib/torch/tvm_module_wrapper/RuntimeModuleWrapperTorch.cc b/src/contrib/torch/tvm_module_wrapper/RuntimeModuleWrapperTorch.cc deleted file mode 100644 index 3159438d7202..000000000000 --- a/src/contrib/torch/tvm_module_wrapper/RuntimeModuleWrapperTorch.cc +++ /dev/null @@ -1,215 +0,0 @@ -/* - * 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 "runtime_bridge.h" - -namespace tvm { -namespace contrib { - -/* - * Convert Torch tensor to DLPack extended tensor. - * The boolean Torch tensor will convert to DLtensor with `is_bool=True` flag. - * @param src Torch tensor - * @return DLPack extended tensor - */ -DLPackTensorExt ToDLPackExt(const at::Tensor& src) { - if (!src.is_contiguous()) { - return ToDLPackExt(src.contiguous()); - } - DLPackTensorExt ret; - if (src.dtype().isScalarType(torch::kBool)) { - auto temp = src.toType(torch::kUInt8); - ret.dl_managed_tensor = at::toDLPack(temp); - ret.is_bool = true; - } else { - ret.dl_managed_tensor = at::toDLPack(src); - ret.is_bool = false; - } - - return ret; -} - -/* - * Convert DLPack extended tensor to Torch tensor. - * @param src DLPack extended tensor - * @return Torch tensor - */ -at::Tensor FromDLPackExt(const DLPackTensorExt& src) { - if (src.is_bool) { - return at::fromDLPack(src.dl_managed_tensor).toType(torch::kBool); - } else { - return at::fromDLPack(src.dl_managed_tensor); - } -} - -/** - * @brief A Torch's module which wraps TVM's OperatorModule Class. - * The basic forward function calling TVM's runtime is provided. - * The TVM module can be serialized/deserialized as a Torch module. - */ -class OperatorModuleWrapper : public torch::jit::CustomClassHolder { - public: - OperatorModuleWrapper() { runtime_module_ = tvm_contrib_torch_get_last_saved_runtime_module(); } - ~OperatorModuleWrapper() { tvm_contrib_torch_free_runtime_module(runtime_module_); } - - void forward(const c10::List& inputs) { - int input_length = inputs.size(); - - std::vector tensors; - - // Torch tensor supports boolean type while DLpack does not, - // we convert Torch tensor to an extension of DLPack tensor - for (int i = 0; i < input_length; ++i) tensors.push_back(ToDLPackExt(inputs[i])); - tvm_contrib_torch_operator_module_forward(this->runtime_module_, tensors.data(), - tensors.size()); - - for (int k = 0; k < input_length; ++k) { - if (tvm_contrib_torch_tensor_ability_of_zero_copy(&tensors[k])) { - // We need to free memory manually - tensors[k].dl_managed_tensor->deleter(tensors[k].dl_managed_tensor); - } else { - // Ownership transferred - inputs[k].copy_(FromDLPackExt(tensors[k])); - } - } - } - - std::string Serialize() { - auto encoding = tvm_contrib_torch_encode(runtime_module_); - auto ret = std::string(encoding); - tvm_contrib_torch_free_encoding(encoding); - return ret; - } - - explicit OperatorModuleWrapper(std::string state) { - runtime_module_ = tvm_contrib_torch_decode(state.c_str()); - } - - private: - /* - * TVM runtime module wrapper - */ - TVMContribTorchRuntimeModule* runtime_module_; -}; - -/** - * @brief A Torch's module which wraps TVM's GraphExecutorFactory Class. - * The basic forward function calling TVM's runtime is provided. - * The TVM module can be serialized/deserialized as a Torch module. - */ -class GraphExecutorFactoryWrapper : public torch::jit::CustomClassHolder { - public: - explicit GraphExecutorFactoryWrapper(TVMContribTorchRuntimeModule* executor_factory) - : executor_factory_(executor_factory), executor_factory_runtime_(nullptr) {} - - ~GraphExecutorFactoryWrapper() { - tvm_contrib_torch_free_runtime_module(executor_factory_); - tvm_contrib_torch_free_runtime_module(executor_factory_runtime_); - } - - GraphExecutorFactoryWrapper() - : GraphExecutorFactoryWrapper(tvm_contrib_torch_get_last_saved_runtime_module()) {} - - std::string Serialize() { - auto encoding = tvm_contrib_torch_encode(executor_factory_); - auto ret = std::string(encoding); - tvm_contrib_torch_free_encoding(encoding); - return ret; - } - - explicit GraphExecutorFactoryWrapper(std::string state) { - executor_factory_ = tvm_contrib_torch_decode(state.c_str()); - executor_factory_runtime_ = nullptr; - } - - c10::List forward(const c10::List& inputs) { - int input_length = inputs.size(); - - TORCH_CHECK(input_length > 0, "Receive empty list of input tensors"); - - std::vector tensors; - - // Torch tensor supports boolean type while DLpack does not, - // we convert Torch tensor to an extension of DLPack tensor - for (int i = 0; i < input_length; ++i) tensors.push_back(ToDLPackExt(inputs[i])); - - DLPackTensorExt* outputs; - if (executor_factory_runtime_ == nullptr) { - executor_factory_runtime_ = tvm_contrib_torch_create_graph_runtime_module( - this->executor_factory_, tensors[0].dl_managed_tensor); - } - auto num_outputs = tvm_contrib_torch_graph_executor_module_forward( - executor_factory_runtime_, tensors.data(), tensors.size(), &outputs); - - c10::List ret; - ret.reserve(num_outputs); - - for (size_t k = 0; k < num_outputs; ++k) { - at::Tensor atTensor = FromDLPackExt(outputs[k]); - ret.emplace_back(atTensor); - } - - for (int k = 0; k < input_length; ++k) { - tensors[k].dl_managed_tensor->deleter(tensors[k].dl_managed_tensor); - } - tvm_contrib_torch_free_dlpack_tensor_ext_array(outputs); - - return ret; - } - - private: - /* - * TVM Graph Executor Factory module wrapper - */ - TVMContribTorchRuntimeModule* executor_factory_; - - /* - * TVM runtime module wrapper - */ - TVMContribTorchRuntimeModule* executor_factory_runtime_; -}; - -TORCH_LIBRARY(tvm_torch, m) { - m.class_("OperatorModuleWrapper") - .def(torch::init<>()) - .def("forward", &OperatorModuleWrapper::forward) - .def_pickle( - [](const c10::intrusive_ptr& self) -> std::string { - return self->Serialize(); - }, - [](std::string state) { return c10::make_intrusive(state); }); - m.class_("GraphExecutorFactoryWrapper") - .def(torch::init<>()) - .def("forward", &GraphExecutorFactoryWrapper::forward) - .def_pickle( - [](const c10::intrusive_ptr& self) -> std::string { - return self->Serialize(); - }, - [](std::string state) { - return c10::make_intrusive(state); - }); -} - -} // namespace contrib -} // namespace tvm diff --git a/src/contrib/torch/tvm_module_wrapper/runtime_bridge.h b/src/contrib/torch/tvm_module_wrapper/runtime_bridge.h deleted file mode 100644 index 58cd53a2840d..000000000000 --- a/src/contrib/torch/tvm_module_wrapper/runtime_bridge.h +++ /dev/null @@ -1,116 +0,0 @@ -/* - * 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 runtime_bridge.h - * \brief Util functions for pytorch tvm interaction. - */ -#ifndef TVM_CONTRIB_TORCH_TVM_MODULE_WRAPPER_RUNTIME_BRIDGE_H_ -#define TVM_CONTRIB_TORCH_TVM_MODULE_WRAPPER_RUNTIME_BRIDGE_H_ - -extern "C" { - -/* - * DLPack data structure extend with `is_bool` flag. - * DLPack haven't support boolean tensor - * (https://github.com/pytorch/pytorch/blob/4618371da56c887195e2e1d16dad2b9686302800/aten/src/ATen/DLConvertor.cpp#L42), - * thus a boolean tensor will be regarded as a UInt8 tensor - * (https://github.com/apache/tvm/blob/de124862714e747764aa8b7f41a90bcb25f3c6a8/python/tvm/_ffi/runtime_ctypes.py#L91). - */ -struct DLPackTensorExt { - DLManagedTensor* dl_managed_tensor; - bool is_bool; -}; - -/* - * A wrapper pointing to TVM runtime module. - */ -struct TVMContribTorchRuntimeModule; - -/* - * Obtain a saved runtime module passed by TVM FFI. - * @return A TVM runtime module wrapper. - */ -TVMContribTorchRuntimeModule* tvm_contrib_torch_get_last_saved_runtime_module(); - -/* - * Delete TVMContribTorchRuntimeModule pointer. - */ -void tvm_contrib_torch_free_runtime_module(TVMContribTorchRuntimeModule* module_ptr); - -/* - * Obtain ExecutorFactory runtime module from ExecutorFactory class. - * @param graph_executor_factory ExecutorFactory class - * @param input_example For obtaining device information - * @return ExecutorFactory TVM runtime module wrapper - */ -TVMContribTorchRuntimeModule* tvm_contrib_torch_create_graph_runtime_module( - TVMContribTorchRuntimeModule* graph_executor_factory, DLManagedTensor* input_example); - -/* - * Forward method for OperatorModuleWrapper. - * @param runtime_module TVM runtime module wrapper - * @param inputs Array pointer of the input tensors - * @param input_size The number of input tensors - */ -void tvm_contrib_torch_operator_module_forward(TVMContribTorchRuntimeModule* runtime_module, - DLPackTensorExt* inputs, size_t input_size); - -/* - * Forward method for GraphExecutorFactoryWrapper. - * @param graph_executor_factory TVM runtime module wrapper - * @param inputs Array pointer of the input tensors - * @param input_size The number of input tensors - * @param outputs The resulting output tensors pointer - * @return The number of output tensors - */ -size_t tvm_contrib_torch_graph_executor_module_forward( - TVMContribTorchRuntimeModule* graph_executor_factory, DLPackTensorExt* inputs, - size_t input_size, DLPackTensorExt** outputs); - -/* - * Encode TVM runtime module. - * @param runtime_module TVM runtime module wrapper - * @return The encoding stream (char array) - */ -char* tvm_contrib_torch_encode(TVMContribTorchRuntimeModule* runtime_module); - -/* - * Decode TVM runtime module. - * @param state The encoding stream (char array) of TVM runtime module - * @return TVM runtime module wrapper - */ -TVMContribTorchRuntimeModule* tvm_contrib_torch_decode(const char* state); - -/* - * Delete DLPackTensorExt pointer. - */ -void tvm_contrib_torch_free_dlpack_tensor_ext_array(DLPackTensorExt*); - -/* - * Delete char array pointer. - */ -void tvm_contrib_torch_free_encoding(char* encoding); - -/* - * Checking if a DLPackTensorExt is boolean or cannot be copied in zero cost. - */ -bool tvm_contrib_torch_tensor_ability_of_zero_copy(DLPackTensorExt*); -} - -#endif // TVM_CONTRIB_TORCH_TVM_MODULE_WRAPPER_RUNTIME_BRIDGE_H_ diff --git a/src/contrib/torch/utils.h b/src/contrib/torch/utils.h deleted file mode 100644 index a98e058ca346..000000000000 --- a/src/contrib/torch/utils.h +++ /dev/null @@ -1,264 +0,0 @@ -/* - * 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_ diff --git a/src/driver/driver_api.cc b/src/driver/driver_api.cc index 1e576bc91002..86c1d44f0108 100644 --- a/src/driver/driver_api.cc +++ b/src/driver/driver_api.cc @@ -24,8 +24,6 @@ #include #include #include -#include -#include #include #include #include @@ -33,8 +31,6 @@ #include #include -#include -#include namespace tvm { @@ -481,7 +477,7 @@ runtime::Module TIRToRuntime(const Map& inputs_arg, // Take the attrs from the first module so the eventual modules have them. // Ideally this would just be one unified module all the way through; IRModule first_module = (*inputs.begin()).second; - IRModule mhost_all = IRModule(Map(), {}, {}, {}, first_module->attrs); + IRModule mhost_all = IRModule(Map(), {}, first_module->attrs); ICHECK(mhost_all.defined()) << "The host module must be defined"; @@ -611,15 +607,7 @@ transform::Sequential MixedModulePassManager(IRModule mixed_mod, Target target) // because the merged allocation site is at the beginning of each device function mixed_pass_list.push_back(tir::transform::MergeSharedMemoryAllocations()); - bool unpacked_api = mixed_mod->GetAttr(tvm::attr::kExecutor) - .value_or(relay::Executor::Create("graph", {})) - ->GetAttr("unpacked-api") - .value_or(Bool(false)); - if (unpacked_api) { - mixed_pass_list.push_back(tir::transform::MakeUnpackedAPI()); - } else { - mixed_pass_list.push_back(tir::transform::MakePackedAPI()); - } + mixed_pass_list.push_back(tir::transform::MakePackedAPI()); mixed_pass_list.push_back(tir::transform::FP8StorageLegalize()); mixed_pass_list.push_back(tir::transform::BF16StorageLegalize()); @@ -635,7 +623,6 @@ TVM_REGISTER_GLOBAL("driver.mixed_mod_passes") transform::Sequential HostModulePassManager(IRModule mixed_mod, Target target_host) { transform::PassContext pass_ctx = transform::PassContext::Current(); - bool enable_debug = pass_ctx->GetConfig("tir.enable_debug", Bool(false)).value(); Array host_pass_list; @@ -655,10 +642,6 @@ transform::Sequential HostModulePassManager(IRModule mixed_mod, Target target_ho host_pass_list.push_back(tir::transform::LowerDeviceStorageAccessInfo()); host_pass_list.push_back(tir::transform::CombineContextCall()); - if (enable_debug) { - host_pass_list.push_back(tir::transform::InstallDebugSpans()); - } - return transform::Sequential(host_pass_list); } diff --git a/src/ir/adt.cc b/src/ir/adt.cc deleted file mode 100644 index 3533c8c514cd..000000000000 --- a/src/ir/adt.cc +++ /dev/null @@ -1,76 +0,0 @@ -/* - * 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 src/ir/adt.cc - * \brief ADT type definitions. - */ -#include -#include -#include - -namespace tvm { - -Constructor::Constructor(String name_hint, tvm::Array inputs, GlobalTypeVar belong_to) { - ObjectPtr n = make_object(); - n->name_hint = std::move(name_hint); - n->inputs = std::move(inputs); - n->belong_to = std::move(belong_to); - data_ = std::move(n); -} - -TVM_REGISTER_NODE_TYPE(ConstructorNode); - -TVM_REGISTER_GLOBAL("ir.Constructor") - .set_body_typed([](String name_hint, tvm::Array inputs, GlobalTypeVar belong_to) { - return Constructor(name_hint, inputs, belong_to); - }); - -TVM_STATIC_IR_FUNCTOR(ReprPrinter, vtable) - .set_dispatch([](const ObjectRef& ref, ReprPrinter* p) { - auto* node = static_cast(ref.get()); - p->stream << "ConstructorNode(" << node->name_hint << ", " << node->inputs << ", " - << node->belong_to << ")"; - }); - -TypeData::TypeData(GlobalTypeVar header, tvm::Array type_vars, - tvm::Array constructors) { - ObjectPtr n = make_object(); - n->header = std::move(header); - n->type_vars = std::move(type_vars); - n->constructors = std::move(constructors); - data_ = std::move(n); -} - -TVM_REGISTER_NODE_TYPE(TypeDataNode); - -TVM_REGISTER_GLOBAL("ir.TypeData") - .set_body_typed([](GlobalTypeVar header, tvm::Array type_vars, - tvm::Array constructors) { - return TypeData(header, type_vars, constructors); - }); - -TVM_STATIC_IR_FUNCTOR(ReprPrinter, vtable) - .set_dispatch([](const ObjectRef& ref, ReprPrinter* p) { - auto* node = static_cast(ref.get()); - p->stream << "TypeDataNode(" << node->header << ", " << node->type_vars << ", " - << node->constructors << ")"; - }); - -} // namespace tvm diff --git a/src/ir/function.cc b/src/ir/function.cc index 59f94201b241..3c92787530a6 100644 --- a/src/ir/function.cc +++ b/src/ir/function.cc @@ -23,7 +23,6 @@ */ #include #include -#include #include #include @@ -37,8 +36,6 @@ TVM_REGISTER_GLOBAL("ir.BaseFuncWithAttr") .set_body_typed([](BaseFunc func, String key, ObjectRef value) -> BaseFunc { if (func->IsInstance()) { return WithAttr(Downcast(std::move(func)), key, value); - } else if (func->IsInstance()) { - return WithAttr(Downcast(std::move(func)), key, value); } else if (func->IsInstance()) { return WithAttr(Downcast(std::move(func)), key, value); } else { @@ -68,8 +65,6 @@ TVM_REGISTER_GLOBAL("ir.BaseFuncWithoutAttr") .set_body_typed([](BaseFunc func, String key) -> BaseFunc { if (func->IsInstance()) { return WithoutAttr(Downcast(std::move(func)), key); - } else if (func->IsInstance()) { - return WithoutAttr(Downcast(std::move(func)), key); } else if (func->IsInstance()) { return WithoutAttr(Downcast(std::move(func)), key); } else { diff --git a/src/ir/memory_pools.cc b/src/ir/memory_pools.cc index f5064af207cc..912ad80b3cce 100644 --- a/src/ir/memory_pools.cc +++ b/src/ir/memory_pools.cc @@ -23,7 +23,6 @@ */ #include -#include namespace tvm { diff --git a/src/ir/module.cc b/src/ir/module.cc index 261fbfe087c6..b27dec5719ab 100644 --- a/src/ir/module.cc +++ b/src/ir/module.cc @@ -33,17 +33,11 @@ namespace tvm { -IRModule::IRModule(tvm::Map functions, - tvm::Map type_definitions, - std::unordered_set import_set, SourceMap source_map, DictAttrs attrs, +IRModule::IRModule(tvm::Map functions, SourceMap source_map, DictAttrs attrs, Map> global_infos) { auto n = make_object(); n->functions = std::move(functions); - n->type_definitions = std::move(type_definitions); - n->global_type_var_map_ = {}; n->global_var_map_ = {}; - n->constructor_tag_map_ = {}; - n->import_set_ = std::move(import_set); n->source_map = source_map; n->attrs = std::move(attrs); n->global_infos = std::move(global_infos); @@ -55,13 +49,6 @@ IRModule::IRModule(tvm::Map functions, n->global_var_map_.Set(kv.first->name_hint, kv.first); } - for (const auto& kv : n->type_definitions) { - // set global typevar map - ICHECK(n->global_type_var_map_.count(kv.first->name_hint) == 0) - << "Duplicate global type definition name " << kv.first->name_hint; - n->global_type_var_map_.Set(kv.first->name_hint, kv.first); - n->RegisterConstructors(kv.first, kv.second); - } data_ = std::move(n); } @@ -78,8 +65,7 @@ bool IRModuleNode::SEqualReduce(const IRModuleNode* other, SEqualReducer equal) if (functions.size() != other->functions.size()) return false; // Update GlobalVar remap if (equal.IsPathTracingEnabled()) { - if ((functions.size() != other->functions.size()) || - (type_definitions.size() != other->type_definitions.size())) { + if (functions.size() != other->functions.size()) { return false; } } @@ -95,23 +81,12 @@ bool IRModuleNode::SEqualReduce(const IRModuleNode* other, SEqualReducer equal) return false; } } - for (const auto& gtv : this->GetGlobalTypeVars()) { - if (other->ContainGlobalTypeVar(gtv->name_hint)) { - if (!equal.DefEqual(gtv, other->GetGlobalTypeVar(gtv->name_hint))) return false; - } else if (!equal.IsPathTracingEnabled()) { - return false; - } - } // Checking functions and type definitions if (!equal(this->functions, other->functions, [](const auto& path) { return path->Attr("functions"); })) { return false; } - if (!equal(this->type_definitions, other->type_definitions, - [](const auto& path) { return path->Attr("type_definitions"); })) { - return false; - } return true; } @@ -143,11 +118,6 @@ void IRModuleNode::SHashReduce(SHashReducer hash_reduce) const { } reduce_temp(); - temp.clear(); - for (const auto& kv : this->type_definitions) { - temp.emplace_back(kv.first->name_hint, kv.first, kv.second); - } - reduce_temp(); hash_reduce(this->attrs); hash_reduce(this->global_infos); } @@ -156,10 +126,6 @@ bool IRModuleNode::ContainGlobalVar(const String& name) const { return global_var_map_.find(name) != global_var_map_.end(); } -bool IRModuleNode::ContainGlobalTypeVar(const String& name) const { - return global_type_var_map_.find(name) != global_type_var_map_.end(); -} - GlobalVar IRModuleNode::GetGlobalVar(const String& name) const { auto it = global_var_map_.find(name); if (it == global_var_map_.end()) { @@ -190,38 +156,8 @@ tvm::Array IRModuleNode::GetGlobalVars() const { return tvm::Array(global_vars); } -GlobalTypeVar IRModuleNode::GetGlobalTypeVar(const String& name) const { - ICHECK(global_type_var_map_.defined()); - auto it = global_type_var_map_.find(name); - ICHECK(it != global_type_var_map_.end()) - << "Cannot find global type var " << name << " in the Module"; - return (*it).second; -} - -Constructor IRModuleNode::GetConstructor(const String& adt, const String& cons) const { - TypeData typeDef = this->LookupTypeDef(adt); - for (Constructor c : typeDef->constructors) { - if (cons.compare(c->name_hint) == 0) { - return c; - } - } - - LOG(FATAL) << adt << " does not contain constructor " << cons; -} - -tvm::Array IRModuleNode::GetGlobalTypeVars() const { - std::vector global_type_vars; - for (const auto& pair : global_type_var_map_) { - global_type_vars.push_back(pair.second); - } - return tvm::Array(global_type_vars); -} - void IRModuleNode::Add(const GlobalVar& var, const BaseFunc& f, bool update) { BaseFunc checked_func = f; - if (const auto* f = runtime::Registry::Get("relay.ir.WarnIfMalformed")) { - (*f)(GetRef(this), checked_func); - } AddUnchecked(var, checked_func); } @@ -238,44 +174,10 @@ void IRModuleNode::AddUnchecked(const GlobalVar& var, const BaseFunc& func) { global_var_map_.Set(var->name_hint, var); } -void IRModuleNode::RegisterConstructors(const GlobalTypeVar& var, const TypeData& type) { - // We hash the global type var name to use as a globally unique prefix for tags. - // The hash will be used as the most significant byte of the tag, with the index of - // the constructor in the less significant bytes - size_t hash = std::hash()(var->name_hint); - int32_t prefix = static_cast(hash & 0xff) << 24; - for (size_t i = 0; i < type->constructors.size(); ++i) { - type->constructors[i]->tag = prefix | static_cast(i); - constructor_tag_map_[type->constructors[i]->tag] = type->constructors[i]; - } -} - -void IRModuleNode::AddTypeDef(const GlobalTypeVar& var, const TypeData& type, bool update) { - // TODO(@jroesch): we have temporarily removed kind checking here, and will consolidate - // to the type checker in follow up PR. - AddTypeDefUnchecked(var, type, update); -} - -void IRModuleNode::AddTypeDefUnchecked(const GlobalTypeVar& var, const TypeData& type, - bool update) { - this->type_definitions.Set(var, type); - if (!update) { - // set global type var map - ICHECK(global_type_var_map_.count(var->name_hint) == 0) - << "Duplicate global type definition name " << var; - } - global_type_var_map_.Set(var->name_hint, var); - RegisterConstructors(var, type); -} - void IRModuleNode::Update(const GlobalVar& var, const BaseFunc& func) { this->Add(var, func, true); } -void IRModuleNode::UpdateTypeDef(const GlobalTypeVar& var, const TypeData& type) { - this->AddTypeDef(var, type, true); -} - void IRModuleNode::UpdateGlobalInfo(const String& name, const Array& info) { this->global_infos.Set(name, info); } @@ -298,28 +200,7 @@ BaseFunc IRModuleNode::Lookup(const String& name) const { return this->Lookup(id); } -TypeData IRModuleNode::LookupTypeDef(const GlobalTypeVar& var) const { - auto it = type_definitions.find(var); - ICHECK(it != type_definitions.end()) << "There is no definition of " << var; - return (*it).second; -} - -TypeData IRModuleNode::LookupTypeDef(const String& name) const { - GlobalTypeVar id = this->GetGlobalTypeVar(name); - return this->LookupTypeDef(id); -} - -Constructor IRModuleNode::LookupTag(const int32_t tag) { - auto it = constructor_tag_map_.find(tag); - ICHECK(it != constructor_tag_map_.end()) << "There is no constructor with the tag " << tag; - return (*it).second; -} - void IRModuleNode::Update(const IRModule& mod) { - if (const auto* f = runtime::Registry::Get("relay.ir.IRModuleUpdateWithRenamer")) { - (*f)(GetRef(this), mod); - return; - } for (auto pair : mod->functions) { // TODO(@jroesch): rename into IRModule. this->AddUnchecked(pair.first, pair.second); @@ -327,15 +208,12 @@ void IRModuleNode::Update(const IRModule& mod) { } IRModule IRModuleNode::ShallowCopy() { - return IRModule(this->functions, this->type_definitions, this->Imports(), this->source_map, - this->attrs, this->global_infos); + return IRModule(this->functions, this->source_map, this->attrs, this->global_infos); } -std::pair IRModule::FromExprInContext( - const RelayExpr& expr, const tvm::Map& global_funcs, - const tvm::Map& type_definitions, - std::unordered_set import_set) { - auto mod = IRModule(global_funcs, type_definitions, std::move(import_set)); +IRModule IRModule::FromExpr(const RelayExpr& expr, + const tvm::Map& global_funcs) { + auto mod = IRModule(global_funcs); String gv_name; // All global definitions must be functions. @@ -346,10 +224,6 @@ std::pair IRModule::FromExprInContext( // Function literal has been annotated with it's required global symbol. gv_name = opt.value(); } - } else if (const auto* f = runtime::Registry::Get("relay.ir.FunctionFromExprInContext")) { - func = (*f)(expr, mod); - } else { - LOG(FATAL) << "`relay.ir.FunctionFromExprInContext` is not registered"; } GlobalVar main_gv; @@ -361,47 +235,14 @@ std::pair IRModule::FromExprInContext( main_gv = global_var_supply->UniqueGlobalFor(gv_name, false); } mod->Add(main_gv, func); - return {mod, main_gv}; -} - -IRModule IRModule::FromExpr(const RelayExpr& expr, const Map& global_funcs, - const Map& type_definitions) { - return FromExprInContext(expr, global_funcs, type_definitions).first; -} - -void IRModuleNode::Import(const String& path) { - static const auto* f = runtime::Registry::Get("relay.parser.ParseModule"); - ICHECK(f != nullptr) << "ValueError: Relay parser is not available"; - if (this->import_set_.count(path) == 0) { - this->import_set_.insert(path); - std::fstream src_file(path, std::fstream::in); - std::string file_contents{std::istreambuf_iterator(src_file), - std::istreambuf_iterator()}; - auto mod_to_import = (*f)(path, file_contents, GetRef(this)); - Update(mod_to_import); - } -} - -void IRModuleNode::ImportFromStd(const String& path) { - auto* f = tvm::runtime::Registry::Get("tvm.relay.std_path"); - ICHECK(f != nullptr) << "The Relay std_path is not set, please register tvm.relay.std_path."; - std::string std_path = (*f)(); - this->Import(std_path + "/" + path); -} - -std::unordered_set IRModuleNode::Imports() const { return this->import_set_; } - -IRModule IRModule::FromText(const String& text, const String& source_path) { - static const auto* f = runtime::Registry::Get("relay.parser.ParseModule"); - ICHECK(f != nullptr) << "ValueError: Relay parser is not available"; - return (*f)(source_path, text, Optional()); + return mod; } TVM_REGISTER_NODE_TYPE(IRModuleNode); TVM_REGISTER_GLOBAL("ir.IRModule") - .set_body_typed([](tvm::Map funcs, tvm::Map types, - tvm::ObjectRef attrs, Map> global_infos) { + .set_body_typed([](tvm::Map funcs, tvm::ObjectRef attrs, + Map> global_infos) { auto dict_attrs = [&attrs]() { if (!attrs.defined()) { return DictAttrs(); @@ -414,7 +255,7 @@ TVM_REGISTER_GLOBAL("ir.IRModule") } }(); - return IRModule(funcs, types, {}, {}, dict_attrs, global_infos); + return IRModule(funcs, {}, dict_attrs, global_infos); }); TVM_REGISTER_GLOBAL("ir.Module_Clone").set_body_typed([](IRModule mod) -> IRModule { @@ -461,26 +302,15 @@ TVM_REGISTER_GLOBAL("ir.Module_Contains") } }); -TVM_REGISTER_GLOBAL("ir.Module_AddDef").set_body_method(&IRModuleNode::AddTypeDef); - TVM_REGISTER_GLOBAL("ir.Module_GetGlobalVar") .set_body_method(&IRModuleNode::GetGlobalVar); TVM_REGISTER_GLOBAL("ir.Module_GetGlobalVars") .set_body_method(&IRModuleNode::GetGlobalVars); -TVM_REGISTER_GLOBAL("ir.Module_GetGlobalTypeVars") - .set_body_method(&IRModuleNode::GetGlobalTypeVars); - TVM_REGISTER_GLOBAL("ir.Module_ContainGlobalVar") .set_body_method(&IRModuleNode::ContainGlobalVar); -TVM_REGISTER_GLOBAL("ir.Module_ContainGlobalTypeVar") - .set_body_method(&IRModuleNode::ContainGlobalTypeVar); - -TVM_REGISTER_GLOBAL("ir.Module_GetGlobalTypeVar") - .set_body_method(&IRModuleNode::GetGlobalTypeVar); - TVM_REGISTER_GLOBAL("ir.Module_Lookup").set_body_typed([](IRModule mod, GlobalVar var) { return mod->Lookup(var); }); @@ -489,18 +319,6 @@ TVM_REGISTER_GLOBAL("ir.Module_Lookup_str").set_body_typed([](IRModule mod, Stri return mod->Lookup(var); }); -TVM_REGISTER_GLOBAL("ir.Module_LookupDef").set_body_typed([](IRModule mod, GlobalTypeVar var) { - return mod->LookupTypeDef(var); -}); - -TVM_REGISTER_GLOBAL("ir.Module_LookupDef_str").set_body_typed([](IRModule mod, String var) { - return mod->LookupTypeDef(var); -}); - -TVM_REGISTER_GLOBAL("ir.Module_LookupTag").set_body_typed([](IRModule mod, int32_t tag) { - return mod->LookupTag(tag); -}); - TVM_REGISTER_GLOBAL("ir.Module_FromExpr").set_body_typed(&IRModule::FromExpr); TVM_REGISTER_GLOBAL("ir.Module_Update").set_body_typed([](IRModule mod, IRModule from) { @@ -515,14 +333,6 @@ TVM_REGISTER_GLOBAL("ir.Module_UpdateGlobalInfo") mod->UpdateGlobalInfo(name, global_info); }); -TVM_REGISTER_GLOBAL("ir.Module_Import").set_body_typed([](IRModule mod, String path) { - mod->Import(path); -}); - -TVM_REGISTER_GLOBAL("ir.Module_ImportFromStd").set_body_typed([](IRModule mod, String path) { - mod->ImportFromStd(path); -}); - TVM_REGISTER_GLOBAL("ir.Module_GetAttrs").set_body_typed([](IRModule mod) -> ObjectRef { return mod->GetAttrs(); }); diff --git a/src/ir/op.cc b/src/ir/op.cc index e80a10f84def..687ad1e83415 100644 --- a/src/ir/op.cc +++ b/src/ir/op.cc @@ -112,45 +112,6 @@ TVM_REGISTER_GLOBAL("ir.RegisterOp").set_body_typed([](String op_name, String de op.describe(descr); }); -// This is exposed FFI api for prototyping using in python. -// Note: it is not full of the C++ type relation, -// since in python side we don't have access to the type reporter, -// and cannot propagate constraints to the inputs, only to the output. -TVM_REGISTER_GLOBAL("ir.OpAddTypeRel") - .set_body_typed([](Op op, String rel_name, runtime::TVMArgValue value) { - auto& reg = OpRegistry::Global()->RegisterOrGet(op->name).set_name(); - if (value.type_code() == kTVMPackedFuncHandle) { - // do an eager copy of the PackedFunc to avoid deleting function from frontend. - PackedFunc fcopy = value; - auto f = [=](const Array& args, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) -> bool { - Array input_types(args.begin(), args.end() - 1); - // call customized relation functions - // *fcopy's signature: function (args: List[Type], attrs: Attrs) -> Type - Type ret_type = fcopy(input_types, attrs); - // when defined ret_type, inference of output type is ok, do type assign - // otherwise, inference failure happens - if (ret_type.defined()) { - // the last argument is output - // TODO(xqdan): support multiple outputs - reporter->Assign(args.back(), ret_type); - return true; - } - return false; - }; - // adjust function call to call conventions of relay type system with TypeReporter - auto type_rel = runtime::TypedPackedFunc&, int, const Attrs&, - const TypeReporter&)>(f); - reg.add_type_rel(rel_name, type_rel); - } else if (value.type_code() == kTVMNullptr) { - // Call relation functions of relay - auto func_name = std::string("tvm.relay.type_relation.") + rel_name; - auto* f = runtime::Registry::Get(func_name); - ICHECK(f != nullptr) << "AddTypeRel error: no type_relation registered."; - reg.add_type_rel(rel_name, *f); - } - }); - TVM_REGISTER_GLOBAL("ir.OpAddArgument") .set_body_typed([](Op op, String name, String type, String description) { auto& reg = OpRegistry::Global()->RegisterOrGet(op->name).set_name(); diff --git a/src/ir/si_builder.cc b/src/ir/si_builder.cc deleted file mode 100644 index c82b963d104d..000000000000 --- a/src/ir/si_builder.cc +++ /dev/null @@ -1,326 +0,0 @@ -/* - * 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 tvm/ir/si_builder.cc - * \brief Implementation for building a source info during rewriting expressions. - */ -#include -#include -#include - -#include - -namespace tvm { - -using RelayExprSet = std::unordered_set; -using PrimExprSet = std::unordered_set; -using StmtSet = std::unordered_set; - -class RelayCollectSpans : public relay::ExprVisitor { - public: - explicit RelayCollectSpans(const RelayExprSet& inputs = {}) : inputs_(inputs) {} - - // From entry to inputs, recursively collect spans. The spans of inputs are included. - Span CollectSpans(const relay::Expr& entry); - - void VisitExpr(const relay::Expr& expr) final; - - private: - Array spans_; - const RelayExprSet& inputs_; -}; - -void RelayCollectSpans::VisitExpr(const relay::Expr& expr) { - if (visit_counter_.count(expr.get())) { - return; - } - if (expr->span.defined()) { - spans_.push_back(expr->span); - } - if (inputs_.find(expr) != inputs_.end()) { - // becuase it returns directly, it should be recorded as visted manually. - visit_counter_.insert({expr.get(), 1}); - return; - } - relay::ExprVisitor::VisitExpr(expr); -} - -Span RelayCollectSpans::CollectSpans(const relay::Expr& entry) { - VisitExpr(entry); - return SequentialSpan(spans_); -} - -class RelayRecursivelyFill : public relay::ExprMutator { - public: - explicit RelayRecursivelyFill(const Span& span, const RelayExprSet& inputs = {}) - : span_(span), inputs_(inputs) {} - - // From entry until inputs, recursively fill spans into expressions. Inputs are not filled. - void Fill(const relay::Expr& entry); - - relay::Expr VisitExpr(const relay::Expr& expr) final; - - private: - const Span& span_; - const RelayExprSet& inputs_; -}; - -relay::Expr RelayRecursivelyFill::VisitExpr(const relay::Expr& expr) { - if (inputs_.find(expr) != inputs_.end()) { - return expr; - } - // Skip op node. Align with python frontend - if (!expr.as()) { - expr->span = span_; - } - - return relay::ExprMutator::VisitExpr(expr); -} - -void RelayRecursivelyFill::Fill(const relay::Expr& entry) { Mutate(entry); } - -class TirCollectSpans : public tir::StmtExprVisitor { - public: - explicit TirCollectSpans(const PrimExprSet& expr_inputs = {}, const StmtSet& stmt_inputs = {}) - : expr_inputs_(expr_inputs), stmt_inputs_(stmt_inputs) {} - - void VisitExpr(const PrimExpr& expr) final; - void VisitStmt(const tir::Stmt& stmt) final; - - bool IsInput(const PrimExpr& expr); - bool IsInput(const tir::Stmt& stmt); - - // From entry to inputs, recursively collect spans. The spans of inputs are included. - Span CollectSpans(const PrimExpr& expr); - // From entry to inputs, recursively collect spans. The spans of inputs are included. - Span CollectSpans(const tir::Stmt& stmt); - - private: - Array spans_; - std::unordered_map visit_counter_; - const PrimExprSet& expr_inputs_; - const StmtSet& stmt_inputs_; -}; - -Span TirCollectSpans::CollectSpans(const PrimExpr& expr) { - operator()(expr); - return SequentialSpan(spans_); -} - -Span TirCollectSpans::CollectSpans(const tir::Stmt& stmt) { - operator()(stmt); - return SequentialSpan(spans_); -} - -bool TirCollectSpans::IsInput(const PrimExpr& expr) { - return expr_inputs_.find(expr) != expr_inputs_.end(); -} - -bool TirCollectSpans::IsInput(const tir::Stmt& stmt) { - return stmt_inputs_.find(stmt) != stmt_inputs_.end(); -} - -void TirCollectSpans::VisitExpr(const PrimExpr& expr) { - if (visit_counter_.count(expr.get())) { - return; - } - if (expr->span.defined()) { - spans_.push_back(expr->span); - } - if (IsInput(expr)) { - // becuase it returns directly, it should be recorded as visted manually. - visit_counter_.insert({expr.get(), 1}); - return; - } - StmtExprVisitor::VisitExpr(expr); -} - -void TirCollectSpans::VisitStmt(const tir::Stmt& stmt) { - if (visit_counter_.count(stmt.get())) { - return; - } - if (stmt->span.defined()) { - spans_.push_back(stmt->span); - } - if (IsInput(stmt)) { - // becuase it returns directly, it should be recorded as visted manually. - visit_counter_.insert({stmt.get(), 1}); - return; - } - StmtExprVisitor::VisitStmt(stmt); -} - -class TirRecursivelyFill : public tir::StmtExprMutator { - public: - TirRecursivelyFill(const Span& span, const PrimExprSet& expr_inputs = {}, - const StmtSet& stmt_inputs = {}) - : span_(span), expr_inputs_(expr_inputs), stmt_inputs_(stmt_inputs) {} - - // From entry until inputs, recursively fill spans into expressions. Inputs are not filled. - tir::Stmt Fill(const tir::Stmt& s) { return operator()(s); } - // From entry until inputs, recursively fill spans into expressions. Inputs are not filled. - PrimExpr Fill(const PrimExpr& e) { return operator()(e); } - - bool IsInput(const PrimExpr& expr); - bool IsInput(const tir::Stmt& stmt); - - PrimExpr VisitExpr(const PrimExpr& expr) final; - tir::Stmt VisitStmt(const tir::Stmt& stmt) final; - - private: - const Span& span_; - const PrimExprSet& expr_inputs_; - const StmtSet& stmt_inputs_; -}; - -tir::Stmt TirRecursivelyFill::VisitStmt(const tir::Stmt& stmt) { - if (IsInput(stmt)) { - return stmt; - } - stmt->span = span_; - return StmtExprMutator::VisitStmt(stmt); -} - -bool TirRecursivelyFill::IsInput(const PrimExpr& expr) { - return expr_inputs_.find(expr) != expr_inputs_.end(); -} - -bool TirRecursivelyFill::IsInput(const tir::Stmt& stmt) { - return stmt_inputs_.find(stmt) != stmt_inputs_.end(); -} - -PrimExpr TirRecursivelyFill::VisitExpr(const PrimExpr& expr) { - if (IsInput(expr)) { - return expr; - } - expr->span = span_; - return StmtExprMutator::VisitExpr(expr); -} - -struct SIBuilder::Impl { - virtual ~Impl() {} - virtual Span Build() const { return Span(); } - virtual void RecursivelyFillSpan(const relay::Expr& entry, const RelayExprSet& inputs) const {} - virtual void RecursivelyFillSpan(const PrimExpr& entry, const PrimExprSet& inputs) const {} - virtual void RecursivelyFillSpan(const tir::Stmt& entry, const PrimExprSet& inputs) const {} - virtual void RecursivelyFillSpan(const tir::Stmt& entry, const StmtSet& inputs) const {} - virtual void CollectSpansSpan(const relay::Expr& entry, const RelayExprSet& inputs) {} - virtual void CollectSpansSpan(const PrimExpr& entry, const PrimExprSet& inputs) {} - virtual void CollectSpansSpan(const tir::Stmt& entry, const PrimExprSet& inputs) {} - virtual void CollectSpansSpan(const tir::Stmt& entry, const StmtSet& inputs) {} -}; - -SIBuilder::~SIBuilder() = default; - -Span SIBuilder::Build() const { return impl_->Build(); } - -template <> -void SIBuilder::RecursivelyFillSpan(const relay::Expr& entry, const RelayExprSet& inputs) const { - impl_->RecursivelyFillSpan(entry, inputs); -} - -template <> -void SIBuilder::RecursivelyFillSpan(const PrimExpr& entry, const PrimExprSet& inputs) const { - impl_->RecursivelyFillSpan(entry, inputs); -} - -void SIBuilder::RecursivelyFillSpan(const tir::Stmt& entry, const PrimExprSet& inputs) const { - impl_->RecursivelyFillSpan(entry, inputs); -} - -void SIBuilder::RecursivelyFillSpan(const tir::Stmt& entry, const StmtSet& inputs) const { - impl_->RecursivelyFillSpan(entry, inputs); -} - -std::unique_ptr SIBuilder::CreateImpl(const Span& span) { - struct Impl : public SIBuilder::Impl { - explicit Impl(const Span& span) : span_(span) {} - Span Build() const final { return span_; } - void RecursivelyFillSpan(const relay::Expr& entry, const RelayExprSet& inputs) const final { - RelayRecursivelyFill(Build(), inputs).Fill(entry); - } - void RecursivelyFillSpan(const PrimExpr& entry, const PrimExprSet& inputs) const final { - TirRecursivelyFill(Build(), inputs).Fill(entry); - } - void RecursivelyFillSpan(const tir::Stmt& entry, const PrimExprSet& inputs) const final { - TirRecursivelyFill(Build(), inputs).Fill(entry); - } - void RecursivelyFillSpan(const tir::Stmt& entry, const StmtSet& inputs) const final { - TirRecursivelyFill(Build(), {}, inputs).Fill(entry); - } - void CollectSpansSpan(const relay::Expr& entry, const RelayExprSet& inputs) final { - span_ = RelayCollectSpans(inputs).CollectSpans(entry); - } - void CollectSpansSpan(const PrimExpr& entry, const PrimExprSet& inputs) final { - span_ = TirCollectSpans(inputs).CollectSpans(entry); - } - void CollectSpansSpan(const tir::Stmt& entry, const PrimExprSet& inputs) final { - span_ = TirCollectSpans(inputs).CollectSpans(entry); - } - void CollectSpansSpan(const tir::Stmt& entry, const StmtSet& inputs) final { - span_ = TirCollectSpans({}, inputs).CollectSpans(entry); - } - - private: - Span span_; - }; - - const bool enable_si_builder = transform::PassContext::Current() - ->GetConfig("ir.enable_si_builder", Bool(false)) - .value(); - - if (enable_si_builder) { - return std::make_unique(span); - } - - return std::make_unique(); -} - -SIBuilder::SIBuilder(const Span& span) : impl_(CreateImpl(span)) {} -SIBuilder::SIBuilder(const Array& spans) : impl_(CreateImpl(SequentialSpan(spans))) {} -SIBuilder::SIBuilder(const std::initializer_list& init) - : impl_(CreateImpl(SequentialSpan(Array(init)))) {} - -template <> -SIBuilder::SIBuilder(const relay::Expr& expr, const Array& inputs) - : impl_(CreateImpl(Span())) { - impl_->CollectSpansSpan(expr, RelayExprSet(inputs.begin(), inputs.end())); -} - -template <> -SIBuilder::SIBuilder(const PrimExpr& expr, const Array& inputs) - : impl_(CreateImpl(Span())) { - impl_->CollectSpansSpan(expr, PrimExprSet(inputs.begin(), inputs.end())); -} - -SIBuilder::SIBuilder(const tir::Stmt& s, const Array& inputs) - : impl_(CreateImpl(Span())) { - impl_->CollectSpansSpan(s, PrimExprSet(inputs.begin(), inputs.end())); -} - -SIBuilder::SIBuilder(const tir::Stmt& s, const Array& inputs) - : impl_(CreateImpl(Span())) { - impl_->CollectSpansSpan(s, StmtSet(inputs.begin(), inputs.end())); -} - -// Register build pipeline related options -TVM_REGISTER_PASS_CONFIG_OPTION("ir.enable_si_builder", Bool); - -} // namespace tvm diff --git a/src/ir/type.cc b/src/ir/type.cc index b61a3df09107..3c648418c6a9 100644 --- a/src/ir/type.cc +++ b/src/ir/type.cc @@ -52,52 +52,19 @@ TVM_REGISTER_GLOBAL("ir.PointerType") return PointerType(element_type, storage_scope); }); -TypeVar::TypeVar(String name, TypeKind kind, Span span) { - ObjectPtr n = make_object(); - n->name_hint = std::move(name); - n->kind = std::move(kind); - n->span = std::move(span); - data_ = std::move(n); -} - -TVM_REGISTER_NODE_TYPE(TypeVarNode); - -TVM_REGISTER_GLOBAL("ir.TypeVar").set_body_typed([](String name, int kind) { - return TypeVar(name, static_cast(kind)); -}); - -GlobalTypeVar::GlobalTypeVar(String name, TypeKind kind, Span span) { - ObjectPtr n = make_object(); - n->name_hint = std::move(name); - n->kind = std::move(kind); - n->span = std::move(span); - data_ = std::move(n); -} - -TVM_REGISTER_NODE_TYPE(GlobalTypeVarNode); - -TVM_REGISTER_GLOBAL("ir.GlobalTypeVar").set_body_typed([](String name, int kind) { - return GlobalTypeVar(name, static_cast(kind)); -}); - -FuncType::FuncType(tvm::Array arg_types, Type ret_type, tvm::Array type_params, - tvm::Array type_constraints, Span span) { +FuncType::FuncType(tvm::Array arg_types, Type ret_type, Span span) { ObjectPtr n = make_object(); n->arg_types = std::move(arg_types); n->ret_type = std::move(ret_type); - n->type_params = std::move(type_params); - n->type_constraints = std::move(type_constraints); n->span = std::move(span); data_ = std::move(n); } TVM_REGISTER_NODE_TYPE(FuncTypeNode); -TVM_REGISTER_GLOBAL("ir.FuncType") - .set_body_typed([](tvm::Array arg_types, Type ret_type, tvm::Array type_params, - tvm::Array type_constraints) { - return FuncType(arg_types, ret_type, type_params, type_constraints); - }); +TVM_REGISTER_GLOBAL("ir.FuncType").set_body_typed([](tvm::Array arg_types, Type ret_type) { + return FuncType(arg_types, ret_type); +}); TupleType::TupleType(Array fields, Span span) { ObjectPtr n = make_object(); @@ -114,30 +81,4 @@ TVM_REGISTER_GLOBAL("ir.TupleType").set_body_typed([](Array fields) { return TupleType(fields); }); -IncompleteType::IncompleteType(TypeKind kind, Span span) { - auto n = make_object(); - n->kind = std::move(kind); - n->span = std::move(span); - data_ = std::move(n); -} - -TVM_REGISTER_NODE_TYPE(IncompleteTypeNode); - -TVM_REGISTER_GLOBAL("ir.IncompleteType").set_body_typed([](int kind) { - return IncompleteType(static_cast(kind)); -}); - -RelayRefType::RelayRefType(Type value, Span span) { - ObjectPtr n = make_object(); - n->value = std::move(value); - n->span = std::move(span); - data_ = std::move(n); -} - -TVM_REGISTER_GLOBAL("ir.RelayRefType").set_body_typed([](Type value) { - return RelayRefType(value); -}); - -TVM_REGISTER_NODE_TYPE(RelayRefTypeNode); - } // namespace tvm diff --git a/src/ir/type_functor.cc b/src/ir/type_functor.cc index 4a69c64fbd3b..1bfc435cff72 100644 --- a/src/ir/type_functor.cc +++ b/src/ir/type_functor.cc @@ -27,21 +27,9 @@ namespace tvm { -void TypeVisitor::VisitType_(const TypeVarNode* op) {} - void TypeVisitor::VisitType_(const TensorTypeNode* op) {} -void TypeVisitor::VisitType_(const IncompleteTypeNode* op) {} - void TypeVisitor::VisitType_(const FuncTypeNode* op) { - for (auto type_param : op->type_params) { - this->VisitType(type_param); - } - - for (auto type_cs : op->type_constraints) { - this->VisitType(type_cs); - } - for (auto arg_type : op->arg_types) { this->VisitType(arg_type); } @@ -54,37 +42,6 @@ void TypeVisitor::VisitType_(const TupleTypeNode* op) { } } -void TypeVisitor::VisitType_(const RelayRefTypeNode* op) { this->VisitType(op->value); } - -void TypeVisitor::VisitType_(const TypeRelationNode* op) { - for (const Type& t : op->args) { - this->VisitType(t); - } -} - -void TypeVisitor::VisitType_(const GlobalTypeVarNode* op) {} - -void TypeVisitor::VisitType_(const TypeCallNode* op) { - this->VisitType(op->func); - for (const Type& t : op->args) { - this->VisitType(t); - } -} - -void TypeVisitor::VisitType_(const TypeDataNode* op) { - this->VisitType(op->header); - for (const auto& v : op->type_vars) { - this->VisitType(v); - } - - for (const auto& c : op->constructors) { - this->VisitType(c->belong_to); - for (const auto& t : c->inputs) { - this->VisitType(t); - } - } -} - void TypeVisitor::VisitType_(const PrimTypeNode* op) {} void TypeVisitor::VisitType_(const PointerTypeNode* op) { this->VisitType(op->element_type); } @@ -100,38 +57,13 @@ Array TypeMutator::MutateArray(Array arr) { return arr.Map([this](const Type& ty) { return VisitType(ty); }); } -Type TypeMutator::VisitType_(const TypeVarNode* op) { return GetRef(op); } - Type TypeMutator::VisitType_(const TensorTypeNode* op) { // TODO(tvm-team) recursively visit to replace Var return GetRef(op); } -Type TypeMutator::VisitType_(const IncompleteTypeNode* op) { return GetRef(op); } - Type TypeMutator::VisitType_(const FuncTypeNode* op) { bool changed = false; - Array type_params; - for (auto type_param : op->type_params) { - auto new_type_param = VisitType(type_param); - changed = changed || !new_type_param.same_as(type_param); - if (auto tin = new_type_param.as()) { - type_params.push_back(tin.value()); - } else { - LOG(FATAL) << new_type_param; - } - } - - Array type_constraints; - for (auto type_cs : op->type_constraints) { - auto new_type_cs = VisitType(type_cs); - changed = changed || !new_type_cs.same_as(type_cs); - if (auto tin = new_type_cs.as()) { - type_constraints.push_back(tin.value()); - } else { - LOG(FATAL) << new_type_cs; - } - } Array new_args = MutateArray(op->arg_types); changed = changed || !new_args.same_as(op->arg_types); @@ -140,7 +72,7 @@ Type TypeMutator::VisitType_(const FuncTypeNode* op) { changed = changed || !new_ret_type.same_as(op->ret_type); if (!changed) return GetRef(op); - return FuncType(new_args, new_ret_type, type_params, type_constraints); + return FuncType(new_args, new_ret_type); } Type TypeMutator::VisitType_(const TupleTypeNode* op) { @@ -152,33 +84,6 @@ Type TypeMutator::VisitType_(const TupleTypeNode* op) { } } -Type TypeMutator::VisitType_(const RelayRefTypeNode* op) { - return RelayRefType(this->VisitType(op->value)); -} - -Type TypeMutator::VisitType_(const TypeRelationNode* type_rel) { - Array new_args = MutateArray(type_rel->args); - if (new_args.same_as(type_rel->args)) { - return GetRef(type_rel); - } else { - return TypeRelation(type_rel->func, new_args, type_rel->num_inputs, type_rel->attrs); - } -} - -Type TypeMutator::VisitType_(const GlobalTypeVarNode* op) { return GetRef(op); } - -Type TypeMutator::VisitType_(const TypeCallNode* op) { - Type new_func = VisitType(op->func); - Array new_args = MutateArray(op->args); - if (new_args.same_as(op->args) && new_func.same_as(op->func)) { - return GetRef(op); - } else { - return TypeCall(new_func, new_args); - } -} - -Type TypeMutator::VisitType_(const TypeDataNode* op) { return GetRef(op); } - Type TypeMutator::VisitType_(const PrimTypeNode* op) { return GetRef(op); } Type TypeMutator::VisitType_(const PointerTypeNode* op) { @@ -191,27 +96,4 @@ Type TypeMutator::VisitType_(const PointerTypeNode* op) { } } -// Implements bind. -class TypeBinder : public TypeMutator { - public: - explicit TypeBinder(const tvm::Map& args_map) : args_map_(args_map) {} - - Type VisitType_(const TypeVarNode* op) override { - auto id = GetRef(op); - auto it = args_map_.find(id); - if (it != args_map_.end()) { - return (*it).second; - } else { - return std::move(id); - } - } - - private: - const tvm::Map& args_map_; -}; - -Type Bind(const Type& type, const tvm::Map& args_map) { - return TypeBinder(args_map).VisitType(type); -} - } // namespace tvm diff --git a/src/ir/type_relation.cc b/src/ir/type_relation.cc deleted file mode 100644 index f038a6678b42..000000000000 --- a/src/ir/type_relation.cc +++ /dev/null @@ -1,69 +0,0 @@ -/* - * 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 src/ir/type_relation.cc - * \brief Type relation - */ -#include -#include -#include -namespace tvm { - -TypeCall::TypeCall(Type func, tvm::Array args) { - ObjectPtr n = make_object(); - n->func = std::move(func); - n->args = std::move(args); - data_ = std::move(n); -} - -TVM_REGISTER_NODE_TYPE(TypeCallNode); - -TVM_REGISTER_GLOBAL("ir.TypeCall").set_body_typed([](Type func, Array type) { - return TypeCall(func, type); -}); - -TVM_STATIC_IR_FUNCTOR(ReprPrinter, vtable) - .set_dispatch([](const ObjectRef& ref, ReprPrinter* p) { - auto* node = static_cast(ref.get()); - p->stream << "TypeCallNode(" << node->func << ", " << node->args << ")"; - }); - -TypeRelation::TypeRelation(TypeRelationFn func, Array args, int num_inputs, Attrs attrs) { - ObjectPtr n = make_object(); - n->func = std::move(func); - n->args = std::move(args); - n->num_inputs = num_inputs; - n->attrs = std::move(attrs); - data_ = std::move(n); -} - -TVM_REGISTER_NODE_TYPE(TypeRelationNode); - -TVM_REGISTER_GLOBAL("ir.TypeRelation") - .set_body_typed([](TypeRelationFn func, Array args, int num_inputs, Attrs attrs) { - return TypeRelation(func, args, num_inputs, attrs); - }); - -TVM_STATIC_IR_FUNCTOR(ReprPrinter, vtable) - .set_dispatch([](const ObjectRef& ref, ReprPrinter* p) { - auto* node = static_cast(ref.get()); - p->stream << "TypeRelationNode(" << node->func->name << ", " << node->args << ")"; - }); -} // namespace tvm diff --git a/src/meta_schedule/utils.h b/src/meta_schedule/utils.h index 61ce62347af7..0d8df31e1f69 100644 --- a/src/meta_schedule/utils.h +++ b/src/meta_schedule/utils.h @@ -600,7 +600,7 @@ class BlockCollector : public tir::StmtVisitor { } else { for (const auto& [gv, base_func] : sch_->mod()->functions) { // `gv->name_hint` is the name of the function - // `base_func` can be PrimFunc or relay::Function + // `base_func` can be PrimFunc or relax::Function if (const auto* func = base_func.as()) { f_collect(GetRef(func), gv->name_hint); } diff --git a/src/node/serialization.cc b/src/node/serialization.cc index 09eb02e10bfa..8b19104dfaeb 100644 --- a/src/node/serialization.cc +++ b/src/node/serialization.cc @@ -26,8 +26,6 @@ #include #include #include -#include -#include #include #include #include @@ -136,21 +134,6 @@ class NodeIndexer : public AttrVisitor { MakeIndex(const_cast(kv.second.get())); } } - } else if (node->IsInstance()) { - auto pre_visit = [this](const relay::LetNode* op) { - MakeNodeIndex(const_cast(static_cast(op))); - MakeIndex(const_cast(static_cast(op->var.get()))); - MakeIndex(const_cast(static_cast(op->value.get()))); - MakeIndex(const_cast(static_cast(op->span.get()))); - MakeIndex(const_cast(static_cast(op->checked_type_.get()))); - if (!op->body.as()) { - MakeIndex(const_cast(static_cast(op->body.get()))); - } - }; - auto post_visit = [](const relay::LetNode* op) {}; - if (!reflection_->GetReprBytes(node, nullptr)) { - relay::ExpandANormalForm(static_cast(node), pre_visit, post_visit); - } } else { // if the node already have repr bytes, no need to visit Attrs. if (!reflection_->GetReprBytes(node, nullptr)) { diff --git a/src/relay/analysis/graph_partitioner.cc b/src/relax/analysis/graph_partitioner.cc similarity index 96% rename from src/relay/analysis/graph_partitioner.cc rename to src/relax/analysis/graph_partitioner.cc index d233d43ad7eb..53f66f42160a 100644 --- a/src/relay/analysis/graph_partitioner.cc +++ b/src/relax/analysis/graph_partitioner.cc @@ -22,7 +22,7 @@ #include namespace tvm { -namespace relay { +namespace relax { DominatorTree DominatorTree::PostDom(support::Arena* arena, const IndexedForwardGraph& graph) { DominatorTree tree; @@ -156,7 +156,7 @@ bool GraphPartitioner::CheckPath(IndexedForwardGraph::Node* src, IndexedForwardG } OpPatternKind CombinePattern(OpPatternKind lhs, OpPatternKind rhs) { - if (lhs > relay::kBroadcast && rhs > relay::kBroadcast) { + if (lhs > kBroadcast && rhs > kBroadcast) { LOG(FATAL) << "Cannot merge two complex group together"; } if (lhs > rhs) return lhs; @@ -226,14 +226,8 @@ size_t GraphPartitioner::CountFusedNodesWithNewChild(IndexedForwardGraph::Node* } size_t GraphPartitioner::CountAdditionalArgs_(const TensorTypeNode* ttype, bool with_strides) { - size_t any_dims = 0; - for (const auto& dim : ttype->shape) { - if (dim.as()) { - any_dims++; - } - } - if (with_strides && any_dims > 0) any_dims += ttype->shape.size(); - return any_dims; + // TODO(@syfeng): need to clean this up + return 0; } size_t GraphPartitioner::CountArgs_(IndexedForwardGraph::Node* src, @@ -244,7 +238,7 @@ size_t GraphPartitioner::CountArgs_(IndexedForwardGraph::Node* src, auto sum = gnode->args_num; visited_groups.insert(gnode->FindRoot()); auto calc_args_number = [this, src, &graph, &visited_groups, - update_postpone](const relay::Expr& arg) -> size_t { + update_postpone](const Expr& arg) -> size_t { if (arg.as()) return 0; auto* node = graph.node_map.at(arg.get()); Group* prev_group = groups_[node->index]->FindRoot(); @@ -339,7 +333,7 @@ void GraphPartitioner::InitGroups(const IndexedForwardGraph& graph) { group_node->pattern = graph_node->pattern; group_node->root_ref = graph_node->ref; // set anchor ref if necessary. - if (group_node->pattern == relay::kOutEWiseFusable) { + if (group_node->pattern == kOutEWiseFusable) { group_node->anchor_ref = graph_node->ref; } group_node->args_num = args_counter(graph_node->ref); @@ -393,12 +387,12 @@ void GraphPartitioner::RunFuse(const IndexedForwardGraph& graph, // if (phase == 2) { // Fuse injective ops into intermediate tuples, if any - if (group_node->pattern > relay::kInjective) continue; + if (group_node->pattern > kInjective) continue; Group* dom_parent_group = groups_[dom_parent_gindex]; Group* dom_root_group = dom_parent_group->FindRoot(); // If dom node group has a tuple as its root, we do not fuse tuple fields into it - if (dom_root_group->pattern == relay::kTuple) continue; - if (dom_parent_group->pattern == kTuple && dom_root_group->pattern <= relay::kInjective) { + if (dom_root_group->pattern == kTuple) continue; + if (dom_parent_group->pattern == kTuple && dom_root_group->pattern <= kInjective) { // Now we know the tuple has been fused into subsequent injective ops auto fcond = [](OpPatternKind kind, bool is_sink) { return kind <= kInjective; }; // dom_root_group can also be tuple, as in inception layers @@ -466,5 +460,5 @@ void GraphPartitioner::RunFuse(const IndexedForwardGraph& graph, // } } -} // namespace relay +} // namespace relax } // namespace tvm diff --git a/src/relay/analysis/graph_partitioner.h b/src/relax/analysis/graph_partitioner.h similarity index 97% rename from src/relay/analysis/graph_partitioner.h rename to src/relax/analysis/graph_partitioner.h index b3f934b972de..6abf10c5a602 100644 --- a/src/relay/analysis/graph_partitioner.h +++ b/src/relax/analysis/graph_partitioner.h @@ -18,14 +18,14 @@ */ /*! - * \file src/relay/analysis/graph_partitioner.h + * \file src/relax/analysis/graph_partitioner.h * \brief The helper function for op fusion. */ -#ifndef TVM_RELAY_ANALYSIS_GRAPH_PARTITIONER_H_ -#define TVM_RELAY_ANALYSIS_GRAPH_PARTITIONER_H_ +#ifndef TVM_RELAX_ANALYSIS_GRAPH_PARTITIONER_H_ +#define TVM_RELAX_ANALYSIS_GRAPH_PARTITIONER_H_ -#include +#include #include #include @@ -34,7 +34,7 @@ #include "../../support/arena.h" namespace tvm { -namespace relay { +namespace relax { using support::LinkedList; using support::LinkNode; @@ -304,6 +304,6 @@ class GraphPartitioner { void RunFuse(const IndexedForwardGraph& graph, const DominatorTree& post_dom_tree, int phase); }; -} // namespace relay +} // namespace relax } // namespace tvm -#endif // TVM_RELAY_ANALYSIS_GRAPH_PARTITIONER_H_ +#endif // TVM_RELAX_ANALYSIS_GRAPH_PARTITIONER_H_ diff --git a/src/relax/analysis/layout_transformation.cc b/src/relax/analysis/layout_transformation.cc index 2e850fa9dee3..25ce1a58f276 100644 --- a/src/relax/analysis/layout_transformation.cc +++ b/src/relax/analysis/layout_transformation.cc @@ -454,7 +454,7 @@ class BlockAnalyzer : public StmtExprVisitor { spatial_dom_.Set(v->var, v->dom); continue; } - if (v->iter_type == kCommReduce) continue; + if (v->iter_type == tir::kCommReduce) continue; LOG(WARNING) << "[LayoutInference] Cannot compute block spatial domain in presence of " "unknown block iter_type : " << v->iter_type; diff --git a/src/relax/analysis/struct_info_analysis.cc b/src/relax/analysis/struct_info_analysis.cc index 6fe8f36020bf..50931cbf38a2 100644 --- a/src/relax/analysis/struct_info_analysis.cc +++ b/src/relax/analysis/struct_info_analysis.cc @@ -66,7 +66,7 @@ class StaticTypeDeriver : public StructInfoFunctor { Array params = op->params.value().Map( [this](const StructInfo& sinfo) { return this->VisitStructInfo(sinfo); }); Type ret = this->VisitStructInfo(op->ret); - return FuncType(params, ret, {}, {}, op->span); + return FuncType(params, ret, op->span); } }; diff --git a/src/relax/analysis/tir_op_pattern_kind.cc b/src/relax/analysis/tir_op_pattern_kind.cc index 44a888d7e6c9..fe7b7bbeb547 100644 --- a/src/relax/analysis/tir_op_pattern_kind.cc +++ b/src/relax/analysis/tir_op_pattern_kind.cc @@ -19,6 +19,7 @@ #include #include +#include #include #include #include @@ -54,7 +55,7 @@ class PatternKindAnalyzer : public StmtExprVisitor { // We only support one buffer store in a block (usually generated by TE compute) // If we have already seen buffer store in the current block, classify as Opaque. if (store_.defined() && !IsSameArray(op->indices, store_.value()->indices)) { - kind_ = relay::kOpaque; + kind_ = kOpaque; return; } store_ = GetRef(op); @@ -82,14 +83,14 @@ class PatternKindAnalyzer : public StmtExprVisitor { // We support exactly one buffer store in a block (usually generated by TE compute) // If we have not seen any store in the current block, classify as Opaque. if (!store_.defined()) { - kind_ = relay::kOpaque; + kind_ = kOpaque; return; } BufferStore store = store_.value(); // Step 3. Checking load store indices pattern - relay::OpPatternKind index_pair_pattern = relay::kElemWise; + OpPatternKind index_pair_pattern = kElemWise; bool has_elem_wise = false; for (const BufferLoad& load : loads_) { // Since elemwise is stricter than broadcast and broadcast is stricter than injective, @@ -99,27 +100,27 @@ class PatternKindAnalyzer : public StmtExprVisitor { // Buffer C and A are elemwise but C and B are broadcast. So the whole block follows // broadcast pattern. if (IsElemwisePattern(store, load)) { - index_pair_pattern = std::max(index_pair_pattern, relay::kElemWise); + index_pair_pattern = std::max(index_pair_pattern, kElemWise); has_elem_wise = true; } else if (IsBroadcastPattern(store, load)) { - index_pair_pattern = std::max(index_pair_pattern, relay::kBroadcast); + index_pair_pattern = std::max(index_pair_pattern, kBroadcast); } else if (IsInjectivePattern(store, load)) { - index_pair_pattern = std::max(index_pair_pattern, relay::kInjective); + index_pair_pattern = std::max(index_pair_pattern, kInjective); } else { - index_pair_pattern = relay::kOpaque; + index_pair_pattern = kOpaque; break; } } // If there is a index pair is kElemWise and others are kBroadcast, we regard it as kElemWise // e.g. A[i, j] = B[i, j] + C[i] - if (index_pair_pattern == relay::kBroadcast && has_elem_wise) { - index_pair_pattern = relay::kElemWise; + if (index_pair_pattern == kBroadcast && has_elem_wise) { + index_pair_pattern = kElemWise; } // If the block index pattern is not opaque, update kind. - if (index_pair_pattern != relay::kOpaque) { + if (index_pair_pattern != kOpaque) { // This rule for softmax: reduce + injective. - if (IsOutputBlock(op) && kind_ == relay::kCommReduce) { - kind_ = relay::kOutEWiseFusable; + if (IsOutputBlock(op) && kind_ == kCommReduce) { + kind_ = kOutEWiseFusable; } else { kind_ = std::max(kind_, index_pair_pattern); } @@ -130,7 +131,7 @@ class PatternKindAnalyzer : public StmtExprVisitor { bool has_reduction = false; Array reduce_vars; for (const IterVar& it : op->iter_vars) { - if (it->iter_type == kCommReduce) { + if (it->iter_type == tir::IterVarType::kCommReduce) { has_reduction = true; reduce_vars.push_back(it->var); } @@ -139,21 +140,21 @@ class PatternKindAnalyzer : public StmtExprVisitor { if (has_reduction) { if (IsFMA(op->body)) { // FMA is regards as kOutEWiseFusable, e.g. Matmul or Conv. - kind_ = std::max(kind_, relay::kOutEWiseFusable); + kind_ = std::max(kind_, kOutEWiseFusable); return; } else { for (size_t i = 0; i < loads_.size(); ++i) { // If it's not a pure reduce, regards as kOutEWiseFusable. // This rule works for pooling for now. if (!IsPureReducePattern(reduce_vars, loads_[i]->indices)) { - kind_ = std::max(kind_, relay::kOutEWiseFusable); + kind_ = std::max(kind_, kOutEWiseFusable); return; } } } - kind_ = std::max(kind_, relay::kCommReduce); + kind_ = std::max(kind_, kCommReduce); } else { - kind_ = relay::kOpaque; + kind_ = kOpaque; } } @@ -335,15 +336,15 @@ class PatternKindAnalyzer : public StmtExprVisitor { /*! \brief The BufferLoad nodes in the current block. */ Array loads_; /*! \brief The result of op pattern. */ - relay::OpPatternKind kind_ = relay::kElemWise; + OpPatternKind kind_ = kElemWise; /*! \brief The buffers from function params. I.e. the input and output buffers. */ std::unordered_set param_buffers_; public: - relay::OpPatternKind GetResult() { return kind_; } + OpPatternKind GetResult() { return kind_; } }; -relay::OpPatternKind AnalyzeOpPatternKind(const PrimFunc& func) { +OpPatternKind AnalyzeOpPatternKind(const PrimFunc& func) { PatternKindAnalyzer analyzer(func); analyzer(func->body); return analyzer.GetResult(); diff --git a/src/relay/backend/contrib/codegen_c/codegen_c.h b/src/relax/backend/contrib/codegen_c/codegen_c.h similarity index 91% rename from src/relay/backend/contrib/codegen_c/codegen_c.h rename to src/relax/backend/contrib/codegen_c/codegen_c.h index cdbfbed8db89..67e97cc99dbf 100644 --- a/src/relay/backend/contrib/codegen_c/codegen_c.h +++ b/src/relax/backend/contrib/codegen_c/codegen_c.h @@ -18,15 +18,14 @@ */ /*! - * \file src/relay/backend/contrib/codegen_c/codegen_c.h + * \file src/relax/backend/contrib/codegen_c/codegen_c.h * \brief The base class for external codegen tools. */ -#ifndef TVM_RELAY_BACKEND_CONTRIB_CODEGEN_C_CODEGEN_C_H_ -#define TVM_RELAY_BACKEND_CONTRIB_CODEGEN_C_CODEGEN_C_H_ +#ifndef TVM_RELAX_BACKEND_CONTRIB_CODEGEN_C_CODEGEN_C_H_ +#define TVM_RELAX_BACKEND_CONTRIB_CODEGEN_C_CODEGEN_C_H_ -#include -#include -#include +#include +#include #include #include @@ -34,7 +33,7 @@ #include namespace tvm { -namespace relay { +namespace relax { namespace contrib { struct Output { @@ -51,24 +50,6 @@ struct GenerateBodyOutput { Array headers; }; -class CSourceModuleCodegenBase { - public: - CSourceModuleCodegenBase() = default; - virtual ~CSourceModuleCodegenBase() = default; - - /*! - * \brief Create a runtime module for the external library. For example, it - * could be a CSourceModule that can be directly compiled and linked together - * with a DSOModule, or a json style module that emitts a json artifact that - * is able to be executed by a customized json runtime. - * - * \param ref The ext_func Relay expression/module to be executed using extern ops. - * - * \return A runtime module. - */ - virtual runtime::Module CreateCSourceModule(const ObjectRef& ref) = 0; -}; - // The base class to generate the declaration functions in C. class CodegenCBase { public: @@ -447,14 +428,8 @@ class CodegenCBase { int indent_{0}; }; -/*! - * \brief A pass to translate all "Primitive" Relay functions with "Compiler=ccompiler" to - * a \p CSourceModule. - */ -transform::Pass CCompilerPass(); - } // namespace contrib -} // namespace relay +} // namespace relax } // namespace tvm -#endif // TVM_RELAY_BACKEND_CONTRIB_CODEGEN_C_CODEGEN_C_H_ +#endif // TVM_RELAX_BACKEND_CONTRIB_CODEGEN_C_CODEGEN_C_H_ diff --git a/src/relax/backend/contrib/cutlass/codegen.cc b/src/relax/backend/contrib/cutlass/codegen.cc index 8ae0036db76d..980243cf8128 100644 --- a/src/relax/backend/contrib/cutlass/codegen.cc +++ b/src/relax/backend/contrib/cutlass/codegen.cc @@ -21,7 +21,6 @@ * \file src/relax/backend/contrib/cutlass/codegen.cc * \brief Implementation of the CUTLASS code generator for Relax. */ -#include "../../../../relay/backend/contrib/cutlass/codegen.h" #include #include @@ -34,22 +33,116 @@ #include #include -#include "../../../../relay/backend/contrib/codegen_c/codegen_c.h" +#include "../codegen_c/codegen_c.h" #include "../utils.h" namespace tvm { namespace relax { namespace contrib { -using namespace relay::contrib::cutlass; +std::string EmitSignature(const std::vector& out, const std::string& func_id, + const std::vector& arg_names) { + std::ostringstream code_stream_; + code_stream_ << "void " << func_id << "_("; + for (const auto& arg_name : arg_names) { + code_stream_ << "DLTensor* " << arg_name << ", "; + } + for (size_t i = 0; i < out.size() - 1; ++i) { + code_stream_ << "DLTensor* out" << i << ", "; + } + code_stream_ << "DLTensor* out" << out.size() - 1 << ")"; + return code_stream_.str(); +} + +runtime::Module Finalize(const std::string& code, const Array& func_names) { + ICHECK(!func_names.empty()) + << "Should only create CUTLASS CSourceModule if there is at least one CUTLASS partition"; + + std::ostringstream default_headers; + default_headers << "#include \n"; + default_headers << "#include \n"; + default_headers << "#include \n"; + default_headers << "#include \n"; + default_headers << "#include \n"; + default_headers << "#include \n"; + default_headers << "#include \n"; + + const auto* pf = runtime::Registry::Get("runtime.CSourceModuleCreate"); + ICHECK(pf != nullptr) << "Cannot find CSource module to create the external runtime module"; + VLOG(1) << "Generated CUTLASS code:" << std::endl << code; + return (*pf)(default_headers.str() + code, "cu", func_names, /*const_vars=*/Array()); +} + +class CodegenResultNode : public Object { + public: + String code; + Array headers; + + void VisitAttrs(AttrVisitor* v) { + v->Visit("code", &code); + v->Visit("headers", &headers); + } + static constexpr const char* _type_key = "contrib.cutlass.CodegenResult"; + TVM_DECLARE_FINAL_OBJECT_INFO(CodegenResultNode, Object); +}; + +class CodegenResult : public ObjectRef { + public: + CodegenResult(String code, Array headers) { + auto n = make_object(); + n->code = std::move(code); + n->headers = std::move(headers); + data_ = std::move(n); + } + + TVM_DEFINE_OBJECT_REF_METHODS(CodegenResult, ObjectRef, CodegenResultNode) +}; + +TVM_REGISTER_NODE_TYPE(CodegenResultNode); + +TVM_REGISTER_GLOBAL("contrib.cutlass.CodegenResult") + .set_body_typed([](String code, Array headers) { + return CodegenResult(code, headers); + }); + +GenerateBodyOutput GenerateBody(const std::string& func_name, const std::string& ext_func_id, + const std::vector& output_types, + const Array& func_args, const Map& attrs, + int* buf_idx) { + // Make function call with input buffers when visiting arguements + ICHECK_GT(func_args.size(), 0); + std::ostringstream decl_stream; + decl_stream << "(" << func_args[0]; + for (size_t i = 1; i < func_args.size(); ++i) { + decl_stream << ", " << func_args[i]; + } + GenerateBodyOutput ret; + for (const auto& out_type : output_types) { + const std::string out = "out" + std::to_string(*buf_idx++); + decl_stream << ", " << out; + Output output; + output.name = out; + output.dtype = out_type; + output.need_copy = false; + ret.outputs.push_back(output); + } + decl_stream << ");"; + + const auto* instantiate_template_func = + runtime::Registry::Get("contrib.cutlass.instantiate_template"); + ICHECK(instantiate_template_func); + + CodegenResult codegen_res = (*instantiate_template_func)(func_name, attrs, func_args); + ret.decl = codegen_res->code; + ret.headers = codegen_res->headers; + + return ret; +} -using Output = relay::contrib::Output; -using GenerateBodyOutput = relay::contrib::GenerateBodyOutput; -using relay::contrib::cutlass::GenerateBody; using OutputType = std::vector; class CodegenCutlass : public relax::MemoizedExprTranslator, - public relay::contrib::CodegenCBase { + public relax::contrib::CodegenCBase { public: CodegenCutlass(const std::string& id, const Map& bindings) : ext_func_id_(id), bindings_(bindings) {} @@ -119,7 +212,7 @@ class CodegenCutlass : public relax::MemoizedExprTranslator, return ret.outputs; } - OutputType VisitExpr_(const FunctionNode* fn) { + OutputType VisitExpr_(const FunctionNode* fn) final { ICHECK(fn->GetAttr(attr::kComposite).defined()) << "JSON runtime only supports composite functions"; // FunctionNode should be handled by the caller. @@ -169,7 +262,7 @@ class CodegenCutlass : public relax::MemoizedExprTranslator, return outputs; } - OutputType VisitExpr_(const SeqExprNode* op) { + OutputType VisitExpr_(const SeqExprNode* op) final { OutputType outputs; for (BindingBlock block : op->blocks) { diff --git a/src/relax/backend/contrib/tensorrt/codegen.cc b/src/relax/backend/contrib/tensorrt/codegen.cc index 5ce6bf5e7d42..4a0a3ea5e12f 100644 --- a/src/relax/backend/contrib/tensorrt/codegen.cc +++ b/src/relax/backend/contrib/tensorrt/codegen.cc @@ -22,6 +22,7 @@ * \brief Implementation of the TensorRT JSON serializer. */ #include +#include // TODO(sunggg): add operator attribute when it's ready // #include #include diff --git a/src/relax/backend/vm/codegen_vm.cc b/src/relax/backend/vm/codegen_vm.cc index ca2d4d4fdb2e..8c0ddeb6c34d 100644 --- a/src/relax/backend/vm/codegen_vm.cc +++ b/src/relax/backend/vm/codegen_vm.cc @@ -33,7 +33,7 @@ #include #include -#include "../../../target/metadata_module.h" +#include "../../../runtime/const_loader_module.h" #include "../../../target/source/codegen_source_base.h" namespace tvm { @@ -428,30 +428,66 @@ IRModule VMCodeGen(ExecBuilder exec_builder, IRModule mod) { TVM_REGISTER_GLOBAL("relax.VMCodeGen").set_body_typed(VMCodeGen); +/*! + * \brief Link the modules together, possibly create a constant module. + * + * \param params The metadata for initialization of all modules. + * \param lib the internal module that is compiled by tvm. + * \param ext_libs The external modules that needs to be imported inside the metadata + * module(s). + * \return The created module. + */ +void LinkModules(ObjectPtr exec, const Map& params, + const tvm::runtime::Module& lib, const Array& ext_libs) { + // query if we need const loader for ext_modules + // Wrap all submodules in the initialization wrapper. + std::unordered_map> const_vars_by_symbol; + for (tvm::runtime::Module mod : ext_libs) { + auto pf_sym = mod.GetFunction("get_symbol"); + auto pf_var = mod.GetFunction("get_const_vars"); + std::vector symbol_const_vars; + if (pf_sym != nullptr && pf_var != nullptr) { + String symbol = pf_sym(); + Array variables = pf_var(); + for (size_t i = 0; i < variables.size(); i++) { + symbol_const_vars.push_back(variables[i].operator std::string()); + } + ICHECK_EQ(const_vars_by_symbol.count(symbol), 0U) << "Found duplicated symbol: " << symbol; + const_vars_by_symbol[symbol] = symbol_const_vars; + } + } + if (!const_vars_by_symbol.empty() || !params.empty()) { + // need runtime const information, run link const loader + std::unordered_map const_var_ndarray; + for (const auto& [name, param] : params) { + const_var_ndarray[name] = param; + } + runtime::Module const_loader_mod = + runtime::ConstLoaderModuleCreate(const_var_ndarray, const_vars_by_symbol); + const_loader_mod.Import(lib); + for (const auto& it : ext_libs) { + const_loader_mod.Import(it); + } + exec->Import(const_loader_mod); + } else { + // directly import the ext_modules as we don't need const loader + exec->Import(lib); + for (const auto& it : ext_libs) { + exec->Import(it); + } + } +} + /*! * \brief Link the libraries together. */ Module VMLink(ExecBuilder builder, Target target, Optional lib, Array ext_libs, Map params) { - // TODO(relax-team) Revisit the param and ext_lib options. ObjectPtr executable = builder->Get(); if (!lib.defined()) { lib = codegen::CSourceModuleCreate(";", "", Array{}); } - std::unordered_map conv_params; - for (const auto& [name, param] : params) { - conv_params[name] = param; - } - Module combined_lib = codegen::CreateMetadataModule( - conv_params, lib.value(), ext_libs, target, - - // TODO(@sunggg): Currently, CRT uses relay-specific executor for uTVM support. - // Before jumping into details, only support cpp runtime for now. - relay::Runtime::Create("cpp"), - relay::Executor::Create("graph"), // TODO(@sunggg): pass arbitrarily executor. CPP runtime - // won't use this anyways. - relay::backend::ExecutorCodegenMetadata()); - executable->Import(combined_lib); + LinkModules(executable, params, lib.value(), ext_libs); return Module(executable); } diff --git a/src/relax/ir/block_builder.cc b/src/relax/ir/block_builder.cc index b8092bbf3a4d..8df9a67a26f6 100644 --- a/src/relax/ir/block_builder.cc +++ b/src/relax/ir/block_builder.cc @@ -29,7 +29,6 @@ #include #include #include -#include #include #include diff --git a/src/relax/ir/expr.cc b/src/relax/ir/expr.cc index 6ace974985a5..ca97744c5125 100644 --- a/src/relax/ir/expr.cc +++ b/src/relax/ir/expr.cc @@ -138,7 +138,7 @@ TVM_REGISTER_GLOBAL("relax.If") return If(cond, true_branch, false_branch, span); }); -Tuple::Tuple(tvm::Array fields, Span span) { +Tuple::Tuple(tvm::Array fields, Span span) { Optional tuple_sinfo = [&]() -> Optional { Array field_sinfo; for (const auto& field : fields) { @@ -163,7 +163,7 @@ Tuple::Tuple(tvm::Array fields, Span span) { TVM_REGISTER_NODE_TYPE(TupleNode); -TVM_REGISTER_GLOBAL("relax.Tuple").set_body_typed([](tvm::Array fields, Span span) { +TVM_REGISTER_GLOBAL("relax.Tuple").set_body_typed([](tvm::Array fields, Span span) { return Tuple(fields, span); }); diff --git a/src/relax/ir/transform.cc b/src/relax/ir/transform.cc index 9f418bff5c6d..ddf252f2cce6 100644 --- a/src/relax/ir/transform.cc +++ b/src/relax/ir/transform.cc @@ -27,7 +27,6 @@ #include #include #include -#include #include namespace tvm { namespace relax { @@ -82,14 +81,6 @@ class FunctionPassNode : public tvm::transform::PassNode { TVM_DECLARE_FINAL_OBJECT_INFO(FunctionPassNode, PassNode); private: - /* - * \brief Check if a function should be skipped for optimization. - * - * \param func The target function to be checked. - * - * \return Return true if the function will be skipped, otherwise false. - */ - bool SkipFunction(const Function& func) const; }; class FunctionPass : public Pass { @@ -145,7 +136,7 @@ IRModule FunctionPassNode::operator()(IRModule mod, const PassContext& pass_ctx) // only picks up relax::Function if (auto* n = it.second.as()) { Function func = GetRef(n); - auto updated_func = SkipFunction(func) ? func : pass_func(func, updated_mod, pass_ctx); + auto updated_func = pass_func(func, updated_mod, pass_ctx); updates.push_back({it.first, updated_func}); } } @@ -165,12 +156,6 @@ IRModule FunctionPassNode::operator()(IRModule mod, const PassContext& pass_ctx) return updated_mod; } -bool FunctionPassNode::SkipFunction(const Function& func) const { - // TODO(@yuchen): will need to revisit in the future - return (func->GetAttr(relay::attr::kCompiler).defined()) || - func->GetAttr(relay::attr::kSkipOptimization, 0) != 0; -} - Pass CreateFunctionPass( const runtime::TypedPackedFunc& pass_func, int opt_level, String name, tvm::Array required, bool traceable) { diff --git a/src/relax/op/distributed/distributed.cc b/src/relax/op/distributed/distributed.cc index 67e11f153511..cdeb537c3d9f 100644 --- a/src/relax/op/distributed/distributed.cc +++ b/src/relax/op/distributed/distributed.cc @@ -100,7 +100,7 @@ StructInfo InferStructInfoCallTIRLocalView(const Call& call, const BlockBuilder& return call->sinfo_args[0]; } -RELAY_REGISTER_OP("relax.dist.call_tir_local_view") +TVM_REGISTER_OP("relax.dist.call_tir_local_view") .set_num_inputs(3) .add_argument("func", "Expr", "The destination-passing-style function.") .add_argument("args", "Tuple", "The input arguments.") diff --git a/src/relax/op/distributed/utils.h b/src/relax/op/distributed/utils.h index 54087639f116..1656df286784 100644 --- a/src/relax/op/distributed/utils.h +++ b/src/relax/op/distributed/utils.h @@ -28,8 +28,6 @@ #include #include #include -#include -#include #include "../op_common.h" diff --git a/src/relax/op/op.cc b/src/relax/op/op.cc index a7d97a59a100..f886b3b4bb1c 100644 --- a/src/relax/op/op.cc +++ b/src/relax/op/op.cc @@ -21,7 +21,6 @@ #include #include #include -#include #include "op_common.h" @@ -102,7 +101,7 @@ StructInfo InferStructInfoCallPurePacked(const Call& call, const BlockBuilder& c } } -RELAY_REGISTER_OP("relax.call_pure_packed") +TVM_REGISTER_OP("relax.call_pure_packed") .set_num_inputs(-1) .add_argument("args", "Array", "The first argument is the function being called. The rest are the " @@ -215,7 +214,7 @@ StructInfo InferStructInfoCallInplacePacked(const Call& call, const BlockBuilder TVM_REGISTER_NODE_TYPE(CallInplacePackedAttrs); -RELAY_REGISTER_OP("relax.call_inplace_packed") +TVM_REGISTER_OP("relax.call_inplace_packed") .set_num_inputs(-1) .set_attrs_type() .add_argument("args", "Array", @@ -562,7 +561,7 @@ void ValidateCallTIR(Call call) { } } -RELAY_REGISTER_OP("relax.call_tir") +TVM_REGISTER_OP("relax.call_tir") .set_num_inputs(3) .add_argument("func", "Expr", "The destination-passing-style function.") .add_argument("args", "Tuple", "The input arguments.") @@ -607,7 +606,7 @@ TVM_REGISTER_GLOBAL("relax.op.call_tir").set_body_typed(MakeCallTIR); TVM_REGISTER_NODE_TYPE(CallTIRWithGradAttrs); -RELAY_REGISTER_OP("relax.call_tir_with_grad") +TVM_REGISTER_OP("relax.call_tir_with_grad") .set_num_inputs(3) .set_attrs_type() .add_argument("func", "Expr", "The destination-passing-style function.") @@ -748,7 +747,7 @@ Expr NormalizeCallTIRInPlace(const BlockBuilder& ctx, Call call) { TVM_REGISTER_NODE_TYPE(CallTIRInplaceAttrs); -RELAY_REGISTER_OP("relax.call_tir_inplace") +TVM_REGISTER_OP("relax.call_tir_inplace") .set_num_inputs(3) .set_attrs_type() .add_argument("func", "Expr", "The destination-passing-style function.") @@ -806,7 +805,7 @@ StructInfo InferStructInfoCallDPSPacked(const Call& call, const BlockBuilder& ct return call->sinfo_args[0]; } -RELAY_REGISTER_OP("relax.call_dps_packed") +TVM_REGISTER_OP("relax.call_dps_packed") .set_num_inputs(2) .add_argument("func", "Expr", "The destination-passing-style function.") .add_argument("args", "Tuple", "The input arguments.") @@ -877,7 +876,7 @@ TVM_REGISTER_GLOBAL("relax.op.null_value").set_body_typed(MakeCallNullValue); // print -RELAY_REGISTER_OP("relax.print") +TVM_REGISTER_OP("relax.print") .set_num_inputs(-1) .add_argument("vals", "Array", "The first value is Python-style format string to use to print. The others " @@ -919,7 +918,7 @@ StructInfo InferAssertStructInfo(const Call& call, const BlockBuilder& ctx) { return ReturnVoidStructInfo(call, ctx); } -RELAY_REGISTER_OP("relax.assert_op") +TVM_REGISTER_OP("relax.assert_op") .set_num_inputs(-1) .add_argument("vals", "Array", "The first value is used as the assertion condition. The second value is " @@ -943,7 +942,7 @@ TVM_REGISTER_GLOBAL("relax.op.assert_op").set_body_typed(MakeAssertOp); // make_closure -RELAY_REGISTER_OP("relax.make_closure") +TVM_REGISTER_OP("relax.make_closure") .set_num_inputs(2) .add_argument("func", "Expr", "The closure.") .add_argument("args", "Tuple", "The captured variables.") @@ -969,7 +968,7 @@ StructInfo InferStructInfoInvokeClosure(const Call& call, const BlockBuilder& ct } } -RELAY_REGISTER_OP("relax.invoke_closure") +TVM_REGISTER_OP("relax.invoke_closure") .set_num_inputs(2) .add_argument("closure", "Expr", "The VMClosure.") .add_argument("args", "Tuple", "The captured variables.") @@ -986,7 +985,7 @@ TVM_REGISTER_GLOBAL("relax.op.invoke_closure").set_body_typed(InvokeClosure); // invoke_pure_closure -RELAY_REGISTER_OP("relax.invoke_pure_closure") +TVM_REGISTER_OP("relax.invoke_pure_closure") .set_num_inputs(2) .add_argument("closure", "Expr", "The VMClosure.") .add_argument("args", "Tuple", "The captured variables.") @@ -1002,7 +1001,7 @@ TVM_REGISTER_GLOBAL("relax.op.invoke_pure_closure").set_body_typed(InvokePureClo // shape_of -RELAY_REGISTER_OP("relax.shape_of") +TVM_REGISTER_OP("relax.shape_of") .set_num_inputs(1) .add_argument("input", "Expr", "The input expression") .set_attr("FInferStructInfo", InferStructInfoShapeOf) @@ -1036,7 +1035,7 @@ StructInfo ReturnTensorToShapeStructInfo(const Call& call, const BlockBuilder& c return ShapeStructInfo(kUnknownNDim); } -RELAY_REGISTER_OP("relax.tensor_to_shape") +TVM_REGISTER_OP("relax.tensor_to_shape") .set_num_inputs(1) .add_argument("input", "Expr", "The input expression") .set_attr("FInferStructInfo", ReturnTensorToShapeStructInfo) @@ -1059,7 +1058,7 @@ StructInfo ReturnShapeToTensorStructInfo(const Call& call, const BlockBuilder& c return TensorStructInfo(ShapeExpr({PrimExpr(ndim)}), DataType::Int(64)); } -RELAY_REGISTER_OP("relax.shape_to_tensor") +TVM_REGISTER_OP("relax.shape_to_tensor") .set_num_inputs(1) .add_argument("input", "Expr", "The input expression") .set_attr("FInferStructInfo", ReturnShapeToTensorStructInfo) @@ -1088,7 +1087,7 @@ StructInfo InferStructInfoAllocateTensor(const Call& call, const BlockBuilder& c return TensorStructInfo(call->args[0], out_dtype); } -RELAY_REGISTER_OP("relax.builtin.alloc_tensor") +TVM_REGISTER_OP("relax.builtin.alloc_tensor") .set_num_inputs(4) .add_argument("shape", "Expr", "The shape of the tensor to allocate.") .add_argument("dtype", "DataTypeImm", "The dtype of the tensor to allocate.") @@ -1112,7 +1111,7 @@ TVM_REGISTER_GLOBAL("relax.op.builtin.alloc_tensor").set_body_typed(MakeAllocTen // memory planning alloc_storage -RELAY_REGISTER_OP("relax.memory.alloc_storage") +TVM_REGISTER_OP("relax.memory.alloc_storage") .set_num_inputs(4) .add_argument("total_space", "Expr", "The total space of the storage to allocate.") .add_argument( @@ -1148,7 +1147,7 @@ StructInfo InferStructInfoMemAllocTensor(const Call& call, const BlockBuilder& c return TensorStructInfo(call->args[2], out_dtype); } -RELAY_REGISTER_OP("relax.memory.alloc_tensor") +TVM_REGISTER_OP("relax.memory.alloc_tensor") .set_num_inputs(4) .add_argument("storage", "Expr", "The storage to allocate the tensor to.") .add_argument("offset", "PrimValue", "Storage offset to allocate the tensor.") @@ -1168,7 +1167,7 @@ TVM_REGISTER_GLOBAL("relax.op.memory.alloc_tensor").set_body_typed(MakeMemAllocT // memory planning kill_storage -RELAY_REGISTER_OP("relax.memory.kill_storage") +TVM_REGISTER_OP("relax.memory.kill_storage") .set_num_inputs(1) .add_argument("storage", "Expr", "The storage to be killed.") .set_attr("FInferStructInfo", ReturnVoidStructInfo) @@ -1184,7 +1183,7 @@ TVM_REGISTER_GLOBAL("relax.op.memory.kill_storage").set_body_typed(MakeMemKillSt // memory planning kill_tensor -RELAY_REGISTER_OP("relax.memory.kill_tensor") +TVM_REGISTER_OP("relax.memory.kill_tensor") .set_num_inputs(1) .add_argument("tensor", "Expr", "The tensor to be killed.") .set_attr("FInferStructInfo", ReturnVoidStructInfo) @@ -1200,7 +1199,7 @@ TVM_REGISTER_GLOBAL("relax.op.memory.kill_tensor").set_body_typed(MakeMemKillTen // vm alloc_storage -RELAY_REGISTER_OP("relax.vm.alloc_storage") +TVM_REGISTER_OP("relax.vm.alloc_storage") .set_num_inputs(4) .add_argument("size", "Expr", "The size of the storage to allocate.") .add_argument("dtype", "DataTypeImm", "The dtype of the tensor to allocate.") @@ -1242,7 +1241,7 @@ StructInfo InferStructInfoVMAllocTensor(const Call& call, const BlockBuilder& ct return TensorStructInfo(out_dtype, kUnknownNDim); } -RELAY_REGISTER_OP("relax.vm.alloc_tensor") +TVM_REGISTER_OP("relax.vm.alloc_tensor") .set_num_inputs(4) .add_argument("storage", "Expr", "The storage to allocate the tensor to.") .add_argument("offset", "PrimValue", "Storage offset to allocate the tensor.") @@ -1278,7 +1277,7 @@ TVM_REGISTER_GLOBAL("relax.op.vm.kill_object").set_body_typed(MakeVMKillObject); // vm call_tir_dyn -RELAY_REGISTER_OP("relax.vm.call_tir_dyn") +TVM_REGISTER_OP("relax.vm.call_tir_dyn") .set_num_inputs(2) .add_argument("func", "Expr", "The destination-passing-style function.") .add_argument("args", "Tuple", @@ -1299,7 +1298,7 @@ StructInfo InferStructInfoStopLiftParams(const Call& call, const BlockBuilder& c return InferStructInfoUnaryArith(call, ctx); } -RELAY_REGISTER_OP("relax.builtin.stop_lift_params") +TVM_REGISTER_OP("relax.builtin.stop_lift_params") .set_num_inputs(1) .add_argument("x", "Expr", "The input data") .set_attr("FInferStructInfo", InferStructInfoStopLiftParams) @@ -1327,7 +1326,7 @@ StructInfo InferToVDeviceStructInfo(const Call& call, const BlockBuilder& ctx) { return TensorStructInfo(data_sinfo->dtype, data_sinfo->ndim, vdev, data_sinfo->span); } -RELAY_REGISTER_OP("relax.to_vdevice") +TVM_REGISTER_OP("relax.to_vdevice") .set_num_inputs(1) .set_attrs_type() .add_argument("data", "Expr", "The input expression to be copied") @@ -1353,7 +1352,7 @@ StructInfo InferHintOnDeviceStructInfo(const Call& call, const BlockBuilder& ctx return data_sinfo; } -RELAY_REGISTER_OP("relax.hint_on_device") +TVM_REGISTER_OP("relax.hint_on_device") .set_num_inputs(1) .set_attrs_type() .add_argument("data", "Expr", "The input expression") diff --git a/src/relax/op/op_common.h b/src/relax/op/op_common.h index eb9caae4b9e1..6e2ef6bd2bef 100644 --- a/src/relax/op/op_common.h +++ b/src/relax/op/op_common.h @@ -27,8 +27,6 @@ #include #include -#include -#include #include #include diff --git a/src/relax/transform/allocate_workspace.cc b/src/relax/transform/allocate_workspace.cc index 05aa8ce5528d..b3f67070a66b 100644 --- a/src/relax/transform/allocate_workspace.cc +++ b/src/relax/transform/allocate_workspace.cc @@ -26,6 +26,7 @@ #include #include #include +#include #include "../op/op_common.h" diff --git a/src/relax/transform/annotate_tir_op_pattern.cc b/src/relax/transform/annotate_tir_op_pattern.cc index b1c1ed29aff3..5ef104d2e3ee 100644 --- a/src/relax/transform/annotate_tir_op_pattern.cc +++ b/src/relax/transform/annotate_tir_op_pattern.cc @@ -33,7 +33,7 @@ tir::PrimFunc AnnotateOpPattern(tir::PrimFunc f) { if (f->HasNonzeroAttr("op_pattern")) { return f; } else { - relay::OpPatternKind kind = AnalyzeOpPatternKind(f); + OpPatternKind kind = AnalyzeOpPatternKind(f); return WithAttr(std::move(f), "op_pattern", Integer(static_cast(kind))); } } diff --git a/src/relax/transform/call_tir_rewrite.cc b/src/relax/transform/call_tir_rewrite.cc index 157bff70cb02..5281a7e28259 100644 --- a/src/relax/transform/call_tir_rewrite.cc +++ b/src/relax/transform/call_tir_rewrite.cc @@ -28,7 +28,6 @@ #include #include -#include "../../relay/transforms/pattern_utils.h" #include "utils.h" namespace tvm { diff --git a/src/relax/transform/few_shot_tuning.cc b/src/relax/transform/few_shot_tuning.cc index 4ad5e2367524..12b18bd2c47b 100644 --- a/src/relax/transform/few_shot_tuning.cc +++ b/src/relax/transform/few_shot_tuning.cc @@ -167,11 +167,9 @@ Pass FewShotTuning(int valid_count, bool benchmark) { result.Set(gv, func); } } - return IRModule(result, // functions - m->type_definitions, // type_definitions - m->import_set_, // import_set - m->source_map, // map - m->attrs); // attrs); + return IRModule(result, // functions + m->source_map, // map + m->attrs); // attrs); }; return CreateModulePass(/*pass_function=*/pass_func, // /*opt_level=*/0, // diff --git a/src/relax/transform/fuse_ops.cc b/src/relax/transform/fuse_ops.cc index 85c739e08353..bcab018b1ee8 100644 --- a/src/relax/transform/fuse_ops.cc +++ b/src/relax/transform/fuse_ops.cc @@ -38,10 +38,9 @@ #include #include -#include -#include "../../relay/analysis/graph_partitioner.h" #include "../../support/arena.h" +#include "../analysis/graph_partitioner.h" #include "tvm/relax/expr.h" #include "utils.h" @@ -88,9 +87,6 @@ namespace relax { - We use an Union-Find data structure to manage the groups. */ -using relay::GraphPartitioner; -using relay::IndexedForwardGraph; -using relay::OpPatternKind; using support::LinkNode; constexpr uint32_t kMaxFusedOps = 256; diff --git a/src/relax/transform/fuse_tir.cc b/src/relax/transform/fuse_tir.cc index fe247645dc24..8fba54628153 100644 --- a/src/relax/transform/fuse_tir.cc +++ b/src/relax/transform/fuse_tir.cc @@ -26,8 +26,6 @@ #include #include -#include "../../relay/analysis/graph_partitioner.h" -#include "../../support/arena.h" #include "../../tir/ir/functor_common.h" namespace tvm { diff --git a/src/relax/transform/lambda_lift.cc b/src/relax/transform/lambda_lift.cc index f45d82129db6..ee19872e5af5 100644 --- a/src/relax/transform/lambda_lift.cc +++ b/src/relax/transform/lambda_lift.cc @@ -400,7 +400,7 @@ class LambdaLifter : public ExprMutator { if (auto it = nested_closure_map_.find(var); it != nested_closure_map_.end()) { Call nested_call = it->second; - Array new_args = call->args; + Array new_args = call->args; for (const auto arg : nested_call->args) { new_args.push_back(arg); } diff --git a/src/relax/transform/merge_composite_functions.cc b/src/relax/transform/merge_composite_functions.cc index 0a3c4ff0a193..8eebeb82db9d 100644 --- a/src/relax/transform/merge_composite_functions.cc +++ b/src/relax/transform/merge_composite_functions.cc @@ -65,8 +65,6 @@ namespace tvm { namespace relax { -using relay::GraphPartitioner; - namespace { using Group = GraphPartitioner::Group; diff --git a/src/relax/transform/meta_schedule.cc b/src/relax/transform/meta_schedule.cc index c54e75e9cb88..06d10919a702 100644 --- a/src/relax/transform/meta_schedule.cc +++ b/src/relax/transform/meta_schedule.cc @@ -169,8 +169,6 @@ Pass MetaScheduleApplyDatabase(Optional work_dir, bool enable_warning = result.Set(gv, base_func); } return IRModule(result, // functions - {}, // type_definitions - {}, // import_set {}, // map mod->attrs); // attrs); }; diff --git a/src/relax/transform/run_codegen.cc b/src/relax/transform/run_codegen.cc index af9ed2fffce2..48237a20ac16 100644 --- a/src/relax/transform/run_codegen.cc +++ b/src/relax/transform/run_codegen.cc @@ -25,8 +25,7 @@ #include #include - -#include +#include #include "../../support/ordered_set.h" #include "utils.h" diff --git a/src/relax/transform/utils.h b/src/relax/transform/utils.h index 55e355b4bac2..67d5fd4875e0 100644 --- a/src/relax/transform/utils.h +++ b/src/relax/transform/utils.h @@ -35,8 +35,8 @@ #include #include -#include "../../relay/analysis/graph_partitioner.h" #include "../../support/array.h" +#include "../analysis/graph_partitioner.h" #include "../op/nn/convolution.h" #include "../op/nn/nn.h" #include "../op/nn/pooling.h" @@ -141,8 +141,7 @@ inline std::string GetExtSymbol(const Function& func) { * \return A new module containing grouped functions. */ IRModule MakeGroupedFunctions( - IRModule mod, - const std::unordered_map& partition, + IRModule mod, const std::unordered_map& partition, bool lift_constants = true, const Array& entry_function_names = {}); /*! diff --git a/src/relay/analysis/annotated_region_set.cc b/src/relay/analysis/annotated_region_set.cc deleted file mode 100644 index ef21604d8a71..000000000000 --- a/src/relay/analysis/annotated_region_set.cc +++ /dev/null @@ -1,242 +0,0 @@ -/* - * 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 "annotated_region_set.h" - -#include -#include - -#include -#include - -namespace tvm { -namespace relay { - -AnnotatedRegion AnnotatedRegionSetNode::GetRegion(const Expr& expr) const { - for (auto candidate : regions_) { - if (candidate->nodes_.find(expr) != candidate->nodes_.end()) { - return candidate; - } - } - return AnnotatedRegion(nullptr); -} - -void AnnotatedRegionSetNode::MergeRegions(AnnotatedRegion src, AnnotatedRegion dest) { - if (dest == src) { - return; - } - - // Merge src to dest and erase src. - dest->nodes_.insert(src->nodes_.begin(), src->nodes_.end()); - for (const auto& input : src->ins_) { - dest->ins_.push_back(input); - } - for (const auto& output : src->outs_) { - dest->outs_.push_back(output); - } - // if any of the outputs of src are inputs of dest, they become internal nodes - // so remove them from outs - std::vector ins_to_remove; - for (const auto& input : dest->ins_) { - auto call = Downcast(input); - auto it = src->nodes_.find(call->args[0]); - if (it != src->nodes_.end()) { - dest->outs_.remove(*it); - ins_to_remove.push_back(input); - } - } - for (const auto& input : ins_to_remove) { - dest->ins_.remove(input); - } - regions_.erase(src); -} - -void AnnotatedRegionSetNode::AddToRegion(AnnotatedRegion dest, const Expr& expr) { - auto src = GetRegion(expr); - if (src.defined()) { - MergeRegions(src, dest); - } else { - dest->nodes_.insert(expr); - } -} - -AnnotatedRegion AnnotatedRegionSetNode::MakeRegion(const std::string& func_name, - const std::string& target) { - auto ret = regions_.emplace(AnnotatedRegion()); - (*ret.first)->id_ = region_id_++; - (*ret.first)->target_ = target; - (*ret.first)->func_name_ = func_name; - return *ret.first; -} - -class AnnotatedRegionSet::Creator : protected MixedModeVisitor { - public: - Creator(const Op& region_begin_op, const Op& region_end_op, - const std::string& func_name = "default") - : begin_op_(region_begin_op), end_op_(region_end_op), func_name_(func_name) {} - - AnnotatedRegionSet Create(const Expr& expr) { - VisitExpr(expr); - return std::move(region_set_); - } - - void AddToArgRegion(Expr expr, Array args) { - // Merge argument regions and add itself to the region. - - // Find the first open region. - AnnotatedRegion region; - for (auto arg : args) { - const CallNode* end = arg.as(); - if (end && end->op == end_op_) { // Ignore closed regions. - continue; - } - - region = region_set_->GetRegion(arg); - if (region.defined()) { - break; - } - } - - // Try to merge open regions. - for (auto arg : args) { - const CallNode* end = arg.as(); - if (end && end->op == end_op_) { // Ignore closed regions. - continue; - } - - auto arg_region = region_set_->GetRegion(arg); - ICHECK_EQ(region.defined(), arg_region.defined()) - << "Arg regions are inconsistent: " << AsText(expr); - if (region.defined() && region != arg_region) { - region_set_->MergeRegions(arg_region, region); - } - } - if (region.defined()) { - region_set_->AddToRegion(region, expr); - } - } - - void VisitExpr_(const CallNode* call) { - auto op_node = call->op.as(); - - if (op_node == nullptr || call->attrs.as() == nullptr) { - AddToArgRegion(GetRef(call), call->args); - } else if (call->op == begin_op_) { - // The annotation node is inserted on edge so it must have only one argument. - ICHECK_EQ(call->args.size(), 1U); - std::string target = call->attrs.as()->compiler; - - // Check if the argument already belongs to a region - auto region = region_set_->GetRegion(GetRef(call)); - ICHECK(!region.defined()); - - // Create a new region. - region = region_set_->MakeRegion(func_name_, target); - region->nodes_.insert(GetRef(call)); - region->ins_.push_back(GetRef(call)); - } else { - ICHECK_EQ(call->op, end_op_); - // The annotation node is inserted on edge so it must have only one argument. - ICHECK_EQ(call->args.size(), 1U); - std::string target = call->attrs.as()->compiler; - - // Check if the argument already belongs to a region - auto region = region_set_->GetRegion(call->args[0]); - if (!region.defined()) { - throw CompileError(ErrorBuilder() - << "Cannot find the corresponding region for end annotation:\n" - << AsText(GetRef(call), false)); - } else { - // If the argument is belonged to a region, it must have the same target. - // Otherwise we should see a region_begin op. - ICHECK_EQ(region->GetTarget(), target); - } - region->nodes_.insert(GetRef(call)); - region->outs_.push_back(GetRef(call)); - } - } - - void VisitExpr_(const TupleNode* op) { AddToArgRegion(GetRef(op), op->fields); } - - void VisitExpr_(const TupleGetItemNode* g) { - Array args = {g->tuple}; - AddToArgRegion(GetRef(g), args); - } - - void VisitExpr_(const LetNode* op) { - Array args = {op->var, op->value, op->body}; - AddToArgRegion(GetRef(op), args); - ExprVisitor::VisitExpr_(op); - } - - void VisitExpr_(const IfNode* op) { - Array args = {op->cond, op->true_branch, op->false_branch}; - AddToArgRegion(GetRef(op), args); - ExprVisitor::VisitExpr_(op); - } - - void VisitExpr_(const RefCreateNode* op) { - Array args = {op->value}; - AddToArgRegion(GetRef(op), args); - ExprVisitor::VisitExpr_(op); - } - - void VisitExpr_(const RefReadNode* op) { - Array args = {op->ref}; - AddToArgRegion(GetRef(op), args); - ExprVisitor::VisitExpr_(op); - } - - void VisitExpr_(const RefWriteNode* op) { - Array args = {op->ref}; - AddToArgRegion(GetRef(op), args); - ExprVisitor::VisitExpr_(op); - } - - private: - /*! \brief The region set being constructed.*/ - AnnotatedRegionSet region_set_; - /*! \brief Region 'begin' annotation operator. */ - const Op begin_op_; - /*! \brief Region 'end' annotation operator. */ - const Op end_op_; - /*! \brief The unique function name that is used to be the name of this region set. */ - const std::string func_name_; -}; - -AnnotatedRegionSet AnnotatedRegionSet::Create(const Expr& expr, const Op& begin, const Op& end, - const std::string& func_name) { - return Creator(begin, end, func_name).Create(expr); -} - -TVM_REGISTER_NODE_TYPE(AnnotatedRegionNode); -TVM_REGISTER_NODE_TYPE(AnnotatedRegionSetNode); - -TVM_REGISTER_GLOBAL("relay.analysis.AnnotatedRegionSet") - .set_body_typed([](Expr expr, Op begin, Op end) { - return AnnotatedRegionSet::Create(expr, begin, end); - }); - -TVM_REGISTER_GLOBAL("relay.analysis.GetRegion") - .set_body_typed([](AnnotatedRegionSet region_set, Expr expr) { - return region_set->GetRegion(expr); - }); - -} // namespace relay -} // namespace tvm diff --git a/src/relay/analysis/annotated_region_set.h b/src/relay/analysis/annotated_region_set.h deleted file mode 100644 index 443bd5ec1da3..000000000000 --- a/src/relay/analysis/annotated_region_set.h +++ /dev/null @@ -1,279 +0,0 @@ -/* - * 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 tvm/relay/transforms/annotated_region_set.h - * \brief Define data structures to extract and manipulate regions from - * a relay function. Regions are denoted by region_begin and region_end - * annotations that exist on all the input and output edges of the region. - */ - -#ifndef TVM_RELAY_ANALYSIS_ANNOTATED_REGION_SET_H_ -#define TVM_RELAY_ANALYSIS_ANNOTATED_REGION_SET_H_ - -#include -#include -#include -#include -#include -#include - -#include -#include -#include -#include -#include - -namespace tvm { -namespace relay { - -class AnnotatedRegion; -class AnnotatedRegionSet; - -class AnnotatedRegionNode : public Object { - public: - void VisitAttrs(AttrVisitor* v) { - v->Visit("id", &id_); - v->Visit("target", &target_); - Array nodes_array(nodes_.begin(), nodes_.end()); - v->Visit("nodes", &nodes_array); - Array args_array(ins_.begin(), ins_.end()); - v->Visit("args", &args_array); - Array rets_array(outs_.begin(), outs_.end()); - v->Visit("rets", &rets_array); - } - - /*! \brief Get the region ID. */ - int GetID() const { return id_; } - - /*! \brief Get the region name. */ - std::string GetName() const { return func_name_; } - - /*! \brief Get the region target. */ - std::string GetTarget() const { return target_; } - - /*! \brief Get the region's inputs. */ - std::list GetInputs() const { return ins_; } - - /*! \brief Get the region's outputs. */ - std::list GetOutputs() const { return outs_; } - - /*! \brief Get the region's nodes. */ - std::unordered_set GetNodes() const { return nodes_; } - - static constexpr const char* _type_key = "relay.AnnotatedRegion"; - TVM_DECLARE_FINAL_OBJECT_INFO(AnnotatedRegionNode, Object); - - protected: - /*! \brief The region ID. */ - int id_{-1}; - /*! \brief The func name. */ - std::string func_name_ = "default"; - /*! \brief The target for this region. */ - std::string target_ = "default"; - /*! \brief The inputs to this region. */ - std::list ins_; - /*! \brief The outputs of this region */ - std::list outs_; - /*! \brief Nodes in this region. */ - std::unordered_set nodes_; - - friend class AnnotatedRegionSet; - friend class AnnotatedRegionSetNode; -}; - -/*! - * \brief An object to hold the properties of a region as used by the - * AnnotatedRegionSet class. This should be considered read-only. - */ -class AnnotatedRegion : public ObjectRef { - public: - AnnotatedRegion() { - auto n = make_object(); - data_ = std::move(n); - } - - /*! - * \brief Construct from an object pointer. - * \param n The object pointer. - */ - explicit AnnotatedRegion(ObjectPtr n) : ObjectRef(n) {} - - /*! \return Mutable pointers to the node. */ - AnnotatedRegionNode* operator->() const { - auto* ptr = get_mutable(); - ICHECK(ptr != nullptr); - return static_cast(ptr); - } -}; - -class AnnotatedRegionSetNode : public Object { - using UnorderedRegionSet = std::unordered_set; - // Create iterator alias for a RegionSet object. - using iterator = UnorderedRegionSet::iterator; - using const_iterator = UnorderedRegionSet::const_iterator; - - public: - /*! \brief Default constructor. */ - AnnotatedRegionSetNode() = default; - - /*! \return The begin iterator */ - iterator begin() { return regions_.begin(); } - /*! \return The end iterator */ - iterator end() { return regions_.end(); } - /*! \return The const begin iterator */ - const_iterator begin() const { return regions_.begin(); } - /*! \return The const end iterator */ - const_iterator end() const { return regions_.end(); } - - /*! - * \brief Get the region that an expression belongs to. - * - * \param expr Which expr to get the region for. - * - * \return A pointer to the region, nullptr if the expression - * doesn't belong to a region. - */ - AnnotatedRegion GetRegion(const Expr& expr) const; - - /*! - * \brief Merge src region into dest region. - * - * \param src The region to merge - will be erased. - * \param dest The region into which src will be merged. - */ - void MergeRegions(AnnotatedRegion src, AnnotatedRegion dest); - - void VisitAttrs(AttrVisitor* v) { - Array regions_array(regions_.begin(), regions_.end()); - v->Visit("regions", ®ions_array); - } - - static constexpr const char* _type_key = "relay.AnnotatedRegionSet"; - TVM_DECLARE_FINAL_OBJECT_INFO(AnnotatedRegionSetNode, Object); - - private: - /*! - * \brief Add an expression to a region. - * - * \param dest The region to add the expression to. - * \param expr The expression. - */ - void AddToRegion(AnnotatedRegion dest, const Expr& expr); - - /*! - * \brief Make a new region for a target. - * - * \return The new region. - */ - AnnotatedRegion MakeRegion(const std::string& func_name, const std::string& target); - - std::unordered_set regions_; - /*! \brief The next region ID to assign. */ - int region_id_{0}; - - friend class AnnotatedRegionSet; -}; - -/*! - * \brief A class to hold a set of regions produced from a relay expression - * that contains 'region_begin' and 'region_end' style annotations. The - * regions should be disjoint. The class provides both a method to construct - * the region set of a given relay expression as well as additional methods - * to update and query regions. - */ -class AnnotatedRegionSet : public ObjectRef { - using UnorderedRegionSet = std::unordered_set; - // Create iterator alias for a RegionSet object. - using iterator = UnorderedRegionSet::iterator; - using const_iterator = UnorderedRegionSet::const_iterator; - - public: - AnnotatedRegionSet() { - auto n = make_object(); - data_ = std::move(n); - } - - /*! - * \brief Construct from an object pointer. - * - * \param n The object pointer. - */ - explicit AnnotatedRegionSet(ObjectPtr n) : ObjectRef(n) {} - - /*! \return The begin iterator. */ - iterator begin() { - auto* n = operator->(); - ICHECK(n); - return n->begin(); - } - /*! \return The end iterator. */ - iterator end() { - auto* n = operator->(); - ICHECK(n); - return n->end(); - } - /*! \return The begin iterator. */ - const_iterator begin() const { - const auto* n = operator->(); - ICHECK(n); - return n->begin(); - } - /*! \return The end iterator. */ - const_iterator end() const { - const auto* n = operator->(); - ICHECK(n); - return n->end(); - } - - /*! \return mutable pointers to the node. */ - AnnotatedRegionSetNode* operator->() const { - auto* ptr = get_mutable(); - ICHECK(ptr != nullptr); - return static_cast(ptr); - } - - /*! \return The region an expression belongs to. */ - AnnotatedRegion operator[](const Expr& expr) { - const auto* n = operator->(); - ICHECK(n); - return n->GetRegion(expr); - } - - /*! \brief Create a RegionSet from a relay expression. - * - * \param expr The relay expr from which to construct the set. - * \param begin Region begin annotation operator. - * \param end Region end annotation operator. - * \param func_name function name - * - * \return The created RegionSet for the expression. - */ - static AnnotatedRegionSet Create(const Expr& expr, const Op& begin, const Op& end, - const std::string& func_name = "default"); - - private: - /*! \brief Helper class to construct a RegionSet from an expr.*/ - class Creator; -}; - -} // namespace relay -} // namespace tvm - -#endif // TVM_RELAY_ANALYSIS_ANNOTATED_REGION_SET_H_ diff --git a/src/relay/analysis/call_graph.cc b/src/relay/analysis/call_graph.cc deleted file mode 100644 index 9d1041f56551..000000000000 --- a/src/relay/analysis/call_graph.cc +++ /dev/null @@ -1,345 +0,0 @@ -/* - * 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 src/relay/analysis/call_graph.cc - * \brief Implementation of APIs to handle the call graph of a Relay module. - */ - -#include "call_graph.h" - -#include -#include -#include - -#include -#include -#include -#include -#include - -#include "../op/call/call.h" - -namespace tvm { -namespace relay { - -CallGraph::CallGraph(IRModule module) { - auto n = make_object(); - n->module = std::move(module); - auto gvar_funcs = n->module->functions; - for (const auto& it : gvar_funcs) { - if (auto func = it.second.as()) { - // Add the global function to gradually build up the call graph. - n->AddToCallGraph(it.first, func.value()); - } - } - data_ = std::move(n); -} - -void CallGraphNode::AddToCallGraph(const GlobalVar& gv, const Function& func) { - ICHECK(func.defined() && gv.defined()); - // Add the current global function as an entry to the call grpah. - CallGraphEntry* cg_node = LookupGlobalVar(gv); - - // Only GlobalVar nodes need to be handled in a function. It indicates that - // the global function of a callee is called by the function that is being - // processed. An edge will be added from the current global function, cg_node, - // to the node that contains the found callee GlobalVarNode. - // - // This is the major overhead for constructing a call graph because the - // post-order visitor will visit each AST node of the current function to - // figure out the dependencies between functions. - PostOrderVisit(func, [&](const Expr& expr) { - // TODO(mbs): Cleanup shapes functions. - if (const auto* call_node = expr.as()) { - CallLoweredProps props = GetCallLoweredProps(call_node); - if (props.lowered_func.defined() && props.attrs.metadata.count("prim_shape_fn_var")) { - // We are implicitly calling the shape function *in addition to* the call target. - CallGraphEntry* callee_cg_node = - LookupGlobalVar(Downcast(props.attrs.metadata["prim_shape_fn_var"])); - cg_node->AddCalledGlobal(callee_cg_node); - } - } else if (auto callee = expr.as()) { - CallGraphEntry* callee_cg_node = LookupGlobalVar(callee.value()); - cg_node->AddCalledGlobal(callee_cg_node); - } - }); -} - -const CallGraphEntry* CallGraphNode::operator[](const GlobalVar& gv) const { - const_iterator cit = call_graph_.find(gv); - ICHECK(cit != call_graph_.end()) - << "GlobalVar " << gv->name_hint << " not found in the call graph!"; - return cit->second.get(); -} - -CallGraphEntry* CallGraphNode::operator[](const GlobalVar& gv) { - const_iterator cit = call_graph_.find(gv); - ICHECK(cit != call_graph_.end()) - << "GlobalVar " << gv->name_hint << " not found in the call graph!"; - return cit->second.get(); -} - -BaseFunc CallGraphNode::GetGlobalFunction(const GlobalVar& var) const { - ICHECK(module->ContainGlobalVar(var->name_hint)) - << "GlobalVar " << var->name_hint << " not found in the current ir module"; - return module->Lookup(var->name_hint); -} - -CallGraphEntry* CallGraphNode::LookupGlobalVar(const GlobalVar& gv) { - ICHECK(gv.defined()); - - // This inserts an element to the call graph if it is not there yet. - auto& call_graph_node = call_graph_[gv]; - if (call_graph_node) return call_graph_node.get(); - - // Create the node for the inserted entry. - call_graph_node = std::make_unique(gv); - return call_graph_node.get(); -} - -void CallGraphNode::Print(std::ostream& os) const { - // Print the call graph in the topological order. - std::vector nodes = TopologicalOrder(); - for (const auto* cgn : nodes) { - cgn->Print(os); - } -} - -GlobalVar CallGraphNode::RemoveGlobalVarFromModule(CallGraphEntry* cg_node, - bool update_call_graph) { - ICHECK(cg_node->empty() || (cg_node->IsRecursive() && cg_node->size() == 1)) - << "Cannot remove global var " << cg_node->GetNameHint() - << " from call graph, because it still calls " << cg_node->size() - << " other global functions"; - - if (update_call_graph) { - // Update the call graph by removing all edges that point to the node - // `cg_node`. - for (auto& it : *this) { - it.second->RemoveAllCallTo(cg_node); - } - } - GlobalVar gv = cg_node->GetGlobalVar(); - call_graph_.erase(gv); - // Update the IR module. - module->Remove(gv); - return gv; -} - -std::vector CallGraphNode::GetEntryGlobals() const { - std::vector ret; - // An entry function in Relay is a function that never called by other - // functions or only called by itself. - for (const auto& it : *this) { - if (it.second->GetRefCount() == 0 || it.second->IsRecursiveEntry()) { - ret.push_back(it.second.get()); - } - } - return ret; -} - -std::vector CallGraphNode::TopologicalOrder() const { - std::vector ret; - // Collect all entry nodes. - std::vector entries = GetEntryGlobals(); - CallGraphEntry::CallGraphEntrySet visited; - - for (const auto& it : entries) { - // Keep tracking the nodes that have been visited. - auto topo = it->TopologicalOrder(&visited); - // Prepend the collected items. The intermediate nodes that are shared by - // multiple entries are guaranteed to be collected when visiting the - // previous entries. Therefore, topological order remains. - ret.insert(ret.begin(), topo.begin(), topo.end()); - } - - // Find out the missing global functions if there are any to help debugging. - if (ret.size() != module->functions.size()) { - for (auto it : module->functions) { - if (visited.find((*this)[it.first]) == visited.end()) { - LOG(WARNING) << "Missing global:" << it.first->name_hint - << " with # refs = " << (*this)[it.first]->GetRefCount(); - } - } - LOG(FATAL) << "Expected " << module->functions.size() << " globals, but received " - << ret.size(); - } - - return ret; -} - -// BSF traversal is used to collect the nodes in a CallGraphEntry. The nodes -// that are visited by previous CallGraphEntry entries can be memoized. This -// helps us to make sure no entry will be visited multiple times when collecting -// the nodes for an entire call graph. -std::vector CallGraphEntry::TopologicalOrder(CallGraphEntrySet* visited) const { - std::vector ret; - std::vector current_nodes; - if (visited->find(this) == visited->end()) { - visited->emplace(this); - current_nodes.emplace_back(const_cast(this)); - } - - std::vector next_nodes; - while (!current_nodes.empty()) { - for (const auto& node : current_nodes) { - ret.push_back(node); - // Iterate through the called entries. - for (auto git = node->begin(); git != node->end(); ++git) { - if (visited->find(git->second) == visited->end()) { - next_nodes.push_back(git->second); - visited->emplace(git->second); - } - } - } - // Update the current level and clean the next level. - current_nodes = next_nodes; - next_nodes.clear(); - } - return ret; -} - -void CallGraphEntry::CleanCallGraphEntries() { - while (!called_globals_.empty()) { - // Decrement the reference counter - called_globals_.back().second->DecRef(); - called_globals_.pop_back(); - } -} - -inline void CallGraphEntry::AddCalledGlobal(CallGraphEntry* cg_node) { - called_globals_.emplace_back(global_, cg_node); - // Increment the reference to indicate that another call site is found for - // the callee in `cg_node`. - cg_node->IncRef(); - // Mark the global function as recursive if it calls itself. - if (global_ == cg_node->GetGlobalVar()) { - cg_node->is_recursive_ = true; - } -} - -// Remove an edge from the current global function to the callee. -void CallGraphEntry::RemoveCallTo(const GlobalVar& callee) { - for (auto it = begin();; ++it) { - ICHECK(it != end()) << "Cannot find global function " << callee->name_hint << " to remove!"; - if (it->second->GetGlobalVar() == callee) { - // Only remove one occurrence of the call site. - it->second->DecRef(); - *it = called_globals_.back(); - called_globals_.pop_back(); - return; - } - } -} - -// Remove all edges from the current global function to the callee. -void CallGraphEntry::RemoveAllCallTo(CallGraphEntry* callee) { - for (uint32_t i = 0, e = size(); i != e;) { - if (called_globals_[i].second == callee) { - callee->DecRef(); - called_globals_[i] = called_globals_.back(); - called_globals_.pop_back(); - --e; - } else { - ++i; - } - } - // Make sure all references to the callee are removed. - ICHECK_EQ(callee->GetRefCount(), 0U) - << "All references to " << callee->GetNameHint() << " should have been removed"; -} - -void CallGraphEntry::Print(std::ostream& os) const { - if (!global_.defined()) { - os << "GlobalVar is not defined\n"; - return; - } - - os << "Call graph node: " << global_->name_hint; - os << " at: " << this << ", #refs = " << GetRefCount() << "\n"; - - for (const auto& it : *this) { - os << " call site: <" << it.first->name_hint << "> calls "; - os << it.second->GetNameHint() << "\n"; - } - os << "\n"; -} - -std::ostream& operator<<(std::ostream& os, const CallGraph& cg) { - cg->Print(os); - return os; -} - -std::ostream& operator<<(std::ostream& os, const CallGraphEntry& cgn) { - cgn.Print(os); - return os; -} - -TVM_REGISTER_NODE_TYPE(CallGraphNode); - -TVM_STATIC_IR_FUNCTOR(ReprPrinter, vtable) - .set_dispatch([](const ObjectRef& ref, ReprPrinter* p) { - auto* node = static_cast(ref.get()); - ICHECK(node); - p->stream << "CallGraph: \n" << GetRef(node); - }); - -TVM_REGISTER_GLOBAL("relay.analysis.CallGraph").set_body_typed([](IRModule module) { - return CallGraph(module); -}); - -TVM_REGISTER_GLOBAL("relay.analysis.PrintCallGraph").set_body_typed([](CallGraph call_graph) { - std::stringstream ss; - ss << call_graph; - return ss.str(); -}); - -TVM_REGISTER_GLOBAL("relay.analysis.GetModule").set_body_typed([](CallGraph call_graph) { - return call_graph->module; -}); - -TVM_REGISTER_GLOBAL("relay.analysis.PrintCallGraphGlobalVar") - .set_body_typed([](CallGraph call_graph, GlobalVar var) { - const auto* entry_node = call_graph[var]; - std::stringstream ss; - ss << *entry_node; - return ss.str(); - }); - -TVM_REGISTER_GLOBAL("relay.analysis.GetRefCountGlobalVar") - .set_body_typed([](CallGraph call_graph, GlobalVar var) { - const auto* entry_node = call_graph[var]; - return static_cast(entry_node->GetRefCount()); - }); - -TVM_REGISTER_GLOBAL("relay.analysis.GetGlobalVarCallCount") - .set_body_typed([](CallGraph call_graph, GlobalVar var) { - const auto* entry_node = call_graph[var]; - return static_cast(entry_node->size()); - }); - -TVM_REGISTER_GLOBAL("relay.analysis.IsRecursive") - .set_body_typed([](CallGraph call_graph, GlobalVar var) { - const auto* entry_node = call_graph[var]; - return entry_node->IsRecursive(); - }); - -} // namespace relay -} // namespace tvm diff --git a/src/relay/analysis/call_graph.h b/src/relay/analysis/call_graph.h deleted file mode 100644 index 091891acd414..000000000000 --- a/src/relay/analysis/call_graph.h +++ /dev/null @@ -1,478 +0,0 @@ -/* - * 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 src/relay/analysis/call_graph.h - * \brief Define data structures for the call graph of a IRModule. It borrows - * the idea how LLVM constructs CallGraph. - * - * https://llvm.org/doxygen/CallGraph_8h_source.html - */ - -#ifndef TVM_RELAY_ANALYSIS_CALL_GRAPH_H_ -#define TVM_RELAY_ANALYSIS_CALL_GRAPH_H_ - -#include -#include -#include -#include - -#include -#include -#include -#include -#include -#include - -namespace tvm { -namespace relay { - -class CallGraphEntry; -class CallGraph; - -class CallGraphNode : public Object { - using CallGraphMap = std::unordered_map>; - // Create iterator alias for a CallGraphNode object. - using iterator = CallGraphMap::iterator; - using const_iterator = CallGraphMap::const_iterator; - - public: - /*! \brief The IR module for creating a CallGraphNode. */ - IRModule module; - - /*! \brief Default constructor. */ - CallGraphNode() {} - - void VisitAttrs(AttrVisitor* v) { v->Visit("module", &module); } - - /*! - * \brief Print the call graph. - * - * \param os The stream for printing. - */ - void Print(std::ostream& os) const; - - /*! \return The begin iterator. */ - iterator begin() { return call_graph_.begin(); } - /*! \return The end iterator. */ - iterator end() { return call_graph_.end(); } - /*! \return The begin iterator. */ - const_iterator begin() const { return call_graph_.begin(); } - /*! \return The end iterator. */ - const_iterator end() const { return call_graph_.end(); } - - /*! - * \brief Get an element from the CallGraphNode using a GlobalVar. - * - * \param gv The GlobalVar used for indexing. - * - * \return The fetched element. - */ - const CallGraphEntry* operator[](const GlobalVar& gv) const; - /*! - * \brief Get an element from the CallGraphNode using a GlobalVar. - * - * \param gv The GlobalVar used for indexing. - * - * \return The fetched element. - */ - CallGraphEntry* operator[](const GlobalVar& gv); - /*! - * \brief Get an element from the CallGraphNode using the global function name. - * - * \param gvar_name The global function name used for indexing. - * - * \return The fetched element. - */ - const CallGraphEntry* operator[](const std::string& gvar_name) const { - return (*this)[module->GetGlobalVar(gvar_name)]; - } - /*! - * \brief Get an element from the CallGraphNode using the global function name. - * - * \param gvar_name The global function name used for indexing. - * - * \return The fetched element. - */ - CallGraphEntry* operator[](const std::string& gvar_name) { - return (*this)[module->GetGlobalVar(gvar_name)]; - } - - /*! - * \brief Get the global function corresponding to the variable. - * - * \param var The global variable. - * - * \return The found global function. - */ - BaseFunc GetGlobalFunction(const GlobalVar& var) const; - - /*! - * \brief Get the entries/root nodes of CallGraphNode. - * - * Entry functions are never referenced by other functions. - * Note these functions can be recursive as well. - * - * \return The list of CallGraphEntry that represent entry nodes. - */ - std::vector GetEntryGlobals() const; - - /*! - * \brief Remove a GlobalVar in a given CallGraphEntry from the current - * IR module. - * - * \param cg_node The CallGraphEntry that contains a global function to be - * removed. - * \param update_call_graph Indicate if we will update the CallGraph as well - * since updating is costly. We are only able to remove a leaf function - * when update_call_graph is disabled because the edges pointing to - * functions being removed are not updated. - * - * \return The GlobalVar removed from the current module. - */ - GlobalVar RemoveGlobalVarFromModule(CallGraphEntry* cg_node, bool update_call_graph = false); - - /*! - * \brief Lookup a GlobalVar for the CallGraphNode. It creates an entry for - * the GlobalVar if it doesn't exist. - * - * \param gv The GlobalVar for query. - * - * \return The queried entry. - */ - CallGraphEntry* LookupGlobalVar(const GlobalVar& gv); - - /*! - * \brief Get the entries from the CallGraphNode in the topological order. - * - * This is useful for various module-level optimizations/analysis. For example, - * inlining requires the correct order of the functions being processed, i.e. - * callee should be always handled before callers. - * - * \return The list of collected entries that are sorted in the topological order. - */ - std::vector TopologicalOrder() const; - - static constexpr const char* _type_key = "relay.CallGraph"; - TVM_DECLARE_FINAL_OBJECT_INFO(CallGraphNode, Object); - - private: - /*! - * \brief Create a CallGraphEntry for a global function and add it to the - * CallGraphNode. - * - * \param gv The global var. - * \param func The global function corresponding to `gv`. - */ - void AddToCallGraph(const GlobalVar& gv, const Function& func); - - /*! \brief A record contains GlobalVar to CallGraphEntry mapping. */ - CallGraphMap call_graph_; - - friend CallGraph; -}; - -/*! - * \brief The class that represents the call graph of a Relay IR module. It also - * provides a variety of utility functions for users to query, view, and update - * a call graph. - */ -class CallGraph : public ObjectRef { - using CallGraphMap = std::unordered_map>; - // Create iterator alias for a CallGraph object. - using iterator = CallGraphMap::iterator; - using const_iterator = CallGraphMap::const_iterator; - - public: - /*! - * \brief Construct a CallGraph from a IR module. - * - * \param module The IR module - */ - explicit CallGraph(IRModule module); - - /*! - * \brief Construct from an object pointer. - * \param n The object pointer. - */ - explicit CallGraph(ObjectPtr n) : ObjectRef(n) {} - - /*! \return The begin iterator. */ - iterator begin() { - auto* n = operator->(); - ICHECK(n); - return n->begin(); - } - /*! \return The end iterator. */ - iterator end() { - auto* n = operator->(); - ICHECK(n); - return n->end(); - } - /*! \return The begin iterator. */ - const_iterator begin() const { - const auto* n = operator->(); - ICHECK(n); - return n->begin(); - } - /*! \return The end iterator. */ - const_iterator end() const { - const auto* n = operator->(); - ICHECK(n); - return n->end(); - } - - /*! - * \brief Get an element from the CallGraph using a GlobalVar. - * - * \param gv The GlobalVar used for indexing. - * - * \return The fetched element. - */ - const CallGraphEntry* operator[](const GlobalVar& gv) const { - const auto* n = operator->(); - ICHECK(n); - return (*n)[gv]; - } - /*! - * \brief Get an element from the CallGraph using a GlobalVar. - * - * \param gv The GlobalVar used for indexing. - * - * \return The fetched element. - */ - CallGraphEntry* operator[](const GlobalVar& gv) { - auto* n = operator->(); - ICHECK(n); - return (*n)[gv]; - } - /*! - * \brief Get an element from the CallGraph using the global function name. - * - * \param gvar_name The global function name used for indexing. - * - * \return The fetched element. - */ - const CallGraphEntry* operator[](const std::string& gvar_name) const { - const auto* n = operator->(); - ICHECK(n); - return (*n)[gvar_name]; - } - /*! - * \brief Get an element from the CallGraph using the global function name. - * - * \param gvar_name The global function name used for indexing. - * - * \return The fetched element. - */ - CallGraphEntry* operator[](const std::string& gvar_name) { - auto* n = operator->(); - ICHECK(n); - return (*n)[gvar_name]; - } - - /*! \return mutable pointers to the node. */ - CallGraphNode* operator->() const { - auto* ptr = get_mutable(); - ICHECK(ptr != nullptr); - return static_cast(ptr); - } - - private: - /*! \brief Overload the << operator to print a call graph. */ - friend std::ostream& operator<<(std::ostream& os, const CallGraph&); -}; - -/*! - * \brief A node in the call graph. It maintains the edges from a caller to - * all callees. - */ -class CallGraphEntry { - public: - using CallGraphEntryPair = std::pair; - using CallGraphEntryVector = std::vector; - using CallGraphEntrySet = std::unordered_set; - // Create iterator alias for a CallGraphEntry object. - using iterator = std::vector::iterator; - using const_iterator = std::vector::const_iterator; - - /*! - * \brief Construct from a GlobalVar. - * - * \param gv The GlobalVar to create a CallGraphEntry. - */ - explicit CallGraphEntry(const GlobalVar& gv) : global_(gv) {} - /*! - * \brief Delete copy constructor. - */ - CallGraphEntry(const CallGraphEntry&) = delete; - /*! \brief Delete assignment. */ - CallGraphEntry& operator=(const CallGraphEntry&) = delete; - - /*! \return The begin iterator */ - iterator begin() { return called_globals_.begin(); } - /*! \return The end iterator */ - iterator end() { return called_globals_.end(); } - /*! \return The const begin iterator */ - const_iterator begin() const { return called_globals_.begin(); } - /*! \return The const end iterator */ - const_iterator end() const { return called_globals_.end(); } - - /*! - * \brief Return if the list of called nodes is empty. - * - * \return true if the list is empty. Otherwise, false. - */ - bool empty() const { return called_globals_.empty(); } - - /*! - * \brief Return the size of the list that represents the nodes are called by - * the current node. - * - * \return The number of called nodes. - */ - uint32_t size() const { return static_cast(called_globals_.size()); } - - /*! - * \brief Fetch the i-th CallGraphEntry from the list of nodes that are called - * by the current function. - * - * \param i The index. - * - * \return The fetched CallGraphEntry. - */ - CallGraphEntry* operator[](size_t i) const { - ICHECK_LT(i, called_globals_.size()) << "Invalid Index"; - return called_globals_[i].second; - } - - /*! - * \brief Print the call graph that is stemmed from the current CallGraphEntry. - * - * \param os The stream for printing. - */ - void Print(std::ostream& os) const; - - /*! - * \brief Return the number of times the global function is referenced. - * - * \return The count. - */ - uint32_t GetRefCount() const { return ref_cnt_; } - - /*! - * \brief Return the GlobalVar stored in the current CallGraphEntry. - * - * \return The GlobalVar. - */ - GlobalVar GetGlobalVar() const { return global_; } - - /*! - * \brief Return the name hint of the GlobalVar stored in the CallGraphEntry. - * - * \return The name hint of the global function. - */ - std::string GetNameHint() const { return global_->name_hint; } - - /*! - * \brief Return if the global function corresponding to the current - * CallGraphEntry is a recursive function. - * - * \return true if it is recursive. Otherwise, false. - */ - bool IsRecursive() const { return is_recursive_; } - - /*! - * \brief Return if the global function corresponding to the current - * CallGraphEntry is both a recursive function and an entry function. This type - * of function only has one reference which is called by itself. - * - * \return true if it is both a recursive function and an entry. Otherwise, false. - */ - bool IsRecursiveEntry() const { return GetRefCount() == 1 && IsRecursive(); } - - /*! - * \brief Return the topological order of the CallGraphEntry. - * - * \param visited A set of CallGraphEntry objects that have been visited. - * - * \return The list of CallGraphEntry that is represented in topological order. - */ - std::vector TopologicalOrder( - CallGraphEntrySet* visited = new CallGraphEntrySet()) const; - - /*! - * \brief Remove all edges from the current CallGraphEntry to any global - * function it calls. - */ - void CleanCallGraphEntries(); - - /*! - * \brief Add a node to the list of nodes that are being called by the current - * global function. - * - * \param cg_node The CallGraphEntry that will be added to the call list. - */ - void AddCalledGlobal(CallGraphEntry* cg_node); - - /*! - * \brief Remove a call edge to the global function from the current - * function. - * - * \param callee The function that is being called. - */ - void RemoveCallTo(const GlobalVar& callee); - - /*! - * \brief Remove all the edges that represent that calls to the global function - * stored in a given CallGraphEntry. - * - * \param callee The function that is being called. - */ - void RemoveAllCallTo(CallGraphEntry* callee); - - private: - /*! \brief Decrement the reference counter by 1. */ - void DecRef() { - ICHECK_GT(ref_cnt_, 0); - --ref_cnt_; - } - /*! \brief Increment the reference counter by 1. */ - void IncRef() { ++ref_cnt_; } - - /*! - * \brief Mark if the global function stored in the CallGraphEntry is - * recursive function. - */ - bool is_recursive_{false}; - /*! \brief Count the number of times the global function is referenced. */ - uint32_t ref_cnt_{0}; - /*! \brief The GlobalVar stored in the current CallGraphEntry. */ - GlobalVar global_; - /*! \brief The list of entries called by the current CallGraphEntry. */ - CallGraphEntryVector called_globals_; - - friend class CallGraph; - /*! \brief Overload the << operator to print a call graph node. */ - friend std::ostream& operator<<(std::ostream& os, const CallGraphEntry&); -}; - -} // namespace relay -} // namespace tvm -#endif // TVM_RELAY_ANALYSIS_CALL_GRAPH_H_ diff --git a/src/relay/analysis/dependency_graph.cc b/src/relay/analysis/dependency_graph.cc deleted file mode 100644 index 91711fa4baa8..000000000000 --- a/src/relay/analysis/dependency_graph.cc +++ /dev/null @@ -1,212 +0,0 @@ -/* - * 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 src/relay/analysis/dependency_graph.cc - * \brief Implementation of dependency graph APIs. - */ -#include "dependency_graph.h" - -#include - -#include -#include - -namespace tvm { -namespace relay { - -// Creator of DependencyGraph -class DependencyGraph::Creator : private MixedModeVisitor { - public: - explicit Creator(support::Arena* arena) : arena_(arena) {} - - DependencyGraph Create(const Expr& body) { - this->VisitExpr(body); - return std::move(graph_); - } - - private: - /*! \brief allocator of all the internal node object */ - support::Arena* arena_; - // The output. - DependencyGraph graph_; - // Update the message stored at the node. - void Depend(DependencyGraph::Node* parent, const Expr& child) { - VisitExpr(child); - - ICHECK_NE(graph_.expr_node.count(child), 0); - - Depend(parent, graph_.expr_node[child]); - } - - void Depend(DependencyGraph::Node* parent, DependencyGraph::Node* child) { - auto* parent_link = arena_->make>(); - parent_link->value = parent; - child->parents.Push(parent_link); - - auto* child_link = arena_->make>(); - child_link->value = child; - parent->children.Push(child_link); - } - - std::unordered_set visited_; - - DependencyGraph::Node* NewNode(bool new_scope) { - auto* ret = arena_->make(); - ret->new_scope = new_scope; - return ret; - } - - void VisitLeaf(const Expr& e) override { - if (visited_.count(e) == 0) { - if (graph_.expr_node.count(e) == 0) { - graph_.expr_node[e] = NewNode(false); - } - visited_.insert(e); - MixedModeVisitor::VisitLeaf(e); - graph_.post_dfs_order.push_back(graph_.expr_node[e]); - } - } - - void VisitExpr_(const CallNode* c) final { - DependencyGraph::Node* n = graph_.expr_node[GetRef(c)]; - Depend(n, c->op); - for (const auto& a : c->args) { - Depend(n, a); - } - } - - void VisitExpr_(const TupleNode* t) final { - DependencyGraph::Node* n = graph_.expr_node[GetRef(t)]; - for (const auto& a : t->fields) { - Depend(n, a); - } - } - - void VisitExpr_(const TupleGetItemNode* t) final { - DependencyGraph::Node* n = graph_.expr_node[GetRef(t)]; - Depend(n, t->tuple); - } - - void VisitExpr_(const RefCreateNode* r) final { - DependencyGraph::Node* n = graph_.expr_node[GetRef(r)]; - Depend(n, r->value); - } - - void VisitExpr_(const RefReadNode* r) final { - DependencyGraph::Node* n = graph_.expr_node[GetRef(r)]; - Depend(n, r->ref); - } - - void VisitExpr_(const RefWriteNode* r) final { - DependencyGraph::Node* n = graph_.expr_node[GetRef(r)]; - Depend(n, r->ref); - Depend(n, r->value); - } - - void VisitExpr_(const IfNode* i) final { - DependencyGraph::Node* n = graph_.expr_node[GetRef(i)]; - DependencyGraph::Node* t = NewNode(true); - DependencyGraph::Node* f = NewNode(true); - Depend(n, i->cond); - Depend(n, t); - Depend(n, f); - Depend(t, i->true_branch); - Depend(f, i->false_branch); - graph_.post_dfs_order.push_back(f); - graph_.post_dfs_order.push_back(t); - } - - void VisitExpr_(const FunctionNode* f) final { - DependencyGraph::Node* n = graph_.expr_node[GetRef(f)]; - DependencyGraph::Node* b = NewNode(true); - Depend(n, b); - for (const auto& p : f->params) { - Depend(b, p); - } - Depend(b, f->body); - graph_.post_dfs_order.push_back(b); - } - - void VisitExpr_(const LetNode* l) final { - std::unordered_map b_map; - auto pre_visit = [&](const LetNode* op) { - Expr e = GetRef(op); - // Derived VisitLeaf - if (visited_.count(e) == 0) { - if (graph_.expr_node.count(e) == 0) { - graph_.expr_node[e] = NewNode(false); - } - visited_.insert(e); - } - DependencyGraph::Node* n = graph_.expr_node[e]; - DependencyGraph::Node* b = NewNode(true); - Depend(n, b); - Depend(b, op->var); - Depend(b, op->value); - b_map[op] = b; - }; - auto post_visit = [&](const LetNode* op) { - ICHECK(b_map.count(op)); - DependencyGraph::Node* b = b_map[op]; - Expr e = GetRef(op); - Depend(b, op->body); - graph_.post_dfs_order.push_back(b); - if (op != l) { - // Base VisitLeaf - this->visit_counter_[op]++; - // Derived VisitLeaf - graph_.post_dfs_order.push_back(graph_.expr_node[e]); - } - }; - ExpandANormalForm(l, pre_visit, post_visit); - } - - void VisitExpr_(const MatchNode* m) final { - DependencyGraph::Node* n = graph_.expr_node[GetRef(m)]; - Depend(n, m->data); - std::vector v; - for (const Clause& c : m->clauses) { - DependencyGraph::Node* b = NewNode(true); - Depend(n, b); - Depend(b, c->rhs); - v.push_back(b); - } - for (auto it = v.rbegin(); it != v.rend(); ++it) { - graph_.post_dfs_order.push_back(*it); - } - } - - void VisitExpr_(const VarNode* v) final {} - - void VisitExpr_(const GlobalVarNode* v) final {} - - void VisitExpr_(const ConstantNode* c) final {} - - void VisitExpr_(const OpNode* o) final {} - - void VisitExpr_(const ConstructorNode* c) final {} -}; - -DependencyGraph DependencyGraph::Create(support::Arena* arena, const Expr& body) { - return Creator(arena).Create(body); -} - -} // namespace relay -} // namespace tvm diff --git a/src/relay/analysis/dependency_graph.h b/src/relay/analysis/dependency_graph.h deleted file mode 100644 index 1de125770c7c..000000000000 --- a/src/relay/analysis/dependency_graph.h +++ /dev/null @@ -1,77 +0,0 @@ -/* - * 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 src/relay/analysis/dependency_graph.h - * \brief create a dependency graph. - */ -#ifndef TVM_RELAY_ANALYSIS_DEPENDENCY_GRAPH_H_ -#define TVM_RELAY_ANALYSIS_DEPENDENCY_GRAPH_H_ - -#include - -#include -#include - -#include "../../support/arena.h" -#include "../transforms/let_list.h" - -namespace tvm { -namespace relay { - -using support::LinkedList; -using support::LinkNode; - -/* DependencyGraph track input and output of an Expr. - * Additionally, dummy scope is created to model scope. - * It allow us to traverse the graph in reverse order. - */ -class DependencyGraph { - public: - /*! \brief A node in the graph. */ - struct Node { - // Determine scope boundaries. Used for calculating scopes, not for - // constructing dependency graph. - bool new_scope = false; - // incoming edges - LinkedList children; - // outgoing edges - LinkedList parents; - }; - - /*! \brief Maps a Relay Expr to its node in the dependency graph. */ - std::unordered_map expr_node; - - /*! \brief The dependency graph in post DFS order. */ - std::vector post_dfs_order; - - /*! - * \brief Create a dependency graph. - * \param arena The arena used for data allocation. - * \param body The body of the expression to create a graph. - */ - static DependencyGraph Create(support::Arena* arena, const Expr& body); - - private: - class Creator; -}; - -} // namespace relay -} // namespace tvm -#endif // TVM_RELAY_ANALYSIS_DEPENDENCY_GRAPH_H_ diff --git a/src/relay/analysis/extract_fake_quantized_ops.cc b/src/relay/analysis/extract_fake_quantized_ops.cc deleted file mode 100644 index d66bbd635480..000000000000 --- a/src/relay/analysis/extract_fake_quantized_ops.cc +++ /dev/null @@ -1,80 +0,0 @@ -/* - * 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 extract_fake_quantized_ops.cc - * \brief Extract fake quantized operators from an IRModule - */ -#include -#include -#include - -#include "../transforms/fake_quantization_to_integer.h" - -namespace tvm { -namespace relay { - -using ExprSet = std::unordered_set; - -class ExtractFakeQuantizedOpsWrapper : private MixedModeVisitor { - public: - Map Extract(const IRModule& m) { - IRModule mod(m); - mod = transform::InferType()(mod); - VisitExpr(mod->Lookup("main")); - - return fake_quantized_op_freqs_; - } - - private: - using MixedModeVisitor::VisitExpr_; - - void VisitExpr_(const CallNode* call_node) override { - if (call_node->op == quantize_op_) { - SubgraphExtractor extractor; - ExprSet subgraph = extractor.GetSubgraph(GetRef(call_node)); - - for (auto expr : subgraph) { - const Op op = Downcast(expr.as()->op); - if (op != dequantize_op_) { - if (fake_quantized_op_freqs_.find(op->name) != fake_quantized_op_freqs_.end()) { - fake_quantized_op_freqs_.Set(op->name, - fake_quantized_op_freqs_.at(op->name).IntValue() + 1); - } else { - fake_quantized_op_freqs_.Set(op->name, 1); - } - } - } - } - } - - Map fake_quantized_op_freqs_; - const Op quantize_op_ = Op::Get("qnn.quantize"); - const Op dequantize_op_ = Op::Get("qnn.dequantize"); -}; - -Map ExtractFakeQuantizedOpsPacked(const IRModule& mod) { - return ExtractFakeQuantizedOpsWrapper().Extract(mod); -} - -TVM_REGISTER_GLOBAL("relay.analysis.ExtractFakeQuantizedOps") - .set_body_typed(ExtractFakeQuantizedOpsPacked); - -} // namespace relay -} // namespace tvm diff --git a/src/relay/analysis/extract_fused_functions.cc b/src/relay/analysis/extract_fused_functions.cc deleted file mode 100644 index e76b54e2d0b7..000000000000 --- a/src/relay/analysis/extract_fused_functions.cc +++ /dev/null @@ -1,83 +0,0 @@ -/* - * 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 extract_fused_functions.cc - * \brief Apply fusion and extract fused primitive functions from an IRModule - */ -#include -#include -#include -#include -#include - -namespace tvm { -namespace relay { - -class FusedFunctionExtractorWrapper : private ExprVisitor { - public: - explicit FusedFunctionExtractorWrapper(const IRModule& mod) : mod_(mod) {} - - IRModule Extract() { - VisitExpr(this->mod_->Lookup("main")); - - auto functions = Map(); - for (auto pair : this->functions) { - functions.Set(GlobalVar(pair.first), pair.second); - } - - this->mod_->functions = functions; - return this->mod_; - } - - private: - const IRModule mod_; - // This is not simply Map because GlobalVar doesn't - // have the desired equals property - Map functions; - - void VisitExpr_(const FunctionNode* n) final { - if (n->HasNonzeroAttr(attr::kPrimitive)) { - // Add function to functions, keyed by function hash string - Function func = Function(n->params, n->body, n->ret_type, n->type_params, n->attrs); - size_t hash_ = tvm::StructuralHash()(func); - this->functions.Set(std::to_string(hash_), func); - } - - ExprVisitor::VisitExpr_(n); - } -}; - -namespace transform { - -Pass ExtractFusedFunctions() { - runtime::TypedPackedFunc pass_func = - [=](IRModule m, PassContext pc) { return FusedFunctionExtractorWrapper(m).Extract(); }; - auto fused_function_extractor_pass = CreateModulePass(pass_func, 1, "ExtractFusedFunctions", {}); - - return Sequential({SimplifyInference(), FuseOps(3), fused_function_extractor_pass}, - "ExtractFusedFunctions"); -} - -TVM_REGISTER_GLOBAL("relay.analysis.ExtractFusedFunctions").set_body_typed(ExtractFusedFunctions); - -} // namespace transform - -} // namespace relay -} // namespace tvm diff --git a/src/relay/analysis/extract_intermediate_expr.cc b/src/relay/analysis/extract_intermediate_expr.cc deleted file mode 100644 index d7466e2729db..000000000000 --- a/src/relay/analysis/extract_intermediate_expr.cc +++ /dev/null @@ -1,88 +0,0 @@ -/* - * 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 extract_intermediate_expr.cc - * \brief Used for extracting Relay Expr - by the expression ID of the main function - that we can see in `print(mod["main"])`. - */ -#include -#include -#include -#include - -namespace tvm { -namespace relay { - -class ExtractIntermediateExprWrapper : private MixedModeVisitor { - public: - explicit ExtractIntermediateExprWrapper(const IRModule& mod, const int expr_id) - : mod_(mod), target_expr_id_(expr_id), counter_(0) {} - - IRModule Extract() { - VisitExpr(this->mod_->Lookup("main")); - - // ensure the target expr_id we want to extract is valid. - ICHECK(target_expr_id_ >= 0 && target_expr_id_ < counter_); - - return IRModule::FromExpr(target_op_, {}); - } - - private: - using MixedModeVisitor::VisitExpr_; - - const IRModule mod_; - /*! \brief the expr id that we want to extract. */ - const int target_expr_id_; - int counter_; - Expr target_op_; - - void VisitExpr_(const CallNode* n) final { - CheckCounterAndIncrease(GetRef(n)); - MixedModeVisitor::VisitExpr_(n); - } - - void VisitExpr_(const TupleNode* n) final { - CheckCounterAndIncrease(GetRef(n)); - MixedModeVisitor::VisitExpr_(n); - } - - void VisitExpr_(const TupleGetItemNode* n) final { - CheckCounterAndIncrease(GetRef(n)); - MixedModeVisitor::VisitExpr_(n); - } - - void CheckCounterAndIncrease(const Expr& expr) { - if (target_expr_id_ == counter_) { - target_op_ = expr; - } - ++counter_; - } -}; - -IRModule ExtractIntermediateExprPacked(const IRModule& mod, const int expr_id) { - return ExtractIntermediateExprWrapper(mod, expr_id).Extract(); -} - -TVM_REGISTER_GLOBAL("relay.analysis.ExtractIntermediateExpr") - .set_body_typed(ExtractIntermediateExprPacked); - -} // namespace relay -} // namespace tvm diff --git a/src/relay/analysis/extract_operators.cc b/src/relay/analysis/extract_operators.cc deleted file mode 100644 index 051c1971f20e..000000000000 --- a/src/relay/analysis/extract_operators.cc +++ /dev/null @@ -1,77 +0,0 @@ -/* - * 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 extract_operators.cc - * \brief Extract unique operators from an IRModule - */ -#include -#include -#include -#include - -namespace tvm { -namespace relay { - -class OperatorExtractorWrapper : private MixedModeVisitor { - public: - explicit OperatorExtractorWrapper(const IRModule& mod) : mod_(mod) {} - - Map Extract() { - VisitExpr(this->mod_->Lookup("main")); - - return operator_freqs_; - } - - private: - using MixedModeVisitor::VisitExpr_; - - const IRModule mod_; - /*! \brief Map of operator to frequency. */ - Map operator_freqs_; - - void VisitExpr_(const CallNode* n) final { - VisitExpr(n->op); - - auto op = n->op.as(); - if (op) { - auto it = operator_freqs_.find(op->name); - ICHECK(it != operator_freqs_.end()) - << "Call's OpNode must be visited and registered before access"; - operator_freqs_.Set(op->name, 1 + operator_freqs_.at(op->name).IntValue()); - } - - MixedModeVisitor::VisitExpr_(n); - } - - void VisitExpr_(const OpNode* n) final { - // NOTE: OpNode is visited only once for every operator kind - // regardless of how many times that op appears in the graph. - operator_freqs_.Set(n->name, 0U); - } -}; - -Map ExtractOperatorsPacked(const IRModule& mod) { - return OperatorExtractorWrapper(mod).Extract(); -} - -TVM_REGISTER_GLOBAL("relay.analysis.ExtractOperators").set_body_typed(ExtractOperatorsPacked); - -} // namespace relay -} // namespace tvm diff --git a/src/relay/analysis/feature.cc b/src/relay/analysis/feature.cc deleted file mode 100644 index f72b4e105749..000000000000 --- a/src/relay/analysis/feature.cc +++ /dev/null @@ -1,153 +0,0 @@ -/* - * 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 feature.cc - * \brief Detect features used in Expr/Module - */ -#include -#include -#include -#include -#include - -#include "../transforms/pass_utils.h" - -namespace tvm { -namespace relay { - -FeatureSet DetectFeature(const Expr& expr) { - if (!expr.defined()) { - return FeatureSet::No(); - } - struct FeatureDetector : ExprVisitor { - std::unordered_set visited_; - FeatureSet fs = FeatureSet::No(); - - void VisitExpr(const Expr& expr) final { - if (visited_.count(expr) == 0) { - visited_.insert(expr); - ExprVisitor::VisitExpr(expr); - } else { - if (!IsAtomic(expr)) { - fs += fGraph; - } - } - } -#define DETECT_CONSTRUCT(CONSTRUCT_NAME, STMT) \ - void VisitExpr_(const CONSTRUCT_NAME##Node* op) final { STMT fs += f##CONSTRUCT_NAME; } -#define DETECT_DEFAULT_CONSTRUCT(CONSTRUCT_NAME) \ - DETECT_CONSTRUCT(CONSTRUCT_NAME, { ExprVisitor::VisitExpr_(op); }) - DETECT_DEFAULT_CONSTRUCT(Var) - DETECT_DEFAULT_CONSTRUCT(GlobalVar) - DETECT_DEFAULT_CONSTRUCT(Constant) - DETECT_DEFAULT_CONSTRUCT(Tuple) - DETECT_DEFAULT_CONSTRUCT(TupleGetItem) - DETECT_CONSTRUCT(Function, { - if (!op->HasNonzeroAttr(attr::kPrimitive)) { - ExprVisitor::VisitExpr_(op); - } - }) - DETECT_DEFAULT_CONSTRUCT(Op) - DETECT_DEFAULT_CONSTRUCT(Call) - DETECT_CONSTRUCT(Let, { - for (const Var& v : FreeVars(op->value)) { - if (op->var == v) { - fs += fLetRec; - } - } - ExprVisitor::VisitExpr_(op); - }) - DETECT_DEFAULT_CONSTRUCT(If) - DETECT_DEFAULT_CONSTRUCT(RefCreate) - DETECT_DEFAULT_CONSTRUCT(RefRead) - DETECT_DEFAULT_CONSTRUCT(RefWrite) - DETECT_DEFAULT_CONSTRUCT(Constructor) - DETECT_DEFAULT_CONSTRUCT(Match) -#undef DETECT_DEFAULT_CONSTRUCT - } fd; - fd(expr); - return fd.fs; -} - -std::string FeatureSet::ToString() const { - std::string ret; - ret += "["; - size_t detected = 0; -#define DETECT_FEATURE(FEATURE_NAME) \ - ++detected; \ - if (bs_[FEATURE_NAME]) { \ - ret += #FEATURE_NAME; \ - ret += ", "; \ - } - DETECT_FEATURE(fVar); - DETECT_FEATURE(fGlobalVar); - DETECT_FEATURE(fConstant); - DETECT_FEATURE(fTuple); - DETECT_FEATURE(fTupleGetItem); - DETECT_FEATURE(fFunction); - DETECT_FEATURE(fOp); - DETECT_FEATURE(fCall); - DETECT_FEATURE(fLet); - DETECT_FEATURE(fIf); - DETECT_FEATURE(fRefCreate); - DETECT_FEATURE(fRefRead); - DETECT_FEATURE(fRefWrite); - DETECT_FEATURE(fConstructor); - DETECT_FEATURE(fMatch); - DETECT_FEATURE(fGraph); - DETECT_FEATURE(fLetRec); -#undef DETECT_FEATURE - ICHECK(detected == feature_count) << "some feature not printed"; - ret += "]"; - return ret; -} - -FeatureSet DetectFeature(const IRModule& mod) { - FeatureSet fs = FeatureSet::No(); - for (const auto& f : mod->functions) { - fs += DetectFeature(f.second); - } - return fs; -} - -Array PyDetectFeature(const Expr& expr, const Optional& mod) { - FeatureSet fs = DetectFeature(expr); - if (mod.defined()) { - fs = fs + DetectFeature(mod.value()); - } - return static_cast>(fs); -} - -TVM_REGISTER_GLOBAL("relay.analysis.detect_feature").set_body_typed(PyDetectFeature); - -void CheckFeature(const Expr& expr, const FeatureSet& fs) { - auto dfs = DetectFeature(expr); - ICHECK(dfs.is_subset_of(fs)) << AsText(expr, false) - << "\nhas unsupported feature: " << (dfs - fs).ToString(); -} - -void CheckFeature(const IRModule& mod, const FeatureSet& fs) { - for (const auto& f : mod->functions) { - CheckFeature(f.second, fs); - } -} - -} // namespace relay -} // namespace tvm diff --git a/src/relay/analysis/get_calibration_data.cc b/src/relay/analysis/get_calibration_data.cc deleted file mode 100644 index 0d99e0a9ecad..000000000000 --- a/src/relay/analysis/get_calibration_data.cc +++ /dev/null @@ -1,201 +0,0 @@ -/* - * 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 src/relay/analysis/get_calibration_data.cc - * - * \brief To get the calibration data, we need to perform two - * steps. First, we need to prepare the module that generates - * the tensor values (GetCalibrateModule). Second, we need to - * generate the mapping between the values and the functions - * (GetCalibrateOutputMap). - */ - -#include -#include -#include - -namespace tvm { -namespace relay { - -/*! - * \brief This function returns a module that will be used by - * the relay graph executor for collecting the calibration data. - * To do that, we first make all inputs and outputs of each - * function into the final output (i.e., the final output is a - * tuple of tensors). Then, we change the compiler attribute of - * each function. Finally, we mark all function to be inlined. - */ - -class Collector : public ExprRewriter { - public: - explicit Collector(const IRModule& module) : module_(module) {} - - Expr Rewrite_(const CallNode* call, const Expr& post) final { - // check if the function implementation is available - // intrinsic functions are excluded for now - if (call->op->IsInstance()) { - auto var = Downcast(call->op); - ICHECK(module_->ContainGlobalVar(var->name_hint)) << "Function " << var << " is not defined"; - // we only handle functions with Compiler attribute set - auto func = Downcast(module_->Lookup(var)); - if (func->GetAttr(attr::kCompiler)) { - // collect all the inputs and outputs - for (const auto& it : call->args) new_outputs_.push_back(it); - new_outputs_.push_back(post); - } - } - return post; - } - - Array GetNewOutputs() { return new_outputs_; } - - private: - const IRModule& module_; - Array new_outputs_; -}; - -Expr FlattenOutputTuple(const Array& exprs) { - Array fields; - for (const auto& it : exprs) { - ICHECK(it->checked_type_.defined()); - if (auto* tn = it->checked_type_.as()) { - // TODO(seanlatias): for now input argument cannot be a tuple - ICHECK(it->IsInstance()); - for (size_t i = 0; i < tn->fields.size(); i++) { - fields.push_back(TupleGetItem(it, i)); - } - } else { - fields.push_back(it); - } - } - return Tuple(fields); -} - -IRModule GetCalibrateModule(IRModule module) { - auto glob_funcs = module->functions; - // module is mutable, hence, we make a copy of it. - module.CopyOnWrite(); - for (const auto& pair : glob_funcs) { - if (auto opt = pair.second.as()) { - // we only collect the outputs for main function - if (pair.first->name_hint == "main") { - auto func = opt.value(); - Collector collector(module); - PostOrderRewrite(func->body, &collector); - auto new_outputs = collector.GetNewOutputs(); - Expr tuple = FlattenOutputTuple(new_outputs); - func = Function(func->params, tuple, tuple->checked_type_, func->type_params, func->attrs); - module->Update(pair.first, func); - } - } - } - // reset the attribute of functions for running graph executor - for (const auto& pair : glob_funcs) { - if (auto opt = pair.second.as()) { - auto func = opt.value(); - if (func->GetAttr(attr::kCompiler)) { - // we need to inline the functions in order to run grpah runtime - func = WithAttr(std::move(func), attr::kInline, tvm::Integer(1)); - // reset the compiler attribute to null for llvm execution - func = WithAttr(std::move(func), attr::kCompiler, NullValue()); - module->Update(pair.first, func); - } - } - } - return module; -} - -/*! - * \brief This function generates the output mapping between - * the calibration data and each function. The key is a - * GlobalVar that corresponds to each function and the value - * is an array of integers. The size of the array is always - * three. The first value is the offset the points to the start. - * The second value is the number of inputs. The third value - * is the number of outputs. - */ - -class OutputMapper : public ExprRewriter { - public: - OutputMapper(Map>* output_map, const IRModule& module, size_t* offset) - : output_map_(output_map), module_(module), offset_(offset) {} - - Expr Rewrite_(const CallNode* call, const Expr& post) final { - if (call->op->IsInstance()) { - auto var = Downcast(call->op); - ICHECK(module_->ContainGlobalVar(var->name_hint)) << "Function " << var << " is not defined"; - ICHECK_EQ(output_map_->count(var), 0) - << "Repeated function call " << var << " is not supported."; - auto func = Downcast(module_->Lookup(var)); - // we only handle functions with Compiler attribute set - if (func->GetAttr(attr::kCompiler)) { - Array info; - // the first value is the offset - info.push_back(Integer(*offset_)); - // the second value is the number of inputs - info.push_back(Integer(call->args.size())); - // the third value is the number of outputs - // we need to check if the output is a tuple - size_t out_size = 1; - if (auto* tn = func->body.as()) { - info.push_back(Integer(tn->fields.size())); - out_size = tn->fields.size(); - } else { - info.push_back(Integer(1)); - } - output_map_->Set(var, info); - // calculate the offset for the next function - *offset_ = *offset_ + call->args.size() + out_size; - } - } - return post; - } - - private: - Map>* output_map_; - const IRModule& module_; - size_t* offset_; -}; - -Map> GetCalibrateOutputMap(const IRModule& module) { - Map> output_map; - size_t offset = 0; - auto glob_funcs = module->functions; - for (const auto& pair : glob_funcs) { - if (const auto* func = pair.second.as()) { - if (pair.first->name_hint == "main") { - OutputMapper output_mapper(&output_map, module, &offset); - PostOrderRewrite(func->body, &output_mapper); - } - } - } - - return output_map; -} - -TVM_REGISTER_GLOBAL("relay.analysis.get_calibrate_module").set_body_typed([](IRModule mod) { - return GetCalibrateModule(mod); -}); - -TVM_REGISTER_GLOBAL("relay.analysis.get_calibrate_output_map") - .set_body_typed([](const IRModule& mod) { return GetCalibrateOutputMap(mod); }); - -} // namespace relay -} // namespace tvm diff --git a/src/relay/analysis/kind_check.cc b/src/relay/analysis/kind_check.cc deleted file mode 100644 index f7a5e7bf2d12..000000000000 --- a/src/relay/analysis/kind_check.cc +++ /dev/null @@ -1,195 +0,0 @@ -/* - * 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 kindchecker.cc - * - * \brief Check that types are well formed by applying "kinding rules". - * - * This pass ensures we do not do things that violate the design of the - * type system when writing down types. - * - * For example tensors are not allowed to contain functions in Relay. - * - * We check this by ensuring the `dtype` field of a Tensor always - * contains a data type such as `int`, `float`, `uint`. - */ -#include -#include -#include - -namespace tvm { -namespace relay { - -using namespace tvm::runtime; - -struct KindChecker : TypeFunctor { - const IRModule& mod; - Optional diag_ctx; - - explicit KindChecker(const IRModule& mod, Optional diag_ctx) - : mod(mod), diag_ctx(diag_ctx) {} - - void EmitFatal(Diagnostic diagnostic) { - if (this->diag_ctx) { - this->diag_ctx.value().EmitFatal(diagnostic); - } else { - LOG(FATAL) << diagnostic->message; - } - } - - void CheckKindMatches(const Type& t, const Type& outer, Kind expected, - const std::string& description) { - Kind k = this->VisitType(t); - if (k != expected) { - EmitFatal(Diagnostic::Error(t->span) - << "Incorrect kind for a " << description << ". Type " << t << " inside " << outer - << " is of kind " << k << " but was expected to be " << expected); - } - } - - Kind VisitType_(const IncompleteTypeNode* op) override { return op->kind; } - - Kind VisitType_(const TypeVarNode* op) override { return op->kind; } - - Kind VisitType_(const GlobalTypeVarNode* op) override { return op->kind; } - - Kind VisitType_(const TensorTypeNode* op) override { return Kind::kType; } - - Kind VisitType_(const TupleTypeNode* op) override { - // tuples should only contain normal types - for (const Type& t : op->fields) { - CheckKindMatches(t, GetRef(op), Kind::kType, "tuple member"); - } - return Kind::kType; - } - - Kind VisitType_(const FuncTypeNode* op) override { - // Func types should only take normal types for arguments - // and only return a normal type. They should also have - // well-formed constraints - FuncType ft = GetRef(op); - for (const Type& t : op->arg_types) { - CheckKindMatches(t, ft, Kind::kType, "function type parameter"); - } - - CheckKindMatches(ft->ret_type, ft, Kind::kType, "function return type"); - - for (const TypeConstraint& tc : op->type_constraints) { - CheckKindMatches(tc, ft, Kind::kConstraint, "function type constraint"); - } - - return Kind::kType; - } - - Kind VisitType_(const RelayRefTypeNode* op) override { - // ref types should only contain normal types - RelayRefType rt = GetRef(op); - CheckKindMatches(op->value, rt, Kind::kType, "ref contents"); - return Kind::kType; - } - - Kind VisitType_(const TypeRelationNode* op) override { - // arguments to type relation should be normal types - for (const Type& t : op->args) { - CheckKindMatches(t, GetRef(op), Kind::kType, "argument to type relation"); - } - return Kind::kConstraint; - } - - Kind VisitType_(const TypeCallNode* op) override { - // type call func should be a global type var, args should be type - TypeCall tc = GetRef(op); - const auto* gtv = op->func.as(); - if (gtv == nullptr) { - EmitFatal(Diagnostic::Error(op->span) - << "The callee in " << tc << " is not a global type var, but is " << op->func); - } - - CheckKindMatches(op->func, tc, Kind::kAdtHandle, "type call function"); - - for (const Type& t : op->args) { - CheckKindMatches(t, tc, Kind::kType, "type call argument"); - } - - // finally we need to check the module to check the number of type params - auto var = GetRef(gtv); - try { - auto data = mod->LookupTypeDef(var); - - if (data->type_vars.size() != op->args.size()) { - EmitFatal(Diagnostic::Error(op->span) - << "Expected " << data->type_vars.size() << "arguments for " << tc << "; got " - << op->args.size()); - } - } catch (const Error& err) { - // TODO(@jroesch): can probably relax to just emit - EmitFatal(Diagnostic::Error(op->span) - << "the type variable : `" << var->name_hint << "` is undefined"); - } - - return Kind::kType; - } - - Kind VisitType_(const TypeDataNode* op) override { - // Constructors can reference the header var, but no other GlobalTypeVars. - // In theory, a TypeData could be nested, so the header scope - // should be tracked recursively, but it is unclear that we need - // to support it. - TypeData td = GetRef(op); - CheckKindMatches(op->header, td, Kind::kAdtHandle, "type data header"); - - for (const auto& var : op->type_vars) { - CheckKindMatches(var, td, Kind::kType, "ADT type var"); - } - - for (const auto& con : op->constructors) { - if (!con->belong_to.same_as(op->header)) { - EmitFatal(Diagnostic::Error(op->span) << con << " has header " << con->belong_to << " but " - << op << " has header " << op->header); - } - - for (const Type& t : con->inputs) { - CheckKindMatches(t, td, Kind::kType, "ADT constructor input"); - } - } - return Kind::kTypeData; - } - - Kind Check(const Type& t) { return this->VisitType(t); } -}; - -Kind KindCheck(const Type& t, const IRModule& mod, Optional diag_ctx) { - KindChecker kc(mod, diag_ctx); - return kc.Check(t); -} - -TVM_REGISTER_GLOBAL("relay.analysis.check_kind").set_body([](TVMArgs args, TVMRetValue* ret) { - if (args.size() == 1) { - *ret = KindCheck(args[0], IRModule({}, {})); - } else if (args.size() == 2) { - *ret = KindCheck(args[0], args[1], Optional()); - } else { - *ret = KindCheck(args[0], args[1], args[2]); - } -}); - -} // namespace relay -} // namespace tvm diff --git a/src/relay/analysis/mac_count.cc b/src/relay/analysis/mac_count.cc deleted file mode 100644 index 29edf55812cc..000000000000 --- a/src/relay/analysis/mac_count.cc +++ /dev/null @@ -1,196 +0,0 @@ -/* - * 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 mac_count.cc - * \brief Pass to roughly count the number of MACs (Multiply-Accumulate) - * operations of a model. Only MACs in CONV and Dense ops are counted. - * This pass is valid after the type infer pass is called, - * otherwise the count is 0. - */ - -#include -#include -#include -#include -#include - -#include "../transforms/pattern_utils.h" - -namespace tvm { -namespace relay { - -namespace mac_count { - -inline int64_t GetCartesianProd(Array arr) { - int64_t ret = 1; - for (size_t i = 0; i < arr.size(); i++) { - const auto* intImm = arr[i].as(); - ret *= static_cast(intImm->value); - } - return ret; -} - -/* - * \brief Preparation function for MAC count. - * \param call_node The call node. - * \return The number of MACs. - */ -using FMacCount = runtime::TypedPackedFunc; - -//---------------------------------------------- -// Per operator defs for MAC count -//---------------------------------------------- - -int64_t ConvMacCount(const Call& call_node) { - if (!call_node->checked_type_.defined()) { - LOG(WARNING) << "The infer type pass should be called before the mac count pass"; - return 0; - } - Array args = call_node->args; - ICHECK_EQ(args.size(), 2) << "The number of input arguments of a CONV 2D node should be 2."; - const auto* conv_2d_attr = call_node->attrs.as(); - const auto* data_type = args[0]->checked_type().as(); - Array data_shape = data_type->shape; - std::string data_layout = conv_2d_attr->data_layout; - int32_t C_ind = Layout(data_layout).IndexOf(LayoutAxis::Get('C')); - int32_t c_ind = Layout(data_layout).IndexOf(LayoutAxis::Get('c')); - ICHECK_NE(C_ind, -1) << "There is no input channel dimension."; - int64_t input_channel = static_cast(data_shape[C_ind].as()->value); - if (c_ind != -1) input_channel *= static_cast(data_shape[c_ind].as()->value); - Array kernel_size = conv_2d_attr->kernel_size; - ICHECK_EQ(kernel_size.size(), 2) << "The dimension of the kernel in Conv 2D should be 2."; - const auto* expr = call_node->checked_type().as(); - Array output_tensor = expr->shape; - ICHECK(output_tensor.size() == 4 || output_tensor.size() == 5) - << "The dimension of the output tensor in Conv 2D should be 4 or 5."; - int64_t count = GetCartesianProd(output_tensor) * GetCartesianProd(kernel_size); - ICHECK_EQ(input_channel % conv_2d_attr->groups, 0) - << "The number of input channels is not divisble by groups."; - count *= input_channel / conv_2d_attr->groups; - return count; -} - -int64_t Conv2dTransposeMacCount(const Call& call_node) { - if (!call_node->checked_type_.defined()) { - LOG(WARNING) << "The infer type pass should be called before the mac count pass"; - return 0; - } - Array args = call_node->args; - ICHECK_EQ(args.size(), 2) - << "The number of input arguments of a CONV 2D Transpose node should be 2."; - const auto* conv_2d_transpose_attr = call_node->attrs.as(); - const auto* data_type = args[0]->checked_type().as(); - Array data_shape = data_type->shape; - std::string data_layout = conv_2d_transpose_attr->data_layout; - int32_t C_ind = Layout(data_layout).IndexOf(LayoutAxis::Get('C')); - int32_t c_ind = Layout(data_layout).IndexOf(LayoutAxis::Get('c')); - ICHECK_NE(C_ind, -1) << "There is no input channel dimension."; - int64_t input_channel = static_cast(data_shape[C_ind].as()->value); - if (c_ind != -1) input_channel *= static_cast(data_shape[c_ind].as()->value); - Array kernel_size = conv_2d_transpose_attr->kernel_size; - ICHECK_EQ(kernel_size.size(), 2) - << "The dimension of the kernel in Conv 2D Transpose should be 2."; - const auto* expr = call_node->checked_type().as(); - Array output_tensor = expr->shape; - ICHECK(output_tensor.size() == 4 || output_tensor.size() == 5) - << "The dimension of the output tensor in Conv 2D Transpose should be 4 or 5."; - int64_t count = GetCartesianProd(output_tensor) * GetCartesianProd(kernel_size); - ICHECK_EQ(input_channel % conv_2d_transpose_attr->groups, 0) - << "The number of input channels is not divisble by groups."; - count *= input_channel / conv_2d_transpose_attr->groups; - return count; -} - -int64_t DenseMacCount(const Call& call_node) { - if (!call_node->checked_type_.defined()) { - LOG(WARNING) << "The infer type pass should be called before the mac count pass"; - return 0; - } - Array args = call_node->args; - ICHECK_EQ(args.size(), 2) << "The number of input arguments of a Dense node should be 2."; - const auto* data_type = args[0]->checked_type().as(); - const auto* weight_type = args[1]->checked_type().as(); - Array data_shape = data_type->shape; - Array weight_shape = weight_type->shape; - ICHECK(data_shape.size() == 2 && weight_shape.size() == 2) - << "The dimension of an input tensor to Dense node should be 2."; - int64_t d1 = static_cast(data_shape[0].as()->value); - int64_t d2 = static_cast(data_shape[1].as()->value); - int64_t d3 = static_cast(weight_shape[0].as()->value); - int64_t d4 = static_cast(weight_shape[1].as()->value); - ICHECK_EQ(d2, d4) << "The dimensions of input arguments do not match."; - int64_t count = d1 * d2 * d3; - return count; -} - -int64_t BatchMatmulMacCount(const Call& call_node) { - if (!call_node->checked_type_.defined()) { - LOG(WARNING) << "The infer type pass should be called before the mac count pass"; - return 0; - } - Array args = call_node->args; - ICHECK_EQ(args.size(), 2); - Array x_shape = args[0]->checked_type().as()->shape; - Array y_shape = args[1]->checked_type().as()->shape; - int64_t batch = x_shape[0].as()->value; - int64_t m = x_shape[1].as()->value; - int64_t k = x_shape[2].as()->value; - int64_t n = y_shape[1].as()->value; - return batch * m * k * n; -} - -RELAY_REGISTER_OP("nn.conv2d").set_attr("FMacCount", ConvMacCount); - -RELAY_REGISTER_OP("nn.conv2d_transpose").set_attr("FMacCount", Conv2dTransposeMacCount); - -RELAY_REGISTER_OP("nn.dense").set_attr("FMacCount", DenseMacCount); - -RELAY_REGISTER_OP("nn.batch_matmul").set_attr("FMacCount", BatchMatmulMacCount); - -class MacCounter : private ExprVisitor { - public: - MacCounter() { count_ = 0; } - static int64_t GetTotalMacNumber(const Expr& expr) { - LOG(INFO) << "This pass only counts MACs in direct conv2d, " - << "conv2d_transpose, dense, and batch_matmul ops"; - MacCounter counter; - counter(expr); - return counter.count_; - } - - private: - void VisitExpr_(const CallNode* call_node) final { - static const auto& fprep = Op::GetAttrMap("FMacCount"); - auto f = fprep.get(call_node->op, nullptr); - if (f != nullptr) count_ += f(GetRef(call_node)); - ExprVisitor::VisitExpr_(call_node); - } - - int64_t count_; -}; - -int64_t GetTotalMacNumber(const Expr& expr) { return MacCounter::GetTotalMacNumber(expr); } - -TVM_REGISTER_GLOBAL("relay.analysis.GetTotalMacNumber").set_body_typed(GetTotalMacNumber); - -} // namespace mac_count -} // namespace relay -} // namespace tvm diff --git a/src/relay/analysis/match_exhaustion.cc b/src/relay/analysis/match_exhaustion.cc deleted file mode 100644 index 9f92ebaa8b47..000000000000 --- a/src/relay/analysis/match_exhaustion.cc +++ /dev/null @@ -1,319 +0,0 @@ -/* - * 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 match_exhaustion.cc - * \brief Checking Relay match expression exhaustiveness. - * - * This file implements a function that checks whether a match - * expression is exhaustive, that is, whether a given match clause - * matches every possible case. This is important for ensuring - * code correctness, since hitting an unmatched case results in a - * dynamic error unless exhaustiveness is checked in advance. - */ -#include -#include -#include -#include - -#include - -namespace tvm { -namespace relay { - -/*! \brief Possible pattern match results */ -enum MatchResult : int { - kMatch = 0, // pattern matches - kClash = 1, // pattern conflicts - kUnspecified = 2, // ambiguous: candidate needs more constructors specified -}; - -class CandidateChecker : public PatternFunctor { - public: - explicit CandidateChecker() {} - - MatchResult Check(const Pattern& pat, const Pattern& candidate) { - return this->VisitPattern(pat, candidate); - } - - // for a constructor pattern, we must ensure that the candidate is - // a ConstructorPattern, that it has the same constructor, and - // that its fields match the subpatterns. - MatchResult VisitPattern_(const PatternConstructorNode* op, const Pattern& cand) override { - auto* ctor_cand = cand.as(); - // attempting to match non-constructor to constructor pattern: need to specify - if (ctor_cand == nullptr) { - return MatchResult::kUnspecified; - } - - // check that constructors match - if (!op->constructor.same_as(ctor_cand->constructor)) { - return MatchResult::kClash; - } - - // now check that subpatterns match - ICHECK_EQ(op->patterns.size(), ctor_cand->patterns.size()); - bool unspecified = false; - for (size_t i = 0; i < op->patterns.size(); i++) { - MatchResult submatch = this->Check(op->patterns[i], ctor_cand->patterns[i]); - // if we have a clash anywhere, then we can return clash - if (submatch == MatchResult::kClash) { - return MatchResult::kClash; - } - if (submatch == MatchResult::kUnspecified) { - unspecified = true; - } - } - // only return unspecified if we have ruled out a clash - if (unspecified) { - return MatchResult::kUnspecified; - } - return MatchResult::kMatch; - } - - MatchResult VisitPattern_(const PatternTupleNode* op, const Pattern& cand) override { - auto* tuple_cand = cand.as(); - // attempting to match non-tuple to constructor pattern: need to specify - if (tuple_cand == nullptr) { - return MatchResult::kUnspecified; - } - - // now check that subpatterns match - ICHECK_EQ(op->patterns.size(), tuple_cand->patterns.size()); - bool unspecified = false; - for (size_t i = 0; i < op->patterns.size(); i++) { - MatchResult submatch = this->Check(op->patterns[i], tuple_cand->patterns[i]); - // if we have a clash anywhere, then we can return clash - if (submatch == MatchResult::kClash) { - return MatchResult::kClash; - } - if (submatch == MatchResult::kUnspecified) { - unspecified = true; - } - } - // only return unspecified if we have ruled out a clash - if (unspecified) { - return MatchResult::kUnspecified; - } - return MatchResult::kMatch; - } - - // wildcard and var patterns always match - MatchResult VisitPattern_(const PatternWildcardNode*, const Pattern&) override { - return MatchResult::kMatch; - } - - MatchResult VisitPattern_(const PatternVarNode*, const Pattern&) override { - return MatchResult::kMatch; - } -}; - -// Returns list of arrays corresponding to Cartesian product of input list. -// Note: CartesianProduct({}) = {{}} -Array> CartesianProduct(Array> fields) { - // the only combination of 0 fields is 0 fields - if (fields.size() == 0) { - return {{}}; - } - - Array field_vals = fields[fields.size() - 1]; - Array> ret; - - // base case: this is the last field left - if (fields.size() == 1) { - for (auto val : field_vals) { - ret.push_back(Array{val}); - } - return ret; - } - - // if we have more fields left, get the sub-candidates by getting - // their cartesian product and appending the elements here onto those - Array> remaining_fields; - for (size_t i = 0; i < fields.size() - 1; i++) { - remaining_fields.push_back(fields[i]); - } - Array> candidates = CartesianProduct(remaining_fields); - for (auto val : field_vals) { - for (auto candidate : candidates) { - candidate.push_back(val); - ret.push_back(candidate); - } - } - return ret; -} - -Array ExpandWildcardsConstructor(const PatternConstructor& clause_ctor, - const Pattern& cand, const IRModule& mod); - -Array ExpandWildcardsTuple(const PatternTuple& clause_tuple, const Pattern& cand, - const IRModule& mod); - -// Expands all wildcards in the candidate pattern once -// Returns a list of all possible expansions. -Array ExpandWildcards(const Pattern& clause_pat, const Pattern& cand, - const IRModule& mod) { - if (auto clause_ctor = clause_pat.as()) { - return ExpandWildcardsConstructor(clause_ctor.value(), cand, mod); - } else if (auto clause_tup = clause_pat.as()) { - return ExpandWildcardsTuple(clause_tup.value(), cand, mod); - } else { - return {cand}; - } -} - -// Expands all wildcards in the candidate pattern once. -// Use the pattern to decide which constructors to insert. -// Returns a list of all possible expansions. -Array ExpandWildcardsConstructor(const PatternConstructor& clause_ctor, - const Pattern& cand, const IRModule& mod) { - auto gtv = Downcast(clause_ctor->constructor->belong_to); - - // for a wildcard node, create constructor nodes with wildcards for all args. - if (cand.as()) { - TypeData td = mod->LookupTypeDef(gtv); - // for each constructor add a candidate. - Array ret; - for (auto constructor : td->constructors) { - Array args; - for (auto inp : constructor->inputs) { - args.push_back(PatternWildcard()); - } - ret.push_back(PatternConstructor(constructor, args)); - } - return ret; - } - - auto ctor_cand = Downcast(cand); - - // expand all fields' wildcards - Array> values_by_field; - for (size_t i = 0; i < ctor_cand->constructor->inputs.size(); i++) { - values_by_field.push_back( - ExpandWildcards(clause_ctor->patterns[i], ctor_cand->patterns[i], mod)); - } - - // generate new candidates using a cartesian product. - auto all_subfields = CartesianProduct(values_by_field); - Array ret; - for (auto subfields : all_subfields) { - ret.push_back(PatternConstructor(ctor_cand->constructor, subfields)); - } - return ret; -} - -// Expands all wildcards in the candidate pattern once. -// Returns a list of all possible expansions. -Array ExpandWildcardsTuple(const PatternTuple& clause_tuple, const Pattern& cand, - const IRModule& mod) { - // for a wildcard node, create tuple with wildcards for all args. - if (cand.as()) { - Array args; - for (auto inp : clause_tuple->patterns) { - args.push_back(PatternWildcard()); - } - return {PatternTuple(args)}; - } - - auto tuple_cand = Downcast(cand); - - // expand all members' patterns - Array> values_by_field; - for (size_t i = 0; i < tuple_cand->patterns.size(); i++) { - values_by_field.push_back( - ExpandWildcards(clause_tuple->patterns[i], tuple_cand->patterns[i], mod)); - } - - // generate new candidates using a cartesian product - auto all_subfields = CartesianProduct(values_by_field); - Array ret; - for (auto subfields : all_subfields) { - ret.push_back(PatternTuple(subfields)); - } - return ret; -} - -/*! - * \brief Finds cases that the match expression does not catch, if any. - * \return Returns a list of cases that are not handled by the match - * expression. - */ -Array UnmatchedCases(const Match& match, const IRModule& mod) { - /* algorithm: - * candidates = { Wildcard } - * while candidates not empty { - * cand = candidates.pop() - * for clause in clauses { - * if clause fails: next clause - * if clause matches candidate: next candidate - * if candidate is not specific enough: - * candidates += expand_possible_wildcards(cand) - * next candidate - * } - * failed_candidates += { cand } - * } - * return failed_candidates - */ - std::stack candidates; - candidates.push(PatternWildcard()); - CandidateChecker checker; - - Array failures; - - while (!candidates.empty()) { - Pattern cand = candidates.top(); - candidates.pop(); - - bool failure = true; - for (auto clause : match->clauses) { - // if the check fails, we move on to the next - MatchResult check = checker.Check(clause->lhs, cand); - if (check == MatchResult::kClash) { - continue; - } - - // either success or we need to generate more candidates; - // either way, we're done with this candidate - failure = false; - if (check == MatchResult::kUnspecified) { - auto new_candidates = ExpandWildcards(clause->lhs, cand, mod); - for (auto candidate : new_candidates) { - candidates.push(candidate); - } - } - break; - } - - if (failure) { - failures.push_back(cand); - } - } - - return failures; -} - -// expose for testing only -TVM_REGISTER_GLOBAL("relay.analysis.unmatched_cases") - .set_body_typed([](const Match& match, const Optional& mod_ref) { - IRModule call_mod = mod_ref.defined() ? mod_ref.value() : IRModule({}, {}); - return UnmatchedCases(match, call_mod); - }); - -} // namespace relay -} // namespace tvm diff --git a/src/relay/analysis/type_solver.cc b/src/relay/analysis/type_solver.cc deleted file mode 100644 index c4fab210acb8..000000000000 --- a/src/relay/analysis/type_solver.cc +++ /dev/null @@ -1,690 +0,0 @@ -/* - * 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 type_solver.cc - * \brief Type solver implementations. - */ -#include "type_solver.h" - -#include -#include -#include -#include - -#include -#include -#include -#include - -namespace tvm { -namespace relay { - -class TypeSolver::Reporter : public TypeReporterNode { - public: - explicit Reporter(TypeSolver* solver) : solver_(solver) {} - - void Assign(const Type& dst, const Type& src) final { solver_->Unify(dst, src, span); } - - bool Assert(const IndexExpr& cond) final { - if (const int64_t* pdiff = tir::as_const_int(cond)) { - return pdiff[0]; - } - return true; - } - - bool AssertEQ(const IndexExpr& lhs, const IndexExpr& rhs) final { - // early warning constant case. - IndexExpr diff = lhs - rhs; - if (const int64_t* pdiff = tir::as_const_int(diff)) { - return pdiff[0] == 0; - } - return true; - } - - TVM_DLL void SetSpan(const Span& span) final { this->span = span; } - - TVM_DLL Span GetSpan() final { return this->span; } - - TVM_DLL DiagnosticContext GetDiagCtx() final { return this->solver_->diag_ctx_; } - - // TVM_DLL void Emit(Diagnostic diagnostic) final { - // return this->solver_-> - // } - - TVM_DLL IRModule GetModule() final { return this->solver_->module_; } - - private: - /*! \brief The span to report unification errors at. */ - mutable Span span; - - TypeSolver* solver_; -}; - -class TypeSolver::AnyChecker : public tir::ExprVisitor { - public: - void VisitExpr_(const AnyNode* op) final { found_ = true; } - - bool Check(const PrimExpr& expr) { - tir::ExprVisitor::VisitExpr(expr); - return found_; - } - - private: - bool found_{false}; -}; - -class TypeSolver::OccursChecker : public TypeVisitor { - public: - explicit OccursChecker(TypeSolver* solver, TypeNode* var) - : solver_(solver), var_(var), found_(false) {} - - bool Check(const Type& t) { - VisitType(t); - return found_; - } - - void VisitType_(const IncompleteTypeNode* op) override { - IncompleteType t = GetRef(op); - TypeNode* node = solver_->GetTypeNode(t); - found_ = found_ || (var_->FindRoot() == node->FindRoot()); - } - - private: - TypeSolver* solver_; - TypeNode* var_; - bool found_; -}; - -class TypeSolver::Unifier : public TypeFunctor { - public: - explicit Unifier(TypeSolver* solver, const Span& span) : solver_(solver), span(span) {} - - Type Unify(const Type& lhs_type, const Type& rhs_type, bool assign_lhs = true, - bool assign_rhs = true) { - // Known limitation - // - handle shape pattern matching - TypeNode* lhs = solver_->GetTypeNode(lhs_type); - TypeNode* rhs = solver_->GetTypeNode(rhs_type); - - // do occur check so we don't create self-referencing structure - if (lhs->FindRoot() == rhs->FindRoot()) { - return lhs->resolved_type; - } - - if (lhs->resolved_type.as()) { - ICHECK(!OccursCheck(lhs, rhs->resolved_type)) - << "Incomplete type " << lhs->resolved_type << " occurs in " << rhs->resolved_type - << ", cannot unify"; - - solver_->MergeFromTo(lhs, rhs); - return rhs->resolved_type; - } else if (rhs->resolved_type.as()) { - ICHECK(!OccursCheck(rhs, lhs->resolved_type)) - << "Incomplete type " << rhs->resolved_type << " occurs in " << lhs->resolved_type - << ", cannot unify"; - solver_->MergeFromTo(rhs, lhs); - return lhs->resolved_type; - } else { - Type resolved = this->VisitType(rhs->resolved_type, lhs->resolved_type); - - if (!resolved.defined()) { - solver_->Emit(Diagnostic::Error(this->span) - << "The Relay type checker is unable to show the following types match.\n" - << "In particular " - << "`" << PrettyPrint(lhs->resolved_type) << "` does not match `" - << PrettyPrint(rhs->resolved_type) << "`"); - return lhs->resolved_type; - } else { - TypeNode* top = solver_->GetTypeNode(resolved); - if (assign_lhs) solver_->MergeFromTo(lhs, top); - if (assign_rhs) solver_->MergeFromTo(rhs, top); - return resolved; - } - } - } - - bool HasAny(const PrimExpr& expr) { - AnyChecker ac; - return ac.Check(expr); - } - - // Checks whether lhs (taken to be a type var) occurs in t, meaning - // there is a recursive equality constraint, which should be rejected. - // N.b.: A tautology like ?a = ?a is okay and should be checked for - // *before* calling this method - // - // See: https://en.wikipedia.org/wiki/Occurs_check - bool OccursCheck(TypeNode* lhs, const Type& t) { - OccursChecker rc(solver_, lhs); - return rc.Check(t); - } - - // default: unify only if structural-equal - Type VisitTypeDefault_(const Object* op, const Type& tn) final { - ObjectRef nr = GetRef(op); - Type t1 = Downcast(nr); - if (!tvm::StructuralEqual()(t1, tn)) { - return Type(nullptr); - } - return t1; - } - - IndexExpr GetShape(const IndexExpr& e) { - IndexExpr ex = e; - while (true) { - auto it = solver_->shape_uf_.find(ex); - if (it == solver_->shape_uf_.end()) { - return ex; - } else { - ex = (*it).second; - } - } - } - - IndexExpr UnifyDim(const IndexExpr& lhs, const IndexExpr& rhs) { - auto ulhs = GetShape(lhs); - auto urhs = GetShape(rhs); - - if (ulhs.same_as(urhs)) { - return ulhs; - } - if (HasAny(ulhs) || HasAny(urhs)) { - return Any(); - } - - auto left_index0 = ulhs.as(); - auto right_index0 = urhs.as(); - if (left_index0 && right_index0) { - solver_->shape_uf_.Set(ulhs, urhs); - return urhs; - } - - auto left_index1 = ulhs.as(); - auto right_index1 = urhs.as(); - if (left_index1 && right_index1) { - solver_->shape_uf_.Set(urhs, ulhs); - return ulhs; - } - - auto left_index2 = ulhs.as(); - auto right_index2 = urhs.as(); - if (left_index2 && right_index2 && left_index2->value == right_index2->value) { - return ulhs; - } - - return tvm::PrimExpr(); - } - - Type VisitType_(const TensorTypeNode* op, const Type& tn) final { - const auto* tt_node = tn.as(); - if (!tt_node) { - return Type(nullptr); - } - - auto tt1 = GetRef(op); - auto tt2 = GetRef(tt_node); - - if (tvm::StructuralEqual()(tt1, tt2)) { - return std::move(tt1); - } - - if (tt1->dtype != tt2->dtype) { - return Type(nullptr); - } - - tvm::Array shape; - if (tt1->shape.size() != tt2->shape.size()) { - this->solver_->Emit(Diagnostic::Error(this->span) - << "tensor type `" << PrettyPrint(tt1) << "` has " << tt1->shape.size() - << " dimensions, while `" << PrettyPrint(tt2) << "` has " - << tt2->shape.size() << " dimensions"); - return Type(nullptr); - } - - std::vector> mismatches; - - ICHECK_EQ(tt1->shape.size(), tt2->shape.size()); - for (size_t i = 0; i < tt1->shape.size(); i++) { - auto dim = UnifyDim(tt1->shape[i], tt2->shape[i]); - if (!dim.defined()) { - // NB: We push an arbitrary dimension here so we can continue error propagation. - shape.push_back(tt1->shape[i]); - tvm::PrimExpr shape1 = tt1->shape[i]; - tvm::PrimExpr shape2 = tt2->shape[i]; - std::tuple tuple = std::make_tuple(i, shape1, shape2); - mismatches.push_back(tuple); - } else { - shape.push_back(dim); - } - } - - if (mismatches.size() != 0) { - auto err = Diagnostic::Error(this->span); - err << "The Relay type checker is unable to show the following types match:\n" - << " " << PrettyPrint(tt1) << "\n" - << " " << PrettyPrint(tt2) << "\n"; - err << "In particular:\n"; - for (auto mismatch : mismatches) { - err << " dimension " << std::get<0>(mismatch) << " conflicts: " << std::get<1>(mismatch) - << " does not match " << std::get<2>(mismatch) << "."; - } - this->solver_->Emit(err); - return Type(nullptr); - } - - return TensorType(shape, tt1->dtype); - } - - Type VisitType_(const TupleTypeNode* op, const Type& tn) final { - const auto* ttn = tn.as(); - if (!ttn || op->fields.size() != ttn->fields.size()) { - return Type(nullptr); - } - - TupleType tt1 = GetRef(op); - TupleType tt2 = GetRef(ttn); - - std::vector new_fields; - for (size_t i = 0; i < tt1->fields.size(); i++) { - Type field = Unify(tt1->fields[i], tt2->fields[i]); - new_fields.push_back(field); - } - return TupleType(new_fields); - } - - Type VisitType_(const FuncTypeNode* op, const Type& tn) final { - const auto* ftn = tn.as(); - if (!ftn || op->arg_types.size() != ftn->arg_types.size() || - op->type_constraints.size() != ftn->type_constraints.size()) { - return Type(nullptr); - } - - // without loss of generality, suppose op->type_params.size() >= ftn->type_params.size(). - if (op->type_params.size() < ftn->type_params.size()) { - return VisitType_(ftn, GetRef(op)); - } - - // remap type vars so they match - Map subst_map; - tvm::Array ft_type_params; - for (size_t i = 0; i < ftn->type_params.size(); ++i) { - subst_map.Set(op->type_params[i], ftn->type_params[i]); - ft_type_params.push_back(op->type_params[i]); - } - - for (size_t i = ftn->type_params.size(); i < op->type_params.size(); ++i) { - subst_map.Set(op->type_params[i], IncompleteType(kType)); - } - - FuncType ft = FuncType(op->arg_types, op->ret_type, ft_type_params, op->type_constraints); - auto ft1 = Downcast(Bind(ft, subst_map)); - auto ft2 = GetRef(ftn); - - Type ret_type = Unify(ft1->ret_type, ft2->ret_type); - - std::vector arg_types; - for (size_t i = 0; i < ft2->arg_types.size(); ++i) { - Type arg_type = Unify(ft1->arg_types[i], ft2->arg_types[i]); - arg_types.push_back(arg_type); - } - - std::vector type_constraints; - for (size_t i = 0; i < ft1->type_constraints.size(); ++i) { - Type unified_constraint = Unify(ft1->type_constraints[i], ft2->type_constraints[i]); - const auto* tcn = unified_constraint.as(); - ICHECK(tcn) << "Two type constraints unified into a non-constraint?" - << ft1->type_constraints[i] << " and " << ft2->type_constraints[i]; - type_constraints.push_back(GetRef(tcn)); - } - - return FuncType(arg_types, ret_type, ft2->type_params, type_constraints); - } - - Type VisitType_(const RelayRefTypeNode* op, const Type& tn) final { - const auto* rtn = tn.as(); - if (!rtn) { - return Type(nullptr); - } - return RelayRefType(Unify(op->value, rtn->value)); - } - - Type VisitType_(const TypeCallNode* op, const Type& tn) override { - const auto* tcn = tn.as(); - if (!tcn || tcn->args.size() != op->args.size()) { - return Type(); - } - - Type func = Unify(op->func, tcn->func); - tvm::Array args; - for (size_t i = 0; i < op->args.size(); i++) { - args.push_back(Unify(op->args[i], tcn->args[i])); - } - return TypeCall(func, args); - } - - private: - TypeSolver* solver_; - Span span; -}; - -class TypeSolver::Resolver : public TypeMutator { - public: - explicit Resolver(TypeSolver* solver) : solver_(solver) {} - - Type Resolve(const Type& t) { - if (!t.defined()) { - return t; - } - return VisitType(t); - } - - Type VisitType_(const IncompleteTypeNode* op) override { - auto* node = solver_->GetTypeNode(GetRef(op)); - return node->resolved_type; - } - - private: - TypeSolver* solver_; -}; - -// It ends up being more compact to simply have TypeFunctor { - public: - explicit Propagator(TypeSolver* solver, const std::unordered_set* rels) - : solver_(solver), rels_(rels) {} - - // adds the relation node to t and all child types of t - void Propagate(const Type& t) { VisitType(t); } - - void UpdateRelSet(const Type& t) { - TypeNode* tnode = solver_->GetTypeNode(t); - for (auto* rel : *rels_) { - tnode->rel_set.insert(rel); - } - } - - void VisitTypeDefault_(const Object* op) override { - ObjectRef nr = GetRef(op); - Type t = Downcast(nr); - UpdateRelSet(t); - } - - void VisitType_(const TupleTypeNode* op) override { - TupleType tt = GetRef(op); - UpdateRelSet(tt); - - for (const Type& t : tt->fields) { - Propagate(t); - } - } - - void VisitType_(const FuncTypeNode* op) override { - FuncType ft = GetRef(op); - UpdateRelSet(ft); - - Propagate(ft->ret_type); - for (auto arg_type : ft->arg_types) { - Propagate(arg_type); - } - - for (auto type_param : ft->type_params) { - Propagate(type_param); - } - - for (auto type_cs : ft->type_constraints) { - Propagate(type_cs); - } - } - - void VisitType_(const TypeCallNode* op) override { - TypeCall tc = GetRef(op); - UpdateRelSet(tc); - - Propagate(tc->func); - for (auto arg : tc->args) { - Propagate(arg); - } - } - - private: - TypeSolver* solver_; - const std::unordered_set* rels_; -}; - -// similarly, we use TypeFunctor so we can use -// the default visitor case to avoid more overrides -class TypeSolver::Merger : public TypeFunctor { - public: - explicit Merger(TypeSolver* solver) : solver_(solver) {} - - // Merges src node to dst, ensures *all* type relations of all - // child nodes of src are transferred to dst. - void Merge(TypeNode* src, TypeNode* dst) { - if (src == dst) return; - dst_ = dst; - VisitType(src->resolved_type); - // set parent at the end so later calls to GetTypeNode go back to src - src->parent = dst; - - // now propagate relations to child nodes, since change to - // a child node should update parent too - Propagator prop(solver_, &dst->rel_set); - prop.Propagate(dst->resolved_type); - } - - // Transfers any relations linked to t to the stored dst. - // Any unresolved relations are added back to the queue, since - // there is now new information - void TransferLinks(const Type& t) { - TypeNode* src = solver_->GetTypeNode(t); - if (src == dst_) return; - for (auto* rel : src->rel_set) { - // if the relation is not yet resolved, add to queue - if (!rel->resolved) { - solver_->AddToQueue(rel); - dst_->rel_set.insert(rel); - } - } - } - - void VisitTypeDefault_(const Object* op) override { - ObjectRef nr = GetRef(op); - Type t = Downcast(nr); - TransferLinks(t); - } - - void VisitType_(const TupleTypeNode* ttn) override { - auto tup = GetRef(ttn); - TransferLinks(tup); - - for (auto field : tup->fields) { - VisitType(field); - } - } - - void VisitType_(const FuncTypeNode* ftn) override { - auto func = GetRef(ftn); - TransferLinks(func); - - VisitType(func->ret_type); - for (auto arg : func->arg_types) { - VisitType(arg); - } - for (auto param : func->type_params) { - VisitType(param); - } - for (auto constraint : func->type_constraints) { - VisitType(constraint); - } - } - - private: - TypeSolver* solver_; - TypeNode* dst_; -}; - -// constructor -TypeSolver::TypeSolver(const GlobalVar& current_func, DiagnosticContext diag_ctx) - : reporter_(make_object(this)), - current_func_(current_func), - diag_ctx_(diag_ctx), - module_(diag_ctx->module) { - ICHECK(module_.defined()); -} - -// destructor -TypeSolver::~TypeSolver() { - // call destructor of all non-POD arena object - for (TypeNode* ptr : type_nodes_) { - ptr->~TypeNode(); - } - for (RelationNode* ptr : rel_nodes_) { - ptr->~RelationNode(); - } -} - -// merge src type node to dst -void TypeSolver::MergeFromTo(TypeNode* src, TypeNode* dst) { - Merger merger(this); - merger.Merge(src, dst); -} - -// Add equality constraint -Type TypeSolver::Unify(const Type& dst, const Type& src, const Span& span, bool assign_lhs, - bool assign_rhs) { - Unifier unifier(this, span); - return unifier.Unify(dst, src, assign_lhs, assign_rhs); -} - -// Add type constraint to the solver. -void TypeSolver::AddConstraint(const TypeConstraint& constraint, const Span& span) { - if (const auto* op = constraint.as()) { - // create a new relation node. - RelationNode* rnode = arena_.make(); - rnode->span = span; - rnode->rel = GetRef(op); - rel_nodes_.push_back(rnode); - // populate the type information. - for (size_t i = 0; i < op->args.size(); ++i) { - // insert link to the type list - LinkNode* tlink = arena_.make>(); - TypeNode* tnode = GetTypeNode(op->args[i]); - tlink->value = tnode; - rnode->type_list.Push(tlink); - // insert type->relation node - std::unordered_set singleton{rnode}; - Propagator prop(this, &singleton); - prop.Propagate(tnode->resolved_type); - } - // add the relation to the working queue. - this->AddToQueue(rnode); - } else { - LOG(FATAL) << "Do not know how to handle constraint type" << constraint->GetTypeKey(); - } -} - -// Resolve a type in the solver context. -Type TypeSolver::Resolve(const Type& type) { - Resolver resolver(this); - auto it = tmap_.find(type); - Type t = (it != tmap_.end()) ? it->second->FindRoot()->resolved_type : type; - return resolver.Resolve(t); -} - -bool TypeSolver::Solve() { - while (!update_queue_.empty()) { - RelationNode* rnode = update_queue_.front(); - const auto& rel = rnode->rel; - update_queue_.pop(); - ICHECK(!rnode->resolved); - // update the relation with given evidence. - Array args; - for (auto* tlink = rnode->type_list.head; tlink != nullptr; tlink = tlink->next) { - args.push_back(Resolve(tlink->value->FindRoot()->resolved_type)); - ICHECK_LE(args.size(), rel->args.size()); - } - - // We need to set this in order to understand where unification - // errors generated by the error reporting are coming from. - reporter_->SetSpan(rnode->span); - - try { - // Call the Type Relation's function. - bool resolved = rel->func(args, rel->num_inputs, rel->attrs, reporter_); - - if (resolved) { - ++num_resolved_rels_; - } - - rnode->resolved = resolved; - } catch (const CompileError& err) { - this->Emit(Diagnostic::Error(rnode->span) << err.what()); - rnode->resolved = false; - } - - // Mark inqueue as false after the function call - // so that rnode itself won't get enqueued again. - rnode->inqueue = false; - } - - // This criterion is not necessarily right for all the possible cases - // TODO(tqchen): We should also count the number of in-complete types. - return num_resolved_rels_ == rel_nodes_.size(); -} - -// Expose type solver only for debugging purposes. -TVM_REGISTER_GLOBAL("relay.analysis._test_type_solver") - .set_body([](runtime::TVMArgs args, runtime::TVMRetValue* ret) { - using runtime::PackedFunc; - using runtime::TypedPackedFunc; - auto module = IRModule({}, {}); - DiagnosticContext diag_ctx = DiagnosticContext::Default(module); - auto dummy_fn_name = GlobalVar("test"); - module->Add(dummy_fn_name, Function({}, Tuple(tvm::Array({})), Type(), {})); - auto solver = std::make_shared(dummy_fn_name, diag_ctx); - - auto mod = [module, solver, diag_ctx](std::string name) -> PackedFunc { - if (name == "Solve") { - return TypedPackedFunc([solver]() { return solver->Solve(); }); - } else if (name == "Unify") { - return TypedPackedFunc([module, solver, diag_ctx](Type lhs, Type rhs) { - auto res = solver->Unify(lhs, rhs, Span()); - DiagnosticContext ctx = diag_ctx; - ctx.Render(); - return res; - }); - } else if (name == "Resolve") { - return TypedPackedFunc([solver](Type t) { return solver->Resolve(t); }); - } else if (name == "AddConstraint") { - return TypedPackedFunc([solver](TypeConstraint c) { - Expr e = Var("dummy_var", IncompleteType(Kind::kType), Span(SourceName(), 0, 0, 0, 0)); - return solver->AddConstraint(c, e->span); - }); - } else { - return PackedFunc(); - } - }; - *ret = runtime::TypedPackedFunc(mod); - }); - -} // namespace relay -} // namespace tvm diff --git a/src/relay/analysis/type_solver.h b/src/relay/analysis/type_solver.h deleted file mode 100644 index 5d32afab6442..000000000000 --- a/src/relay/analysis/type_solver.h +++ /dev/null @@ -1,224 +0,0 @@ -/* - * 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 type_solver.h - * \brief Solver logic for type inference. - */ -#ifndef TVM_RELAY_ANALYSIS_TYPE_SOLVER_H_ -#define TVM_RELAY_ANALYSIS_TYPE_SOLVER_H_ - -#include -#include -#include -#include - -#include -#include -#include -#include - -#include "../../support/arena.h" - -namespace tvm { -namespace relay { - -using support::LinkedList; -using support::LinkNode; - -/*! - * \brief Interface of type solver used in type inference. - * - * TypeSolver works on a list of constraints among incomplete types. - * The user will populate the constraints by AddConstraint and Assign. - * Then we can call Solve to trying to resolve the unknown. - * - * This can be viewed as "type program(computational graph)" of types, where - * the type constraint are operators of the graph and the incomplete - * types are intermediate value of the graph. - * If all the input types are concretely known, we should be able to - * just run a forward pass on the "type program" to get all the types. - * - * The list of constraints representation means we are storing it as a bipartite - * graph instead of a DAG. This is because some constraints might go both direction. - * TypeSolver could take advantage of bidirectional constraints to deduce input - * value given output ones. Never-the-less, we should keep in mind that - * there is a "forward direction" that the TypeSolver should take advantage of. - */ -class TypeSolver { - public: - TypeSolver(const GlobalVar& current_func, DiagnosticContext diag_ctx); - ~TypeSolver(); - /*! - * \brief Add a type constraint to the solver. - * \param constraint The constraint to be added. - * \param location The location at which the constraint was incurred. - */ - void AddConstraint(const TypeConstraint& constraint, const Span& span); - /*! - * \brief Resolve type to the solution type in the solver. - * \param type The type to be resolved. - * \return The resolved type. - */ - Type Resolve(const Type& type); - /*! - * \brief Start to solve the types using the current known information. - * \return Whether all the incomplete types has been fully resolved. - */ - bool Solve(); - /*! - * \brief Unify lhs and rhs. - * \param lhs The left operand. - * \param rhs The right operand - * \param location The location at which the unification problem arose. - */ - Type Unify(const Type& lhs, const Type& rhs, const Span& span, bool assign_lhs = true, - bool assign_rhs = true); - /*! - * \brief Report a diagnostic. - * \param diag The diagnostic to report. - */ - void Emit(const Diagnostic& diag) { diag_ctx_.Emit(diag); } - - private: - class AnyChecker; - class OccursChecker; - class Unifier; - class Resolver; - class Propagator; - class Merger; - class Reporter; - struct TypeNode; - struct RelationNode; - // Internally the solver maintains a bipartite graph of Relation and Types. - // All the object in the structure is managed by a arena allocator - // which releases the memory upon distruction of the type solver. - /*! - * \brief type node struct - * TypeNode implements a union-find data structure(via parent) - * that can unifies the same types to the name resolved_type. - * - * It also contains collection of links to related Relations, - * which is stored in rel_set. - */ - struct TypeNode { - /*! \brief The final resolved type */ - Type resolved_type; - /*! \brief type node in the union find algorithm */ - TypeNode* parent{nullptr}; - /*! \brief set of relations that is related to this type node */ - std::unordered_set rel_set; - - /*! - * \brief Find the root type node, perform path compression - * \return The root type node. - */ - TypeNode* FindRoot() { - // fast path - if (this->parent == nullptr) return this; - // slow path with path compression. - TypeNode* root = this; - while (root->parent != nullptr) { - root = root->parent; - } - for (TypeNode* p = this; p != root;) { - TypeNode* parent = p->parent; - p->parent = root; - p = parent; - } - return root; - } - }; - - /*! \brief relation node */ - struct RelationNode { - /*! \brief Whether the relation is in the queue to be solved */ - bool inqueue{false}; - /*! \brief Whether the relation is resolved */ - bool resolved{false}; - /*! \brief The corresponding type relation */ - TypeRelation rel; - /*! \brief list types to this relation */ - LinkedList type_list; - /*! \brief The location this type relation originated from. */ - Span span; - }; - - /*! \brief A simple union find between shapes. */ - tvm::Map shape_uf_; - /*! \brief List of all allocated type nodes */ - std::vector type_nodes_; - /*! \brief List of all allocated relation nodes */ - std::vector rel_nodes_; - /*! \brief Number of resolved relations */ - size_t num_resolved_rels_{0}; - /*! \brief map from types to type nodes. */ - std::unordered_map tmap_; - /*! \brief Internal queue to update the relation */ - std::queue update_queue_; - /*! \brief allocator of all the internal node obhect*/ - support::Arena arena_; - /*! \brief Reporter that reports back to self */ - TypeReporter reporter_; - /*! \brief The global representing the current function. */ - GlobalVar current_func_; - /*! \brief The diagnostic context. */ - DiagnosticContext diag_ctx_; - /*! \brief The module. */ - IRModule module_; - - /*! - * \brief GetTypeNode that is corresponds to t. - * if it do not exist, create a new one. - * \return The type node. - */ - TypeNode* GetTypeNode(const Type& t) { - auto it = tmap_.find(t); - if (it != tmap_.end()) { - return it->second->FindRoot(); - } else { - TypeNode* n = arena_.make(); - type_nodes_.push_back(n); - n->resolved_type = t; - tmap_[t] = n; - return n; - } - } - /*! - * \brief Add relation node rel to the update queue - * \param rel The relation node - */ - void AddToQueue(RelationNode* rel) { - if (rel->inqueue) return; - ICHECK(!rel->resolved); - rel->inqueue = true; - update_queue_.push(rel); - } - - /*! - * \brief Merge rhs type node to lhs - * \param src The source operand - * \param dst The dst operand. - */ - void MergeFromTo(TypeNode* src, TypeNode* dst); -}; - -} // namespace relay -} // namespace tvm -#endif // TVM_RELAY_ANALYSIS_TYPE_SOLVER_H_ diff --git a/src/relay/analysis/util.cc b/src/relay/analysis/util.cc deleted file mode 100644 index 96db7d762cae..000000000000 --- a/src/relay/analysis/util.cc +++ /dev/null @@ -1,512 +0,0 @@ -/* - * 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 util.cc - * - * \brief Utility functions for Relay. - */ -#include -#include -#include -#include -#include -#include -#include - -#include "../transforms/pass_utils.h" - -namespace tvm { -namespace relay { - -template -struct InsertionSet { - std::unordered_set set; - std::vector data; - void Insert(const T& t) { - if (set.count(t) == 0) { - set.insert(t); - data.push_back(t); - } - } -}; - -class TypeVarTVisitor : public TypeVisitor { - public: - TypeVarTVisitor(InsertionSet* type_vars, InsertionSet* bound_type_vars) - : type_vars_(type_vars), bound_type_vars_(bound_type_vars) {} - - void VisitType_(const TypeVarNode* tp) final { - TypeVar var = GetRef(tp); - type_vars_->Insert(var); - } - - void VisitType_(const FuncTypeNode* f) final { - for (auto type_param : f->type_params) { - type_vars_->Insert(type_param); - bound_type_vars_->Insert(type_param); - } - TypeVisitor::VisitType_(f); - } - - private: - InsertionSet* type_vars_; - InsertionSet* bound_type_vars_; -}; - -class TypeVarEVisitor : private MixedModeVisitor { - public: - explicit TypeVarEVisitor(const IRModule& mod) : mod_(mod) {} - - Array CollectFree() { - Array ret; - for (const auto& v : type_vars_.data) { - if (bound_type_vars_.set.count(v) == 0) { - ret.push_back(v); - } - } - return ret; - } - - Array CollectBound() { - Array ret; - for (const auto& v : bound_type_vars_.data) { - ret.push_back(v); - } - return ret; - } - - Array CollectAll() { - Array ret; - for (const auto& v : type_vars_.data) { - ret.push_back(v); - } - return ret; - } - - Array Free(const Expr& expr) { - VisitExpr(expr); - return CollectFree(); - } - - Array Free(const Type& type) { - VisitType(type); - return CollectFree(); - } - - Array Bound(const Expr& expr) { - VisitExpr(expr); - return CollectBound(); - } - - Array Bound(const Type& type) { - VisitType(type); - return CollectBound(); - } - - Array All(const Expr& expr) { - VisitExpr(expr); - return CollectAll(); - } - - Array All(const Type& type) { - VisitType(type); - return CollectAll(); - } - - using MixedModeVisitor::VisitExpr_; - - void VisitExpr_(const FunctionNode* f) final { - for (const auto& tp : f->type_params) { - type_vars_.Insert(tp); - bound_type_vars_.Insert(tp); - } - ExprVisitor::VisitExpr_(f); - } - - void VisitExpr_(const LetNode* op) final { - auto pre_visit = [this](const LetNode* op) { - this->VisitExpr(op->var); - this->VisitExpr(op->value); - }; - auto post_visit = [this](const LetNode* op) { - this->VisitExpr(op->body); - this->visit_counter_[op] += 1; - }; - ExpandANormalForm(op, pre_visit, post_visit); - } - - void VisitExpr_(const ConstructorNode* cn) final { - // for constructors, type vars will be bound in the module - auto data = mod_->LookupTypeDef(cn->belong_to); - for (const auto& tv : data->type_vars) { - type_vars_.Insert(tv); - bound_type_vars_.Insert(tv); - } - ExprVisitor::VisitExpr_(cn); - } - - void VisitType(const Type& t) final { - TypeVarTVisitor(&type_vars_, &bound_type_vars_).VisitType(t); - } - - private: - InsertionSet type_vars_; - InsertionSet bound_type_vars_; - const IRModule& mod_; -}; - -class VarVisitor : protected MixedModeVisitor, protected PatternVisitor { - public: - Array Free(const Expr& expr) { - this->VisitExpr(expr); - Array ret; - for (const auto& v : vars_.data) { - if (bound_vars_.set.count(v) == 0) { - ret.push_back(v); - } - } - return ret; - } - - Array Collect() { - Array ret; - for (const auto& v : bound_vars_.data) { - ret.push_back(v); - } - return ret; - } - - Array Bound(const Expr& expr) { - this->VisitExpr(expr); - return Collect(); - } - - Array Bound(const Pattern& pat) { - this->VisitPattern(pat); - return Collect(); - } - - Array All(const Expr& expr) { - this->VisitExpr(expr); - Array ret; - for (const auto& v : vars_.data) { - ret.push_back(v); - } - return ret; - } - - void MarkBounded(const Var& v) { - bound_vars_.Insert(v); - vars_.Insert(v); - } - - using MixedModeVisitor::VisitExpr_; - - void VisitExpr_(const VarNode* var) final { vars_.Insert(GetRef(var)); } - - void VisitExpr_(const FunctionNode* op) final { - for (const auto& param : op->params) { - MarkBounded(param); - } - VisitExpr(op->body); - } - - void VisitExpr_(const LetNode* op) final { - Expr let = GetRef(op); - while (auto let_node = let.as()) { - MarkBounded(let_node->var); - VisitExpr(let_node->value); - let = let_node->body; - } - VisitExpr(let); - } - - void VisitPattern(const Pattern& p) final { PatternVisitor::VisitPattern(p); } - - void VisitPattern_(const PatternVarNode* op) final { MarkBounded(op->var); } - - private: - InsertionSet vars_; - InsertionSet bound_vars_; -}; - -tvm::Array FreeTypeVars(const Expr& expr, const IRModule& mod) { - return TypeVarEVisitor(mod).Free(expr); -} - -tvm::Array FreeTypeVars(const Type& type, const IRModule& mod) { - return TypeVarEVisitor(mod).Free(type); -} - -tvm::Array BoundTypeVars(const Expr& expr, const IRModule& mod) { - return TypeVarEVisitor(mod).Bound(expr); -} - -tvm::Array BoundTypeVars(const Type& type, const IRModule& mod) { - return TypeVarEVisitor(mod).Bound(type); -} - -tvm::Array AllTypeVars(const Expr& expr, const IRModule& mod) { - return TypeVarEVisitor(mod).All(expr); -} - -tvm::Array AllTypeVars(const Type& type, const IRModule& mod) { - return TypeVarEVisitor(mod).All(type); -} - -tvm::Array FreeVars(const Expr& expr) { return VarVisitor().Free(expr); } - -tvm::Array BoundVars(const Expr& expr) { return VarVisitor().Bound(expr); } - -tvm::Array BoundVars(const Pattern& pat) { return VarVisitor().Bound(pat); } - -tvm::Array AllVars(const Expr& expr) { return VarVisitor().All(expr); } - -TVM_REGISTER_GLOBAL("relay.analysis.free_vars").set_body_typed(FreeVars); - -TVM_REGISTER_GLOBAL("relay.analysis.bound_vars").set_body([](TVMArgs args, TVMRetValue* ret) { - ObjectRef x = args[0]; - if (x.as()) { - *ret = BoundVars(Downcast(x)); - } else { - *ret = BoundVars(Downcast(x)); - } -}); - -TVM_REGISTER_GLOBAL("relay.analysis.all_vars").set_body_typed(AllVars); - -TVM_REGISTER_GLOBAL("relay.analysis.free_type_vars").set_body([](TVMArgs args, TVMRetValue* ret) { - ObjectRef x = args[0]; - IRModule mod = args[1]; - if (x.as()) { - *ret = FreeTypeVars(Downcast(x), mod); - } else { - *ret = FreeTypeVars(Downcast(x), mod); - } -}); - -TVM_REGISTER_GLOBAL("relay.analysis.bound_type_vars").set_body([](TVMArgs args, TVMRetValue* ret) { - ObjectRef x = args[0]; - IRModule mod = args[1]; - if (x.as()) { - *ret = BoundTypeVars(Downcast(x), mod); - } else { - *ret = BoundTypeVars(Downcast(x), mod); - } -}); - -TVM_REGISTER_GLOBAL("relay.analysis.all_type_vars").set_body([](TVMArgs args, TVMRetValue* ret) { - ObjectRef x = args[0]; - IRModule mod = args[1]; - if (x.as()) { - *ret = AllTypeVars(Downcast(x), mod); - } else { - *ret = AllTypeVars(Downcast(x), mod); - } -}); - -class DtypeCollector : protected ExprVisitor, protected TypeVisitor { - public: - void VisitExpr(const Expr& expr) final { - if (expr->checked_type_.defined()) { - TypeVisitor::VisitType(expr->checked_type()); - } - ExprVisitor::VisitExpr(expr); - } - - void VisitType_(const TensorTypeNode* op) final { dtypes_.insert(DLDataType2String(op->dtype)); } - - Array All(const Expr& expr) { - VisitExpr(expr); - - Array res; - for (const auto& dtype : dtypes_) { - res.push_back(String(dtype)); - } - return res; - } - - private: - std::unordered_set dtypes_; -}; - -tvm::Array AllDtypes(const Expr& expr) { return DtypeCollector().All(expr); } - -TVM_REGISTER_GLOBAL("relay.analysis.all_dtypes").set_body_typed(AllDtypes); - -/*! - * \brief Get reference counter of each internal ExprNode in body. - * \param body The body expression. - * \return The reference count mapping. - */ -std::unordered_map GetExprRefCount(const Expr& body) { - class ExprRefCounter : private MixedModeVisitor { - public: - std::unordered_map Get(const Expr& body) { - this->VisitExpr(body); - return std::move(this->visit_counter_); - } - }; - return ExprRefCounter().Get(body); -} - -template -bool IsNDArrayAllGreaterEqual(const runtime::NDArray& tensor, T value) { - ICHECK_EQ(tensor->device.device_type, kDLCPU); - ICHECK(tensor->strides == nullptr); - ICHECK_EQ(tensor->byte_offset, 0); - const T* data = static_cast(tensor->data); - int64_t num_elems = 1; - for (int i = 0; i < tensor->ndim; ++i) { - num_elems *= tensor->shape[i]; - } - - for (int64_t i = 0; i < num_elems; i++) { - if (*data < value) { - return false; - } - data++; - } - return true; -} - -bool IsAllPositiveConstant(const Expr& expr) { - // Cache the operators that are checked recursively to reduce lookup overhead. - static const auto& expand_dims_op = Op::Get("expand_dims"); - static const auto& reshape_op = Op::Get("reshape"); - static const auto& transpose_op = Op::Get("transpose"); - static const auto& squeeze_op = Op::Get("squeeze"); - static const auto& repeat_op = Op::Get("repeat"); - - // peel through a few common transform ops. - if (const auto* constant = expr.as()) { - const auto& tensor = constant->data; - const auto& dtype = tensor->dtype; - if (dtype.lanes != 1) { - return false; - } else if (dtype.code == kDLFloat && dtype.bits == 32) { - return IsNDArrayAllGreaterEqual(tensor, 0); - } else if (dtype.code == kDLFloat && dtype.bits == 64) { - return IsNDArrayAllGreaterEqual(tensor, 0); - } else if (dtype.code == kDLInt && dtype.bits == 8) { - return IsNDArrayAllGreaterEqual(tensor, 0); - } else if (dtype.code == kDLInt && dtype.bits == 32) { - return IsNDArrayAllGreaterEqual(tensor, 0); - } else if (dtype.code == kDLUInt && dtype.bits == 8) { - return IsNDArrayAllGreaterEqual(tensor, 0); - } else if (dtype.code == kDLUInt && dtype.bits == 32) { - return IsNDArrayAllGreaterEqual(tensor, 0); - } else { - return false; - } - } else if (const auto* op = expr.as()) { - // tail recursion. - if (op->op == expand_dims_op || op->op == reshape_op || op->op == transpose_op || - op->op == squeeze_op || op->op == repeat_op) { - return IsAllPositiveConstant(op->args[0]); - } else { - return false; - } - } else { - return false; - } -} - -Type TypeSubst(const Type& type, const TypeVar& tvar, const Type& subst) { - return TypeSubst(type, tvm::Map({{tvar, subst}})); -} - -Expr TypeSubst(const Expr& expr, const TypeVar& tvar, const Type& subst) { - return TypeSubst(expr, tvm::Map({{tvar, subst}})); -} - -Type TypeSubst(const Type& type, const tvm::Map& subst_map) { - return Bind(type, subst_map); -} - -Expr TypeSubst(const Expr& expr, const tvm::Map& subst_map) { - class TypeSubstMutator : public ExprMutator, public PatternMutator { - public: - explicit TypeSubstMutator(const tvm::Map& subst_map) : subst_map_(subst_map) {} - Type VisitType(const Type& t) final { return TypeSubst(t, subst_map_); } - Var VisitVar(const Var& v) final { return Downcast(VisitExpr(v)); } - - Pattern VisitPattern(const Pattern& p) final { return PatternMutator::VisitPattern(p); } - - Clause VisitClause(const Clause& c) final { - Pattern pat = VisitPattern(c->lhs); - return Clause(pat, VisitExpr(c->rhs)); - } - - private: - const tvm::Map& subst_map_; - }; - ICHECK(WellFormed(expr)); - auto ret = TypeSubstMutator(subst_map).VisitExpr(expr); - ICHECK_EQ(FreeVars(expr).size(), FreeVars(ret).size()); - ICHECK(WellFormed(ret)); - return ret; -} - -struct IsDynamicVisitor : public TypeVisitor { - bool is_dyn{false}; - void VisitType_(const TensorTypeNode* tt) { - for (auto dim : tt->shape) { - if (dim.as() == nullptr) { - is_dyn = true; - break; - } - } - } -}; - -bool IsDynamic(const Type& ty) { - IsDynamicVisitor v; - v.VisitType(ty); - return v.is_dyn; -} - -TVM_REGISTER_GLOBAL("relay.ir.IsDynamic").set_body_typed(IsDynamic); - -bool IsDataDependent(const CallNode* call) { - static auto tshape_data_dependent = Op::GetAttrMap("TShapeDataDependent"); - Op op = Downcast(call->op); - - if (!tshape_data_dependent.count(op)) { - return false; - } - - if (op->name == "strided_slice") { - if (const auto* attrs = call->attrs.as()) { - if (attrs->begin && attrs->end && attrs->strides) { - // not data dependent if begin, end and strides exist - return false; - } - } - } - - for (auto req : tshape_data_dependent[op]) { - if (req->value != 0) return true; - } - return false; -} -} // namespace relay -} // namespace tvm diff --git a/src/relay/analysis/well_formed.cc b/src/relay/analysis/well_formed.cc deleted file mode 100644 index d8a5bb8e4f65..000000000000 --- a/src/relay/analysis/well_formed.cc +++ /dev/null @@ -1,165 +0,0 @@ -/* - * 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 well_formed.cc - * \brief check that expression is well formed. - */ -#include -#include -#include -#include - -#include - -namespace tvm { -namespace relay { - -//! brief make sure each Var is bound at most once in a scope. -class WellFormedChecker : private MixedModeVisitor, PatternVisitor { - public: - Optional diag_ctx; - Span occurs_in; - - explicit WellFormedChecker(const Optional& ctx) : diag_ctx(ctx) {} - - bool well_formed = true; - - void Illformed(Diagnostic diag) { - well_formed = false; - if (diag_ctx) { - diag_ctx.value().Emit(diag); - } else { - LOG(INFO) << "The IR is not well formed with: " << diag->message; - } - } - - std::vector> scope; - std::unordered_set current_bound; - std::unordered_set total_bound; - std::unordered_set free; - - struct Scope { - WellFormedChecker* wfc; - explicit Scope(WellFormedChecker* wfc) : wfc(wfc) { wfc->scope.push_back({{}}); } - ~Scope() { - ICHECK_GE(wfc->scope.size(), 0); - for (const Var& v : wfc->scope.back()) { - ICHECK_GE(wfc->current_bound.count(v), 0); - wfc->current_bound.erase(v); - } - wfc->scope.pop_back(); - } - }; - - void Bound(const Var& v) { - if (current_bound.count(v) != 0 || total_bound.count(v) != 0 || free.count(v) != 0) { - Illformed(Diagnostic::Error(v->span) << "The variable " << v->name_hint() - << " is bound more than once, this is not valid IR"); - } - ICHECK_GE(scope.size(), 0); - scope.back().insert(v); - current_bound.insert(v); - total_bound.insert(v); - } - - using MixedModeVisitor::VisitExpr_; - - void VisitExpr_(const VarNode* op) final { - Var v = GetRef(op); - if (current_bound.count(v) == 0) { - if (total_bound.count(v) != 0) { - Illformed(Diagnostic::Error(v->span) << "the variable " << v->name_hint() - << "is bound more then once, this is not valid IR"); - } else { - free.insert(v); - } - } - } - - void VisitExpr_(const LetNode* l) final { - std::vector scopes; - Expr let = GetRef(l); - while (auto let_node = let.as()) { - scopes.push_back(new Scope(this)); - // we do letrec only for FunctionNode, - // but shadowing let in let binding is likely programming error, and we should forbidden it. - Bound(let_node->var); - CheckWellFormed(let_node->value); - let = let_node->body; - } - CheckWellFormed(let); - while (!scopes.empty()) { - delete scopes.back(); - scopes.pop_back(); - } - } - - void VisitExpr_(const FunctionNode* f) final { - Scope s(this); - for (const Var& param : f->params) { - Bound(param); - } - CheckWellFormed(f->body); - } - - void VisitExpr_(const CallNode* call) final { - ICHECK(call->op.defined()); - - for (auto arg : call->args) { - ICHECK(arg.defined()); - } - - // ICHECK(call->attrs.defined()); - ICHECK(call->type_args.defined()); - MixedModeVisitor::VisitExpr_(call); - } - - void VisitClause(const Clause& c) final { - Scope s(this); - VisitPattern(c->lhs); - VisitExpr(c->rhs); - } - - void VisitPattern(const Pattern& p) final { PatternVisitor::VisitPattern(p); } - - void VisitVar(const Var& v) final { Bound(v); } - - public: - bool CheckWellFormed(const Expr& e) { - if (auto v = e.as()) { - VisitExpr_(v); - } else { - // this->occurs_in = e->span; - VisitExpr(e); - } - return well_formed; - } -}; - -bool WellFormed(const Expr& e, Optional diag_ctx) { - return WellFormedChecker(diag_ctx).CheckWellFormed(e); -} - -TVM_REGISTER_GLOBAL("relay.analysis.well_formed").set_body_typed([](Expr e) { - return WellFormed(e); -}); - -} // namespace relay -} // namespace tvm diff --git a/src/relay/backend/annotate_used_memory.cc b/src/relay/backend/annotate_used_memory.cc deleted file mode 100644 index 8e7ab68cbac9..000000000000 --- a/src/relay/backend/annotate_used_memory.cc +++ /dev/null @@ -1,236 +0,0 @@ -/* - * 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 src/relay/backend/annotate_used_memory.cc - * \brief Analyzes the used memory at the callsite of primitive functions. - */ - -#include -#include -#include - -#include -#include - -#include "../transforms/device_aware_visitors.h" -#include "../transforms/pass_utils.h" -#include "./liveness_analysis.h" -#include "./utils.h" - -namespace tvm { -namespace relay { -namespace backend { - -/*! - * \brief Annotates the minimum required memory of each primitive function callsite by analyzing - * the liveness of the input/output tensors at each function callsite and calculating the total - * amount of memory these tensors require. This is added as a "used_memory" annotation to the - * function in question as a list of the number of bytes for each callsite. In addition, the - * containing function is annotated with an "io_used_memory" annotation which refers to the total - * memory required for the IO tensors. - * - * Note: This pass does not support dynamic shapes, it is the users responsibility to check this - * pass isn't applied where dynamic shapes may be input. - * - * A simple example: - * - * Before: - * \verbatim - * def @main(%input: Tensor[(1, 2, 2, 4), int8]) -> Tensor[(1, 2, 2, 4), int8] { - * let %x_0 = fn (%x: Tensor[(1, 2, 2, 4), int8], Primitive=1) -> Tensor[(1, 2, 2, 4), int8] { - * nn.max_pool2d(%x, pool_size=[1, 1], padding=[0, 0, 0, 0]) - * }; - * let %x_1 = %x_0(%input); - * %x_1 - * } - * \endverbatim - * - * After: - * \verbatim - * def @main(%input: Tensor[(1, 2, 2, 4), int8], io_used_memory=32) -> Tensor[(1, 2, 2, 4), int8] { - * let %x_0: fn (%x: Tensor[(1, 2, 2, 4), int8], Primitive=1, used_memory=[32]) -> Tensor[(1, 2, - * 2, 4), int8] { - * nn.max_pool2d(%x, pool_size=[1, 1], padding=[0, 0, 0, 0]) - * }; - * let %x_1: Tensor[(1, 2, 2, 4), int8] = %x_0(%input); - * %x_1 - * } - * \endverbatim - * - * Note that in the simple example above io_used_memory and used_memory are the same since there - * is only one primitive function. - */ -class AnnotateUsedMemoryMutator : public transform::DeviceAwareExprMutator { - public: - AnnotateUsedMemoryMutator(const IRModule& module, const transform::ControlFlowGraph& cfg, - const transform::LivenessAnalysis& lva) - : DeviceAwareExprMutator(module), control_flow_graph_(cfg), liveness_(lva) {} - - /*! - * \brief Mutates the input function. In addition, an "io_used_memory" annotation is - * added to the input function which refers to the total size required for the IO - * tensors. - */ - Function operator()(const Function& func) { - uint64_t io_used_memory = 0; - - // Inputs - for (const Var& param : func->params) { - Type type = param->checked_type(); - ICHECK(type.defined()) << "InferType pass should be run before AnnotateUsedMemory."; - ICHECK(!IsDynamic(type)) << "AnnotateUsedMemory does not support dynamic shapes."; - io_used_memory += CalculateRelayExprSizeBytes(type); - } - - // Outputs - Type type = func->body->checked_type(); - ICHECK(type.defined()) << "InferType pass should be run before AnnotateUsedMemory."; - ICHECK(!IsDynamic(type)) << "AnnotateUsedMemory does not support dynamic shapes."; - io_used_memory += CalculateRelayExprSizeBytes(type); - - Expr new_func_body = VisitExpr(func->body); - Function new_func = WithFields(func, func->params, new_func_body); - return WithAttr(std::move(new_func), "io_used_memory", - tvm::IntImm(tvm::DataType::UInt(64), io_used_memory)); - } - - /*! - * \brief Establish which let bindings have primitive function values. - */ - std::pair PreVisitLetBinding_(const Var& var, const Expr& value) override { - if (const auto* func_node = value.as()) { - ICHECK(func_node->attrs.HasNonzeroAttr(attr::kPrimitive)) - << "Expect top-level functions to be primitive."; - let_bound_prim_func_.insert(var); - } - return DeviceAwareExprMutator::PreVisitLetBinding_(var, value); - } - - /*! - * \brief Visit let nodes and perform one of two actions depending on their value: - * - * 1. CallNode - Calculate "used_memory" annotation value at the callsite of - * primitive functions. - * - * 2. FunctionNode - Annotate functions with "used_memory" annotation based on the - * previous analysis at the callsite. - * - */ - Expr PostVisitLet_(const LetNode* pre_let_node, const LetNode* post_let_node) override { - Var let_var = post_let_node->var; - Expr let_value = IgnoreOnDevice(post_let_node->value); - - if (let_value->IsInstance()) { - Call callsite = Downcast(let_value); - if (CheckPrimitiveFunctionCall(callsite)) { - Var call_op = Downcast(callsite->op); - - // Find all the vars that are live at the callsite. This is done by merging the - // in and out varset's and then removing the var that references the primitive - // function itself since we don't want this included in the calculation. - const transform::ControlFlowGraph::NodePtr cfg_node = - control_flow_graph_.let_map.at(GetRef(pre_let_node)); - transform::VarSet live_tensors = liveness_.live_in.at(cfg_node); - const transform::VarSet& live_out = liveness_.live_out.at(cfg_node); - live_tensors.insert(live_out.begin(), live_out.end()); - live_tensors.erase(call_op); - - // Calculate size of live tensors and store to allow annotation when the function - // gets visited. - uint64_t used_memory = 0; - for (const auto& var : live_tensors) { - Type type = var->checked_type(); - ICHECK(type.defined()) << "InferType pass should be run before AnnotateUsedMemory."; - ICHECK(!IsDynamic(type)) << "AnnotateUsedMemory does not support dynamic shapes."; - used_memory += CalculateRelayExprSizeBytes(type); - } - IntImm annotation(DataType::UInt(64), used_memory); - used_memory_annotations_[call_op].push_back(annotation); - } - } else if (let_value->IsInstance()) { - Function func = Downcast(let_value); - ICHECK(used_memory_annotations_.find(let_var) != used_memory_annotations_.end()) - << "Could not find used_memory value for primitive function bound at " - << let_var->name_hint(); - Array used_memory = used_memory_annotations_[let_var]; - used_memory_annotations_.erase(let_var); - - Function new_func = WithAttr(std::move(func), "used_memory", - Array(used_memory.rbegin(), used_memory.rend())); - return Let(let_var, new_func, post_let_node->body, post_let_node->span); - } - - return DeviceAwareExprMutator::PostVisitLet_(pre_let_node, post_let_node); - } - - private: - /*! - * \brief Check if a call is a primitive function callsite. - */ - bool CheckPrimitiveFunctionCall(const Call& callsite) { - if (auto var = callsite->op.as()) { - if (let_bound_prim_func_.find(var.value()) != let_bound_prim_func_.end()) { - return true; - } - } - return false; - } - - /*! \brief Control flow graph representation of the main function. */ - transform::ControlFlowGraph control_flow_graph_; - /*! \brief Liveness analysis of the main function. */ - transform::LivenessAnalysis liveness_; - /*! \brief Var's that reference primitive functions. */ - std::unordered_set let_bound_prim_func_; - /*! \brief Stores the calculated uint64 used_memory values so they can be annotated on the - * relevant function. */ - std::unordered_map, ObjectPtrHash, ObjectPtrEqual> used_memory_annotations_; -}; - -} // namespace backend - -namespace transform { - -Pass AnnotateUsedMemory() { - runtime::TypedPackedFunc pass_func = [=](IRModule mod, - PassContext ctx) { - GlobalVar gv = mod->GetGlobalVar("main"); - Function main_func = Downcast(mod->Lookup("main")); - - // Perform liveness analysis to determine what tensors are 'live' at each functions callsite. - support::Arena arena; - ControlFlowGraph cfg = ControlFlowGraph::Create(&arena, main_func); - UseDefAnalysis use_def = UseDefAnalysis::Analyze(cfg); - LivenessAnalysis lva = LivenessAnalysis::Analyze(cfg, use_def); - - auto new_main_func = backend::AnnotateUsedMemoryMutator(mod, cfg, lva)(main_func); - if (!new_main_func.same_as(main_func)) { - mod->Update(gv, new_main_func); - } - return mod; - }; - return CreateModulePass(pass_func, 0, "AnnotateUsedMemory", {"ToANormalForm", "InferType"}); -} - -TVM_REGISTER_GLOBAL("relay._transform.AnnotateUsedMemory").set_body_typed(AnnotateUsedMemory); - -} // namespace transform -} // namespace relay -} // namespace tvm diff --git a/src/relay/backend/aot/aot_lower_main.cc b/src/relay/backend/aot/aot_lower_main.cc deleted file mode 100644 index c7752c08053f..000000000000 --- a/src/relay/backend/aot/aot_lower_main.cc +++ /dev/null @@ -1,866 +0,0 @@ -/* - * 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 src/relay/backend/aot/aot_lower_main.cc - * \brief Lower the Relay main func into an AOT TIR main func. - */ -#include "./aot_lower_main.h" - -#include -#include -#include - -#include "../../op/call/call.h" -#include "../../op/memory/device_copy.h" -#include "../../op/memory/memory.h" -#include "../../transforms/device_aware_visitors.h" -#include "../name_transforms.h" -#include "../utils.h" - -namespace tvm { -namespace relay { -namespace backend { -namespace aot { - -/*! - * \brief Looks at the expressions in a given function and produces an Expr to - * StorageInfo map by assigning one or more StorageInfos to the expressions that - * require storage. - * - * This pass is leveraged by AOTMainLowerer to perform an initial naive allocation - * for tensors in the Relay main function. The resulting storage map is then lowered - * into TIR allocations by AOTMainLowerer where the allocation can be subsequently - * optimized by later passes (e.g. USMP). - */ -class ExprAllocator : public transform::DeviceAwareExprVisitor { - public: - ExprAllocator() : transform::DeviceAwareExprVisitor(Optional()) {} - - // run the visitor on a global function. - void Run(const Function& func) { VisitExpr(func); } - - std::vector GetReturnSIDs() const { return return_sids_; } - - StorageMap GetStorageMap() const { return expr_storage_map_; } - - using ExprVisitor::VisitExpr_; - - void DeviceAwareVisitExpr_(const CallNode* call_node) final { - Array args; - - CallLoweredProps call_lowered_props = GetCallLoweredProps(call_node); - if (call_lowered_props.lowered_func.defined()) { - args = call_lowered_props.arguments; - } else { // Relay functions that have not been lowered and lowered extern functions - args = call_node->args; - if (call_node->op.as()) { // Lowered extern function - ICHECK(!(call_node->attrs.defined())) << "Extern functions should have null attributes."; - } else { // Relay function which has not been lowered yet - ICHECK(call_node->op.as()) - << "Expected the call to be to a lowered primfunc, a lowered extern function or a " - "unlowered Relay function."; - } - } - CreateStorage(call_node); - for (const Expr& arg : args) { - VisitExpr(arg); - } - AssignReturnSID(GetRef(call_node)); - } - - void DeviceAwareVisitExpr_(const FunctionNode* func_node) final { - if (function_nesting() > 1) { - // Do not recurse into sub functions. - return; - } - for (const auto& param : func_node->params) { - CreateStorage(param.get()); - } - VisitExpr(func_node->body); - } - - void PreVisitLetBinding_(const Var& var, const Expr& value) final { - VisitExpr(value); - StorageInfo si = GetStorage(value); - expr_storage_map_[var] = si; - } - - void VisitExpr_(const ConstantNode* op) final { - CreateStorage(op); - AssignReturnSID(GetRef(op)); - } - - void VisitExpr_(const VarNode* op) final { AssignReturnSID(GetRef(op)); } - - void VisitExpr_(const TupleNode* op) final { - std::vector storage_ids; - std::vector virtual_devices; - std::vector storage_sizes_in_bytes; - Expr expr = GetRef(op); - for (Expr field : op->fields) { - auto sid = GetStorage(field); - storage_ids.insert(storage_ids.end(), sid->storage_ids.begin(), sid->storage_ids.end()); - virtual_devices.insert(virtual_devices.end(), sid->virtual_devices.begin(), - sid->virtual_devices.end()); - storage_sizes_in_bytes.insert(storage_sizes_in_bytes.end(), - sid->storage_sizes_in_bytes.begin(), - sid->storage_sizes_in_bytes.end()); - } - expr_storage_map_[expr] = StorageInfo(storage_ids, virtual_devices, storage_sizes_in_bytes); - AssignReturnSID(expr); - } - - void VisitExpr_(const TupleGetItemNode* op) final { - Expr expr = GetRef(op); - auto sids = GetStorage(op->tuple); - ICHECK_LT(static_cast(op->index), sids->storage_ids.size()); - expr_storage_map_[expr] = - StorageInfo({sids->storage_ids[op->index]}, {sids->virtual_devices[op->index]}, - {sids->storage_sizes_in_bytes[op->index]}); - AssignReturnSID(expr); - } - - void VisitExpr_(const IfNode* op) final { LOG(FATAL) << "'If' is not supported."; } - - private: - /*! - * \brief Assign the expression's storage IDs as the return storage IDs. - * \note This is called when visiting every expression on the understanding - * that the returned expression will be visited last. - */ - void AssignReturnSID(const Expr& e) { - if (expr_storage_map_.find(e) != expr_storage_map_.end()) { - StorageInfo& sinfo = expr_storage_map_[e]; - return_sids_.clear(); - for (auto sid : sinfo->storage_ids) { - return_sids_.push_back(sid); - } - } - } - - /*! - * \brief Get the necessary storage for the expression. - * \param expr The expression. - * \return The corresponding token. - */ - StorageInfo GetStorage(const Expr& expr) { - // See through "on_device" calls. - Expr true_expr = IgnoreOnDevice(expr); - VisitExpr(true_expr); - auto it = expr_storage_map_.find(true_expr); - ICHECK(it != expr_storage_map_.end()) << "Could not find " << true_expr->GetTypeKey() << " " - << PrettyPrint(true_expr) << " in storage device map"; - return it->second; - } - - /*! - * \brief Create storage for the expression. - */ - void CreateStorage(const ExprNode* op) { - Expr expr = GetRef(op); - return CreateStorage(expr, GetVirtualDevice(expr)); - } - - /*! - * \brief Create storage to hold the result of evaluating \p expr in \p virtual_device. - */ - void CreateStorage(const Expr& expr, const VirtualDevice& virtual_device) { - ICHECK(!virtual_device->IsFullyUnconstrained()) - << "invalid virtual device for expr:" << std::endl - << PrettyPrint(expr); - std::vector storage_ids; - std::vector virtual_devices; - std::vector storage_sizes_in_bytes; - for (const auto& ttype : FlattenTupleType(expr->checked_type())) { - storage_ids.push_back(next_available_sid_++); - virtual_devices.push_back(virtual_device); - storage_sizes_in_bytes.push_back(GetMemorySizeBytes(ttype->shape, ttype->dtype)); - } - expr_storage_map_[expr] = StorageInfo(std::move(storage_ids), std::move(virtual_devices), - std::move(storage_sizes_in_bytes)); - } - - /*! \brief Map between Exprs and StorageInfos */ - StorageMap expr_storage_map_; - /*! \brief The next available storage ID to be used */ - int next_available_sid_{0}; - /*! \brief The storage IDs that correspond to return values */ - std::vector return_sids_; -}; - -std::tuple> CreateStorage(const Function& func) { - ExprAllocator expr_allocator; - expr_allocator.Run(func); - return std::make_tuple(expr_allocator.GetStorageMap(), expr_allocator.GetReturnSIDs()); -} - -class AOTMainLowerer : public MixedModeVisitor { - public: - AOTMainLowerer(tvm::CompilationConfig config, CallType call_type) - : config_(config), call_type_(call_type) {} - - IRModule Lower(IRModule mod, String mod_name) { - VLOG_CONTEXT << "AOT"; - IRModule lowered_mod = GetRef(mod.CopyOnWrite()); - - auto lowered_main = lowered_mod->Lookup("main"); - auto lowered_main_func = Downcast(lowered_main); - - // Assign StorageInfo to all the Relay exprs and get the return SIDs - std::tie(expr_storage_map_, return_sid_) = CreateStorage(lowered_main_func); - - for (auto input : lowered_main_func->params) { - input_vars_.push_back(input); - std::string input_name = tvm::runtime::SanitizeName(input->name_hint()); - // We don't want the compiler changing input names in the - // event of a sanitization collision. Therefore, enforcing - // the var created to use the input_name strictly. - CreateIOVar(input, input_name, /*use_unique_name = */ false); - } - - // Define the storage allocator ids - for (auto kv : expr_storage_map_) { - for (auto sid : kv.second->storage_ids) { - // The buffer_var is created with storage_scope to be global.workspace to be serviced by - // TVMBackendAllocWorkspace(TVMBAW) calls, explicitly. The reasoning being the executor - // allocates should be serviced by TVMBAWs as the data could be accessed by many devices and - // should not be lowered to the stack. For more details please refer to the discussion here: - // https://github.com/apache/tvm/issues/9022 - tir::Var buffer_var(MakeString("sid_", sid), - PointerType(PrimType(DataType::Int(8)), "global.workspace")); - sids_table_[sid] = buffer_var; - } - } - - // Create output vars for the TIR main func - // If output tensor names were provided use them - if (auto opt = lowered_main->GetAttr>("output_tensor_names")) { - Array output_tensor_names = opt.value(); - Expr output_expr = lowered_main_func->body; - if (output_expr->checked_type()->IsInstance()) { - TupleType output_tuple_type = Downcast(output_expr->checked_type()); - for (unsigned i = 0; i < output_tuple_type->fields.size(); i++) { - // AoT Executor Codegen does not create these names, - // thus should be used as they are provided. - CreateIOVar(output_tuple_type->fields[i], output_tensor_names[i], - /*use_unique_name = */ false); - } - } else { - // AoT Executor Codegen does not create these names, - // thus should be used as they are provided. - CreateIOVar(lowered_main_func->body, output_tensor_names[0], /*use_unique_name = */ false); - } - } else { - // If output tensor names are not provided we will generate output(x) - // where x is a counter to create unique names. - if (lowered_main_func->body->checked_type()->IsInstance()) { - CreateIOVar(lowered_main_func->body, "output"); - } else { - CreateIOVar(lowered_main_func->body, "output", /*use_unique_name = */ false); - } - } - - CollectDeviceVariables(lowered_mod->GetAttr>("device_contexts") - .value_or(Map())); - VisitExpr(lowered_main_func->body); - - // Remove the Relay main and replace it with the lowered TIR version - lowered_mod->Remove(lowered_mod->GetGlobalVar("main")); - auto tir_main_func = CreateMainFunc(mod_name); - lowered_mod->Update(GlobalVar(runtime::symbol::tvm_module_main), tir_main_func); - lowered_mod = tir::transform::RemoveNoOp()(lowered_mod); - return lowered_mod; - } - - void VisitExpr_(const CallNode* call_node) override { - OnDeviceProps on_device_props = GetOnDeviceProps(call_node); - if (on_device_props.body.defined()) { - VisitExpr(on_device_props.body); - return; - } - - DeviceCopyProps device_copy_props = GetDeviceCopyProps(call_node); - CallLoweredProps call_lowered_props = GetCallLoweredProps(call_node); - - if (device_copy_props.body.defined()) { - // TODO(mbs): device_copy cleaunp - // Suspect treating as no-op is better since already built into the StorageInfo? - LOG(FATAL) << "The AOT executor does not currently support device_copy"; - } - - // At this point we should only see calls of the form call_lowered(@callee, (args...)), - // where @callee can be a PrimFunc we've compiled or an external function supplied via - // some other mechanism. - ICHECK(call_lowered_props.lowered_func.defined()) - << "AOT does not support calling Relay functions. Attempting to call:" << std::endl - << PrettyPrint(GetRef(call_node)); - for (const auto& arg : call_lowered_props.arguments) { - // Evaluate the args - VisitExpr(arg); - } - CreateFuncCall(call_lowered_props, GetRef(call_node)); - } - - void VisitExpr_(const VarNode* op) override { - Expr expr = GetRef(op); - StorageInfo& sinfo = expr_storage_map_[expr]; - - // Let bound vars refer to a value, so these should not be considered "output" vars. - if (let_bound_vars_.find(GetRef(op)) != let_bound_vars_.end()) { - return; - } - - // If the Var node is an output node we need to copy the content of the variable to the output - // It's safe to check the SID here because Var StorageToken are never reallocated - auto output_iter = std::find(return_sid_.begin(), return_sid_.end(), sinfo->storage_ids[0]); - if (output_iter != return_sid_.end()) { - int output_index = std::distance(return_sid_.begin(), output_iter); - auto var_expr = FindExpr(expr); - CopyToOutput(GetBufferVarForIO(input_vars_.size() + output_index), var_expr[0], - /*pack_input*/ false, sinfo->storage_sizes_in_bytes[0]); - } - } - - void VisitExpr_(const ConstantNode* op) override { - Expr expr = GetRef(op); - ICHECK(expr_storage_map_.find(expr) != expr_storage_map_.end()) - << "Storage map did not contain constant expr " << PrettyPrint(expr); - StorageInfo& sinfo = expr_storage_map_[expr]; - std::stringstream ss; - ss << "constant_" << constant_map_.size(); - - tir::Var constant(ss.str(), PointerType(PrimType(DataType(op->data->dtype)))); - constant_map_[constant] = op; - auto sid = sinfo->storage_ids[0]; - sids_table_[sid] = constant; - - // If the Constant node is an output node we need to copy the content of the parameter to the - // output. A node can only produce a single output - auto output_iter = std::find(return_sid_.begin(), return_sid_.end(), sid); - if (output_iter != return_sid_.end()) { - int output_index = std::distance(return_sid_.begin(), output_iter); - auto param_handle = tvm::tir::Call(DataType::Handle(), tvm::tir::builtin::lookup_param(), - {tir::StringImm(ss.str())}); - CopyToOutput(GetBufferVarForIO(input_vars_.size() + output_index), constant, - /* pack_input */ false, sinfo->storage_sizes_in_bytes[0]); - } - } - - void VisitExpr_(const TupleNode* op) override { - for (auto field : op->fields) { - VisitExpr(field); - } - } - - void VisitExpr_(const LetNode* op) override { - auto pre_visit = [this](const LetNode* op) { - let_bound_vars_.insert(op->var); - this->VisitExpr(op->value); - }; - auto post_visit = [this](const LetNode* op) { - this->VisitExpr(op->body); - this->visit_counter_[op] += 1; - }; - ExpandANormalForm(op, pre_visit, post_visit); - } - - void VisitExpr_(const TupleGetItemNode* op) override { VisitExpr(op->tuple); } - void VisitExpr_(const OpNode* op) override { - if (GetRef(op) != CallLoweredOp() && GetRef(op) != OnDeviceOp()) { - LOG(FATAL) << "All OpNodes except for call_lowered should have been expanded"; - } - } - void VisitExpr_(const IfNode* op) override { - LOG(FATAL) << "All GlobalVarNodes should be removed before AOT executor's Codegen is called"; - } - void VisitExpr_(const FunctionNode* op) override { - ICHECK(op->GetAttr(attr::kCompiler).defined()) - << "FunctionNode only supported by custom codegen"; - } - void VisitExpr_(const RefCreateNode* op) override { - LOG(FATAL) << "AOT executor does not support references (found RefCreateNode)"; - } - void VisitExpr_(const RefReadNode* op) override { - LOG(FATAL) << "AOT executor does not support references (found RefReadNode)"; - } - void VisitExpr_(const RefWriteNode* op) override { - LOG(FATAL) << "AOT executor does not support references (found RefWriteNode)"; - } - void VisitExpr_(const ConstructorNode* op) override { - LOG(FATAL) << "AOT executor does not support ADTs (found ConstructorNode)"; - } - void VisitExpr_(const MatchNode* op) override { - LOG(FATAL) << "AOT executor does not support matching (found MatchNode)"; - } - - private: - /*! - * \brief Create the main PrimFunc to execute the graph. - * \note The packed function calls don't pack their arguments. The AOT - * runner function needs to be legalized by the LegalizePackedCalls pass. - */ - tir::PrimFunc CreateMainFunc(String mod_name) { - tir::Stmt body = tir::SeqStmt::Flatten(stmts_); - // Allocate the sids - std::unordered_map allocated; - std::vector> sids_to_allocate; - - for (auto kv : expr_storage_map_) { - // Only allocate sids that are needed - const bool is_input = - (std::find(input_vars_.begin(), input_vars_.end(), kv.first) != input_vars_.end()); - if (is_input) { - continue; - } - - for (unsigned int i = 0; i < kv.second->storage_ids.size(); i++) { - sids_to_allocate.push_back( - std::make_pair(kv.second->storage_ids[i], kv.second->storage_sizes_in_bytes[i])); - } - } - - // Sort the SID allocation to make output deterministic - std::sort(sids_to_allocate.begin(), sids_to_allocate.end()); - - for (auto p : sids_to_allocate) { - int sid = p.first; - int size = p.second; - - if (std::find(return_sid_.begin(), return_sid_.end(), sid) != return_sid_.end()) { - continue; - } - - // Make sure it hasn't already been allocated, this can happen - // with let-bound var/value pairs. - if (allocated.find(sid) != allocated.end()) { - continue; - } - - allocated[sid] = constant_map_.count(sids_table_[sid]); - - // TODO(giuseros): we should allocate this once outside the PrimFunc - // so we don't pay the price of allocation for every inference - if (!allocated[sid]) { - PointerType ptype = Downcast(sids_table_[sid]->type_annotation); - DataType element_type = Downcast(ptype->element_type)->dtype; - body = tir::Allocate(sids_table_[sid], element_type, {size}, tir::const_true(), body); - } - allocated[sid] = true; - } - - for (auto kv : constant_map_) { - auto buffer_var = kv.first; - auto dtype = DataType(kv.second->data->dtype); - - int ndim = kv.second->data->ndim; - Array extents; - - for (int i = 0; i < ndim; i++) { - int shape = kv.second->data->shape[i]; - extents.push_back(tir::make_const(DataType::Int(32), shape, Span())); - } - body = tir::AllocateConst(buffer_var, dtype, extents, kv.second->data, body); - } - - // Define the PrimFunc attributes - Map dict_attrs; - String run_func_name = runtime::get_name_mangled(mod_name, runtime::symbol::tvm_module_main); - dict_attrs.Set("global_symbol", run_func_name); - dict_attrs.Set("runner_function", Bool(true)); - dict_attrs.Set(tvm::attr::kTarget, config_->host_target); - Array input_vars = - Array(main_signature_.begin(), main_signature_.begin() + input_vars_.size()); - dict_attrs.Set("input_vars", input_vars); - Array output_vars = - Array(main_signature_.begin() + input_vars_.size(), - main_signature_.begin() + input_vars_.size() + return_sid_.size()); - dict_attrs.Set("output_vars", output_vars); - Array device_names; - for (const auto& it : devices_) { - device_names.push_back(it.first); - } - dict_attrs.Set("devices", device_names); - - tir::Stmt device_activations = GenerateAllDeviceHook("Activate"); - tir::Stmt device_deactivations = GenerateAllDeviceHook("Deactivate"); - tir::Stmt final_body = tir::SeqStmt({device_activations, body, device_deactivations}); - - // Make the PrimFunc - return tir::PrimFunc(main_signature_, final_body, VoidType(), main_buffer_map_, - DictAttrs(dict_attrs)); - } - - /*! - * \brief Collects device context variables for passing to operators - */ - void CollectDeviceVariables(const Map& device_contexts) { - Map target_contexts; - TargetKindAttrMap target_attr_map = tvm::TargetKind::GetAttrMap("use_device_api"); - - for (const auto& it : device_contexts) { - const GlobalVar& global_var = it.first; - const std::string device_context_name = it.second; - - Optional target_kind = tvm::TargetKind::Get(device_context_name); - if (!target_kind || !target_attr_map.count(target_kind.value())) { - return; - } - if (target_attr_map[target_kind.value()]) { - std::string context_name = tvm::runtime::SanitizeName(device_context_name); - tir::Var device_context_var("device_context_" + context_name, DataType::Handle()); - - auto pair = target_contexts.find(target_kind.value()); - if (pair != target_contexts.end()) { - device_context_var = (*pair).second; - } else { - main_signature_.push_back(device_context_var); - devices_.Set(context_name, device_context_var); - target_contexts.Set(target_kind.value(), device_context_var); - } - - device_contexts_.Set(global_var, device_context_var); - } - } - } - - /*! - * \brief Return a vector of variables that represents the sids for the given Relay Expr - */ - std::vector PackSid(Expr expr) { - std::vector buffer_vars; - - ICHECK(expr_storage_map_.find(expr) != expr_storage_map_.end()) - << "Storage map did not contain constant expr " << PrettyPrint(expr); - StorageInfo& sinfo = expr_storage_map_[expr]; - - // Note that an expression can have multiple sids associated with it - // e.g., returning multiple values from a function - for (auto sid : sinfo->storage_ids) { - // Determine if an sid is an output buffer - auto output_iter = std::find(return_sid_.begin(), return_sid_.end(), sid); - if (output_iter != return_sid_.end()) { - int output_index = std::distance(return_sid_.begin(), output_iter); - buffer_vars.push_back(GetBufferVarForIO(input_vars_.size() + output_index)); - continue; - } - - auto sid_value = sids_table_[sid]; - buffer_vars.push_back(sid_value); - } - return buffer_vars; - } - - /*! - * \brief Given an expression return the variable(s) associated with that expression - */ - std::vector FindExpr(Expr arg) { - auto input_iter = std::find(input_vars_.begin(), input_vars_.end(), arg); - if (input_iter != input_vars_.end()) { - // Input variable - int main_index = std::distance(input_vars_.begin(), input_iter); - return {GetBufferVarForIO(main_index)}; - } else { - // Storage identifier (i.e., intermediate memory) - return PackSid(arg); - } - } - - void PushArgs(const Expr& expr, const std::vector& sids, Array* args) { - const TupleNode* t = expr.as(); - if (t != nullptr) { - CHECK_EQ(sids.size(), t->fields.size()) << "Relay tuple does not map 1:1 into TIR; AOT can't " - "handle this type of Relay Expr in a CallNode."; - } - - args->insert(args->end(), sids.begin(), sids.end()); - } - - /*! - * \brief Wraps a call_extern with a tvm_check_return annotation if required otherwise - * returns the passed Call - */ - tir::Call AddCheckReturn(tir::Call existing_call) { - Array args = {tir::make_const(DataType::Int(32, 1), 0, Span()), - tir::make_const(DataType::Int(32, 1), -1, Span()), existing_call}; - return tir::Call(DataType::Int(32), tir::builtin::tvm_check_return(), args); - } - - /*! - * \brief Create a function call - * \param call_lowered_props The lowered function and the arguments to call it with - * \param result_expr The call we got func and args from (so as to recover the storage - * ids to hold the result). - */ - void CreateFuncCall(CallLoweredProps call_lowered_props, const Expr& result_expr) { - std::string func_name = call_lowered_props.lowered_func->name_hint; - tvm::Array args{tvm::tir::StringImm(func_name)}; - std::vector create_func_call_stmts; - - // Pack the inputs - for (const Expr& arg : call_lowered_props.arguments) { - auto sids = FindExpr(arg); - PushArgs(arg, sids, &args); - } - - // Pack the return(s) value. A call node can produce multiple outputs - auto result_expr_sid = PackSid(result_expr); - PushArgs(result_expr, result_expr_sid, &args); - - GlobalVar global_var = call_lowered_props.lowered_func; - bool has_c_device_api_context = device_contexts_.count(global_var) != 0; - tir::Var device_context; - tir::Stmt func_call; - - switch (call_type_) { - case CallType::kUnpacked: { - // call_extern calling convention with optional context - if (has_c_device_api_context) { - device_context = device_contexts_.Get(global_var).value(); - args.push_back(device_context); - } - func_call = tir::Evaluate(AddCheckReturn( - tvm::tir::Call(DataType::Int(32), tvm::tir::builtin::call_extern(), args))); - break; - } - case CallType::kCPacked: { - if (has_c_device_api_context) { - device_context = device_contexts_.Get(global_var).value(); - args.push_back(device_context); - } else { - // NOTE: LowerTVMBuiltin expects some device_context placeholder. - args.push_back(tir::make_zero(DataType::Handle())); - } - func_call = tir::Evaluate( - tvm::tir::Call(DataType::Int(32), tvm::tir::builtin::tvm_call_cpacked(), args)); - create_func_call_stmts.push_back(func_call); - break; - } - case CallType::kPacked: { - // call_packed does not accept a device context. - CHECK(!has_c_device_api_context) << "CallType::kPacked does not accept a device context"; - func_call = tir::Evaluate(AddCheckReturn( - tvm::tir::Call(DataType::Int(32), tvm::tir::builtin::tvm_call_packed(), args))); - create_func_call_stmts.push_back(func_call); - break; - } - default: - ICHECK(false) << "Unknown CallType: " << call_type_; - } - - ICHECK(func_call.defined()) << "Must define func_call"; - - if (has_c_device_api_context) { - func_call = tir::SeqStmt(Array({ - GenerateDeviceHook(device_context, "Open"), - func_call, - GenerateDeviceHook(device_context, "Close"), - })); - } - - tir::Stmt body = tir::SeqStmt::Flatten(func_call); - stmts_.push_back(body); - } - - /*! - * \brief Copy a variable to the output. This function is mainly used in edge cases - * when we want to return an input or a parameter. - * TODO(giuseros): we should try to avoid unnecessary copy to the output, e.g., in a - * copy-on-write fashion. - */ - void CopyToOutput(PrimExpr out, PrimExpr in, bool pack_input, size_t size) { - // Define intermediate DLTensor to load/store the data - tir::Buffer tmp_read = - tir::decl_buffer({IntImm(DataType::UInt(64), size)}, DataType::UInt(8), "tmp_read"); - tir::Buffer tmp_write = - tir::decl_buffer({IntImm(DataType::UInt(64), size)}, DataType::UInt(8), "tmp_write"); - te::Var loop_idx("i", DataType::Int(32)); - auto retval_i = tir::BufferLoad(tmp_read, {loop_idx}); - // Copy the variable from the input to the output - tir::Stmt copy = tir::For( - loop_idx, 0, tir::make_const(DataType::Int(32, 1), size, Span()), tir::ForKind::kSerial, - tir::BufferStore(tmp_write, tir::Let(tmp_read->data, in, retval_i), {loop_idx})); - stmts_.push_back(tir::LetStmt(tmp_write->data, out, copy)); - } - - /*! - * \brief Generates a call to a given hook for all Devices found for C Device API - * \param hook Name of hook to generate statements for - * \return Statement with function calls for each device - */ - tir::Stmt GenerateAllDeviceHook(const String& hook) { - std::vector device_hooks; - for (const auto& it : devices_) { - const String& device_name = it.first; - const tir::Var& context = it.second; - Array sections = {"Device", device_name, hook}; - String device_hook_name = ToCFunctionStyle(PrefixName(sections)); - - tir::Evaluate device_hook( - AddCheckReturn(tvm::tir::Call(DataType::Int(32), tvm::tir::builtin::call_extern(), - {tvm::tir::StringImm(device_hook_name), context}))); - device_hooks.push_back(device_hook); - } - return tir::SeqStmt::Flatten(device_hooks); - } - - /*! - * \brief Generates a call to a given hook for a single Device function - * \param context Device context to call hook on - * \param hook Name of hook to generate statements for - * \return Statement with function call to Device API - */ - tir::Stmt GenerateDeviceHook(const tir::Var& context, const String& hook) { - const auto& it = std::find_if(std::begin(devices_), std::end(devices_), [&](const auto& it) { - return it.second->name_hint == context->name_hint; - }); - const String& device_name = (*it).first; - Array sections = {"Device", device_name, hook}; - String device_hook = ToCFunctionStyle(PrefixName(sections)); - - return tir::Evaluate( - AddCheckReturn(tir::Call(DataType::Int(32), tvm::tir::builtin::call_extern(), - {tvm::tir::StringImm(device_hook), context}))); - } - - /*! - * \brief Utility function to string together different arguments - */ - template - std::string MakeString(Args const&... args) { - std::ostringstream ss; - using List = int[]; - (void)List{0, ((void)(ss << args), 0)...}; - - return ss.str(); - } - - /*! - * \brief Access IO vars using the buffer vars and - * not the actual var. - */ - tir::Var GetBufferVarForIO(int index) { return main_buffer_map_[main_signature_[index]]->data; } - - /*! - * \brief Create tir::Var for input/output while updating the buffer_maps. - * \param expr The expression to evaluate. - * \param original_name The name of the tir::Var. - * \param use_unique_name Whether to generate a new unique name where a name conflicts. - */ - void CreateIOVar(const Expr& expr, const std::string& original_name, - bool use_unique_name = true) { - CreateIOVar(expr->checked_type(), original_name, use_unique_name); - } - - /*! - * \brief Create tir::Var for input/output while updating the buffer_maps. - * \param expr The expression to evaluate. - * \param original_name The name of the tir::Var. - * \param use_unique_name Whether to generate a new unique name where a name conflicts. - */ - void CreateIOVar(const Type& type, const std::string& original_name, - bool use_unique_name = true) { - if (type->IsInstance()) { - TupleType tuple_type = Downcast(type); - for (unsigned i = 0; i < tuple_type->fields.size(); i++) { - CreateIOVar(tuple_type->fields[i], original_name); - } - } else { - std::string name = original_name; - if (use_unique_name) { - name = GetUniqueIOVarName(original_name); - } - tir::Var var = tir::Var(name, DataType::Handle()); - main_signature_.push_back(var); - auto tensor_type = type.as(); - ICHECK(tensor_type) << "Expected TensorType node but was " << type->GetTypeKey(); - DataType elem_type = tensor_type->dtype; - tir::Var buffer_var = - tir::Var(name + "_buffer_var", PointerType(PrimType(elem_type), "global")); - tir::Buffer buffer = tir::Buffer(buffer_var, elem_type, tensor_type->shape, {}, 0, - name + "_buffer", 16, 1, tir::BufferType::kDefault); - main_buffer_map_.Set(var, buffer); - } - } - - /*! - * \brief Create a unique name for I/O Var - */ - std::string GetUniqueIOVarName(std::string name) { - if (io_var_names_.find(name) == io_var_names_.end()) { - io_var_names_[name] = 1; - return name + std::to_string(io_var_names_[name] - 1); - } else { - io_var_names_[name] = io_var_names_[name] + 1; - return name + std::to_string(io_var_names_[name] - 1); - } - } - - /*! \brief list of input expressions (i.e., variable passed by the user) */ - std::vector input_vars_; - /*! \brief map of device contexts variables */ - Map devices_; - /*! \brief map of GlobalVars to C Device API contexts */ - Map device_contexts_; - /*! \brief input and output variables belonging to the main function signature */ - Array main_signature_; - /*! \brief input and output variables belonging to the main function signature */ - Map main_buffer_map_; - /*! \brief All available targets. */ - CompilationConfig config_; - /*! - * \brief The type of kernel call to be emitted. - * See CallType for more documentation. - */ - CallType call_type_; - std::unordered_map - constant_map_; - /*! \brief plan memory of device result */ - StorageMap expr_storage_map_; - /*! \brief mapping sid -> tir::Var */ - std::unordered_map sids_table_; - /*! \brief the set of statements that make the program */ - std::vector stmts_; - /*! \brief the list of return sids (note that the function might return more then one output */ - std::vector return_sid_; - /*! \brief This is per IO var name counter to aid the generating unique names */ - std::unordered_map io_var_names_; - /*! \brief A set of variables that are let bound. */ - std::unordered_set let_bound_vars_; -}; - -Pass AOTLowerMain(String mod_name, tvm::CompilationConfig config, CallType call_type) { - runtime::TypedPackedFunc pass_func = - [=](IRModule module, transform::PassContext ctx) { - return AOTMainLowerer(config, call_type).Lower(module, mod_name); - }; - - return tvm::transform::CreateModulePass(pass_func, 0, "AOTLowerMain", {"InferType"}); -} - -TVM_REGISTER_GLOBAL("relay.backend.aot.AOTLowerMain") - .set_body_typed([](const String& mod_name, const tvm::CompilationConfig& config, - int call_type) { - return AOTLowerMain(mod_name, config, static_cast(call_type)); - }); - -} // namespace aot -} // namespace backend -} // namespace relay -} // namespace tvm diff --git a/src/relay/backend/aot/aot_lower_main.h b/src/relay/backend/aot/aot_lower_main.h deleted file mode 100644 index 8981e7d7434f..000000000000 --- a/src/relay/backend/aot/aot_lower_main.h +++ /dev/null @@ -1,58 +0,0 @@ -/* - * 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. - */ -#ifndef TVM_RELAY_BACKEND_AOT_AOT_LOWER_MAIN_H_ -#define TVM_RELAY_BACKEND_AOT_AOT_LOWER_MAIN_H_ - -#include -#include - -#include -#include -#include - -#include "../utils.h" - -namespace tvm { -namespace relay { -namespace backend { -namespace aot { - -using StorageMap = - std::unordered_map; - -/*! \brief Exposed for testing, part of the implementation of AOTLowerMain */ -std::tuple> CreateStorage(const Function& func); - -/*! \brief Lower the Relay main function into TIR for use with the AOT executor. - * - * This pass expects that all operators have already been lowered to TIR and - * so only Calls to 'call_lowered' are present in main. - * - * \param mod_name The name of the module. - * \param config The compilation config. - * \param call_type The call type to use when calling functions. - */ -transform::Pass AOTLowerMain(String mod_name, tvm::CompilationConfig config, CallType call_type); - -} // namespace aot -} // namespace backend -} // namespace relay -} // namespace tvm - -#endif // TVM_RELAY_BACKEND_AOT_AOT_LOWER_MAIN_H_ diff --git a/src/relay/backend/aot/create_executor_metadata.cc b/src/relay/backend/aot/create_executor_metadata.cc deleted file mode 100644 index 8ad3566880fa..000000000000 --- a/src/relay/backend/aot/create_executor_metadata.cc +++ /dev/null @@ -1,86 +0,0 @@ -/* - * 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 src/relay/backend/aot/create_executor_metadata.cc - * \brief Create the ExecutorCodegenMetadata from a compiled IRModule. - */ - -#include "./create_executor_metadata.h" - -#include "../utils.h" - -namespace tvm { -namespace relay { -namespace backend { -namespace aot { - -ExecutorCodegenMetadata CreateExecutorMetadata(const IRModule& mod, String mod_name, - Executor executor, Integer workspace_byte_alignment, - Integer constant_byte_alignment) { - // Get relevant executor config information - std::string interface_api = executor->GetAttr("interface-api").value_or("packed"); - bool unpacked_api = executor->GetAttr("unpacked-api").value_or(Bool(false)); - // Get the input vars - auto tir_main_func = Downcast(mod->Lookup(runtime::symbol::tvm_module_main)); - Array inputs = tir_main_func->GetAttr>("input_vars").value(); - Array input_tensor_types; - for (const auto& input : inputs) { - auto buffer = tir_main_func->buffer_map.Get(input).value(); - input_tensor_types.push_back(TensorType(buffer->shape, buffer->dtype)); - } - // Extract USMP metadata to pass onto metadata sources - Map pool_var_info; - std::vector pool_vars; - Optional> allocated_pool_infos = - tir_main_func->GetAttr>(tvm::attr::kPoolArgs); - if (allocated_pool_infos) { - for (const tir::usmp::AllocatedPoolInfo& allocated_pool_info : allocated_pool_infos.value()) { - int pool_var_index = allocated_pool_info->pool_var_idx.value()->value; - pool_vars.push_back(tir_main_func->params[pool_var_index]); - pool_var_info.Set(tir_main_func->params[pool_var_index], allocated_pool_info); - } - } - Map io_pool_allocations = - mod->GetAttr>(tvm::attr::kIOTensorPoolAllocations) - .value_or({}); - - Array outputs = tir_main_func->GetAttr>("output_vars").value(); - Array output_tensor_types; - std::vector output_var_names; - for (const auto& output : outputs) { - auto buffer = tir_main_func->buffer_map.Get(output).value(); - output_tensor_types.push_back(TensorType(buffer->shape, buffer->dtype)); - output_var_names.push_back(output->name_hint); - } - auto devices = tir_main_func->GetAttr>("devices").value_or({}); - - return ExecutorCodegenMetadata(inputs, input_tensor_types, output_var_names, output_tensor_types, - pool_vars, devices, runtime::kTvmExecutorAot, mod_name, - interface_api, unpacked_api, workspace_byte_alignment, - constant_byte_alignment, pool_var_info, io_pool_allocations); -} - -TVM_REGISTER_GLOBAL("relay.backend.aot.CreateExecutorMetadata") - .set_body_typed(CreateExecutorMetadata); - -} // namespace aot -} // namespace backend -} // namespace relay -} // namespace tvm diff --git a/src/relay/backend/aot/create_executor_metadata.h b/src/relay/backend/aot/create_executor_metadata.h deleted file mode 100644 index 5657aa02809c..000000000000 --- a/src/relay/backend/aot/create_executor_metadata.h +++ /dev/null @@ -1,50 +0,0 @@ -/* - * 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. - */ -#ifndef TVM_RELAY_BACKEND_AOT_CREATE_EXECUTOR_METADATA_H_ -#define TVM_RELAY_BACKEND_AOT_CREATE_EXECUTOR_METADATA_H_ - -#include -#include -#include - -#include "../utils.h" - -namespace tvm { -namespace relay { -namespace backend { -namespace aot { - -/*! \brief Create ExecutorCodegenMetadata needed for AOT execution. - * \param mod The module. - * \param mod_name The module name. - * \param executor The executor configuration. - * \param workspace_byte_alignment The alignment of the workspace pool. - * \param constant_byte_alignment The alignment of the constant pool. - * \return The ExecutorCodegenMetadata. - */ -ExecutorCodegenMetadata CreateExecutorMetadata(const IRModule& mod, String mod_name, - Executor executor, Integer workspace_byte_alignment, - Integer constant_byte_alignment); - -} // namespace aot -} // namespace backend -} // namespace relay -} // namespace tvm - -#endif // TVM_RELAY_BACKEND_AOT_CREATE_EXECUTOR_METADATA_H_ diff --git a/src/relay/backend/aot/create_function_metadata.cc b/src/relay/backend/aot/create_function_metadata.cc deleted file mode 100644 index 2ef5e495abca..000000000000 --- a/src/relay/backend/aot/create_function_metadata.cc +++ /dev/null @@ -1,124 +0,0 @@ -/* - * 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 src/relay/backend/aot/create_function_metadata.cc - * \brief Create FunctionInfo metadata from a lowered TIR module. - */ -#include "./create_function_metadata.h" - -#include -#include -#include -#include -#include -#include -#include -#include -#include -#include -#include -#include - -#include "../utils.h" - -namespace tvm { -namespace relay { -namespace backend { -namespace aot { - -/*! - * \brief Calculate FunctionInfo for all the PrimFuncs in a module. - */ -Map CalculateFunctionInfos(const IRModule& mod, - Integer workspace_byte_alignment, - Integer constant_byte_alignment) { - Map function_metadata; - for (const auto& kv : mod->functions) { - GlobalVar global_var = kv.first; - BaseFunc base_func = kv.second; - if (base_func->IsInstance()) { - tir::PrimFunc pfunc = Downcast(base_func); - Optional tgt_opt = pfunc->GetAttr(tvm::attr::kTarget); - ICHECK(tgt_opt) << "Target must be defined for all primfuncs."; - Target tgt = tgt_opt.value(); - // Determine the size of input/output buffers - auto params = pfunc->params; - int64_t total_io_bytes = 0; - for (const auto& param : params) { - if (pfunc->buffer_map.find(param) != pfunc->buffer_map.end()) { - auto buffer = pfunc->buffer_map[param]; - total_io_bytes += GetMemorySizeBytes(buffer->shape, buffer->dtype); - } - } - const auto& ws = CalculateWorkspaceBytes(pfunc, workspace_byte_alignment); - const auto& cs = CalculateConstantBytes(pfunc, constant_byte_alignment); - backend::FunctionInfo finfo{ - {{tgt, ws}}, {{tgt, total_io_bytes}}, {{tgt, cs}}, {{tgt, pfunc}}, {}}; - function_metadata.Set(global_var->name_hint, finfo); - } - } - return function_metadata; -} - -Map CreateFunctionMetadata(const IRModule& mod, - Integer workspace_byte_alignment, - Integer constant_byte_alignment) { - // First calculate the FunctionInfos from the buffers that are explicitly allocated - auto function_metadata = - CalculateFunctionInfos(mod, workspace_byte_alignment, constant_byte_alignment); - // Now adjust the FunctionInfo for the main func to also include PoolInfo allocations - // made by the USMP. - Optional> allocated_pool_infos = - mod->GetAttr>(tvm::attr::kPoolArgs); - backend::FunctionInfo main_func_info = - function_metadata.Get(runtime::symbol::tvm_module_main).value(); - if (allocated_pool_infos) { - for (const tir::usmp::AllocatedPoolInfo& allocated_pool_info : allocated_pool_infos.value()) { - for (const auto& tgt : allocated_pool_info->pool_info->targets) { - VLOG(1) << "USMP requires target " << tgt->ToDebugString() << " to have pool size " - << allocated_pool_info->allocated_size->value; - size_t size = allocated_pool_info->allocated_size->value; - if (allocated_pool_info->pool_info->IsInstance()) { - size += main_func_info->constant_sizes.count(tgt) - ? main_func_info->constant_sizes[tgt]->value - : 0; - main_func_info->constant_sizes.Set(tgt, size); - } else if (allocated_pool_info->pool_info->IsInstance()) { - size += main_func_info->workspace_sizes.count(tgt) - ? main_func_info->workspace_sizes[tgt]->value - : 0; - main_func_info->workspace_sizes.Set(tgt, size); - } else { - LOG(FATAL) << "Unknown pool type: " << allocated_pool_info->pool_info->GetTypeKey(); - } - } - } - } - function_metadata.Set(runtime::symbol::tvm_module_main, main_func_info); - return function_metadata; -} - -TVM_REGISTER_GLOBAL("relay.backend.aot.CreateFunctionMetadata") - .set_body_typed(CreateFunctionMetadata); - -} // namespace aot -} // namespace backend -} // namespace relay -} // namespace tvm diff --git a/src/relay/backend/aot/create_function_metadata.h b/src/relay/backend/aot/create_function_metadata.h deleted file mode 100644 index 8c7bf8753496..000000000000 --- a/src/relay/backend/aot/create_function_metadata.h +++ /dev/null @@ -1,49 +0,0 @@ -/* - * 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. - */ -#ifndef TVM_RELAY_BACKEND_AOT_CREATE_FUNCTION_METADATA_H_ -#define TVM_RELAY_BACKEND_AOT_CREATE_FUNCTION_METADATA_H_ - -#include -#include -#include - -#include "../utils.h" - -namespace tvm { -namespace relay { -namespace backend { -namespace aot { - -/*! \brief Create FunctionInfo metadata for all the PrimFuncs in a module lowered - * for AOT execution. - * \param mod The module. - * \param workspace_byte_alignment The alignment of the workspace pool. - * \param constant_byte_alignment The alignment of the constant pool. - * \return A map between function names and FunctionInfos. - */ -Map CreateFunctionMetadata(const IRModule& mod, - Integer workspace_byte_alignment, - Integer constant_byte_alignment); - -} // namespace aot -} // namespace backend -} // namespace relay -} // namespace tvm - -#endif // TVM_RELAY_BACKEND_AOT_CREATE_FUNCTION_METADATA_H_ diff --git a/src/relay/backend/aot_executor_codegen.cc b/src/relay/backend/aot_executor_codegen.cc deleted file mode 100644 index f698c654d6d8..000000000000 --- a/src/relay/backend/aot_executor_codegen.cc +++ /dev/null @@ -1,1455 +0,0 @@ -/* - * 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 src/relay/backend/aot_executor_codegen.cc - * \brief AOT executor codegen - */ - -#include -#include -#include -#include -#include -#include -#include -#include -#include -#include -#include -#include -#include -#include -#include -#include - -#include -#include -#include -#include - -#include "../../target/source/codegen_source_base.h" -#include "../../tir/transforms/ir_utils.h" -#include "../op/annotation/annotation.h" -#include "../op/call/call.h" -#include "../op/memory/device_copy.h" -#include "../transforms/device_aware_visitors.h" -#include "./name_transforms.h" -#include "./te_compiler.h" -#include "./utils.h" - -namespace tvm { -namespace relay { -namespace backend { - -using StorageMap = - std::unordered_map; - -/** - * This is an on demand allocator for AOT. A new temporary - * (storage allocator identifier) is allocated for each operation. - */ -class AOTOnDemandAllocator : public transform::DeviceAwareExprVisitor { - public: - AOTOnDemandAllocator() : transform::DeviceAwareExprVisitor(Optional()) {} - - // run the visitor on a global function. - void Run(const Function& func) { VisitExpr(func); } - - std::vector GetReturnIds() const { return return_ids_; } - std::vector GetReturnTtypes() const { return return_ttypes_; } - - StorageMap GetStorageMap() const { return storage_device_map_; } - - using ExprVisitor::VisitExpr_; - - void VisitExpr_(const ConstantNode* op) final { - CreateStorage(op); - AssignReturnSid(GetRef(op)); - } - - void DeviceAwareVisitExpr_(const CallNode* call_node) final { - // AOTOnDemandAllocator is run both before and after lowering, so we need to handle the case - // where the op of the call is a generic function - - Expr func; - Array args; - - CallLoweredProps call_lowered_props = GetCallLoweredProps(call_node); - if (call_lowered_props.lowered_func.defined()) { - func = call_lowered_props.lowered_func; - args = call_lowered_props.arguments; - } else { // Relay functions that have not been lowered and lowered extern functions - func = call_node->op; - args = call_node->args; - if (call_node->op.as()) { // Lowered extern function - ICHECK(!(call_node->attrs.defined())) << "Extern functions should have null attributes."; - } else { // Relay function which has not been lowered yet - ICHECK(call_node->op.as()) - << "Expected the call to be to a lowered primfunc, a lowered extern function or a " - "unlowered Relay function."; - } - } - VisitExpr(func); - CreateStorage(call_node); - for (const Expr& arg : args) { - VisitExpr(arg); - } - AssignReturnSid(GetRef(call_node)); - } - - void VisitExpr_(const VarNode* op) final { AssignReturnSid(GetRef(op)); } - - void DeviceAwareVisitExpr_(const FunctionNode* func_node) final { - if (function_nesting() > 1) { - // do not recurse into sub functions. - return; - } - if (func_node->HasNonzeroAttr(attr::kPrimitive)) { - // No storage needed for primitive functions. - return; - } - for (const auto& param : func_node->params) { - CreateStorage(param.get()); - } - VisitExpr(func_node->body); - } - - void VisitExpr_(const GlobalVarNode* op) final { - // Do nothing. - } - - void VisitExpr_(const OpNode* op) final { - // Do nothing. - } - - void VisitExpr_(const TupleNode* op) final { - std::vector storage_ids; - std::vector virtual_devices; - std::vector storage_sizes_in_bytes; - Expr expr = GetRef(op); - for (Expr field : op->fields) { - auto sid = GetStorage(field); - storage_ids.insert(storage_ids.end(), sid->storage_ids.begin(), sid->storage_ids.end()); - virtual_devices.insert(virtual_devices.end(), sid->virtual_devices.begin(), - sid->virtual_devices.end()); - storage_sizes_in_bytes.insert(storage_sizes_in_bytes.end(), - sid->storage_sizes_in_bytes.begin(), - sid->storage_sizes_in_bytes.end()); - } - storage_device_map_[expr] = StorageInfo(storage_ids, virtual_devices, storage_sizes_in_bytes); - AssignReturnSid(expr); - } - - void VisitExpr_(const TupleGetItemNode* op) final { - Expr expr = GetRef(op); - auto sids = GetStorage(op->tuple); - ICHECK_LT(static_cast(op->index), sids->storage_ids.size()); - storage_device_map_[expr] = - StorageInfo({sids->storage_ids[op->index]}, {sids->virtual_devices[op->index]}, - {sids->storage_sizes_in_bytes[op->index]}); - AssignReturnSid(expr); - } - - void VisitExpr_(const IfNode* op) final { LOG(FATAL) << "if is not supported."; } - - void PreVisitLetBinding_(const Var& var, const Expr& value) final { - VisitExpr(value); - StorageInfo si = GetStorage(value); - storage_device_map_[var] = si; - } - - private: - void AssignReturnSid(Expr e) { - if (storage_device_map_.find(e) != storage_device_map_.end()) { - StorageInfo& sinfo = storage_device_map_[e]; - return_ids_.clear(); - for (auto sid : sinfo->storage_ids) { - return_ids_.push_back(sid); - } - return_ttypes_.clear(); - return_ttypes_ = FlattenTupleType(e->checked_type()); - } - } - /*! - * \brief ceil(size/word_size) to get number of words. - * \param size The original size. - * \param word_size The element size. - */ - static size_t DivRoundUp(size_t size, size_t word_size) { - return (size + word_size - 1) / word_size; - } - /*! - * \brief Get the memory requirement. - * \param prototype The prototype token. - * \return The required memory size. - * - * TODO(mbs): Cf CalculateRelayExprSizeBytes in utils.cc, GetMemorySize is graph_plan_memory.cc - */ - size_t GetMemorySizeBytes(const TensorType& ttype) { - size_t size = 1; - for (IndexExpr dim : ttype->shape) { - const int64_t* pval = tir::as_const_int(dim); - ICHECK(pval != nullptr) << "Cannot allocate memory symbolic tensor shape " << ttype->shape; - ICHECK_GE(*pval, 0) << "Cannot allocate memory for tensor with negative shape" << *pval; - size *= static_cast(pval[0]); - } - size *= DivRoundUp(ttype->dtype.bits() * ttype->dtype.lanes(), 8); - return size; - } - /*! - * \brief Get the necessary storage for the expression. - * \param expr The expression. - * \return The corresponding token. - */ - StorageInfo GetStorage(const Expr& expr) { - // See through "on_device" calls. - Expr true_expr = IgnoreOnDevice(expr); - VisitExpr(true_expr); - auto it = storage_device_map_.find(true_expr); - ICHECK(it != storage_device_map_.end()) << "Could not find " << true_expr->GetTypeKey() << " " - << PrettyPrint(true_expr) << " in storage device map"; - return it->second; - } - - /*! - * \brief Create storage for the expression. - */ - void CreateStorage(const ExprNode* op) { - Expr expr = GetRef(op); - return CreateStorage(expr, GetVirtualDevice(expr)); - } - - /*! - * \brief Create storage to hold the result of evaluating \p expr in \p virtual_device. - */ - void CreateStorage(const Expr& expr, const VirtualDevice& virtual_device) { - ICHECK(!virtual_device->IsFullyUnconstrained()) - << "invalid virtual device for expr:" << std::endl - << PrettyPrint(expr); - std::vector storage_ids; - std::vector virtual_devices; - std::vector storage_sizes_in_bytes; - for (const auto& ttype : FlattenTupleType(expr->checked_type())) { - storage_ids.push_back(next_available_sid_++); - virtual_devices.push_back(virtual_device); - storage_sizes_in_bytes.push_back(GetMemorySizeBytes(ttype)); - } - storage_device_map_[expr] = StorageInfo(std::move(storage_ids), std::move(virtual_devices), - std::move(storage_sizes_in_bytes)); - } - - /*! \brief mapping of expression -> storageInfo */ - StorageMap storage_device_map_; - /*! \brief current id of the temporary allocated */ - int next_available_sid_{0}; - /*! \brief the set of intermediate tensors that are return variables */ - std::vector return_ids_; - /*! \brief the data types of the return values */ - std::vector return_ttypes_; -}; - -/*! \brief Code generator for AOT executor */ -class AOTExecutorCodegen : public MixedModeVisitor { - protected: - /*! \brief Describes the type of kernel call emitted. */ - enum CallType { - /*! - * \brief Emit PackedFunc calls bound just-in-time using TVMBackend* functions. - * - * When this type is selected, assumes all operators must be called via TVMFuncCall. Given the - * implementation of TVMFuncCall in the C++ runtime, this in practice implies that those - * functions are of type TVMBackendPackedCFunc. - * - * The following code is emitted at call sites to call a function named `func`: - * void* func_ptr = TVMBackendGetFuncFromEnv("func"); - * TVMFuncCall(func_ptr, values, tcodes, num_args, ret_values, ret_tcodes) - * - * The arguments given to the tir::Call node are encoded into `values`, `tcodes`, and `num_args` - * by LowerTVMBuiltin TIR transform. - * - * If `resource_handle` is passed to `func`, it is determined by TVMFuncCall (often, - * `resource_handle` is registered with the C++ runtime to provide a `this` equivalent when - * `func` is implemented in C). - * - * Compatible with both C++ and C runtimes, implemented with the C runtime only. - */ - kPacked, // Emit tir.call_packed and wrap all arguments in DLTensor. - - /*! - * \brief Directly call a TVMBackendPackedCFunc named according to the tir::Call. - * - * When this type is selected, assumes all operators are implemented in functions of type - * `TVMBackendPackedCFunc` and should be called directly. That is, presumes at the time of - * downstream compilation that there is a symbol named after the 0th arg to tir::Call of - * type `TVMBackendPackedCFunc`. This situation should occur when target_host == target. - * - * The following code is emitted at call sites to call a function named `func`: - * func(values, tcodes, num_args, ret_values, ret_tcodes, resource_handle) - * - * The arguments given to the tir::Call node are encoded into `values`, `tcodes`, and `num_args` - * by LowerTVMBuiltin TIR transform. - * - * `resource_handle` is encoded as the final argument to the tir::Call node. In practice, it is - * always the device context parameter when not null. At present, the implementation does not - * support forwarding device context parameters to CPacked. - * - * Compatible with the C runtime and C++ runtime (so long as target_host == target). Implemented - * in the same scenarios. - */ - kCPacked, // Emit tir.call_cpacked and wrap all arguments in DLTensor. - - /*! \brief Directly call a function accepting the `data` arrays as args. - * - * When this type is selected, assumes all operaotrs are implemented in C functions whose - * arguments are 1-to-1 with those in the tir::Call. DLTensor arguments are encoded as just the - * `data` parameters (i.e. no DLTensor object is passed along). - * - * The following code is emitted at call sites to a function named `func`: - * func(void* arg0, void* arg1, ..., void* argN) // no resource_handle - * -or- - * func(void* arg0, void* arg1, ..., void* argN, void* resource_handle) // with resource_handle - * - * `resource_handle` is encoded as the final argument to the tir::Call node. In practice, it is - * always the device context parameter when not null. - * - * Compatible with the C runtime and C++ runtime (so long as target_host == target). Implemented - * with the C runtime only. - */ - kUnpacked, // Emit tir.call_extern passing only the `data` part of DLTensors. - }; - - /*! - * \brief Return a vector of variables that represents the sids for the given Relay Expr - */ - std::vector PackSid(Expr expr) { - std::vector buffer_vars; - - ICHECK(storage_device_map_.find(expr) != storage_device_map_.end()) - << "Storage map did not contain constant expr " << PrettyPrint(expr); - StorageInfo& sinfo = storage_device_map_[expr]; - - // Note that an expression can have multiple sids associated with it - // e.g., returning multiple values from a function - for (auto sid : sinfo->storage_ids) { - // Determine if an sid is an output buffer - auto output_iter = std::find(return_sid_.begin(), return_sid_.end(), sid); - if (output_iter != return_sid_.end()) { - int output_index = std::distance(return_sid_.begin(), output_iter); - buffer_vars.push_back(GetBufferVarForIO(input_vars_.size() + output_index)); - continue; - } - - auto sid_value = sids_table_[sid]; - buffer_vars.push_back(sid_value); - } - return buffer_vars; - } - - /*! - * brief Given an expression return the variable(s) associated with that expression - */ - std::vector FindExpr(Expr arg) { - auto input_iter = std::find(input_vars_.begin(), input_vars_.end(), arg); - if (input_iter != input_vars_.end()) { - // Input variable - int main_index = std::distance(input_vars_.begin(), input_iter); - return {GetBufferVarForIO(main_index)}; - } else { - // Storage identifier (i.e., intermediate memory) - return PackSid(arg); - } - } - - /*! - * \brief Reverse lookup the device name in devices_ map. - * \param device_context Value in devices_ to find. - * \return Key matching device_context in devices_. - */ - std::string FindDeviceName(tir::Var device_context) { - for (std::pair kv : devices_) { - if (kv.second->name_hint == device_context->name_hint) { - return kv.first; - } - } - ICHECK(false) << "Did not find a device name associated with " << device_context; - return ""; - } - - void PushArgs(const Expr& expr, const std::vector& sids, Array* args) { - const TupleNode* t = expr.as(); - if (t != nullptr) { - CHECK_EQ(sids.size(), t->fields.size()) << "Relay tuple does not map 1:1 into TIR; AOT can't " - "handle this type of Relay Expr in a CallNode."; - } - - args->insert(args->end(), sids.begin(), sids.end()); - } - - /* - * Wraps a call_extern with a tvm_check_return annotation if required otherwise - * returns the passed Call - */ - tir::Call AddCheckReturn(tir::Call existing_call) { - Array args = {tir::make_const(DataType::Int(32, 1), 0, Span()), - tir::make_const(DataType::Int(32, 1), -1, Span()), existing_call}; - return tir::Call(DataType::Int(32), tir::builtin::tvm_check_return(), args); - } - - /*! - * brief Create a function call - * \param call_lowered_props The lowered function and the arguments to call it with - * \param result_expr The call we got func and args from (so as to recover the storage - * ids to hold the result). - */ - void CreateFuncCall(CallLoweredProps call_lowered_props, const Expr& result_expr) { - std::string func_name = call_lowered_props.lowered_func->name_hint; - tvm::Array args{tvm::tir::StringImm(func_name)}; - std::vector create_func_call_stmts; - - // Pack the inputs - for (const Expr& arg : call_lowered_props.arguments) { - if (params_by_expr_.find(arg) != params_by_expr_.end()) { - auto param_handle = tvm::tir::Call(DataType::Handle(), tvm::tir::builtin::lookup_param(), - {tir::StringImm(params_by_expr_[arg])}); - // NOTE: this cast looks like a no-op, but is required for compilation downstream. - // Because DataType::Handle has default bits=64, but CodeGenC does not observe this field, - // adding this cast forces the codegen to insert the cast. In this case, a cast is required - // because param_handle is actually code-generated as `const void*`, and the `const` piece - // needs to be removed. - args.push_back(tvm::tir::Cast(DataType::Handle(32, 1), param_handle)); - } else { - auto sids = FindExpr(arg); - PushArgs(arg, sids, &args); - } - } - - // Pack the return(s) value. A call node can produce multiple outputs - auto result_expr_sid = PackSid(result_expr); - PushArgs(result_expr, result_expr_sid, &args); - - GlobalVar global_var = call_lowered_props.lowered_func; - bool has_c_device_api_context = device_contexts_.count(global_var) != 0; - tir::Var device_context; - tir::Stmt func_call; - - switch (call_type_) { - case CallType::kUnpacked: { - // call_extern calling convention with optional context - if (has_c_device_api_context) { - device_context = device_contexts_.Get(global_var).value(); - - // call_extern has no further legalization steps, and - // requires the number of arguments to match exactly. For - // internal calls, conditionally append the device context. - bool requires_device_context = [&]() -> bool { - Optional opt = num_arguments_.Get(global_var); - if (!opt.defined()) { - // For external calls, we must trust that the user has - // supplied a kernel that accepts a device_context - // argument. - return true; - } - int num_callee_params = opt.value()->value; - int num_args = call_lowered_props.arguments.size(); - if (num_callee_params == num_args) { - return false; - } else if (num_callee_params == num_args + 1) { - return true; - } else { - LOG(FATAL) << "Callee " << global_var << " requires " << num_callee_params - << ", but is called with " << num_args << " arguments."; - } - }(); - if (requires_device_context) { - args.push_back(device_context); - } - } - func_call = tir::Evaluate(AddCheckReturn( - tvm::tir::Call(DataType::Int(32), tvm::tir::builtin::call_extern(), args))); - break; - } - case CallType::kCPacked: { - if (has_c_device_api_context) { - device_context = device_contexts_.Get(global_var).value(); - args.push_back(device_context); - } else { - // NOTE: LowerTVMBuiltin expects some device_context placeholder. - args.push_back(tir::make_zero(DataType::Handle())); - } - func_call = tir::Evaluate( - tvm::tir::Call(DataType::Int(32), tvm::tir::builtin::tvm_call_cpacked(), args)); - create_func_call_stmts.push_back(func_call); - break; - } - case CallType::kPacked: { - // call_packed does not accept a device context. - CHECK(!has_c_device_api_context) << "CallType::kPacked does not accept a device context"; - func_call = tir::Evaluate(AddCheckReturn( - tvm::tir::Call(DataType::Int(32), tvm::tir::builtin::tvm_call_packed(), args))); - create_func_call_stmts.push_back(func_call); - break; - } - default: - ICHECK(false) << "Unknown CallType: " << call_type_; - } - - ICHECK(func_call.defined()) << "Must define func_call"; - - if (has_c_device_api_context) { - func_call = tir::SeqStmt(Array({ - GenerateDeviceHook(device_context, "Open"), - func_call, - GenerateDeviceHook(device_context, "Close"), - })); - } - - tir::Stmt body = tir::SeqStmt::Flatten(func_call); - stmts_.push_back(body); - } - - /*! - * \brief Copy a variable to the output. This function is mainly used in edge cases - * when we want to return an input or a parameter. - * TODO(giuseros): we should try to avoid unnecessary copy to the output, e.g., in a - * copy-on-write fashion. - */ - void CopyToOutput(PrimExpr out, PrimExpr in, bool pack_input, size_t size) { - std::vector let_nest; - - // Define intermediate DLTensor to load/store the data - tir::Buffer tmp_read = - tir::decl_buffer({IntImm(DataType::UInt(64), size)}, DataType::UInt(8), "tmp_read"); - tir::Buffer tmp_write = - tir::decl_buffer({IntImm(DataType::UInt(64), size)}, DataType::UInt(8), "tmp_write"); - - // Re-use in/out as the buffer var, if possible - if (auto opt = out.as()) { - tmp_write.CopyOnWrite()->data = opt.value(); - } else { - let_nest.push_back(tir::LetStmt(tmp_write->data, out, tir::Evaluate(0))); - } - if (auto opt = in.as()) { - tmp_read.CopyOnWrite()->data = opt.value(); - } else { - let_nest.push_back(tir::LetStmt(tmp_read->data, in, tir::Evaluate(0))); - } - - // Copy the variable from the input to the output - te::Var loop_idx("i", DataType::Int(32)); - tir::Stmt copy = tir::BufferStore(tmp_write, tir::BufferLoad(tmp_read, {loop_idx}), {loop_idx}); - copy = tir::For(loop_idx, 0, tir::make_const(DataType::Int(32, 1), size, Span()), - tir::ForKind::kSerial, copy); - copy = tir::MergeNest(let_nest, copy); - - stmts_.push_back(copy); - } - - /* - * \brief Collects device context variables for passing to operators - */ - void CollectDeviceVariables(const Map& device_contexts) { - Map target_contexts; - TargetKindAttrMap target_attr_map = tvm::TargetKind::GetAttrMap("use_device_api"); - - for (const auto& it : device_contexts) { - const GlobalVar& global_var = it.first; - const std::string device_context_name = it.second; - - Optional target_kind = tvm::TargetKind::Get(device_context_name); - if (!target_kind || !target_attr_map.count(target_kind.value())) { - return; - } - if (target_attr_map[target_kind.value()]) { - std::string context_name = tvm::runtime::SanitizeName(device_context_name); - tir::Var device_context_var("device_context_" + context_name, DataType::Handle()); - - auto pair = target_contexts.find(target_kind.value()); - if (pair != target_contexts.end()) { - device_context_var = (*pair).second; - } else { - main_signature_.push_back(device_context_var); - devices_.Set(context_name, device_context_var); - target_contexts.Set(target_kind.value(), device_context_var); - } - - device_contexts_.Set(global_var, device_context_var); - } - } - } - - /** - * \brief Generates a call to a given hook for all Devices found for C Device API - * \param Name of hook to generate statements for - * \return Statement with function calls for each device - */ - tir::Stmt GenerateAllDeviceHook(const String& hook) { - std::vector device_hooks; - for (const auto& it : devices_) { - const String& device_name = it.first; - const tir::Var& context = it.second; - Array sections = {"Device", device_name, hook}; - String device_hook_name = ToCFunctionStyle(PrefixName(sections)); - - tir::Evaluate device_hook( - AddCheckReturn(tvm::tir::Call(DataType::Int(32), tvm::tir::builtin::call_extern(), - {tvm::tir::StringImm(device_hook_name), context}))); - device_hooks.push_back(device_hook); - } - return tir::SeqStmt::Flatten(device_hooks); - } - - /** - * \brief Generates a call to a given hook for a single Device function - * \param Var Device context to call hook on - * \param Name of hook to generate statements for - * \return Statement with function call to Device API - */ - tir::Stmt GenerateDeviceHook(const tir::Var& context, const String& hook) { - const auto& it = std::find_if(std::begin(devices_), std::end(devices_), [&](const auto& it) { - return it.second->name_hint == context->name_hint; - }); - const String& device_name = (*it).first; - Array sections = {"Device", device_name, hook}; - String device_hook = ToCFunctionStyle(PrefixName(sections)); - - return tir::Evaluate( - AddCheckReturn(tir::Call(DataType::Int(32), tvm::tir::builtin::call_extern(), - {tvm::tir::StringImm(device_hook), context}))); - } - - /*! - * Utility function to string together different arguments - */ - template - std::string MakeString(Args const&... args) { - std::ostringstream ss; - using List = int[]; - (void)List{0, ((void)(ss << args), 0)...}; - - return ss.str(); - } - - void VisitExpr_(const CallNode* call_node) override { - OnDeviceProps on_device_props = GetOnDeviceProps(call_node); - if (on_device_props.body.defined()) { - VisitExpr(on_device_props.body); - return; - } - - DeviceCopyProps device_copy_props = GetDeviceCopyProps(call_node); - CallLoweredProps call_lowered_props = GetCallLoweredProps(call_node); - - if (device_copy_props.body.defined()) { - // TODO(mbs): device_copy cleaunp - // Suspect treating as no-op is better since already built into the StorageInfo? - LOG(FATAL) << "The AOT executor does not currently support device_copy"; - } - - // At this point we should only see calls of the form call_lowered(@callee, (args...)), - // where @callee can be a PrimFunc we've compiled or an external function supplied via - // some other mechanism. - ICHECK(call_lowered_props.lowered_func.defined()) - << "AOT does not support calling Relay functions. Attempting to call:" << std::endl - << PrettyPrint(GetRef(call_node)); - for (const auto& arg : call_lowered_props.arguments) { - // Evaluate the args - VisitExpr(arg); - } - CreateFuncCall(call_lowered_props, GetRef(call_node)); - } - - void VisitExpr_(const VarNode* op) override { - Expr expr = GetRef(op); - StorageInfo& sinfo = storage_device_map_[expr]; - - // Let bound vars refer to a value, so these should not be considered "output" vars. - if (let_bound_vars_.find(GetRef(op)) != let_bound_vars_.end()) { - return; - } - - // If the Var node is an output node we need to copy the content of the variable to the output - // It's safe to check the SID here because Var StorageToken are never reallocated - auto output_iter = std::find(return_sid_.begin(), return_sid_.end(), sinfo->storage_ids[0]); - if (output_iter != return_sid_.end()) { - int output_index = std::distance(return_sid_.begin(), output_iter); - if (params_by_expr_.find(expr) != params_by_expr_.end()) { - auto param_handle = tvm::tir::Call(DataType::Handle(), tvm::tir::builtin::lookup_param(), - {tir::StringImm(params_by_expr_[expr])}); - CopyToOutput(GetBufferVarForIO(input_vars_.size() + output_index), param_handle, - /*pack_input*/ false, sinfo->storage_sizes_in_bytes[0]); - } else { - auto var_expr = FindExpr(expr); - CopyToOutput(GetBufferVarForIO(input_vars_.size() + output_index), var_expr[0], - /*pack_input*/ false, sinfo->storage_sizes_in_bytes[0]); - } - } - } - - void VisitExpr_(const ConstantNode* op) override { - Expr expr = GetRef(op); - ICHECK(storage_device_map_.find(expr) != storage_device_map_.end()) - << "Storage map did not contain constant expr " << PrettyPrint(expr); - StorageInfo& sinfo = storage_device_map_[expr]; - std::stringstream ss; - ss << "constant_" << constant_map_.size(); - - tir::Var constant(ss.str(), PointerType(PrimType(DataType(op->data->dtype)))); - constant_map_[constant] = op; - auto sid = sinfo->storage_ids[0]; - sids_table_[sid] = constant; - - // If the Constant node is an output node we need to copy the content of the parameter to the - // output. A node can only produce a single output - auto output_iter = std::find(return_sid_.begin(), return_sid_.end(), sid); - if (output_iter != return_sid_.end()) { - int output_index = std::distance(return_sid_.begin(), output_iter); - auto param_handle = tvm::tir::Call(DataType::Handle(), tvm::tir::builtin::lookup_param(), - {tir::StringImm(ss.str())}); - CopyToOutput(GetBufferVarForIO(input_vars_.size() + output_index), constant, - /* pack_input */ false, sinfo->storage_sizes_in_bytes[0]); - } - } - - void VisitExpr_(const TupleNode* op) override { - for (auto field : op->fields) { - VisitExpr(field); - } - } - - void VisitExpr_(const LetNode* op) override { - auto pre_visit = [this](const LetNode* op) { - let_bound_vars_.insert(op->var); - this->VisitExpr(op->value); - }; - auto post_visit = [this](const LetNode* op) { - this->VisitExpr(op->body); - this->visit_counter_[op] += 1; - }; - ExpandANormalForm(op, pre_visit, post_visit); - } - - void VisitExpr_(const TupleGetItemNode* op) override { VisitExpr(op->tuple); } - void VisitExpr_(const OpNode* op) override { - if (GetRef(op) != CallLoweredOp() && GetRef(op) != OnDeviceOp()) { - LOG(FATAL) << "All OpNodes except for call_lowered should have been expanded"; - } - } - void VisitExpr_(const IfNode* op) override { - LOG(FATAL) << "All GlobalVarNodes should be removed before AOT executor's Codegen is called"; - } - void VisitExpr_(const FunctionNode* op) override { - ICHECK(op->GetAttr(attr::kCompiler).defined()) - << "FunctionNode only supported by custom codegen"; - } - void VisitExpr_(const RefCreateNode* op) override { - LOG(FATAL) << "AOT executor does not support references (found RefCreateNode)"; - } - void VisitExpr_(const RefReadNode* op) override { - LOG(FATAL) << "AOT executor does not support references (found RefReadNode)"; - } - void VisitExpr_(const RefWriteNode* op) override { - LOG(FATAL) << "AOT executor does not support references (found RefWriteNode)"; - } - void VisitExpr_(const ConstructorNode* op) override { - LOG(FATAL) << "AOT executor does not support ADTs (found ConstructorNode)"; - } - void VisitExpr_(const MatchNode* op) override { - LOG(FATAL) << "AOT executor does not support matching (found MatchNode)"; - } - - // Create the main PrimFunc to execute the graph. Please note that - // the packed function calls don't pack their arguments. The AOT - // runner function needs to be legalized by the LegalizePackedCalls pass. - tir::PrimFunc CreateMainFunc(String mod_name, unsigned int relay_params) { - tir::Stmt body = tir::SeqStmt::Flatten(stmts_); - // Allocate the sids - std::unordered_map allocated; - - for (auto kv : storage_device_map_) { - // Only allocate sids that are needed - const bool is_input = - (std::find(input_vars_.begin(), input_vars_.end(), kv.first) != input_vars_.end()); - const bool is_param = (params_by_expr_.find(kv.first) != params_by_expr_.end()); - if (is_input || is_param) { - continue; - } - - for (unsigned int i = 0; i < kv.second->storage_ids.size(); i++) { - int size = kv.second->storage_sizes_in_bytes[i]; - int sid = kv.second->storage_ids[i]; - - if (std::find(return_sid_.begin(), return_sid_.end(), sid) != return_sid_.end()) { - continue; - } - - // Make sure it hasn't already been allocated, this can happen - // with let-bound var/value pairs. - if (allocated.find(sid) != allocated.end()) { - continue; - } - - allocated[sid] = constant_map_.count(sids_table_[sid]); - - // TODO(giuseros): we should allocate this once outside the PrimFunc - // so we don't pay the price of allocation for every inference - if (!allocated[sid]) { - PointerType ptype = Downcast(sids_table_[sid]->type_annotation); - DataType element_type = Downcast(ptype->element_type)->dtype; - body = tir::Allocate(sids_table_[sid], element_type, {size}, tir::const_true(), body); - } - allocated[sid] = true; - } - } - - for (auto kv : constant_map_) { - auto buffer_var = kv.first; - auto dtype = DataType(kv.second->data->dtype); - - int ndim = kv.second->data->ndim; - Array extents; - - for (int i = 0; i < ndim; i++) { - int shape = kv.second->data->shape[i]; - extents.push_back(tir::make_const(DataType::Int(32), shape, Span())); - } - body = tir::AllocateConst(buffer_var, dtype, extents, kv.second->data, body); - } - - // Define the PrimFunc attributes - Map dict_attrs; - String run_func_name = runtime::get_name_mangled(mod_name, runtime::symbol::tvm_module_main); - dict_attrs.Set("global_symbol", run_func_name); - dict_attrs.Set("runner_function", Bool(true)); - dict_attrs.Set(tvm::attr::kTarget, config_->host_target); - - tir::Stmt device_activations = GenerateAllDeviceHook("Activate"); - tir::Stmt device_deactivations = GenerateAllDeviceHook("Deactivate"); - tir::Stmt final_body = tir::SeqStmt({device_activations, body, device_deactivations}); - - // Make the PrimFunc - return tir::PrimFunc(main_signature_, final_body, VoidType(), main_buffer_map_, - DictAttrs(dict_attrs)); - } - - /*! - * \brief Access IO vars using the buffer vars and - * not the actual var. - */ - tir::Var GetBufferVarForIO(int index) { return main_buffer_map_[main_signature_[index]]->data; } - - /*! - * \brief Create tir::Var for input/output while updating the buffer_maps. - * - * \param expr The expression to evaluate. - * \param original_name The name of the tir::Var. - * \param use_unique_name Whether to generate a new unique name where a name conflicts. - */ - void CreateIOVar(const Expr& expr, const std::string& original_name, - bool use_unique_name = true) { - CreateIOVar(expr->checked_type(), original_name, use_unique_name); - } - - /*! - * \brief Create tir::Var for input/output while updating the buffer_maps. - * - * \param expr The expression to evaluate. - * \param original_name The name of the tir::Var. - * \param use_unique_name Whether to generate a new unique name where a name conflicts. - */ - void CreateIOVar(const Type& type, const std::string& original_name, - bool use_unique_name = true) { - if (type->IsInstance()) { - TupleType tuple_type = Downcast(type); - for (unsigned i = 0; i < tuple_type->fields.size(); i++) { - CreateIOVar(tuple_type->fields[i], original_name); - } - } else { - std::string name = original_name; - if (use_unique_name) { - name = GetUniqueIOVarName(original_name); - } - tir::Var var = tir::Var(name, DataType::Handle()); - main_signature_.push_back(var); - auto tensor_type = type.as(); - ICHECK(tensor_type) << "Expected TensorType node but was " << type->GetTypeKey(); - DataType elem_type = tensor_type->dtype; - tir::Var buffer_var = - tir::Var(name + "_buffer_var", PointerType(PrimType(elem_type), "global")); - tir::Buffer buffer = tir::Buffer(buffer_var, elem_type, tensor_type->shape, {}, 0, - name + "_buffer", 16, 1, tir::BufferType::kDefault); - main_buffer_map_.Set(var, buffer); - io_tensor_types_.Set(var, Downcast(type)); - } - } - - /*! - * \brief Create a unique name for I/O Var - */ - std::string GetUniqueIOVarName(std::string name) { - if (io_var_names_.find(name) == io_var_names_.end()) { - io_var_names_[name] = 1; - return name; - } else { - io_var_names_[name] = io_var_names_[name] + 1; - return name + std::to_string(io_var_names_[name]); - } - } - - /*! - * \brief Calculate workspace sizes for PrimFuncs in the IRModule - */ - Map CalculateWorkspaceSizes( - const IRModule& lowered_mod, const Map& function_metadata) { - Integer workspace_byte_alignment = GetModuleWorkspaceByteAlignment(lowered_mod); - Map updated_function_metadata; - for (const auto& kv : lowered_mod->functions) { - GlobalVar global_var = kv.first; - BaseFunc base_func = kv.second; - if (base_func->IsInstance()) { - tir::PrimFunc pfunc = Downcast(base_func); - Target tgt = pfunc->GetAttr(tvm::attr::kTarget).value(); - const auto& ws = CalculateWorkspaceBytes(pfunc, workspace_byte_alignment); - if (function_metadata.count(global_var->name_hint)) { - updated_function_metadata.Set(global_var->name_hint, - function_metadata[global_var->name_hint]); - updated_function_metadata[global_var->name_hint]->workspace_sizes.Set(tgt, ws); - } else { - FunctionInfo finfo{{{tgt, ws}}, {}, {}, {{tgt, pfunc}}, {}}; - updated_function_metadata.Set(global_var->name_hint, finfo); - } - } - } - return updated_function_metadata; - } - - /*! - * \brief Run USMP to plan memory for lowered IRModule. - */ - IRModule PlanMemoryWithUSMP(const IRModule& mod) { - VLOG(1) << "Planning memory with USMP for module:" << std::endl << PrettyPrint(mod); - Integer workspace_byte_alignment = GetModuleWorkspaceByteAlignment(mod); - IRModule lowered_mod = mod->ShallowCopy(); - lowered_mod = tir::transform::UnifiedStaticMemoryPlanner()(lowered_mod); - function_metadata_ = CalculateWorkspaceSizes(lowered_mod, function_metadata_); - Optional> allocated_pool_infos = - lowered_mod->GetAttr>(tvm::attr::kPoolArgs); - backend::FunctionInfo main_func_info = - lowered_mod->GetAttr("main_func_info").value(); - main_func_info->workspace_sizes.clear(); - if (allocated_pool_infos) { - for (const tir::usmp::AllocatedPoolInfo& allocated_pool_info : allocated_pool_infos.value()) { - for (const auto& tgt : allocated_pool_info->pool_info->targets) { - VLOG(1) << "USMP requires target " << tgt->ToDebugString() << " to have pool size " - << allocated_pool_info->allocated_size->value; - size_t size = allocated_pool_info->allocated_size->value; - if (allocated_pool_info->pool_info->IsInstance()) { - size += main_func_info->constant_sizes.count(tgt) - ? main_func_info->constant_sizes[tgt]->value - : 0; - main_func_info->constant_sizes.Set(tgt, size); - } else if (allocated_pool_info->pool_info->IsInstance()) { - size += main_func_info->workspace_sizes.count(tgt) - ? main_func_info->workspace_sizes[tgt]->value - : 0; - main_func_info->workspace_sizes.Set(tgt, size); - } else { - LOG(FATAL) << "Unknown pool type: " << allocated_pool_info->pool_info->GetTypeKey(); - } - } - } - } - function_metadata_.Set(runtime::symbol::tvm_module_main, main_func_info); - return lowered_mod; - } - - /*! - * \brief Run StorageRewrite to plan memory for lowered IRModule. - */ - IRModule PlanMemoryWithStorageRewrite(const IRModule& mod) { - Integer workspace_byte_alignment = GetModuleWorkspaceByteAlignment(mod); - IRModule lowered_mod = mod->ShallowCopy(); - function_metadata_ = CalculateWorkspaceSizes(lowered_mod, function_metadata_); - // Running StorageRewrite just on the main function - tir::PrimFunc tir_main_func = - Downcast(lowered_mod->Lookup(::tvm::runtime::symbol::tvm_module_main)); - IRModule main_func_mod; - main_func_mod->Update(lowered_mod->GetGlobalVar(::tvm::runtime::symbol::tvm_module_main), - tir_main_func); - main_func_mod = tir::transform::StorageRewrite()(main_func_mod); - lowered_mod->Update(lowered_mod->GetGlobalVar(::tvm::runtime::symbol::tvm_module_main), - main_func_mod->Lookup(::tvm::runtime::symbol::tvm_module_main)); - tir_main_func = - Downcast(lowered_mod->Lookup(::tvm::runtime::symbol::tvm_module_main)); - // Use the PrimFunc to calculate the workspace required to service the allocates - Integer main_workspace_size_bytes = - CalculateWorkspaceBytes(tir_main_func, workspace_byte_alignment); - backend::FunctionInfo main_func_info = - lowered_mod->GetAttr("main_func_info").value(); - main_func_info->workspace_sizes.Set(config_->host_target, main_workspace_size_bytes); - function_metadata_.Set(runtime::symbol::tvm_module_main, main_func_info); - return lowered_mod; - } - - /*! - * \brief Gets module workspace alignment from supplied executor or defaults to 16 - */ - Integer GetModuleWorkspaceByteAlignment(const IRModule& mod) { - Executor executor_config = mod->GetAttr(tvm::attr::kExecutor).value(); - return executor_config->GetAttr("workspace-byte-alignment").value_or(16); - } - - /*! - * \brief Gets module constant alignment from supplied executor or defaults to 16 - */ - Integer GetModuleConstantByteAlignment(const IRModule& mod) { - Executor executor_config = mod->GetAttr(tvm::attr::kExecutor).value(); - return executor_config->GetAttr("constant-byte-alignment").value_or(16); - } - - protected: - /*! \brief mod */ - runtime::Module* mod_; - /*! \brief list of input expressions (i.e., variable passed by the user) */ - std::vector input_vars_; - /*! \brief map of device contexts variables */ - Map devices_; - /*! \brief map of GlobalVars to C Device API contexts */ - Map device_contexts_; - /*! \brief map of GlobalVars to the number of arguments they require */ - Map num_arguments_; - /*! \brief input and output variables belonging to the main function signature */ - Array main_signature_; - /*! \brief input and output variables belonging to the main function signature */ - Map main_buffer_map_; - /*! \brief maps input and output variables to TensorType which describe them */ - Map io_tensor_types_; - /*! \brief All available targets. */ - CompilationConfig config_; - /*! - * \brief The type of kernel call to be emitted. - * See CallType for more documentation. - */ - CallType call_type_; - - /*! - * \brief parameters (i.e. ConstantNodes found in the graph). - * These are take as inputs to the GraphRuntime. - * Maps param name to a pair of storage_id and NDArray. At runtime, the storage_id can be - * used to lookup the parameter. - */ - std::unordered_map params_; - /*! \brief mapping between expression and parameters */ - Map params_by_expr_; - /*! \brief mapping between parameter names ("p0", "p1", etc..) and storage identifiers*/ - std::unordered_map param_storage_ids_; - std::unordered_map - constant_map_; - - /*! \brief plan memory of device result */ - StorageMap storage_device_map_; - /*! \brief mapping sid -> tir::Var */ - std::unordered_map sids_table_; - /*! \brief lowered funcs */ - Map function_metadata_; - /*! \brief the set of statements that make the program */ - std::vector stmts_; - /*! \brief the list of return sids (note that the function might return more then one output */ - std::vector return_sid_; - /*! \brief This is per IO var name counter to aid the generating unique names */ - std::unordered_map io_var_names_; - /*! \brief A set of variables that are let bound. */ - std::unordered_set let_bound_vars_; - - public: - AOTExecutorCodegen(runtime::Module* mod, const Array& targets) - : mod_(mod), config_(transform::PassContext::Current(), targets) {} - - LoweredOutput Codegen(IRModule mod, relay::Function func, String mod_name) { - VLOG_CONTEXT << "AOT"; - - Runtime runtime_config = mod->GetAttr(tvm::attr::kRuntime).value(); - Integer workspace_byte_alignment = GetModuleWorkspaceByteAlignment(mod); - - Executor executor_config = mod->GetAttr(tvm::attr::kExecutor).value(); - std::string interface_api = - executor_config->GetAttr("interface-api").value_or("packed"); - bool unpacked_api = executor_config->GetAttr("unpacked-api").value_or(Bool(false)); - - // Validate choice of unpacked_api and use_call_cpacked_ - if (runtime_config->name == kTvmRuntimeCrt) { - if (unpacked_api == true) { - call_type_ = CallType::kUnpacked; - } else if (unpacked_api == false && interface_api == "packed") { - call_type_ = CallType::kCPacked; - } else { - CHECK(interface_api == "packed" || unpacked_api == true) - << "Either need interface_api == \"packed\" (got: " << interface_api - << ") or unpacked-api == true (got: " << unpacked_api << ") when targeting c runtime"; - ICHECK(false) << "Unhandled executor option config: interface-api=" << interface_api - << ", unpacked-api=" << unpacked_api; - } - } else if (runtime_config->name == kTvmRuntimeCpp) { - if (unpacked_api == false && interface_api == "packed") { - call_type_ = CallType::kCPacked; - } else { - CHECK(static_cast(unpacked_api) == false && interface_api == "packed") - << "Need unpacked-api == false (got: " << unpacked_api - << ") and interface-api == \"packed\" (got: " << interface_api - << ") when targeting c++ runtime"; - ICHECK(false) << "Unhandled executor option config: interface-api=" << interface_api - << ", unpacked-api=" << unpacked_api; - } - } else { - ICHECK(false) << "runtime_config (" << runtime_config->name - << ") is not one of the expected values"; - } - - mod = transform::ToANormalForm()(mod); - mod = transform::InferType()(mod); - mod = transform::AnnotateUsedMemory()(mod); - - IRModule lowered_mod = - tec::LowerTE(mod_name, config_, [this, workspace_byte_alignment](BaseFunc func) { - // We need to maintain the constant map for external - // functions so we pass this processing function which - // allows us to process each function as we lower it. - if (func->GetAttr(attr::kCompiler).defined()) { - UpdateConstants(func, ¶ms_); - } - - // TODO(@areusch, @jroesch): We should refactor this to - // execute as a further pass, instead writing data to the - // lowering process directly. - tec::UpdateFunctionMetadata(func, this->function_metadata_, workspace_byte_alignment); - })(mod); - - transform::PassContext pass_ctx = transform::PassContext::Current(); - bool enable_remove_reshapes = - pass_ctx->GetConfig("relay.remove_standalone_reshapes.enable", Bool(true)).value(); - if (enable_remove_reshapes) { - lowered_mod = transform::RemoveStandaloneReshapes()(lowered_mod); - } - auto lowered_main = lowered_mod->Lookup("main"); - auto lowered_main_func = Downcast(lowered_main); - - // Post-lowering storage map for writing main func - AOTOnDemandAllocator final_aot_allocator; - final_aot_allocator.Run(lowered_main_func); - storage_device_map_ = final_aot_allocator.GetStorageMap(); - - // TODO(@electriclilies, @jroesch, @Mousius): remove UpdateMainWorkspaceSize - StaticMemoryPlan memory_plan(storage_device_map_); - backend::FunctionInfo func_info = - tec::UpdateMainWorkspaceSize(lowered_mod, config_, memory_plan->expr_to_storage_info); - lowered_mod = WithAttr(lowered_mod, "main_func_info", func_info); - - for (auto input : lowered_main_func->params) { - input_vars_.push_back(input); - std::string input_name = SanitizeName(input->name_hint()); - // We dont want the compiler changing input names in the - // event of a sanitization collision. Therefore, enforcing - // the var created to use the input_name strictly. - CreateIOVar(input, input_name, /*use_unique_name = */ false); - } - - // Define the storage allocator ids - for (auto kv : storage_device_map_) { - for (auto sid : kv.second->storage_ids) { - // The buffer_var is created with storage_scope to be global.workspace to be serviced by - // TVMBackendAllocWorkspace(TVMBAW) calls, explicitly. The reasoning being the executor - // allocates should be serviced by TVMBAWs as the data could be accessed by many devices and - // should not be lowered to the stack. For more details please refer to the discussion here: - // https://github.com/apache/tvm/issues/9022 - te::Var buffer_var(MakeString("sid_", sid), - PointerType(PrimType(DataType::Int(8)), "global.workspace")); - sids_table_[sid] = buffer_var; - } - } - - // Retrieve the return sids - return_sid_ = final_aot_allocator.GetReturnIds(); - // Insert outputs to main func signature - // If output tensor names were provided use them - if (auto opt = func->GetAttr>("output_tensor_names")) { - Array output_tensor_names = opt.value(); - Expr output_expr = lowered_main_func->body; - if (output_expr->checked_type()->IsInstance()) { - TupleType output_tuple_type = Downcast(output_expr->checked_type()); - for (unsigned i = 0; i < output_tuple_type->fields.size(); i++) { - // AoT Executor Codegen does not create these names, - // thus should be used as they are provided. - CreateIOVar(output_tuple_type->fields[i], output_tensor_names[i], - /*use_unique_name = */ false); - } - } else { - // AoT Executor Codegen does not create these names, - // thus should be used as they are provided. - CreateIOVar(lowered_main_func->body, output_tensor_names[0], /*use_unique_name = */ false); - } - } else { - // If output tensor names are not provided we will generate output(x) - // where x is a counter to create unique names. - CreateIOVar(lowered_main_func->body, "output"); - } - - CollectDeviceVariables(lowered_mod->GetAttr>("device_contexts").value()); - num_arguments_ = [&]() -> Map { - Map arg_count; - for (const auto& [gvar, func] : lowered_mod->functions) { - if (const auto* prim_func = func.as()) { - arg_count.Set(gvar, prim_func->params.size()); - } - } - return arg_count; - }(); - VisitExpr(lowered_main_func->body); - - // Create the runner function. Please note that the function is not legal yet - // because the packed calls arguments are not wrapped in TVMValues. To make this happen we need - // to run the LegalizePackedCalls pass. - LoweredOutput ret; - - // Collect any constants extracted by external codegen. - ret.params = std::unordered_map(); - Map const_name_to_constant = - lowered_mod->GetAttr>(tvm::attr::kConstNameToConstant) - .value_or({}); - for (const auto& kv : const_name_to_constant) { - ICHECK(ret.params.emplace(kv.first, kv.second).second); - } - - // Collect any constants extracted during lowering. - for (const auto& kv : params_) { - ICHECK(ret.params.emplace(kv.first, kv.second).second); - } - - // AoT Executor codegen works completely on TIR beyond this point, hence removing relay main - // function and replacing it with its TIR version. We should try to make this a Pass. - lowered_mod->Remove(lowered_mod->GetGlobalVar("main")); - auto tir_main_func = CreateMainFunc(mod_name, lowered_main_func->params.size()); - // Extract additional information around main TIR PrimFunc arguments - Array devices = ListDevices(); - const auto main_func_params_end_iterator = - tir_main_func->params.begin() + tir_main_func->params.size(); - const auto outputs_begin_iterator = - main_func_params_end_iterator - return_sid_.size() - devices.size(); - Array inputs = Array(tir_main_func->params.begin(), outputs_begin_iterator); - Array input_tensor_types; - for (auto i : inputs) { - input_tensor_types.push_back(io_tensor_types_[i]); - } - Array outputs = - Array(outputs_begin_iterator, main_func_params_end_iterator - devices.size()); - - lowered_mod->Update(GlobalVar(::tvm::runtime::symbol::tvm_module_main), tir_main_func); - // Parallel for loops are not supported in AoT codegen. - lowered_mod = tir::transform::ConvertForLoopsToSerial()(lowered_mod); - - // Check USMP option - bool enable_usmp = false; - if (runtime_config->name == kTvmRuntimeCrt) { - enable_usmp = true; - } - if (pass_ctx->GetConfig(kUSMPEnableOption) != nullptr) { - enable_usmp = pass_ctx->GetConfig(kUSMPEnableOption, Bool(false)).value(); - } - - if (enable_usmp) { - lowered_mod = PlanMemoryWithUSMP(lowered_mod); - } else { - lowered_mod = PlanMemoryWithStorageRewrite(lowered_mod); - } - ret.function_metadata = std::move(function_metadata_); - - // Legalize AOT if needed. This means that all the packed calls - // need to be wrapped in TVMValues (unless unpacked_api is set) - if (call_type_ == CallType::kCPacked || call_type_ == CallType::kPacked) { - auto pack_calls = tir::transform::LegalizePackedCalls(); - lowered_mod = pack_calls(lowered_mod); - } - - // Collect any runtime modules generated by external codegen. - ret.external_mods = - lowered_mod->GetAttr>(tvm::attr::kExternalMods).value_or({}); - - // This is the point where we separate the functions in the module by target - VLOG(1) << "lowered module:" << std::endl << PrettyPrint(lowered_mod); - ret.lowered_funcs = tec::GetPerTargetModules(lowered_mod); - VLOG(1) << "per-target modules:"; - for (const auto& kv : ret.lowered_funcs) { - VLOG(1) << "target:" << std::endl - << kv.first->ToDebugString() << std::endl - << "maps to:" << std::endl - << PrettyPrint(kv.second); - } - - // Extract USMP metadata to pass onto metadata sources - Map pool_var_info; - std::vector pool_vars; - tir_main_func = - Downcast(lowered_mod->Lookup(::tvm::runtime::symbol::tvm_module_main)); - Optional> allocated_pool_infos = - tir_main_func->GetAttr>(tvm::attr::kPoolArgs); - if (allocated_pool_infos) { - for (const tir::usmp::AllocatedPoolInfo& allocated_pool_info : allocated_pool_infos.value()) { - int pool_var_index = allocated_pool_info->pool_var_idx.value()->value; - pool_vars.push_back(tir_main_func->params[pool_var_index]); - pool_var_info.Set(tir_main_func->params[pool_var_index], allocated_pool_info); - } - } - Map io_pool_allocations = - lowered_mod - ->GetAttr>(tvm::attr::kIOTensorPoolAllocations) - .value_or({}); - - std::vector output_var_names; - if (auto opt = func->GetAttr>("output_tensor_names")) { - Array output_tensor_names = opt.value(); - for (size_t i = 0; i < output_tensor_names.size(); ++i) { - output_var_names.push_back(output_tensor_names[i]); - } - } - - // If output names have not been specified then generate default output names - if (output_var_names.size() == 0) { - if (return_sid_.size() == 1) { - output_var_names.push_back(String("output")); - } else { - for (size_t i = 0; i < return_sid_.size(); ++i) { - output_var_names.push_back(String("output" + std::to_string(i))); - } - } - } - - Array output_tensor_types{final_aot_allocator.GetReturnTtypes()}; - - ret.metadata = ExecutorCodegenMetadata( - inputs, input_tensor_types, output_var_names, output_tensor_types, pool_vars, devices, - runtime::kTvmExecutorAot, mod_name, interface_api, unpacked_api, - GetModuleWorkspaceByteAlignment(mod), GetModuleConstantByteAlignment(mod), pool_var_info, - io_pool_allocations); - return ret; - } - - /*! - * \brief Get list of devices found - * \return List of devices - */ - Array ListDevices() { - std::vector device_names(devices_.size()); - std::transform(devices_.begin(), devices_.end(), device_names.begin(), - [](const auto& it) -> String { return it.first; }); - return device_names; - } -}; // namespace backend - -class AOTExecutorCodegenModule : public runtime::ModuleNode { - public: - AOTExecutorCodegenModule() {} - virtual PackedFunc GetFunction(const String& name, const ObjectPtr& sptr_to_self) { - if (name == "init") { - return PackedFunc([sptr_to_self, this](TVMArgs args, TVMRetValue* rv) { - ICHECK_EQ(args.num_args, 2) << "The expected of arguments are: " - << "runtime::Module mod and Array targets"; - void* mod = args[0]; - Array targets = args[1]; - init(mod, targets); - }); - } else if (name == "codegen") { - return PackedFunc([sptr_to_self, this](TVMArgs args, TVMRetValue* rv) { - IRModule mod = args[0]; - Function func = args[1]; - String mod_name = args[2]; - this->output_ = this->codegen_->Codegen(mod, func, mod_name); - }); - } else if (name == "list_params_name") { - return PackedFunc( - [sptr_to_self, this](TVMArgs args, TVMRetValue* rv) { *rv = list_params_name(); }); - } else if (name == "get_param_by_name") { - return PackedFunc([sptr_to_self, this](TVMArgs args, TVMRetValue* rv) { - String key = args[0]; - *rv = get_param_by_name(key); - }); - } else if (name == "get_irmodule") { - return PackedFunc( - [sptr_to_self, this](TVMArgs args, TVMRetValue* rv) { *rv = get_irmodule(); }); - } else if (name == "get_external_modules") { - return PackedFunc( - [sptr_to_self, this](TVMArgs args, TVMRetValue* rv) { *rv = get_external_modules(); }); - } else if (name == "get_function_metadata") { - return PackedFunc([sptr_to_self, this](TVMArgs args, TVMRetValue* rv) { - *rv = this->output_.function_metadata; - }); - } else if (name == "get_devices") { - return PackedFunc([sptr_to_self, this](TVMArgs args, TVMRetValue* rv) { - *rv = this->codegen_->ListDevices(); - }); - } else if (name == "get_executor_codegen_metadata") { - return PackedFunc( - [sptr_to_self, this](TVMArgs args, TVMRetValue* rv) { *rv = output_.metadata; }); - } else { - return PackedFunc([](TVMArgs args, TVMRetValue* rv) {}); - } - } - - const char* type_key() const final { return "RelayGraphRuntimeCodegenModule"; } - - /*! \brief Get the property of the runtime module .*/ - int GetPropertyMask() const final { return runtime::ModulePropertyMask::kRunnable; } - - private: - void init(void* mod, const Array& targets) { - codegen_ = - std::make_shared(reinterpret_cast(mod), targets); - } - - Array list_params_name() { - Array ret; - for (const auto& kv : this->output_.params) { - ret.push_back(kv.first); - } - return ret; - } - - runtime::NDArray get_param_by_name(String key) { - auto it = this->output_.params.find(key); - CHECK(it != this->output_.params.end()) << "no such parameter " << key; - return (*it).second; - } - - Array get_external_modules() { return output_.external_mods; } - - Map get_irmodule() { return this->output_.lowered_funcs; } - - std::shared_ptr codegen_; - LoweredOutput output_; -}; - -runtime::Module CreateAOTExecutorCodegenMod() { - auto ptr = make_object(); - return runtime::Module(ptr); -} - -TVM_REGISTER_GLOBAL("relay.build_module._AOTExecutorCodegen") - .set_body([](TVMArgs args, TVMRetValue* rv) { *rv = CreateAOTExecutorCodegenMod(); }); - -} // namespace backend -} // namespace relay -} // namespace tvm diff --git a/src/relay/backend/build_module.cc b/src/relay/backend/build_module.cc deleted file mode 100644 index 83c252d831c5..000000000000 --- a/src/relay/backend/build_module.cc +++ /dev/null @@ -1,523 +0,0 @@ -/* - * 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 relay/backend/build_module.cc - * \brief Code generation for TVM's graph executor. - */ -#include -#include -#include -#include -#include -#include -#include -#include -#include -#include -#include - -#include - -#include "../../driver/internal_driver_api.h" -#include "../../target/func_registry_generator.h" -#include "../../target/metadata_module.h" -#include "../../target/source/codegen_source_base.h" -#include "te_compiler.h" -#include "utils.h" - -namespace tvm { -namespace relay { -namespace transform { -Pass LabelOps(); -} -namespace backend { - -using namespace tvm::relay::transform; - -/*! - * \brief Output of building module - */ -struct BuildOutput { - std::string graph_json; - runtime::Module mod; - std::unordered_map params; -}; - -struct ExecutorCodegen { - void Init(runtime::Module* m, const Array& raw_targets) { - CallFunc("init", m, raw_targets); - } - - void Codegen(IRModule mod, const Function& func, String mod_name) { - CallFunc("codegen", mod, func, mod_name); - } - - virtual void UpdateOutput(BuildOutput* ret) = 0; - - Map GetFunctionMetadata() { - return CallFunc>("get_function_metadata", nullptr); - } - - std::unordered_map GetParams() { - std::unordered_map ret; - auto names = CallFunc>("list_params_name", nullptr); - for (const auto& expr : names) { - // Implicit cast from runtime::String to std::string - std::string key = expr; - ret[key] = CallFunc("get_param_by_name", key); - } - return ret; - } - - Array GetExternalModules() { - return CallFunc>("get_external_modules", nullptr); - } - - Map GetIRModule() { - return CallFunc>("get_irmodule", nullptr); - } - - Array ListDevices() { return CallFunc>("get_devices"); } - - relay::backend::ExecutorCodegenMetadata GetExecutorCodegenMetadata() { - return CallFunc("get_executor_codegen_metadata"); - } - virtual ~ExecutorCodegen() {} - - protected: - tvm::runtime::Module mod; - template - R CallFunc(const std::string& name, Args... args) { - auto pf = mod.GetFunction(name, false); - return pf(std::forward(args)...); - } - template - void CallFunc(const std::string& name, Args... args) { - auto pf = mod.GetFunction(name, false); - pf(std::forward(args)...); - return; - } -}; - -struct AOTCodegen : ExecutorCodegen { - AOTCodegen() { - auto pf = GetPackedFunc("relay.build_module._AOTExecutorCodegen"); - mod = (*pf)(); - } - - void UpdateOutput(BuildOutput* ret) override { ret->graph_json = ""; } - - ~AOTCodegen() {} -}; - -/*! - * \brief GraphCodegen module wrapper - * - */ -struct GraphCodegen : ExecutorCodegen { - GraphCodegen() { - auto pf = GetPackedFunc("relay.build_module._GraphExecutorCodegen"); - mod = (*pf)(); - } - void UpdateOutput(BuildOutput* ret) override { ret->graph_json = GetGraphJSON(); } - - std::string GetGraphJSON() { return CallFunc("get_graph_json", nullptr); } - - ~GraphCodegen() {} -}; - -/*! - * \brief Executor codegen factory function - */ -std::unique_ptr MakeExecutorCodegen(String executor_str) { - std::unique_ptr ret; - if (executor_str == runtime::kTvmExecutorGraph) { - ret = std::make_unique(); - } else if (executor_str == runtime::kTvmExecutorAot) { - ret = std::make_unique(); - } else { - CHECK(false) << "Executor " << executor_str << " not supported"; - } - return ret; -} - -/*! - * \brief Relay build module - * - */ -class RelayBuildModule : public runtime::ModuleNode { - public: - RelayBuildModule() = default; - - /*! - * \brief Get member function to front-end - * \param name The name of the function. - * \param sptr_to_self The pointer to the module node. - * \return The corresponding member function. - */ - PackedFunc GetFunction(const String& name, const ObjectPtr& sptr_to_self) final { - if (name == "get_graph_json") { - return PackedFunc( - [sptr_to_self, this](TVMArgs args, TVMRetValue* rv) { *rv = this->GetGraphJSON(); }); - } else if (name == "get_module") { - return PackedFunc( - [sptr_to_self, this](TVMArgs args, TVMRetValue* rv) { *rv = this->GetModule(); }); - } else if (name == "build") { - return PackedFunc([sptr_to_self, this](TVMArgs args, TVMRetValue* rv) { - ICHECK_EQ(args.num_args, 8); - this->Build(args[0], args[1], args[2], args[3], args[4], args[5], args[6], args[7]); - }); - } else if (name == "list_params") { - return PackedFunc( - [sptr_to_self, this](TVMArgs args, TVMRetValue* rv) { *rv = this->ListParamNames(); }); - } else if (name == "get_params") { - return PackedFunc( - [sptr_to_self, this](TVMArgs args, TVMRetValue* rv) { *rv = this->GetParams(); }); - } else if (name == "set_params") { - return PackedFunc([sptr_to_self, this](TVMArgs args, TVMRetValue* rv) { - Map params = args[0]; - for (const auto& kv : params) { - this->SetParam(kv.first, kv.second->data); - } - }); - } else if (name == "get_devices") { - return PackedFunc([sptr_to_self, this](TVMArgs args, TVMRetValue* rv) { - *rv = this->executor_codegen_->ListDevices(); - }); - } else if (name == "get_irmodule") { - return PackedFunc([sptr_to_self, this](TVMArgs args, TVMRetValue* rv) { - *rv = this->executor_codegen_->GetIRModule(); - }); - } else if (name == "get_external_modules") { - return PackedFunc([sptr_to_self, this](TVMArgs args, TVMRetValue* rv) { - *rv = this->executor_codegen_->GetExternalModules(); - }); - } else if (name == "get_function_metadata") { - return PackedFunc([sptr_to_self, this](TVMArgs args, TVMRetValue* rv) { - *rv = this->executor_codegen_->GetFunctionMetadata(); - }); - } else if (name == "get_executor_codegen_metadata") { - return PackedFunc([sptr_to_self, this](TVMArgs args, TVMRetValue* rv) { - *rv = this->executor_codegen_->GetExecutorCodegenMetadata(); - }); - } else if (name == "optimize") { - return PackedFunc([sptr_to_self, this](TVMArgs args, TVMRetValue* rv) { - ICHECK_EQ(args.num_args, 2); - *rv = this->Optimize(args[0], args[1]); - }); - } else { - LOG(FATAL) << "Unknown packed function: " << name; - return PackedFunc([sptr_to_self, name](TVMArgs args, TVMRetValue* rv) {}); - } - } - - /*! - * \brief Get the GraphJSON for runtime - * - * \return const std::string graph_json - */ - const std::string& GetGraphJSON() { return ret_.graph_json; } - - /*! - * \brief Get the Module object - * - * \return runtime::Module - */ - runtime::Module GetModule() { return ret_.mod; } - - /*! - * \brief List all paramter names - * - * \return Array names of params - */ - Array ListParamNames() { - Array ret; - for (const auto& kv : params_) { - ret.push_back(kv.first); - } - return ret; - } - - /*! - * \brief Get params dictionary - * - * \return Map params dictionary - */ - Map GetParams() { - Map ret; - for (const auto& kv : ret_.params) { - ret.Set(kv.first, Constant(kv.second)); - } - return ret; - } - - /*! - * \brief Set the parameters - * - * \param name name of parameter - * \param data_in input DLTensor - */ - void SetParam(const std::string& name, runtime::NDArray data_in) { params_[name] = data_in; } - - /*! - * \brief type key - * - * \return const char* - */ - const char* type_key() const final { return "RelayBuildModule"; } - - /*! \brief Get the property of the runtime module .*/ - int GetPropertyMask() const final { return runtime::ModulePropertyMask::kRunnable; } - - /*! - * \brief Build relay IRModule for graph executor - * - * \param mod Relay IRModule - * \param raw_targets List of available targets for kernels. - * \param executor Executor to target - * \param runtime Runtime to codegen for - * \param mod_name Name of the module - */ - void Build(IRModule mod, const Array& raw_targets, const tvm::Target& target_host, - const Executor& executor, const Runtime& runtime, - const WorkspaceMemoryPools& workspace_memory_pools, - const ConstantMemoryPools& constant_memory_pools, const String mod_name) { - VLOG_CONTEXT << "Build"; - executor_ = executor; - runtime_ = runtime; - workspace_memory_pools_ = workspace_memory_pools; - constant_memory_pools_ = constant_memory_pools; - config_ = CompilationConfig(PassContext::Current(), raw_targets); - VLOG(1) << "Using compilation config:" << std::endl << config_; - BuildRelay(std::move(mod), mod_name); - } - - protected: - /*! - * \brief Optimize a Relay IRModule. - * - * \param relay_module The input IRModule where optmization will be applied on. - * \param raw_targets List of available targets for kernels. - * - * \return relay::IRModule The updated Relay IR module after optimization. - */ - IRModule Optimize(IRModule relay_module, const Array& raw_targets) { - VLOG_CONTEXT << "Optimize"; - config_ = CompilationConfig(PassContext ::Current(), raw_targets); - VLOG(1) << "Using compilation config:" << std::endl << config_; - return OptimizeImpl(std::move(relay_module)); - } - - IRModule OptimizeImpl(IRModule relay_module) { - ICHECK(relay_module.defined()) << "The IRModule must be defined for the Relay compiler."; - - backend::BindParamsInModule(relay_module, params_); - - Array pass_seqs = - GetPassPrefix(/*is_homogenous=*/config_->primitive_targets.size() == 1, /*is_vm=*/false); - transform::PassContext pass_ctx = PassContext::Current(); - - if (config_->optional_homogeneous_target.defined()) { - // This pass currently only supports the homogeneous case. - pass_seqs.push_back(transform::SplitArgs( - config_->optional_homogeneous_target->GetAttr("max_function_args", 0) - .value() - .IntValue())); - } - - // Always plan devices so the remaining passes don't need to distinguish homogeneous vs - // hetrogenous execution. - pass_seqs.push_back(transform::PlanDevices(config_)); - - // Fuse the operations if it is needed. - pass_seqs.push_back(transform::FuseOps()); - - // Create a sequential pass and perform optimizations. - transform::Pass seq = transform::Sequential(pass_seqs); - if (config_->optional_homogeneous_target.defined()) { - With tctx(config_->optional_homogeneous_target); - relay_module = seq(relay_module); - } else { - relay_module = seq(relay_module); - } - - // Do layout rewrite for auto-scheduler. - if (backend::IsAutoSchedulerEnabled() && config_->optional_homogeneous_target.defined()) { - Pass major_pass = transform::AutoSchedulerLayoutRewrite(); - bool enable_layout_rewrite_targets = - config_->optional_homogeneous_target->GetTargetDeviceType() == kDLCPU || - config_->optional_homogeneous_target->GetAttr("device", "") == "mali"; - if (enable_layout_rewrite_targets && pass_ctx.PassEnabled(major_pass->Info())) { - With tctx(config_->optional_homogeneous_target); - relay_module = major_pass(relay_module); - // Defuse ops to fold constants, then fuse them again - relay_module = transform::DefuseOps()(relay_module); - relay_module = transform::FoldConstant()(relay_module); - relay_module = transform::FuseOps()(relay_module); - } - } - if (backend::IsMetaScheduleEnabled() && config_->optional_homogeneous_target.defined()) { - Pass major_pass = transform::MetaScheduleLayoutRewrite(); - bool enable_layout_rewrite_targets = - config_->optional_homogeneous_target->GetTargetDeviceType() == kDLCPU || - config_->optional_homogeneous_target->GetAttr("device", "") == "mali"; - if (enable_layout_rewrite_targets && pass_ctx.PassEnabled(major_pass->Info())) { - With tctx(config_->optional_homogeneous_target); - relay_module = major_pass(relay_module); - // Defuse ops to fold constants, then fuse them again - relay_module = transform::DefuseOps()(relay_module); - relay_module = transform::FoldConstant()(relay_module); - relay_module = transform::FuseOps()(relay_module); - } - } - - relay_module = transform::InferType()(relay_module); - - // Inline the functions that have been lifted by the module scope. - // - // TODO(@zhiics) Note that we need to be careful about the subgraphs with - // global function calls. We should make sure that these callees are also - // inline functions. However, this should be very unlikely for accelerators - // and vendor-provided libraries. So we don't handle for now. - relay_module = transform::Inline()(relay_module); - relay_module = transform::InferType()(relay_module); - relay_module = transform::LabelOps()(relay_module); - relay_module = transform::AnnotateMemoryScope()(relay_module); - - ICHECK(relay_module.defined()); - - return relay_module; - } - - /*! - * \brief Compile a Relay IR module to runtime module. - * - * \param relay_module The Relay IR module. - * \param params The parameters. - */ - void BuildRelay(IRModule relay_module, const String& mod_name) { - // Relay IRModule -> IRModule optimizations. - IRModule module = WithAttrs( - relay_module, {{tvm::attr::kExecutor, executor_}, {tvm::attr::kRuntime, runtime_}}); - relay_module = OptimizeImpl(std::move(module)); - - // Get the updated function and new IRModule to build. - // Instead of recreating the IRModule, we should look at the differences between this and the - // incoming IRModule to see if we can just pass (IRModule, Function) to the code generator. - Function func = Downcast(relay_module->Lookup("main")); - IRModule func_module = WithAttrs(IRModule::FromExpr(func), - {{tvm::attr::kExecutor, executor_}, - {tvm::attr::kRuntime, runtime_}, - {tvm::attr::kWorkspaceMemoryPools, workspace_memory_pools_}, - {tvm::attr::kConstantMemoryPools, constant_memory_pools_}}); - - // Generate code for the updated function. - executor_codegen_ = MakeExecutorCodegen(executor_->name); - executor_codegen_->Init(nullptr, config_->primitive_targets); - executor_codegen_->Codegen(func_module, func, mod_name); - executor_codegen_->UpdateOutput(&ret_); - ret_.params = executor_codegen_->GetParams(); - - auto lowered_funcs = executor_codegen_->GetIRModule(); - - // No need to build for external functions. - Target ext_dev("ext_dev"); - if (lowered_funcs.find(ext_dev) != lowered_funcs.end()) { - lowered_funcs.Set(ext_dev, IRModule()); - } - - const Target& host_target = config_->host_virtual_device->target; - const runtime::PackedFunc* pf = runtime::Registry::Get("codegen.LLVMModuleCreate"); - // When there is no lowered_funcs due to reasons such as optimization. - if (lowered_funcs.size() == 0) { - if (host_target->kind->name == "llvm") { - CHECK(pf != nullptr) << "Unable to create empty module for llvm without llvm codegen."; - // If we can decide the target is LLVM, we then create an empty LLVM module. - ret_.mod = (*pf)(host_target->str(), "empty_module"); - } else { - // If we cannot decide the target is LLVM, we create an empty CSourceModule. - // The code content is initialized with ";" to prevent complaining - // from CSourceModuleNode::SaveToFile. - ret_.mod = tvm::codegen::CSourceModuleCreate(";", "", Array{}); - } - } else { - ret_.mod = tvm::TIRToRuntime(lowered_funcs, host_target); - } - - auto ext_mods = executor_codegen_->GetExternalModules(); - ret_.mod = tvm::codegen::CreateMetadataModule(ret_.params, ret_.mod, ext_mods, host_target, - runtime_, executor_, - executor_codegen_->GetExecutorCodegenMetadata()); - // Remove external params which were stored in metadata module. - for (tvm::runtime::Module mod : ext_mods) { - auto pf_var = mod.GetFunction("get_const_vars"); - if (pf_var != nullptr) { - Array variables = pf_var(); - for (size_t i = 0; i < variables.size(); i++) { - auto it = ret_.params.find(variables[i].operator std::string()); - if (it != ret_.params.end()) { - VLOG(1) << "constant '" << variables[i] << "' has been captured in external module"; - ret_.params.erase(it); - } - } - } - } - } - - protected: - std::unique_ptr executor_codegen_; - /*! \brief Executor to build for */ - Executor executor_; - /*! \brief Runtime to codegen for */ - Runtime runtime_; - /*! \brief Workspace memory pools to codegen for */ - WorkspaceMemoryPools workspace_memory_pools_; - /*! \brief Constant memory pools to codegen for */ - ConstantMemoryPools constant_memory_pools_; - /*! \brief parameters */ - std::unordered_map params_; - /*! \brief building output */ - BuildOutput ret_; - /*! \brief Collects all the targets and scopes we need during compilation. */ - CompilationConfig config_; -}; - -runtime::Module RelayBuildCreate() { - auto exec = make_object(); - return runtime::Module(exec); -} - -TVM_REGISTER_GLOBAL("relay.build_module._BuildModule").set_body([](TVMArgs args, TVMRetValue* rv) { - *rv = RelayBuildCreate(); -}); - -TVM_REGISTER_GLOBAL("relay.build_module.BindParamsByName") - .set_body([](TVMArgs args, TVMRetValue* rv) { - Map params = args[1]; - std::unordered_map params_; - for (const auto& kv : params) { - params_[kv.first] = kv.second->data; - } - *rv = relay::backend::BindParamsByName(args[0], params_); - }); - -} // namespace backend -} // namespace relay -} // namespace tvm diff --git a/src/relay/backend/contrib/arm_compute_lib/codegen.cc b/src/relay/backend/contrib/arm_compute_lib/codegen.cc deleted file mode 100644 index 3f11e63c7391..000000000000 --- a/src/relay/backend/contrib/arm_compute_lib/codegen.cc +++ /dev/null @@ -1,426 +0,0 @@ -/* - * 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 src/relay/backend/contrib/arm_compute_lib/codegen.cc - * \brief Implementation of the Relay -> ACL JSON serializer. - */ -#include -#include -#include -#include - -#include -#include -#include - -#include "../../utils.h" -#include "../codegen_json/codegen_json.h" - -namespace tvm { -namespace relay { -namespace contrib { - -/*! - * \brief Generates an ACLModule from a relay expression. This "compilation" - * does not require ACL since the actual conversion using ACL APIs is - * deferred until creation of the runtime. This step simply serializes the - * relay program into a JSON string. - */ -class ACLJSONSerializer : public backend::contrib::JSONSerializer { - using JSONGraphNode = tvm::runtime::json::JSONGraphNode; - using JSONGraphNodeEntry = tvm::runtime::json::JSONGraphNodeEntry; - - public: - ACLJSONSerializer(const std::string& symbol, const Expr& expr) : JSONSerializer(symbol, expr) {} - - /*! - * \brief A series of operators that form a composite - * convolution. Supports both nn.conv2d and qnn.conv2d. - */ - struct CompositeConvNode { - const CallNode* pad = nullptr; - const CallNode* conv = nullptr; - const CallNode* bias = nullptr; - const CallNode* activation = nullptr; - const CallNode* requantize = nullptr; - }; - - /*! - * \brief A series of operators that form a composite - * dense layer. Supports both nn.dense and qnn.dense. - */ - struct CompositeDenseNode { - const CallNode* dense = nullptr; - const CallNode* bias = nullptr; - const CallNode* requantize = nullptr; - }; - - /*! - * \brief Visit call nodes and generate appropriate JSON node. - * - * \param cn The current call node. - * \return A list of graph entry nodes. - */ - std::vector VisitExpr_(const CallNode* cn) override { - if (cn->op.as()) { - return JSONSerializer::VisitExpr_(cn); - } - if (!cn->op.as()) { - LOG(FATAL) << "Arm Compute Library JSON runtime does not support calls to " - << cn->op->GetTypeKey(); - } - auto fn = cn->op.as(); - auto comp = fn->GetAttr(attr::kComposite); - ICHECK(comp.defined()) << "Arm Compute Library JSON runtime only supports composite functions."; - const std::string name = comp.value(); - std::shared_ptr json_node; - if (name == "arm_compute_lib.conv2d" || name == "arm_compute_lib.qnn_conv2d") { - json_node = CreateCompositeConvJSONNode(cn); - } else if (name == "arm_compute_lib.dense" || name == "arm_compute_lib.qnn_dense") { - json_node = CreateCompositeDenseJSONNode(cn); - } else if (name == "arm_compute_lib.avg_pool2d") { - json_node = CreateCompositeAvgPool2DJSONNode(cn); - } else if (name == "arm_compute_lib.l2_pool2d") { - json_node = CreateCompositeL2Pool2DJSONNode(cn); - } else if (name == "arm_compute_lib.concatenate") { - return AddCommonSingleJSONNode(cn, "concatenate"); - } else { - LOG(FATAL) << "Unrecognized Arm Compute Library pattern: " << name; - } - return AddNode(json_node, GetRef(cn)); - } - - private: - /*! - * \brief Extract convolution nodes from a composite function. - * - * \param cn The call node of the composite function. - * \return Extracted composite convolution nodes. - */ - static CompositeConvNode UnpackCompositeConvolution(const CallNode* cn) { - CompositeConvNode nodes{}; - const auto* fn = cn->op.as(); - ICHECK(fn); - - // Traverse composite convolution function from child to parent - const auto* current_call = fn->body.as(); - if (backend::IsOp(current_call, "qnn.requantize")) { - nodes.requantize = current_call; - current_call = current_call->args[0].as(); - } - if (backend::IsOp(current_call, "nn.relu")) { - nodes.activation = current_call; - current_call = current_call->args[0].as(); - } - if (backend::IsOp(current_call, "add")) { - nodes.bias = current_call; - current_call = current_call->args[0].as(); - } - // Enforce a convolution node exists at this point during traversal - if (nodes.requantize) { - ICHECK(backend::IsOp(current_call, "qnn.conv2d")); - } else { - ICHECK(backend::IsOp(current_call, "nn.conv2d")); - } - nodes.conv = current_call; - if (!current_call->args.empty() && current_call->args[0]->IsInstance()) { - current_call = current_call->args[0].as(); - if (backend::IsOp(current_call, "nn.pad")) { - nodes.pad = current_call; - } - } - return nodes; - } - - /*! - * \brief Create a JSON representation of a composite convolution. - * - * \param cn The call to be represented. - * \return A JSON representation of a specific operator. - */ - std::shared_ptr CreateCompositeConvJSONNode(const CallNode* cn) { - CompositeConvNode nodes = UnpackCompositeConvolution(cn); - - const auto* conv_attr = nodes.conv->attrs.as(); - ICHECK(conv_attr); - - std::string name; - std::string name_prefix = "nn"; - - // Distinguish between normal and depth-wise convolution - if (conv_attr->channels.defined() && - tvm::tir::ExprDeepEqual()(conv_attr->channels, conv_attr->groups) && - conv_attr->groups != 1) { - name = "depthwise_conv2d"; - ICHECK(conv_attr->kernel_layout == "IHWO") - << "Kernel layout must be IHWO, has the module been pre-processed correctly?"; - } else { - name = "conv2d"; - ICHECK(conv_attr->kernel_layout == "OHWI") - << "Kernel layout must be OHWI, has the module been pre-processed correctly?"; - } - - // Inputs must be added in the same order they appear in the relay graph. - std::vector inputs; - inputs.push_back(VisitExpr(cn->args[0])[0]); - inputs.push_back(VisitExpr(nodes.conv->args[1])[0]); - if (nodes.requantize) { - name_prefix = "qnn"; - inputs.push_back(VisitExpr(nodes.conv->args[2])[0]); // input zero-point - inputs.push_back(VisitExpr(nodes.conv->args[3])[0]); // kernel zero-point - inputs.push_back(VisitExpr(nodes.conv->args[4])[0]); // input scale - inputs.push_back(VisitExpr(nodes.conv->args[5])[0]); // kernel scale - } - if (nodes.bias) { - inputs.push_back(VisitExpr(nodes.bias->args[1])[0]); - } - if (nodes.requantize) { - inputs.push_back(VisitExpr(nodes.requantize->args[3])[0]); // output scale - inputs.push_back(VisitExpr(nodes.requantize->args[4])[0]); // output zero-point - } - - auto json_node = std::make_shared(name_prefix + "." + name, "kernel", inputs, 1); - SetCallNodeAttribute(json_node, nodes.conv); - - // Override attributes - if (nodes.pad) { - const auto* pad_attr = nodes.pad->attrs.as(); - ICHECK(pad_attr); - auto p = pad_attr->pad_width; - // Convert to TVM layout for now, conversion to ACL layout takes place in runtime. - // Standard convolution pad layout for TVM: top, left, bottom, right. - std::vector padding = {std::to_string(p[1][0].as()->value), - std::to_string(p[2][0].as()->value), - std::to_string(p[1][1].as()->value), - std::to_string(p[2][1].as()->value)}; - std::vector padding_attr; - padding_attr.emplace_back(padding); - json_node->SetAttr("padding", padding_attr); - } - if (nodes.activation) { - std::vector activation_type = {"relu"}; - std::vector act_attr; - act_attr.emplace_back(activation_type); - json_node->SetAttr("activation_type", act_attr); - } - return json_node; - } - - /*! - * \brief Extract dense nodes from a composite function. - * - * \param cn The call node of the composite function. - * \return Extracted composite convolution nodes. - */ - static CompositeDenseNode UnpackCompositeDense(const CallNode* cn) { - CompositeDenseNode nodes{}; - const auto* fn = cn->op.as(); - ICHECK(fn); - - // Traverse composite dense function from child to parent - const auto* current_call = fn->body.as(); - if (backend::IsOp(current_call, "qnn.requantize")) { - nodes.requantize = current_call; - current_call = current_call->args[0].as(); - } - if (backend::IsOp(current_call, "add")) { - nodes.bias = current_call; - current_call = current_call->args[0].as(); - } - - // Enforce a dense node exists at this point during traversal - if (nodes.requantize) { - ICHECK(backend::IsOp(current_call, "qnn.dense")); - } else { - ICHECK(backend::IsOp(current_call, "nn.dense")); - } - nodes.dense = current_call; - return nodes; - } - - /*! - * \brief Create a JSON representation of a composite dense (fully-connected) operator. - * - * \param cn The call to be represented. - * \return A JSON representation of a specific operator. - */ - std::shared_ptr CreateCompositeDenseJSONNode(const CallNode* cn) { - CompositeDenseNode nodes = UnpackCompositeDense(cn); - std::string name = "nn.dense"; - - // Inputs must be added in the same order they appear in the relay graph. - std::vector inputs; - inputs.push_back(VisitExpr(cn->args[0])[0]); - inputs.push_back(VisitExpr(nodes.dense->args[1])[0]); - if (nodes.requantize) { - name = "qnn.dense"; - inputs.push_back(VisitExpr(nodes.dense->args[2])[0]); // input zero-point - inputs.push_back(VisitExpr(nodes.dense->args[3])[0]); // weight zero-point - inputs.push_back(VisitExpr(nodes.dense->args[4])[0]); // input scale - inputs.push_back(VisitExpr(nodes.dense->args[5])[0]); // weight scale - } - if (nodes.bias) { - inputs.push_back(VisitExpr(nodes.bias->args[1])[0]); - } - if (nodes.requantize) { - inputs.push_back(VisitExpr(nodes.requantize->args[3])[0]); // output scale - inputs.push_back(VisitExpr(nodes.requantize->args[4])[0]); // output zero-point - } - - auto json_node = std::make_shared(name, "kernel", inputs, 1); - SetCallNodeAttribute(json_node, nodes.dense); - return json_node; - } - - /*! - * \brief Create a JSON representation of a composite (global) average pooling operator. - * - * A composite function is only created when using the int8/uint8 datatype for these operators. - * - * \param cn The call to be represented. - * \return A JSON representation of a specific operator. - */ - std::shared_ptr CreateCompositeAvgPool2DJSONNode(const CallNode* cn) { - const auto* fn = cn->op.as(); - ICHECK(fn); - const auto* cast = fn->body.as(); - ICHECK(cast); - const auto* avg_pool = cast->args[0].as(); - ICHECK(avg_pool); - const auto* avg_pool_op = avg_pool->op.as(); - ICHECK(avg_pool_op); - const std::string name = avg_pool_op->name; - - std::vector inputs; - inputs.push_back(VisitExpr(cn->args[0])[0]); - auto json_node = std::make_shared(name, "kernel", inputs, 1); - SetCallNodeAttribute(json_node, avg_pool); - return json_node; - } - - /*! - * \brief Create a JSON representation of a composite L2 pooling operator. - * - * \note Relay does not have an operator for L2 pooling, instead we can create - * an equivalent from power(2) + nn.avg_pool2d + sqrt. - * - * \param cn The call to be represented. - * \return A JSON representation of a specific operator. - */ - std::shared_ptr CreateCompositeL2Pool2DJSONNode(const CallNode* cn) { - const std::string name = "nn.l2_pool2d"; - const auto* fn = cn->op.as(); - ICHECK(fn); - const auto* sqrt = fn->body.as(); - ICHECK(sqrt); - const auto* avg_pool = sqrt->args[0].as(); - ICHECK(avg_pool); - const auto* pow = avg_pool->args[0].as(); - ICHECK(pow); - const auto* exponent = pow->args[1].as(); - ICHECK(exponent); - ICHECK_EQ(*static_cast(exponent->data->data), 2) << "Exponent must be 2 for L2 pooling"; - - std::vector inputs; - inputs.push_back(VisitExpr(cn->args[0])[0]); - auto json_node = std::make_shared(name, "kernel", inputs, 1); - SetCallNodeAttribute(json_node, avg_pool); - return json_node; - } - - /*! - * \brief Create a JSON representation of a single operator. - * \param cn The call to be represented. - * \param name The name of the operator. - * \return A list of graph entry nodes. - */ - std::vector AddCommonSingleJSONNode(const CallNode* cn, std::string name) { - std::vector inputs; - for (const auto& arg : cn->args) { - auto res = VisitExpr(arg); - inputs.insert(inputs.end(), res.begin(), res.end()); - } - auto node = std::make_shared(name, /* name_ */ - "kernel", /* op_type_ */ - inputs, 1 /* num_outputs_ */); - - const auto* fn = cn->op.as(); - ICHECK(fn); - const auto* callNode = fn->body.as(); - ICHECK(callNode); - SetCallNodeAttribute(node, callNode); - return AddNode(node, GetRef(cn)); - } -}; - -/*! - * \brief Create a runtime module for ACL. - * - * This consists of a series of "serialized functions" which each represent a - * sub-graph to be computed by ACL and will each be executed independently from - * one another. Each function consists of serialized JSON describing the sub-graph - * and serialized constant tensors. - * - * \note The ACL runtime module only supports a single operator per - * sub-graph currently. - * - * \param ref The ext_func Relay expression/module to be executed using extern ops. - * \return A runtime module. - */ -runtime::Module ACLCompiler(const ObjectRef& ref) { - ICHECK(ref->IsInstance()) << "The input ref is expected to be a Relay function."; - Function func = Downcast(ref); - std::string func_name = backend::GetExtSymbol(func); - - ACLJSONSerializer serializer(func_name, func); - serializer.serialize(); - std::string graph_json = serializer.GetJSON(); - - // Note that serializer.const_name_to_constant() is ignored. Instead the TECompiler invokes - // a callback which calls backend::UpdateConstants to capture the map before the function - // 'disappears' into lowered form, on the assumption the visit order and thus constant - // names match those generated by the JSONSerializer. - - const auto* pf = runtime::Registry::Get("runtime.arm_compute_lib_runtime_create"); - ICHECK(pf != nullptr) << "Cannot find JSON runtime module to create"; - runtime::Module lib = (*pf)(func_name, graph_json, serializer.const_names()); - return lib; -} - -TVM_REGISTER_GLOBAL("relay.ext.arm_compute_lib").set_body_typed(ACLCompiler); - -/*! - * \brief Check whether ACL graph executor is used. - * - * \return True if ACL graph executor is enabled, False if not. - */ -inline constexpr bool IsACLRuntimeEnabled() { -#if TVM_GRAPH_EXECUTOR_ARM_COMPUTE_LIB - return true; -#else - return false; -#endif -} - -TVM_REGISTER_GLOBAL("relay.op.is_arm_compute_runtime_enabled").set_body_typed(IsACLRuntimeEnabled); - -} // namespace contrib -} // namespace relay -} // namespace tvm diff --git a/src/relay/backend/contrib/bnns/codegen.cc b/src/relay/backend/contrib/bnns/codegen.cc deleted file mode 100644 index 3791773ad67d..000000000000 --- a/src/relay/backend/contrib/bnns/codegen.cc +++ /dev/null @@ -1,219 +0,0 @@ -/* - * 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 - * \brief Implementation of BNNS codegen APIs. - */ - -#include -#include -#include -#include - -#include -#include - -#include "../../../../runtime/contrib/json/json_node.h" -#include "../../utils.h" -#include "../codegen_json/codegen_json.h" - -namespace tvm { -namespace relay { -namespace contrib { - -using namespace backend; - -/*! - * \brief Retrieve the expected "root" op nested inside a fused call, such as conv2d in - * relu(add(conv2d)) - * \param call A Relay call node. Typically nn.relu when called the first time. - * \param max_depth The maximum number of calls before the root op, counting from current_call. - * \param root_name The name of expected "root" op in this fused call. - * \return A CallNode corresponding to the root op - */ -inline const CallNode* FindCallWithName(const CallNode* current_call, int max_depth, - const std::string& root_name) { - ICHECK(current_call && max_depth >= 0); - - if (max_depth == 0) { - ICHECK(current_call && IsOp(current_call, root_name)); - return current_call; - } - if (IsOp(current_call, root_name)) { - return current_call; - } - - ICHECK_GT(current_call->args.size(), 0); - - const auto* next_call = current_call->args[0].as(); - return FindCallWithName(next_call, max_depth - 1, root_name); -} - -class BNNSJSONSerializer : public backend::contrib::JSONSerializer { - using JSONGraphNode = tvm::runtime::json::JSONGraphNode; - using JSONGraphNodeEntry = tvm::runtime::json::JSONGraphNodeEntry; - - public: - BNNSJSONSerializer(const std::string& symbol, const Expr& expr) : JSONSerializer(symbol, expr) {} - - std::vector VisitExpr_(const CallNode* cn) override { - Expr expr = GetRef(cn); - std::string name; - const CallNode* call = cn; - if (const auto* op_node = cn->op.as()) { - name = op_node->name; - } else if (const auto* fn = cn->op.as()) { - auto comp = fn->GetAttr(attr::kComposite); - ICHECK(comp.defined()) << "BNNS JSON runtime only supports composite functions."; - name = comp.value(); - - auto body = fn->body.as(); - if (name == "bnns.conv2d_bias_relu") { - auto add_op_type = IsOp(body->args[0].as(), "add") ? "add" : "nn.bias_add"; - call = GetRootCall(body, 2, {"nn.conv2d", add_op_type, "nn.relu"}); - } else if (name == "bnns.conv2d_bias") { - auto add_op_type = IsOp(body, "add") ? "add" : "nn.bias_add"; - call = GetRootCall(body, 1, {"nn.conv2d", add_op_type}); - } else if (name == "bnns.conv2d_relu") { - call = GetRootCall(body, 1, {"nn.conv2d", "nn.relu"}); - ICHECK(call->op.as()) << "Not op node"; - } else if (name == "bnns.conv2d_bias_sigmoid") { - auto add_op_type = IsOp(body->args[0].as(), "add") ? "add" : "nn.bias_add"; - call = GetRootCall(body, 2, {"nn.conv2d", add_op_type, "sigmoid"}); - ICHECK(call->op.as()) << "Not op node"; - } else if (name == "bnns.conv2d_sigmoid") { - call = GetRootCall(body, 1, {"nn.conv2d", "sigmoid"}); - ICHECK(call->op.as()) << "Not op node"; - } else if (name == "bnns.dense_bias") { - call = GetRootCall(fn->body.as(), 1, {"nn.dense", "add"}); - } else if (name == "bnns.dense_bias_gelu") { - call = FindCallWithName(fn->body.as(), 10, "nn.dense"); - } else { - LOG(FATAL) << "Unrecognized BNNS pattern: " << name; - } - } else { - LOG(FATAL) << "BNNS JSON runtime does not support calls to " << cn->op->GetTypeKey(); - } - - std::vector inputs; - for (const auto& arg : cn->args) { - auto res = VisitExpr(arg); - inputs.insert(inputs.end(), res.begin(), res.end()); - } - auto node = std::make_shared(name, /* name_ */ - "kernel", /* op_type_ */ - inputs, 1 /* num_outputs_ */); - SetCallNodeAttribute(node, call); - return AddNode(node, GetRef(cn)); - } -}; - -/*! - * \brief The external compiler/codegen tool. It takes a Relay expression/module and - * compile it into a runtime module. - */ -runtime::Module BNNSCompiler(const ObjectRef& ref) { - ICHECK(ref->IsInstance()); - auto func = Downcast(ref); - auto func_name = GetExtSymbol(func); - BNNSJSONSerializer serializer(func_name, func); - serializer.serialize(); - std::string graph_json = serializer.GetJSON(); - - // Note that serializer.const_name_to_constant() is ignored. Instead the TECompiler invokes - // a callback which calls backend::UpdateConstants to capture the map before the function - // 'disappears' into lowered form, on the assumption the visit order and thus constant - // names match those generated by the JSONSerializer. - - const auto* pf = runtime::Registry::Get("runtime.BNNSJSONRuntimeCreate"); - ICHECK(pf != nullptr) << "Cannot find JSON runtime module to create"; - auto mod = (*pf)(func_name, graph_json, serializer.const_names()); - return mod; -} - -TVM_REGISTER_GLOBAL("relay.ext.bnns").set_body_typed(BNNSCompiler); - -/** - * \brief A helper to expand the params by adding ones which used by BNNS runtime - * for a given expression. Same as default ConstantUpdater but skip constant from - * essential BNNS composed function ops. - */ -struct BNNSConstantUpdater : public ConstantUpdater { - public: - BNNSConstantUpdater(const std::string& symbol, - std::unordered_map* params, - const std::vector& skip_mask) - : ConstantUpdater(symbol, params), skip_mask_(skip_mask) {} - using ConstantUpdater::VisitExpr_; - - /**! - * Like an original implementation but avoid visiting of body nodes - * for BNNS specific composite primitives. - */ - void VisitExpr_(const FunctionNode* op) final { - this->VisitSpan(op->span); - for (auto param : op->params) { - this->VisitExpr(param); - } - - if (!isBNNSSpecificCompositeFunc(op)) { - this->VisitExpr(op->body); - } - } - - private: - bool isBNNSSpecificCompositeFunc(const FunctionNode* op) { - auto comp = op->GetAttr(attr::kComposite); - if (!comp) return false; - - auto comp_name = comp.value(); - - bool is_match = false; - for (const auto& mask : skip_mask_) { - if (std::string(comp_name).substr(0, mask.size()) == mask) { - is_match = true; - break; - } - } - return is_match; - } - - std::vector skip_mask_; -}; - -Map BNNSConstantUpdaterFunc(Expr expr, std::string symbol) { - std::vector bnns_composite_filter = {"bnns."}; - - // Visit all suitable constant nodes - std::unordered_map res; - BNNSConstantUpdater const_updater(symbol, &res, bnns_composite_filter); - const_updater(expr); - - // Convert to tvm::Map - Map ret; - for (const auto& kvp : res) ret.Set(kvp.first, kvp.second); - return ret; -} - -TVM_REGISTER_GLOBAL("relay.ext.bnns.constant_updater").set_body_typed(BNNSConstantUpdaterFunc); - -} // namespace contrib -} // namespace relay -} // namespace tvm diff --git a/src/relay/backend/contrib/clml/codegen.cc b/src/relay/backend/contrib/clml/codegen.cc deleted file mode 100644 index 83c6ac31e328..000000000000 --- a/src/relay/backend/contrib/clml/codegen.cc +++ /dev/null @@ -1,483 +0,0 @@ -/* - * 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 src/relay/backend/contrib/clml/codegen.cc - * \brief Implementation of the Relay -> CLML JSON serializer. - */ -#include -#include -#include -#include - -#include -#include -#include -#include - -#include "../../utils.h" -#include "../codegen_json/codegen_json.h" - -namespace tvm { - -constexpr const char* kCLMLTargetVersion = "relay.ext.clml.target_version"; -TVM_REGISTER_PASS_CONFIG_OPTION(kCLMLTargetVersion, Integer); - -namespace relay { -namespace contrib { - -/*! - * \brief Generates an CLMLModule from a relay expression. This "compilation" - * does not require CLML since the actual conversion using CLML APIs is - * deferred until creation of the runtime. This step simply serializes the - * relay program into a JSON string. - */ -class CLMLJSONSerializer : public backend::contrib::JSONSerializer { - using JSONGraphNode = tvm::runtime::json::JSONGraphNode; - using JSONGraphNodeEntry = tvm::runtime::json::JSONGraphNodeEntry; - - public: - CLMLJSONSerializer(const std::string& symbol, const Expr& expr) - : JSONSerializer(symbol, expr), clml_symbol_(symbol) {} - - /*! - * \brief A series of operators that form a composite - * convolution. Supports nn.conv2d - */ - struct CompositeConvNode { - const CallNode* pad = nullptr; - const CallNode* conv = nullptr; - const CallNode* bn = nullptr; - const CallNode* bias = nullptr; - const CallNode* activation = nullptr; - std::string act_type; - }; - - /*! - * \brief Visit call nodes and generate appropriate JSON node. - * - * \param cn The current call node. - * \return A list of graph entry nodes. - */ - std::vector VisitExpr_(const CallNode* cn) override { - if (cn->op.as()) { - return JSONSerializer::VisitExpr_(cn); - } - if (!cn->op.as()) { - LOG(FATAL) << "CLML JSON runtime does not support calls to " << cn->op->GetTypeKey(); - } - auto fn = cn->op.as(); - auto comp = fn->GetAttr(attr::kComposite); - ICHECK(comp.defined()) << "CLML JSON runtime only supports composite functions."; - const std::string name = comp.value(); - std::shared_ptr json_node; - if (name == "clml.conv2d" || name == "clml.pad_conv2d" || name == "clml.conv2d_transpose") { - json_node = CreateCompositeConvJSONNode(cn); - } else if (name == "clml.batch_norm") { - json_node = CreateBatchNormJSONNode(cn); - } else if (name == "clml.dense1d" || name == "clml.dense2d") { - json_node = CreateDenseJSONNode(cn); - } else if (name == "clml.pad") { - json_node = CreatePadJSONNode(cn); - } else if (name == "clml.concat") { - json_node = CreateConcatJSONNode(cn); - } else { - json_node = CreateGenericJSONNode(cn); - } - return AddNode(json_node, GetRef(cn)); - } - - /*! - * \brief Visit call nodes and generate ordered params. - * - * \param cn The current constant node. - * \return A list of graph entry nodes. - */ - std::vector VisitExpr_(const ConstantNode* cn) override { - std::string name = "clml_" + clml_symbol_ + "_const_" + std::to_string(clml_params_.size()); - clml_params_.push_back(name); - clml_params_map_[name] = cn->data; - auto node = std::make_shared(name, "const" /* op_type_ */); - return AddNode(node, GetRef(cn)); - } - - Array GetParams() const { return clml_params_; } - Map GetParamsMap() const { - return Map(clml_params_map_); - } - - private: - std::string clml_symbol_; - Array clml_params_; - std::unordered_map clml_params_map_; - /*! - * \brief Extract convolution nodes from a composite function. - * - * \param cn The call node of the composite function. - * \return Extracted composite convolution nodes. - */ - static CompositeConvNode UnpackCompositeConvolution(const CallNode* cn) { - CompositeConvNode nodes{}; - - const auto* fn = cn->op.as(); - ICHECK(fn); - // Traverse composite convolution function from child to parent - const auto* current_call = fn->body.as(); - if (fn->body.as()) { - auto tuple_item = fn->body.as(); - current_call = tuple_item->tuple.as(); - } else { - current_call = fn->body.as(); - } - if (backend::IsOp(current_call, "nn.relu")) { - nodes.activation = current_call; - nodes.act_type = "relu"; - if (current_call->args[0].as()) { - auto tuple_item = current_call->args[0].as(); - current_call = tuple_item->tuple.as(); - } else { - current_call = current_call->args[0].as(); - } - } else if (backend::IsOp(current_call, "clip")) { - nodes.activation = current_call; - nodes.act_type = "relu6"; - if (current_call->args[0].as()) { - auto tuple_item = current_call->args[0].as(); - current_call = tuple_item->tuple.as(); - } else { - current_call = current_call->args[0].as(); - } - } - if (backend::IsOp(current_call, "nn.batch_norm")) { - nodes.bn = current_call; - current_call = current_call->args[0].as(); - } - if (backend::IsOp(current_call, "add") || backend::IsOp(current_call, "nn.bias_add")) { - nodes.bias = current_call; - current_call = current_call->args[0].as(); - } - // Enforce a convolution node exists at this point during traversal - if (!backend::IsOp(current_call, "nn.conv2d") && - !backend::IsOp(current_call, "nn.conv2d_transpose")) { - LOG(FATAL) << "Can't find primary op in Convolution node"; - } - nodes.conv = current_call; - if (!current_call->args.empty() && current_call->args[0]->IsInstance()) { - current_call = current_call->args[0].as(); - if (backend::IsOp(current_call, "nn.pad")) { - nodes.pad = current_call; - } - } - return nodes; - } - - /*! - * \brief Create a JSON representation of a composite convolution. - * - * \param cn The call to be represented. - * \return A JSON representation of a specific operator. - */ - std::shared_ptr CreateCompositeConvJSONNode(const CallNode* cn) { - CompositeConvNode nodes = UnpackCompositeConvolution(cn); - - std::string name; - std::string name_prefix = "nn"; - if (backend::IsOp(nodes.conv, "nn.conv2d")) { - const auto* conv_attr = nodes.conv->attrs.as(); - ICHECK(conv_attr); - if (conv_attr->channels.defined() && - tvm::tir::ExprDeepEqual()(conv_attr->channels, conv_attr->groups) && - conv_attr->groups != 1) { - name = "depthwise_conv2d"; - ICHECK(conv_attr->kernel_layout == "IOHW") - << "Kernel layout must be IHWO, has the module been pre-processed correctly?"; - } else { - name = "conv2d"; - ICHECK(conv_attr->kernel_layout == "OIHW") - << "Kernel layout must be OHWI, has the module been pre-processed correctly?"; - } - } else if (backend::IsOp(nodes.conv, "nn.conv2d_transpose")) { - name = "conv2d_transpose"; - const auto* conv_transpose_attr = nodes.conv->attrs.as(); - ICHECK(conv_transpose_attr); - ICHECK(conv_transpose_attr->kernel_layout == "OIHW") - << "Kernel layout must be OHWI, has the module been pre-processed correctly?"; - } - - // Inputs must be added in the same order they appear in the relay graph. - std::vector inputs; - - inputs.push_back(VisitExpr(cn->args[0])[0]); - inputs.push_back(VisitExpr(nodes.conv->args[1])[0]); - if (nodes.bias) { - inputs.push_back(VisitExpr(nodes.bias->args[1])[0]); - } - // Deal with Batchnorm Fusing here - if (nodes.bn) { - inputs.push_back(VisitExpr(nodes.bn->args[1])[0]); - inputs.push_back(VisitExpr(nodes.bn->args[2])[0]); - inputs.push_back(VisitExpr(nodes.bn->args[3])[0]); - inputs.push_back(VisitExpr(nodes.bn->args[4])[0]); - } - - auto json_node = std::make_shared(name_prefix + "." + name, "kernel", inputs, 1); - SetCallNodeAttribute(json_node, nodes.conv); - - if (nodes.bn) { - const auto* bn_attr = nodes.bn->attrs.as(); - std::vector bn_any_attr; - std::vector bn_args = { - std::to_string(bn_attr->axis), std::to_string(bn_attr->epsilon), - std::to_string(bn_attr->center), std::to_string(bn_attr->scale)}; - bn_any_attr.emplace_back(bn_args); - json_node->SetAttr("batchnorm", bn_any_attr); - } - - // Override attributes - if (nodes.pad) { - const auto* pad_attr = nodes.pad->attrs.as(); - ICHECK(pad_attr); - auto p = pad_attr->pad_width; - // Standard convolution pad layout for TVM: dimension wise pair of pre and post padding. - // CLML takes dimension wise pre-padding followed by dimension wise post-padding. - std::vector padding = {std::to_string(p[2][0].as()->value), - std::to_string(p[3][0].as()->value), - std::to_string(p[2][1].as()->value), - std::to_string(p[3][1].as()->value)}; - std::vector padding_attr; - padding_attr.emplace_back(padding); - json_node->SetAttr("padding", padding_attr); - } - - if (nodes.activation) { - std::vector activation_type = {nodes.act_type}; - std::vector act_attr; - act_attr.emplace_back(activation_type); - json_node->SetAttr("activation_type", act_attr); - } - return json_node; - } - - /*! - * \brief Create a JSON representation of a Batchnorm operator. - * - * \param cn The call to be represented. - * \return A JSON representation of a specific operator. - */ - std::shared_ptr CreateBatchNormJSONNode(const CallNode* cn) { - const auto* fn = cn->op.as(); - ICHECK(fn); - const auto* tuple_item = fn->body.as(); - ICHECK(tuple_item); - const auto* bn = tuple_item->tuple.as(); - ICHECK(bn); - const auto* bn_op = bn->op.as(); - ICHECK(bn_op); - const std::string name = bn_op->name; - - std::vector inputs; - inputs.push_back(VisitExpr(cn->args[0])[0]); - inputs.push_back(VisitExpr(bn->args[1])[0]); - inputs.push_back(VisitExpr(bn->args[2])[0]); - inputs.push_back(VisitExpr(bn->args[3])[0]); - inputs.push_back(VisitExpr(bn->args[4])[0]); - auto json_node = std::make_shared(name, "kernel", inputs, 1); - SetCallNodeAttribute(json_node, bn); - return json_node; - } - - /*! - * \brief Create a JSON representation of a Concat operator. - * - * \param cn The call to be represented. - * \return A JSON representation of a specific operator. - */ - std::shared_ptr CreateConcatJSONNode(const CallNode* cn) { - const auto* fn = cn->op.as(); - ICHECK(fn); - const auto* concat = fn->body.as(); - - ICHECK(backend::IsOp(concat, "concatenate")); - const auto* concat_op = concat->op.as(); - ICHECK(concat_op); - const std::string name = concat_op->name; - - std::vector inputs; - for (auto arg : cn->args) { - inputs.push_back(VisitExpr(arg)[0]); - } - - auto json_node = std::make_shared(name, "kernel", inputs, 1); - SetCallNodeAttribute(json_node, concat); - return json_node; - } - - /*! - * \brief Create a JSON representation of a Dense operator. - * - * \param cn The call to be represented. - * \return A JSON representation of a specific operator. - */ - std::shared_ptr CreateDenseJSONNode(const CallNode* cn) { - const auto* fn = cn->op.as(); - ICHECK(fn); - const auto* dense = fn->body.as(); - const CallNode* bias = nullptr; - - if (backend::IsOp(dense, "add") || backend::IsOp(dense, "nn.bias_add")) { - bias = dense; - dense = dense->args[0].as(); - } - ICHECK(backend::IsOp(dense, "nn.dense")); - const auto* dense_op = dense->op.as(); - ICHECK(dense_op); - const std::string name = dense_op->name; - - std::vector inputs; - inputs.push_back(VisitExpr(cn->args[0])[0]); - inputs.push_back(VisitExpr(dense->args[1])[0]); - if (bias) { - inputs.push_back(VisitExpr(bias->args[1])[0]); - } - auto json_node = std::make_shared(name, "kernel", inputs, 1); - SetCallNodeAttribute(json_node, dense); - return json_node; - } - - /*! - * \brief Create a JSON representation of a Pad operator. - * - * \param cn The call to be represented. - * \return A JSON representation of a specific operator. - */ - std::shared_ptr CreatePadJSONNode(const CallNode* cn) { - const auto* fn = cn->op.as(); - ICHECK(fn); - const auto* pad = fn->body.as(); - const auto* pad_op = pad->op.as(); - ICHECK(pad_op); - const std::string name = pad_op->name; - - std::vector inputs; - inputs.push_back(VisitExpr(cn->args[0])[0]); - - auto json_node = std::make_shared(name, "kernel", inputs, 1); - - const auto* pad_attr = pad->attrs.as(); - ICHECK(pad_attr); - auto p = pad_attr->pad_width; - // TVM padding format: Dimension wise pair of pre and post padding. - // CLML padding format: Dimension wise pre padding followed by dimension wise post padding. - std::vector padding = {std::to_string(p[2][0].as()->value), - std::to_string(p[2][1].as()->value), - std::to_string(p[3][0].as()->value), - std::to_string(p[3][1].as()->value)}; - std::vector padding_attr; - padding_attr.emplace_back(padding); - json_node->SetAttr("pad_width", padding_attr); - - std::vector pad_mode = {pad_attr->pad_mode}; - std::vector pad_mode_attr; - pad_mode_attr.emplace_back(pad_mode); - json_node->SetAttr("pad_mode", pad_mode_attr); - - return json_node; - } - - std::shared_ptr CreateGenericJSONNode(const CallNode* cn) { - const auto* fn = cn->op.as(); - ICHECK(fn); - const auto* node = fn->body.as(); - - const auto* node_op = node->op.as(); - ICHECK(node_op); - const std::string name = node_op->name; - - std::vector inputs; - unsigned int i = 0; - for (i = 0; i < cn->args.size(); i++) { - inputs.push_back(VisitExpr(cn->args[i])[0]); - } - for (unsigned int j = i; j < node->args.size(); j++) { - inputs.push_back(VisitExpr(node->args[j])[0]); - } - auto json_node = std::make_shared(name, "kernel", inputs, 1); - SetCallNodeAttribute(json_node, node); - return json_node; - } -}; - -/*! - * \brief Create a runtime module for CLML. - * - * This consists of a series of "serialized functions" which each represent a - * sub-graph to be computed by CLML and will each be executed independently from - * one another. Each function consists of serialized JSON describing the sub-graph - * and serialized constant tensors. - * - * \note The CLML runtime module only supports a single operator per - * sub-graph currently. - * - * \param ref The ext_func Relay expression/module to be executed using extern ops. - * \return A runtime module. - */ -runtime::Module CLMLCompiler(const ObjectRef& ref) { - ICHECK(ref->IsInstance()) << "The input ref is expected to be a Relay function."; - Function func = Downcast(ref); - std::string func_name = backend::GetExtSymbol(func); - - CLMLJSONSerializer serializer(func_name, func); - serializer.serialize(); - std::string graph_json = serializer.GetJSON(); - auto param_names = serializer.GetParams(); - const auto* pf = runtime::Registry::Get("runtime.clml_runtime_create"); - ICHECK(pf != nullptr) << "Cannot find CLML runtime module to create"; - runtime::Module lib = (*pf)(func_name, graph_json, param_names); - return lib; -} - -TVM_REGISTER_GLOBAL("relay.ext.clml").set_body_typed(CLMLCompiler); - -/*! - * \brief Check whether CLML graph runtime is used. - * - * \return True if CLML graph runtime is enabled, False if not. - */ -inline constexpr bool IsCLMLRuntimeEnabled() { -#if TVM_GRAPH_EXECUTOR_CLML - return true; -#else - return false; -#endif -} - -TVM_REGISTER_GLOBAL("relay.op.is_clml_runtime_enabled").set_body_typed(IsCLMLRuntimeEnabled); - -Map CLMLConstantUpdater(Expr func, std::string symbol) { - CLMLJSONSerializer serializer(symbol, func); - serializer.serialize(); - auto pmap = serializer.GetParamsMap(); - return pmap; -} - -TVM_REGISTER_GLOBAL("relay.ext.clml.constant_updater").set_body_typed(CLMLConstantUpdater); - -} // namespace contrib -} // namespace relay -} // namespace tvm diff --git a/src/relay/backend/contrib/clml/target.cc b/src/relay/backend/contrib/clml/target.cc deleted file mode 100644 index c7f22c1315c8..000000000000 --- a/src/relay/backend/contrib/clml/target.cc +++ /dev/null @@ -1,41 +0,0 @@ -/* - * 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 src/relay/backend/contrib/clml/target.cc - * \brief Registers the "clml" external codegen TargetKind. - */ - -#include - -namespace tvm { -namespace relay { -namespace contrib { - -/*! - * \brief This external codegen target can use the CLML library linked into the TVM runtime. - * - Patterns and custom compiler: python/tvm/relay/op/contrib/clml.py - * - Runtime: src/runtime/contrib/clml/clml_runtime.cc - */ -TVM_REGISTER_TARGET_KIND("clml", kDLOpenCL) - .set_attr(tvm::attr::kIsExternalCodegen, Bool(true)); - -} // namespace contrib -} // namespace relay -} // namespace tvm diff --git a/src/relay/backend/contrib/codegen_c/codegen.cc b/src/relay/backend/contrib/codegen_c/codegen.cc deleted file mode 100644 index de41807431b0..000000000000 --- a/src/relay/backend/contrib/codegen_c/codegen.cc +++ /dev/null @@ -1,401 +0,0 @@ -/* - * 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 "../../../transforms/compiler_function_utils.h" -#include "../../utils.h" -#include "codegen_c.h" - -namespace tvm { -namespace relay { -namespace contrib { - -/*! \brief Return the "ccompiler" Target instance to use to guide compilation. */ -Target GetCCompilerTarget() { - Target target = Target::Current(/*allow_not_defined=*/true); - if (!target.defined() || target->kind->name != "ccompiler") { - // Use the default compilation options if no specific "ccompiler" target was given - // in the overall targets list. In that case target_hooks.cc will invoke the custom pass - // without pushing any target instance onto the implicit target stack. - target = Target("ccompiler"); - } - return target; -} - -/*! - * \brief Emits C/C++ code for a single function. - * - * For testing and demonstration only, only a few binary operators are supported. - */ -class CodegenC : public backend::MemoizedExprTranslator>, public CodegenCBase { - public: - CodegenC(std::unordered_map* const_name_to_constant, - Array* const_names, bool* needs_extra_headers, std::string ext_func_id) - : const_name_to_constant_(const_name_to_constant), - const_names_(const_names), - needs_extra_headers_(needs_extra_headers), - ext_func_id_(std::move(ext_func_id)) {} - - /*! - * \brief Emit the source code that invokes C compiler compatible wrappers. - * - * \return The emitted code. - */ - std::string JIT(const std::vector& out) override { - // Write function macros - for (auto decl : func_decl_) { - code_stream_ << decl << "\n"; - } - return JitImpl(ext_func_id_, ext_func_args_, buf_decl_, ext_func_body_, const_array_name_, out); - } - - private: - std::vector VisitExprDefault_(const Object* op) override { - LOG(FATAL) << "C codegen doesn't support: " << op->GetTypeKey(); - } - - std::vector VisitExpr_(const VarNode* node) override { - ext_func_args_.push_back(GetRef(node)); - Output output; - output.name = node->name_hint(); - return {output}; - } - - std::vector VisitExpr_(const TupleNode* node) override { - std::vector outs; - for (auto field : node->fields) { - auto res = VisitExpr(field); - ICHECK_EQ(res.size(), 1U) << "Do not support tuple nest"; - outs.push_back(res[0]); - } - return outs; - } - - std::vector VisitExpr_(const TupleGetItemNode* op) override { - auto res = VisitExpr(op->tuple); - ICHECK_GT(res.size(), static_cast(op->index)); - - // Only keep the item we want for the child node. - // FIXME(@comaniac): The other items should still be requried for the primary outputs. - return {res[op->index]}; - } - - std::vector VisitExpr_(const ConstantNode* cn) override { - // Remember we'll need some extra headers to support the runtime constants array. - *needs_extra_headers_ = true; - - std::ostringstream decl_stream; - std::ostringstream buf_stream; - - Output output; - // Get const: static_cast(gcc_0_consts[0]->data) - size_t const_id = const_name_to_constant_->size(); - output.name = CreateDataReference(ext_func_id_, const_id); - const auto* type_node = cn->checked_type().as(); - ICHECK(type_node); - const auto& dtype = GetDtypeString(type_node); - - // Generate the global variable for needed ndarrays - if (const_array_name_.empty()) { - *needs_extra_headers_ = true; - const_array_name_ = CreateNDArrayPool(ext_func_id_); - std::string checker = CreateInitChecker(ext_func_id_); - ext_func_body_.insert(ext_func_body_.begin(), checker); - } - - ICHECK(dtype == "float" || dtype == "int") << "Only float and int are supported for now."; - output.dtype = dtype; - - std::string const_var_name = CreateConstVar(ext_func_id_, const_id); - const_name_to_constant_->emplace(const_var_name, cn->data); - const_names_->push_back(const_var_name); - - return {output}; - } - - std::vector VisitExpr_(const CallNode* call) override { - std::ostringstream macro_stream; - std::ostringstream decl_stream; - std::ostringstream buf_stream; - - std::string func_name = ext_func_id_ + "_" + std::to_string(func_idx++); - - // Make function declaration - macro_stream << "CSOURCE_BINARY_OP_" << call->args.size() << "D(" << func_name << ", "; - - if (backend::IsOp(call, "add")) { - macro_stream << "+"; - } else if (backend::IsOp(call, "subtract")) { - macro_stream << "-"; - } else if (backend::IsOp(call, "multiply")) { - macro_stream << "*"; - } else { - LOG(FATAL) << "Unrecognized op"; - } - - auto in_shape = backend::GetShape(call->args[0]->checked_type()); - for (size_t i = 0; i < in_shape.size(); ++i) { - macro_stream << ", " << in_shape[i]; - } - - const auto* type_node = call->checked_type().as(); - ICHECK(type_node); - const auto& dtype = GetDtypeString(type_node); - macro_stream << ", " << dtype; - - macro_stream << ");"; - func_decl_.push_back(macro_stream.str()); - - // Make function call when visiting arguments - bool first = true; - decl_stream << func_name << "("; - for (size_t i = 0; i < call->args.size(); ++i) { - auto res = VisitExpr(call->args[i]); - for (auto out : res) { - if (!first) { - decl_stream << ", "; - } - first = false; - decl_stream << out.name; - } - } - - std::string out = "buf_" + std::to_string(buf_idx_++); - auto out_shape = backend::GetShape(call->checked_type()); - int out_size = 1; - for (size_t i = 0; i < out_shape.size(); ++i) { - out_size *= out_shape[i]; - } - buf_stream << dtype << "* " << out << " = (" << dtype << "*)malloc(4 * " << out_size << ");"; - buf_decl_.push_back(buf_stream.str()); - - decl_stream << ", " << out << ");"; - ext_func_body_.push_back(decl_stream.str()); - - // Update output buffer - // Note C codegen only handles TensorType. Therefore, we don't flatten - // tuples and only return a single vaule. - Output output; - output.name = out; - output.dtype = dtype; - output.need_copy = true; - output.size = out_size; - return {output}; - } - - /*! - * \brief The accumulated constant name to constant mapping. Shared between all generated - * functions. - */ - std::unordered_map* const_name_to_constant_; - /*! \brief The accumulated constant names, in the order they were generated. */ - Array* const_names_; - /*! - * \brief Set to true if the ndarray and packed function headers are required to declare and - * manage the constants array. - */ - bool* needs_extra_headers_; - /*! \brief Name of the global function currently being compiled. */ - std::string ext_func_id_; - - /*! \brief The index of the next available wrapped C function. */ - int func_idx = 0; - /*! \brief The index of the next available allocated buffers. */ - int buf_idx_ = 0; - /*! \brief The arguments of a C compiler compatible function. */ - Array ext_func_args_; - /*! \brief The statements of a C compiler compatible function. */ - std::vector ext_func_body_; - /*! \brief The array declared to store the constant values. */ - std::string const_array_name_; - /*! \brief The declaration statements of a C compiler compatible function. */ - std::vector func_decl_; - /*! \brief The declaration statements of buffers. */ - std::vector buf_decl_; -}; - -/*! \brief Emits C/C++ code for a module. */ -class CodegenCModule { - public: - CodegenCModule(Target target, IRModule mod) : target_(std::move(target)), mod_(std::move(mod)) {} - - runtime::Module CreateCSourceModule() { - for (const auto& kv : mod_->functions) { - if (const auto* function_node = GetCCompilerFunctionNode(kv.second)) { - GenCFunc(GetRef(function_node)); - } - } - return Finalize(); - } - - /*! \brief Returns the accumulated constant name to constant mapping. */ - const std::unordered_map& const_name_to_constant() const { - return const_name_to_constant_; - } - - private: - /*! \brief Emits the standard C/C++ header into \p os. */ - void EmitPreamble(std::ostringstream& os) { - // Custom header, if any. - Optional header = target_->GetAttr("header"); - if (header.defined() && !header.value().empty()) { - os << header.value().c_str() << "\n"; - } - - // Standard includes. - os << "#include \n"; - os << "#include \n"; - os << "#include \n"; - os << "#include \n"; - os << "#include \n"; - - if (needs_extra_headers_) { - // This segment would be generated in C++ because of the usage - // of tvm::runtime::Array. This is not ideal, but this to demonstrate - // constant copying process used packed imports in other external - // codegen. Moreover, in microTVM we dont expect this part to be generated. - os << "#ifdef __cplusplus\n"; - os << "#include \n"; - os << "#include \n"; - os << "#endif\n"; - } - - // Define some macros to help operator implementations. - const char* operator_macro = R"op_macro( - #define CSOURCE_BINARY_OP_1D(p_ID_, p_OP_, p_DIM1_, p_DTYPE) \ - void p_ID_(p_DTYPE* a, p_DTYPE* b, p_DTYPE* out) { \ - for (int64_t i = 0; i < p_DIM1_; ++i) { \ - out[i] = a[i] p_OP_ b[i]; \ - } \ - } - - #define CSOURCE_BINARY_OP_2D(p_ID_, p_OP_, p_DIM1_, p_DIM2_, p_DTYPE) \ - void p_ID_(p_DTYPE* a, p_DTYPE* b, p_DTYPE* out) { \ - for (int64_t i = 0; i < p_DIM1_; ++i) { \ - for (int64_t j = 0; j < p_DIM2_; ++j) { \ - int64_t k = i * p_DIM2_ + j; \ - out[k] = a[k] p_OP_ b[k]; \ - } \ - } \ - } - )op_macro"; - - os << operator_macro << "\n\n"; - } - - void GenCFunc(const Function& function) { - ICHECK(function.defined()) << "Input error: expect a Relay function."; - std::string ext_func_id = backend::GetExtSymbol(function); - CodegenC builder(&const_name_to_constant_, &const_names_, &needs_extra_headers_, ext_func_id); - std::vector out = builder.VisitExpr(function->body); - code_stream_ << builder.JIT(out); - func_names_.push_back(ext_func_id); - } - - /*! \brief Returns function if it is tagged with "Compiler=ccompiler". */ - static const FunctionNode* GetCCompilerFunctionNode(const Expr& expr) { - if (const auto* function_node = expr.as()) { - Optional opt_compiler = function_node->GetAttr(attr::kCompiler); - if (opt_compiler.defined() && opt_compiler.value() == "ccompiler") { - return function_node; - } - } - return nullptr; - } - - runtime::Module Finalize() { - std::ostringstream os; - EmitPreamble(os); - os << code_stream_.str(); - std::string code = os.str(); - - VLOG(1) << "CodegenCModule generated:" << std::endl << code; - - // Create a CSource module - const auto* pf = runtime::Registry::Get("runtime.CSourceModuleCreate"); - ICHECK(pf != nullptr) << "Cannot find csource module to create the external runtime module"; - return (*pf)(code, "c", func_names_, const_names_); - } - - /*! \brief "ccompiler" Target with compilation options to use. */ - Target target_; - /*! \brief Module we are compiling. */ - IRModule mod_; - - /*! \brief True if we need to include the ndarray and packed function headers. */ - bool needs_extra_headers_ = false; - /*! \brief The accumulated constant name to constant mapping. */ - std::unordered_map const_name_to_constant_; - /*! \brief The accumulated constant names, in the order they were generated. */ - Array const_names_; - /*! \brief The accumulated function names. */ - Array func_names_; - /*! - * \brief The accumulated code stream containing all function definitions. - * (Does not include the preamble.) - */ - std::ostringstream code_stream_; -}; - -/*! \brief The actual translation pass. */ -tvm::transform::Pass CCompilerImpl() { - auto pass_func = [=](IRModule mod, const tvm::transform::PassContext& pass_ctx) { - VLOG(1) << "CCompilerImpl input:" << std::endl << PrettyPrint(mod); - Target target = GetCCompilerTarget(); - - // Emit the C/C++ code and package it as a CSourceModule. - CodegenCModule codegen(target, mod); - runtime::Module runtime_mod = codegen.CreateCSourceModule(); - - // Capture the new runtime module. - Array external_mods = - mod->GetAttr>(tvm::attr::kExternalMods).value_or({}); - external_mods.push_back(runtime_mod); - - // Capture the new constants. - Map const_name_to_constant = - mod->GetAttr>(tvm::attr::kConstNameToConstant).value_or({}); - for (const auto& kv : codegen.const_name_to_constant()) { - ICHECK_EQ(const_name_to_constant.count(kv.first), 0); - const_name_to_constant.Set(kv.first, kv.second); - } - - return WithAttrs(mod, {{tvm::attr::kExternalMods, external_mods}, - {tvm::attr::kConstNameToConstant, const_name_to_constant}}); - }; - return tvm::transform::CreateModulePass(pass_func, 0, "CCompilerImpl", {}); -} - -tvm::transform::Pass CCompilerPass() { - return transform::Sequential( - {transform::OutlineCompilerFunctionsWithExistingGlobalSymbols("ccompiler"), CCompilerImpl(), - transform::MarkCompilerFunctionsAsExtern("ccompiler")}); -} - -} // namespace contrib -} // namespace relay -} // namespace tvm diff --git a/src/relay/backend/contrib/codegen_c/target.cc b/src/relay/backend/contrib/codegen_c/target.cc deleted file mode 100644 index cd1e0283df28..000000000000 --- a/src/relay/backend/contrib/codegen_c/target.cc +++ /dev/null @@ -1,43 +0,0 @@ -/* - * 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 "./codegen_c.h" - -namespace tvm { -namespace relay { -namespace contrib { - -/*! - * \brief This demonstration external codegen target emits C/C++ for compilation by the native c - * compiler on CPU. - * - Patterns: None, functions must be explicitly marked as "Primitive" and "Compiler=ccompiler". - * - Custom compiler: relay/backend/contrib/codegen_c/codegen.cc - */ -TVM_REGISTER_TARGET_KIND("ccompiler", kDLCPU) - .set_attr(tvm::attr::kIsExternalCodegen, Bool(true)) - .set_attr(tvm::attr::kRelayToTIR, CCompilerPass()) - // Value is prepended to every output CModule. - .add_attr_option("header", String("")); - -} // namespace contrib -} // namespace relay -} // namespace tvm diff --git a/src/relay/backend/contrib/codegen_json/codegen_json.h b/src/relay/backend/contrib/codegen_json/codegen_json.h deleted file mode 100644 index 350a1275ae27..000000000000 --- a/src/relay/backend/contrib/codegen_json/codegen_json.h +++ /dev/null @@ -1,377 +0,0 @@ -/* - * 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 relay/backend/contrib/codegen_json.h - * \brief Utilities for json codegen and runtime - */ -#ifndef TVM_RELAY_BACKEND_CONTRIB_CODEGEN_JSON_CODEGEN_JSON_H_ -#define TVM_RELAY_BACKEND_CONTRIB_CODEGEN_JSON_CODEGEN_JSON_H_ - -#include -#include -#include -#include - -#include -#include -#include -#include -#include -#include -#include - -#include "../../../../runtime/contrib/json/json_node.h" -#include "../../../../runtime/contrib/json/json_runtime.h" -#include "../../utils.h" - -namespace tvm { -namespace relay { -namespace backend { -namespace contrib { - -using namespace tvm::runtime::json; - -using ShapeVector = std::vector>; -using TypeVector = std::vector; -using JSONGraphObjectPtr = std::shared_ptr; - -/*! - * \brief Helper class to extract all attributes of a certain op and save them - * into text format. - */ -class OpAttrExtractor : public AttrVisitor { - public: - explicit OpAttrExtractor(JSONGraphObjectPtr node) : node_(node) {} - - template ::value>> - std::string Fp2String(const T value) { - std::ostringstream out; - out.precision(std::numeric_limits::max_digits10); - out << value; - return out.str(); - } - - void SetNodeAttr(const char* key, const std::vector& value) { - std::vector attr; - attr.emplace_back(value); - node_->SetAttr(key, attr); - } - - void Visit(const char* key, double* value) final { SetNodeAttr(key, {Fp2String(*value)}); } - - void Visit(const char* key, int64_t* value) final { SetNodeAttr(key, {std::to_string(*value)}); } - - void Visit(const char* key, uint64_t* value) final { SetNodeAttr(key, {std::to_string(*value)}); } - - void Visit(const char* key, int* value) final { SetNodeAttr(key, {std::to_string(*value)}); } - - void Visit(const char* key, bool* value) final { SetNodeAttr(key, {std::to_string(*value)}); } - - void Visit(const char* key, std::string* value) final { SetNodeAttr(key, {*value}); } - - void Visit(const char* key, DataType* value) final { - if (!value->is_void()) { - SetNodeAttr(key, {runtime::DLDataType2String(*value)}); - } else { - SetNodeAttr(key, {""}); - } - } - - void Visit(const char* key, runtime::ObjectRef* value) final { - if (const auto* an = (*value).as()) { - std::vector attr; - for (size_t i = 0; i < an->size(); ++i) { - if (const auto* im = (*an)[i].as()) { - attr.push_back(std::to_string(im->value)); - } else if (const auto* fm = (*an)[i].as()) { - attr.push_back(Fp2String(fm->value)); - } else if (const auto* str = (*an)[i].as()) { - String s = GetRef(str); - attr.push_back(s); - } else { - LOG(FATAL) << "Not supported type: " << (*an)[i]->GetTypeKey(); - } - } - SetNodeAttr(key, attr); - } else if (!(*value).defined()) { // Skip NullValue - SetNodeAttr(key, std::vector{""}); - } else if (const auto* im = (*value).as()) { - SetNodeAttr(key, std::vector{std::to_string(im->value)}); - } else if (const auto* fm = (*value).as()) { - SetNodeAttr(key, std::vector{Fp2String(fm->value)}); - } else if (const auto* str = (*value).as()) { - String s = GetRef(str); - SetNodeAttr(key, std::vector{s}); - } else { - LOG(FATAL) << "Not yet supported type: " << (*value)->GetTypeKey() << ": " << *value; - } - } - - void Visit(const char* key, runtime::NDArray* value) final { - LOG(FATAL) << "NDArray is not allowed in op attribute"; - } - - void Visit(const char* key, void** value) final { - LOG(FATAL) << "void pointer is not allowed in op attribute"; - } - - void Extract(Object* node) { - if (node) { - reflection_->VisitAttrs(node, this); - } - } - - private: - JSONGraphObjectPtr node_; - ReflectionVTable* reflection_ = ReflectionVTable::Global(); -}; - -/*! \brief Serialize a Relay expression to JSON. */ -class JSONSerializer : public MemoizedExprTranslator> { - public: - /*! - * \brief Constructor - * - * \param symbol The symbol that represents the graph being converted. - * \param expr The Relay expression to be converted to the JSON form. - */ - JSONSerializer(std::string symbol, Expr expr) - : symbol_(std::move(symbol)), func_(std::move(expr)) {} - - void serialize() { - relay::Function func = Downcast(func_); - // First we convert all the parameters into input nodes. - for (const auto& param : func->params) { - auto node_ptr = std::make_shared(param->name_hint(), "input" /* op_type_ */); - memo_[param] = AddNode(node_ptr, param); - } - heads_ = VisitExpr(func->body); - } - - /*! - * \brief Returns the accumulated map from constant names to the NDArray they must be bound to - * at runtime. Also referred to a 'params' elsewhere in the code. - */ - const std::unordered_map& const_name_to_constant() const { - return const_name_to_constant_; - } - - /*! - * \brief Return the constant names in order they were encountered during translation. - */ - const Array& const_names() const { return const_names_; } - - /*!\brief Return the generated json. */ - std::string GetJSON() { - std::ostringstream os; - dmlc::JSONWriter writer(&os); - Save(&writer); - return os.str(); - } - - protected: - /*! - * \brief Add a node to graph. - * - * \param node A graph node. It is a shared pointer. Some attributes of it - * will be added, i.e. shape and type. These attributes are attached to - * the JSON graph in the end. - * \param expr The relay expression. - * \return A list of graph entry nodes. It the relay expr is a tuple type, we - * will flatten it. - */ - std::vector AddNode(JSONGraphObjectPtr node, const Expr& expr) { - auto checked_type = expr->checked_type(); - auto node_id = nodes_.size(); - nodes_.push_back(node); - std::vector ret; - ShapeVector shape; - TypeVector dtype; - // Flatten tuple node. - if (const auto* tuple_type = checked_type.as()) { - for (size_t i = 0; i < tuple_type->fields.size(); ++i) { - const auto* tensor_type = tuple_type->fields[i].as(); - ICHECK(tensor_type) << "Expect TensorType, but received: ." - << tuple_type->fields[i]->GetTypeKey(); - ret.push_back(JSONGraphNodeEntry(node_id, i)); - shape.emplace_back(GetIntShape(tensor_type->shape)); - dtype.emplace_back(DType2String(tensor_type->dtype)); - } - node->SetNumOutput(tuple_type->fields.size()); - } else { - const auto* tensor_type = checked_type.as(); - ICHECK(tensor_type) << "Expect TensorType, but received: " << checked_type->GetTypeKey(); - shape.emplace_back(GetIntShape(tensor_type->shape)); - dtype.emplace_back(DType2String(tensor_type->dtype)); - ret.push_back(JSONGraphNodeEntry(node_id, 0)); - } - std::vector shape_attrs; - shape_attrs.emplace_back(shape); - node->SetAttr("shape", shape_attrs); - - std::vector type_attrs; - type_attrs.emplace_back(dtype); - node->SetAttr("dtype", type_attrs); - return ret; - } - - void SetCallNodeAttribute(JSONGraphObjectPtr node, const CallNode* cn) { - if (cn->op.as()) { - OpAttrExtractor extractor(node); - const Object* call_attr = cn->attrs.get(); - extractor.Extract(const_cast(call_attr)); - } else if (const auto* fn = cn->op.as()) { - auto pattern = fn->GetAttr(attr::kPartitionedFromPattern); - ICHECK(pattern.defined()); - std::vector values; - values.push_back(pattern.value()); - std::vector attr; - attr.emplace_back(values); - node->SetAttr("PartitionedFromPattern", attr); - } - } - - std::vector VisitExprDefault_(const Object* op) { - LOG(FATAL) << "JSON runtime currently doesn't support " << op->GetTypeKey(); - } - - std::vector VisitExpr_(const VarNode* vn) { - ICHECK(memo_.count(GetRef(vn))); - return memo_[GetRef(vn)]; - } - - std::vector VisitExpr_(const ConstantNode* constant_node) { - std::string name = symbol_ + "_const_" + std::to_string(const_names_.size()); - VLOG(1) << "Will require parameter '" << name - << "' to be supplied by the ConstLoaderModule at runtime"; - ICHECK_EQ(const_name_to_constant_.count(name), 0); - const_name_to_constant_.emplace(name, constant_node->data); - const_names_.push_back(name); - auto node = std::make_shared(name, /*op_type=*/"const"); - return AddNode(node, GetRef(constant_node)); - } - - std::vector VisitExpr_(const TupleNode* tn) { - std::vector fields; - for (const auto& field : tn->fields) { - auto ref = VisitExpr(field); - fields.insert(fields.end(), ref.begin(), ref.end()); - } - return fields; - } - - std::vector VisitExpr_(const CallNode* cn) { - Expr expr = GetRef(cn); - std::string name; - if (const auto* op_node = cn->op.as()) { - name = op_node->name; - } else if (const auto* fn = cn->op.as()) { - auto comp = fn->GetAttr(attr::kComposite); - ICHECK(comp.defined()) << "JSON runtime only supports composite functions."; - name = comp.value(); - } else { - LOG(FATAL) << "JSON runtime does not support calls to " << cn->op->GetTypeKey(); - } - - std::vector inputs; - for (const auto& arg : cn->args) { - auto res = VisitExpr(arg); - inputs.insert(inputs.end(), res.begin(), res.end()); - } - auto node = std::make_shared(name, /* name_ */ - "kernel", /* op_type_ */ - inputs, 1 /* num_outputs_ */); - SetCallNodeAttribute(node, cn); - return AddNode(node, GetRef(cn)); - } - - std::vector VisitExpr_(const LetNode* ln) { - ICHECK_EQ(memo_.count(ln->var), 0); - memo_[ln->var] = VisitExpr(ln->value); - return VisitExpr(ln->body); - } - - std::vector VisitExpr_(const TupleGetItemNode* gtn) { - auto vtuple = VisitExpr(gtn->tuple); - return {vtuple[gtn->index]}; - } - - std::vector VisitExpr_(const FunctionNode* fn) { - ICHECK(fn->GetAttr(attr::kComposite).defined()) - << "JSON runtime only supports composite functions"; - // FunctionNode should be handled by the caller. - return {}; - } - - /*! - * \brief Save to JSON graph - * - * \param writer A json writer - */ - void Save(dmlc::JSONWriter* writer) { - std::vector arg_nodes; - for (size_t i = 0; i < nodes_.size(); ++i) { - auto node = nodes_[i]; - if (node->IsLeaf()) { - arg_nodes.push_back(i); - } - } - size_t num_entry = 0; - std::vector node_row_ptr{0}; - for (auto node : nodes_) { - num_entry += node->GetNumOutput(); - node_row_ptr.push_back(num_entry); - } - writer->BeginObject(); - writer->WriteObjectKeyValue("symbol", symbol_); - writer->WriteObjectKeyValue("nodes", nodes_); - writer->WriteObjectKeyValue("arg_nodes", arg_nodes); - writer->WriteObjectKeyValue("heads", heads_); - writer->WriteObjectKeyValue("node_row_ptr", node_row_ptr); - writer->EndObject(); - } - - private: - /*! \brief The symbol that represents the json graph. */ - std::string symbol_; - /*! \brief The function to be serialized. */ - const Expr func_; - /*! \brief JSON graph nodes. */ - std::vector nodes_; - /*! \brief Output of the JSON graph. */ - std::vector heads_; - /*! - * \brief A map from constant names to NDArrays for each Constant encountered during - * translation to JSON. The JSON will record only the constant name. The actual NDArray must - * be made available at runtime from a ConstLoaderModule. - */ - std::unordered_map const_name_to_constant_; - /*! - * \brief The domain of the above map, but in order the constants were encountered during - * translation. - */ - Array const_names_; -}; - -} // namespace contrib -} // namespace backend -} // namespace relay -} // namespace tvm -#endif // TVM_RELAY_BACKEND_CONTRIB_CODEGEN_JSON_CODEGEN_JSON_H_ diff --git a/src/relay/backend/contrib/constant_transforms.cc b/src/relay/backend/contrib/constant_transforms.cc deleted file mode 100644 index 45669b5ef271..000000000000 --- a/src/relay/backend/contrib/constant_transforms.cc +++ /dev/null @@ -1,52 +0,0 @@ -/* - * 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 "constant_transforms.h" - -#include - -#include "../../transforms/fold_constant.h" -#include "../../transforms/pattern_utils.h" -#include "../../transforms/simplify_expr.h" - -/*! - * \file src/relay/backend/contrib/constant_transforms.cc - * \brief Transforms applied to constant operations during codegen for BYOC backends. - */ - -namespace tvm { -namespace relay { -namespace contrib { - -Constant TransposeWeights(const Constant& data, const std::string& source_layout, - const std::string& target_layout) { - Array transpose_matrix; - for (const char& c : target_layout) { - int pos = source_layout.find(c); - transpose_matrix.push_back(pos); - } - Expr transpose = MakeTranspose(data, transpose_matrix); - transpose = InferType(transform::FoldConstantExpr(transpose)); - Constant transposed_data = Downcast(transpose); - return transposed_data; -} - -} // namespace contrib -} // namespace relay -} // namespace tvm diff --git a/src/relay/backend/contrib/constant_transforms.h b/src/relay/backend/contrib/constant_transforms.h deleted file mode 100644 index f642564115b6..000000000000 --- a/src/relay/backend/contrib/constant_transforms.h +++ /dev/null @@ -1,50 +0,0 @@ -/* - * 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 src/relay/backend/contrib/constant_transforms.h - * \brief Transforms applied to constant operations during codegen for BYOC backends. - */ - -#ifndef TVM_RELAY_BACKEND_CONTRIB_CONSTANT_TRANSFORMS_H_ -#define TVM_RELAY_BACKEND_CONTRIB_CONSTANT_TRANSFORMS_H_ - -#include - -#include - -namespace tvm { -namespace relay { -namespace contrib { - -/*! - *\brief Transpose weights from `source_layout` to `target_layout` - * - * \param data The constant expression to transpose. - * \param source_layout The current layout of the constant e.g. "OHWI". - * \param target_layout The target layout of the constant e.g. "HWIO". - */ -Constant TransposeWeights(const Constant& data, const std::string& source_layout, - const std::string& target_layout); - -} // namespace contrib -} // namespace relay -} // namespace tvm - -#endif // TVM_RELAY_BACKEND_CONTRIB_CONSTANT_TRANSFORMS_H_ diff --git a/src/relay/backend/contrib/cublas/target.cc b/src/relay/backend/contrib/cublas/target.cc deleted file mode 100644 index 45d3aaa314ae..000000000000 --- a/src/relay/backend/contrib/cublas/target.cc +++ /dev/null @@ -1,44 +0,0 @@ -/* - * 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 src/relay/backend/contrib/cudnn/target.cc - * \brief Registers the "cublas" external codegen TargetKind. - */ - -#include - -namespace tvm { -namespace relay { -namespace contrib { - -/*! - * \brief This external codegen target can use the CuBLAS library linked into the TVM runtime. - * - Patterns and custom compiler: python/tvm/relay/op/contrib/cublas.py - * - Custom schedules: python/tvm/contrib/cublas.py - * - Runtime: src/runtime/contrib/cublas/cublas.cc - * - * CuBLAS can also be used via the "-libs=cublas" Target option. - */ -TVM_REGISTER_TARGET_KIND("cublas", kDLCUDA) - .set_attr(tvm::attr::kIsExternalCodegen, Bool(true)); - -} // namespace contrib -} // namespace relay -} // namespace tvm diff --git a/src/relay/backend/contrib/cudnn/target.cc b/src/relay/backend/contrib/cudnn/target.cc deleted file mode 100644 index 1f1117391209..000000000000 --- a/src/relay/backend/contrib/cudnn/target.cc +++ /dev/null @@ -1,42 +0,0 @@ -/* - * 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 src/relay/backend/contrib/cudnn/target.cc - * \brief Registers the "cudnn" external codegen TargetKind. - */ - -#include - -namespace tvm { -namespace relay { -namespace contrib { - -/*! - * \brief This external codegen target can use the CuDNN library linked into the TVM runtime. - * - Patterns and custom compiler: python/tvm/relay/op/contrib/cudnn.py - * - Custom schedules: python/tvm/contrib/cudnn.py - * - Runtime: src/runtime/contrib/cudnn/ *.cc - */ -TVM_REGISTER_TARGET_KIND("cudnn", kDLCUDA) - .set_attr(tvm::attr::kIsExternalCodegen, Bool(true)); - -} // namespace contrib -} // namespace relay -} // namespace tvm diff --git a/src/relay/backend/contrib/cutlass/codegen.cc b/src/relay/backend/contrib/cutlass/codegen.cc deleted file mode 100644 index 354b25e509a6..000000000000 --- a/src/relay/backend/contrib/cutlass/codegen.cc +++ /dev/null @@ -1,442 +0,0 @@ -/* - * 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 src/relay/backend/contrib/cutlass/codegen.cc - * \brief The 'custom' compilation pass for CUTLASS (invoked by the RelayToTIRTargetHook pass). - */ - -#include "codegen.h" - -#include -#include -#include -#include -#include -#include - -#include -#include -#include -#include -#include - -#include "../../../transforms/compiler_function_utils.h" -#include "../../utils.h" -#include "../codegen_c/codegen_c.h" - -namespace tvm { -namespace relay { -namespace contrib { -namespace cutlass { - -std::string EmitSignature(const std::vector& out, const std::string& func_id, - const std::vector& arg_names) { - std::ostringstream code_stream_; - code_stream_ << "void " << func_id << "_("; - for (const auto& arg_name : arg_names) { - code_stream_ << "DLTensor* " << arg_name << ", "; - } - for (size_t i = 0; i < out.size() - 1; ++i) { - code_stream_ << "DLTensor* out" << i << ", "; - } - code_stream_ << "DLTensor* out" << out.size() - 1 << ")"; - return code_stream_.str(); -} - -runtime::Module Finalize(const std::string& code, const Array& func_names) { - ICHECK(!func_names.empty()) - << "Should only create CUTLASS CSourceModule if there is at least one CUTLASS partition"; - - std::ostringstream default_headers; - default_headers << "#include \n"; - default_headers << "#include \n"; - default_headers << "#include \n"; - default_headers << "#include \n"; - default_headers << "#include \n"; - default_headers << "#include \n"; - default_headers << "#include \n"; - - const auto* pf = runtime::Registry::Get("runtime.CSourceModuleCreate"); - ICHECK(pf != nullptr) << "Cannot find CSource module to create the external runtime module"; - VLOG(1) << "Generated CUTLASS code:" << std::endl << code; - return (*pf)(default_headers.str() + code, "cu", func_names, /*const_vars=*/Array()); -} - -class CodegenResultNode : public Object { - public: - String code; - Array headers; - - void VisitAttrs(AttrVisitor* v) { - v->Visit("code", &code); - v->Visit("headers", &headers); - } - static constexpr const char* _type_key = "contrib.cutlass.CodegenResult"; - TVM_DECLARE_FINAL_OBJECT_INFO(CodegenResultNode, Object); -}; - -class CodegenResult : public ObjectRef { - public: - CodegenResult(String code, Array headers) { - auto n = make_object(); - n->code = std::move(code); - n->headers = std::move(headers); - data_ = std::move(n); - } - - TVM_DEFINE_OBJECT_REF_METHODS(CodegenResult, ObjectRef, CodegenResultNode) -}; - -TVM_REGISTER_NODE_TYPE(CodegenResultNode); - -TVM_REGISTER_GLOBAL("contrib.cutlass.CodegenResult") - .set_body_typed([](String code, Array headers) { - return CodegenResult(code, headers); - }); - -GenerateBodyOutput GenerateBody(const std::string& func_name, const std::string& ext_func_id, - const std::vector& output_types, - const Array& func_args, const Map& attrs, - int* buf_idx) { - // Make function call with input buffers when visiting arguements - ICHECK_GT(func_args.size(), 0); - std::ostringstream decl_stream; - decl_stream << "(" << func_args[0]; - for (size_t i = 1; i < func_args.size(); ++i) { - decl_stream << ", " << func_args[i]; - } - GenerateBodyOutput ret; - for (const auto& out_type : output_types) { - const std::string out = "out" + std::to_string(*buf_idx++); - decl_stream << ", " << out; - Output output; - output.name = out; - output.dtype = out_type; - output.need_copy = false; - ret.outputs.push_back(output); - } - decl_stream << ");"; - - const auto* instantiate_template_func = - runtime::Registry::Get("contrib.cutlass.instantiate_template"); - ICHECK(instantiate_template_func); - - CodegenResult codegen_res = (*instantiate_template_func)(func_name, attrs, func_args); - ret.decl = codegen_res->code; - ret.headers = codegen_res->headers; - - return ret; -} - -namespace { - -/*! \brief Return the "cutlass" Target instance to use to guide compilation. */ -Target GetCutlassTarget() { - Target target = Target::Current(/*allow_not_defined=*/true); - if (!target.defined() || target->kind->name != "cutlass") { - // Use the default CUTLASS compilation options if no specific "cutlass" target was given - // in the overall targets list. In that case target_hooks.cc will invoke the custom pass - // without pushing any target instance onto the implicit target stack. - target = Target("cutlass"); - } - return target; -} - -class CodegenCutlass : public backend::MemoizedExprTranslator>, - public CodegenCBase { - public: - CodegenCutlass(const std::string& id, const Map& attrs) { - this->ext_func_id_ = id; - this->attrs_ = attrs; - } - - std::vector VisitExprDefault_(const Object* op) final { - LOG(FATAL) << "Cutlass codegen doesn't support: " << op->GetTypeKey(); - } - - std::vector VisitExpr_(const VarNode* node) final { - ext_func_args_.push_back(GetRef(node)); - Output output; - output.name = node->name_hint(); - return {output}; - } - - std::vector VisitExpr_(const CallNode* call) final { - const auto* func = call->op.as(); - ICHECK(func) << "Only composite function is supported for CUTLASS."; - GenerateBodyOutput ret = GenerateCompositeFunctionCall(func, call); - ext_func_body_.push_back(ret.decl); - headers_ = ret.headers; - return ret.outputs; - } - - std::string JIT(const std::vector& out) { - std::vector arg_names; - for (const auto& arg : ext_func_args_) { - arg_names.push_back(arg->name_hint()); - } - - code_stream_ << EmitSignature(out, ext_func_id_, arg_names) << "{\n"; - - this->EnterScope(); - - // Function body - for (auto decl : buf_decl_) { - this->PrintIndents(); - code_stream_ << decl << "\n"; - } - code_stream_ << "\n"; - for (auto stmt : ext_func_body_) { - this->PrintIndents(); - code_stream_ << stmt << "\n"; - } - - this->ExitScope(); - code_stream_ << "}\n"; - - this->GenerateBackendCFunc(ext_func_id_, ext_func_args_, /*const_arr_name=*/"", out, true); - return code_stream_.str(); - } - - Array GetHeaders() { return headers_; } - - private: - Array GetArgumentNames(const CallNode* call) { - Array arg_names; - for (size_t i = 0; i < call->args.size(); ++i) { - auto res = VisitExpr(call->args[i]); - for (const auto& out : res) { - arg_names.push_back(out.name); - } - } - return arg_names; - } - - // Is node `x` an ancestor of `y`? - bool IsAncestor(const CallNode* x, const CallNode* y) { - if (x == y) return true; - for (auto arg : y->args) { - const CallNode* arg_ptr = arg.as(); - if (arg_ptr && IsAncestor(x, arg_ptr)) return true; - } - return false; - } - - GenerateBodyOutput GenerateCompositeFunctionCall(const FunctionNode* callee, - const CallNode* caller) { - const auto pattern_name_opt = callee->GetAttr(attr::kComposite); - ICHECK(pattern_name_opt.defined()) << "Only functions with composite attribute are supported."; - const std::string pattern_name = pattern_name_opt.value(); - - if (pattern_name.find("conv2d") != std::string::npos && - pattern_name.find("residual") != std::string::npos) { - const CallNode* current_call = callee->body.as(); - bool has_relu = current_call->args.size() == 1; - const CallNode* binop = has_relu ? current_call->args[0].as() : current_call; - ICHECK(binop->args.size() == 2); - // Figure out which of the first or second argument corresponds to the residual input - // The root conv2d call can be reached via the other input of the binary op - int residual_index; - if (binop->args[1].as()) { - residual_index = 1; - } else if (binop->args[0].as()) { - residual_index = 0; - } else { - const CallNode* lhs = binop->args[0].as(); - const CallNode* rhs = binop->args[1].as(); - ICHECK(lhs && rhs); - // The residual input should be an ancestor of the non-residual input - residual_index = IsAncestor(rhs, lhs) ? 1 : 0; - } - const auto residual_input = binop->args[residual_index]; - auto call_args = GetArgumentNames(caller); - auto func_args = call_args; - if (call_args.size() == 3) { - // TODO(masahi): This code assumes that there is always a bias_add in a residual block. - for (size_t i = 0; i < call_args.size(); ++i) { - if (callee->params[i] == residual_input) { - auto residual_input_name = call_args[i]; - func_args.push_back(residual_input_name); - } - } - } else { - ICHECK_EQ(func_args.size(), 4) << "Residual block fusion expects 4 input tensors: data, " - "weight, bias, and residual tensor."; - } - return GenerateBody(caller, pattern_name, func_args, attrs_); - } else { - return GenerateBody(caller, pattern_name, attrs_); - } - - LOG(FATAL) << "Unknown composite function: " << pattern_name; - } - - GenerateBodyOutput GenerateBody(const CallNode* call, const std::string& func_name, - const Array& func_args, - const Map& attrs) { - std::vector out_types; - if (call->checked_type()->IsInstance()) { - auto type_node = call->checked_type().as(); - for (auto field : type_node->fields) { - ICHECK(field->IsInstance()); - out_types.push_back(field); - } - } else if (call->checked_type()->IsInstance()) { - ICHECK(call->checked_type()->IsInstance()); - out_types.push_back(call->checked_type()); - } else { - LOG(FATAL) << "Unrecognized type node: " << AsText(call->checked_type(), false); - } - - std::vector out_types_str; - for (const auto& out_type : out_types) { - out_types_str.push_back(GetDtypeString(out_type.as())); - } - - return cutlass::GenerateBody(func_name, ext_func_id_, out_types_str, func_args, attrs, - &buf_idx_); - } - - GenerateBodyOutput GenerateBody(const CallNode* call, const std::string& func_name, - const Map& attrs) { - auto func_args = GetArgumentNames(call); - return GenerateBody(call, func_name, func_args, attrs); - } - - /*! \brief The id of the external cutlass ext_func. */ - std::string ext_func_id_; - /*! \brief The attrs of the external cutlass ext_func. */ - Map attrs_; - /*! - * \brief The index to track the output buffer. Each kernel will redirect the - * output to a buffer that may be consumed by other kernels. - */ - int buf_idx_{0}; - /*! \brief The arguments used by a wrapped function that calls CUTLASS kernels. */ - Array ext_func_args_; - /*! \brief Statement of the function that will be compiled using CUTLASS kernels. */ - std::vector ext_func_body_; - /*! \brief The declaration of intermediate buffers. */ - std::vector buf_decl_; - /*! \brief Required header-file names. */ - Array headers_; -}; // class CodegenCutlass - -class CutlassModuleCodegen { - public: - explicit CutlassModuleCodegen(IRModule mod) : mod_(std::move(mod)) {} - - runtime::Module CreateCSourceModule() { - for (const auto& entry : mod_->functions) { - if (const auto* function_node = GetCutlassFunctionNode(entry.second)) { - GenCutlassFunc(GetRef(function_node)); - } - } - return Finalize(code_stream_.str(), func_names_); - } - - private: - void GenCutlassFunc(const Function& function) { - ICHECK(function.defined()) << "Input error: expect a Relay function."; - - // Record the external symbol for runtime lookup. - Optional opt_global_symbol = function->GetAttr(tvm::attr::kGlobalSymbol); - ICHECK(opt_global_symbol.defined()) - << "CUTLASS functions must have a " << tvm::attr::kGlobalSymbol << " attribute"; - std::string sid = opt_global_symbol.value(); - if (std::find(func_names_.begin(), func_names_.end(), sid) != func_names_.end()) { - // Already emitted. - return; - } - func_names_.push_back(sid); - - const auto* attrs = function->attrs.as(); - ICHECK(attrs != nullptr); - const auto dict = attrs->dict; - CodegenCutlass builder(sid, dict); - VLOG(1) << "Creating cutlass C code for '" << sid << "' from:\n" << PrettyPrint(function); - auto out = builder.VisitExpr(function->body); - auto code = builder.JIT(out); - for (const auto& header : builder.GetHeaders()) { - code_stream_ << "#include <" << header << ">\n"; - } - code_stream_ << "\n" + code; - } - - /*! - * \brief Returns \p expr as function if it is a \p Function with "Compiler" attribute - * value "cutlass". - */ - static const FunctionNode* GetCutlassFunctionNode(const Expr& expr) { - if (const auto* function_node = expr.as()) { - Optional opt_compiler = function_node->GetAttr(attr::kCompiler); - if (opt_compiler.defined() && opt_compiler.value() == "cutlass") { - return function_node; - } - } - return nullptr; - } - - /*! \brief Module we are compiling. */ - IRModule mod_; - /*! \brief The accumulated code stream that will be compiled by NVCC */ - std::ostringstream code_stream_; - /*! \brief The accumulated function names. */ - Array func_names_; -}; // CutlassModuleCodegen - -/*! - * \brief A small shim to redirect to the 'relay.ext.cutlass.compile_for_cutlass' Python - * function which does the main CUTLASS training, c-code generation and compilation steps. - */ -tvm::transform::Pass CompileForCutlassImpl() { - auto pass_func = [=](IRModule mod, const tvm::transform::PassContext& pass_ctx) { - VLOG(1) << "CompileForCutlass input:" << std::endl << PrettyPrint(mod); - const auto* pf = runtime::Registry::Get("relay.ext.cutlass.compile_for_cutlass"); - ICHECK(pf != nullptr) << "Cannot find compile_for_cutlass function"; - Target target = GetCutlassTarget(); - runtime::Module runtime_mod = (*pf)(mod, target); - Array external_mods = - mod->GetAttr>(tvm::attr::kExternalMods).value_or({}); - external_mods.push_back(runtime_mod); - return WithAttr(mod, tvm::attr::kExternalMods, external_mods); - }; - return tvm::transform::CreateModulePass(pass_func, 0, "CompileForCutlass", {}); -} - -runtime::Module CreateCSourceModule(const IRModule& mod) { - VLOG(1) << "Creating CUTLASS CSource module from:" << std::endl << PrettyPrint(mod); - return CutlassModuleCodegen(mod).CreateCSourceModule(); -} - -} // namespace - -TVM_REGISTER_GLOBAL("relay.ext.cutlass.create_c_source_module").set_body_typed(CreateCSourceModule); - -tvm::transform::Pass CompileForCutlass() { - return transform::Sequential( - {transform::OutlineCompilerFunctionsWithExistingGlobalSymbols("cutlass"), - CompileForCutlassImpl(), transform::MarkCompilerFunctionsAsExtern("cutlass")}); -} - -} // namespace cutlass -} // namespace contrib -} // namespace relay -} // namespace tvm diff --git a/src/relay/backend/contrib/cutlass/codegen.h b/src/relay/backend/contrib/cutlass/codegen.h deleted file mode 100644 index 03b8e6afbddc..000000000000 --- a/src/relay/backend/contrib/cutlass/codegen.h +++ /dev/null @@ -1,69 +0,0 @@ -/* - * 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 src/relay/backend/contrib/cutlass/codegen.h - * \brief The 'custom' compilation pass for CUTLASS (invoked by the RelayToTIRTargetHook pass). - */ - -#ifndef TVM_RELAY_BACKEND_CONTRIB_CUTLASS_CODEGEN_H_ -#define TVM_RELAY_BACKEND_CONTRIB_CUTLASS_CODEGEN_H_ - -#include - -#include -#include - -#include "../codegen_c/codegen_c.h" - -namespace tvm { -namespace relay { -namespace contrib { -namespace cutlass { - -/*! - * \brief Returns the pass which replaces all calls to "Primitive" functions with "Compiler" - * attribute of "cutlass" with an call to an extern, and binds a \p runtime::StaticLibrary - * to the IRModule's "external_mods" attribute containing compiled implementations of - * those functions using the CUTLASS C++ template library. - */ -transform::Pass CompileForCutlass(); - -// The rest is sparsely documented since they are exposed only for code sharing between Relay -// and Relax backend implementations. - -/*! \brief Emit the function signature for a kernel */ -std::string EmitSignature(const std::vector& out, - const std::string& func_id, const std::vector& arg_names); - -/*! \brief Generate the body of the kernel */ -GenerateBodyOutput GenerateBody(const std::string& func_name, const std::string& ext_func_id, - const std::vector& output_types, - const Array& func_args, const Map& attrs, - int* buf_idx); - -/*! \brief Create a C-source module from the given kernel string */ -runtime::Module Finalize(const std::string& code, const Array& func_names); - -} // namespace cutlass -} // namespace contrib -} // namespace relay -} // namespace tvm - -#endif // TVM_RELAY_BACKEND_CONTRIB_CUTLASS_CODEGEN_H_ diff --git a/src/relay/backend/contrib/cutlass/target.cc b/src/relay/backend/contrib/cutlass/target.cc deleted file mode 100644 index ea040f6ff56a..000000000000 --- a/src/relay/backend/contrib/cutlass/target.cc +++ /dev/null @@ -1,74 +0,0 @@ -/* - * 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 src/relay/backend/contrib/cutlass/target.cc - * \brief Registers the "cutlass" external codegen TargetKind. - */ - -#include - -#include "./codegen.h" - -namespace tvm { -namespace relay { -namespace contrib { -namespace cutlass { - -/*! - * \brief This external codegen target can use the CUTLASS template library included in - * TVM's 3rdparty/cutlass. - * - Patterns: python/tvm/relay/op/contrib/cutlass.py - * - Custom compiler: python/tvm/contrib/cutlass/build.py, - * src/relay/backend/contrib/cutlass/codegen.cc - */ -TVM_REGISTER_TARGET_KIND("cutlass", kDLCUDA) - .set_attr(tvm::attr::kIsExternalCodegen, runtime::Bool(true)) - .set_attr("RelayToTIR", CompileForCutlass()) - // An integer specifying the compute capability. For example, 75 for Turing and - // 80 or 86 for Ampere. - .add_attr_option("sm", runtime::Int(80)) - // Whether to use slower but very accurate (compared to tf32) 3xtf32 mode for - // fp32 inputs on tensorcore. - .add_attr_option("use_3xtf32", runtime::Bool(true)) - // Split factor candidates for split-K GEMM. If split-K > 1, the GEMM K-loop is computed in - // parallel across split-K blocks, and a separate global reduction kernel is launched to - // accumulate partial reductions. The profiler will pick the best split-k factor from the - // given candidate list. Note that the larger split-K factor requires a larger workspace. - // Currently, parallel split-k has been tested only for wgrad. For GEMM and other conv2d - // kinds, split_k_slices is ignored. - .add_attr_option>("split_k_slices", Array{runtime::Int(1)}) - // When True, profile all kernel variants with smaller alignments than the largest possible. - .add_attr_option("profile_all_alignments", runtime::Bool(false)) - // Whether to profile all candidate kernels, or stop profiling after the first applicable kernel - // is found. - .add_attr_option("find_first_valid", runtime::Bool(false)) - // Whether to compile profiler executables for different kernels in parallel. - .add_attr_option("use_multiprocessing", runtime::Bool(false)) - // Number of threads to use during compilation, or -1 to use number of cpus. - .add_attr_option("threads", runtime::Int(-1)) - // Whether to replace sigmoid with tanh. - .add_attr_option("use_fast_math", runtime::Bool(false)) - // A temporary directory where intermediate compiled artifacts will be stored. - .add_attr_option("tmp_dir", String("./tmp")); - -} // namespace cutlass -} // namespace contrib -} // namespace relay -} // namespace tvm diff --git a/src/relay/backend/contrib/dnnl/codegen.cc b/src/relay/backend/contrib/dnnl/codegen.cc deleted file mode 100644 index 3b7bc8f10d50..000000000000 --- a/src/relay/backend/contrib/dnnl/codegen.cc +++ /dev/null @@ -1,629 +0,0 @@ -/* - * 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 src/relay/backend/contrib/dnnl/codegen.cc - * \brief Implementation of DNNL codegen APIs. - */ - -#include -#include -#include -#include -#include -#include - -#include -#include -#include - -#include "../../utils.h" -#include "comp_op_matcher.h" - -#ifdef USE_JSON_RUNTIME -#include "../../../../runtime/contrib/json/json_node.h" -#include "../codegen_json/codegen_json.h" -#else -#include "../codegen_c/codegen_c.h" -#endif - -namespace tvm { -namespace relay { -namespace contrib { - -using namespace backend; - -/*! - * \brief Replace var expr which bind with args of call node - * - * \param args vector of expression (contains vars or constant nodes) - * \param cn call node which describe mapping of internal body vars with args - * \return updated vector of expressions - */ -static tvm::Array BindToCallNodeArgs(const std::vector& args, const CallNode* cn) { - tvm::Array res; - for (const auto& arg : args) { - if (arg->IsInstance()) { - res.push_back(arg); - } else { - auto body_params = cn->op.as()->params; - auto found = std::find(body_params.begin(), body_params.end(), arg); - ICHECK(found != body_params.end()); - auto idx = std::distance(body_params.begin(), found); - res.push_back(cn->args[idx]); - } - } - return res; -} - -#ifndef USE_JSON_RUNTIME // C source runtime -inline size_t GetShape1DSize(const Type& type) { - const auto shape = GetShape(type); - return std::accumulate(shape.begin(), shape.end(), 1, std::multiplies()); -} - -inline std::string GetShapeString(std::vector shape) { - std::string v = "std::vector{"; - for (auto s : shape) { - v += std::to_string(s) + ","; - } - v += "}"; - return v; -} - -std::vector Conv2d(const CallNode* call) { - std::vector args; - const auto* conv2d_attr = call->attrs.as(); - ICHECK(conv2d_attr); - - auto ishape = GetShape(call->args[0]->checked_type()); - auto wshape = GetShape(call->args[1]->checked_type()); - - // Args: N, C, H, W - for (auto s : ishape) { - args.push_back(std::to_string(s)); - } - - // Args: O, G, Ph0, Pw0, Ph1, Pw1, Kh, Kw, Sh, Sw - args.push_back(std::to_string(wshape[0])); - args.push_back(std::to_string(conv2d_attr->groups)); - args.push_back(std::to_string(conv2d_attr->padding[0].as()->value)); - args.push_back(std::to_string(conv2d_attr->padding[1].as()->value)); - args.push_back(std::to_string(conv2d_attr->padding[2].as()->value)); - args.push_back(std::to_string(conv2d_attr->padding[3].as()->value)); - args.push_back(std::to_string(wshape[2])); - args.push_back(std::to_string(wshape[3])); - args.push_back(std::to_string(conv2d_attr->strides[0].as()->value)); - args.push_back(std::to_string(conv2d_attr->strides[1].as()->value)); - - return args; -} - -std::vector Dense(const CallNode* call) { - std::vector args; - auto ishape = GetShape(call->args[0]->checked_type()); - auto wshape = GetShape(call->args[1]->checked_type()); - - // Args: N, C, O - args.push_back(std::to_string(ishape[0])); - args.push_back(std::to_string(ishape[1])); - args.push_back(std::to_string(wshape[0])); - - return args; -} - -std::vector Relu(const CallNode* call) { - std::vector args; - auto ishape = GetShape(call->args[0]->checked_type()); - // Args: N, C, H, W - args.push_back(GetShapeString(ishape)); - return args; -} - -std::vector BatchNorm(const CallNode* call) { - std::vector args; - const auto* bn_attr = call->attrs.as(); - auto ishape = GetShape(call->args[0]->checked_type()); - - // Args: N, C, H, W - for (auto s : ishape) { - args.push_back(std::to_string(s)); - } - - // Args: epsilon - args.push_back(std::to_string(bn_attr->epsilon)); - - return args; -} - -// should comply with src/runtime/contrib/dnnl/dnnl.cc -#define DNNL_BINARY_ADD 0 -#define DNNL_BINARY_MUL 1 - -std::vector Add(const CallNode* call) { - std::vector args; - auto ishape = GetShape(call->args[0]->checked_type()); - args.push_back(std::to_string(DNNL_BINARY_ADD)); - // Args: H, W - args.push_back(GetShapeString(ishape)); - return args; -} - -std::vector Multiply(const CallNode* call) { - std::vector args; - auto ishape = GetShape(call->args[0]->checked_type()); - args.push_back(std::to_string(DNNL_BINARY_MUL)); - // Args: H, W - args.push_back(GetShapeString(ishape)); - return args; -} - -// TODO(@zhiics, @comaniac): This is a basic implementation. We should implement -// all utilities and make a base class for users to implement. -class CodegenDNNL : public MemoizedExprTranslator>, public CodegenCBase { - public: - explicit CodegenDNNL(const std::string& id) { this->ext_func_id_ = id; } - - std::vector VisitExprDefault_(const Object* op) final { - LOG(FATAL) << "DNNL codegen doesn't support: " << op->GetTypeKey(); - } - - std::vector VisitExpr_(const VarNode* node) final { - ext_func_args_.push_back(GetRef(node)); - Output output; - output.name = node->name_hint(); - return {output}; - } - - std::vector VisitExpr_(const TupleNode* node) final { - std::vector outs; - for (auto field : node->fields) { - auto res = VisitExpr(field); - ICHECK_EQ(res.size(), 1U) << "Do not support tuple nest"; - outs.push_back(res[0]); - } - return outs; - } - - std::vector VisitExpr_(const TupleGetItemNode* op) final { - auto res = VisitExpr(op->tuple); - ICHECK_GT(res.size(), static_cast(op->index)); - - // Only keep the item we want for the child node. - // FIXME(@comaniac): The other items should still be requried for the primary outputs. - return {res[op->index]}; - } - - std::vector VisitExpr_(const ConstantNode* cn) final { - Output output; - // Get const: static_cast(dnnl_0_consts[0]->data) - output.name = CreateDataReference(ext_func_id_, const_idx_); - output.dtype = "float"; - - // Generate the global variable for needed ndarrays - if (const_array_name_.empty()) { - const_array_name_ = CreateNDArrayPool(ext_func_id_); - std::string checker = CreateInitChecker(ext_func_id_); - ext_func_body_.insert(ext_func_body_.begin(), checker); - } - - // Give the ndarray a unique name to ease the initialization of it at - // runtime. - std::string const_symbol = "dnnl_" + ext_func_id_; - std::string const_var_name = CreateConstVar(const_symbol, const_idx_); - const_vars_.push_back(const_var_name); - const_idx_++; - - const auto* type_node = cn->checked_type().as(); - ICHECK(type_node); - ICHECK_EQ(GetDtypeString(type_node), "float") << "Only float is supported for now."; - - return {output}; - } - - std::vector VisitExpr_(const CallNode* call) final { - GenerateBodyOutput ret; - if (const auto* func = call->op.as()) { - ret = GenerateCompositeFunctionCall(func, call); - } else { - ret = GenerateOpCall(call); - } - - buf_decl_.insert(buf_decl_.end(), ret.buffers.begin(), ret.buffers.end()); - ext_func_body_.push_back(ret.decl); - return ret.outputs; - } - - std::string JIT(const std::vector& out) { - return JitImpl(ext_func_id_, ext_func_args_, buf_decl_, ext_func_body_, const_array_name_, out); - } - - private: - std::vector GetArgumentNames(const CallNode* call) { - std::vector arg_names; - for (size_t i = 0; i < call->args.size(); ++i) { - auto res = VisitExpr(call->args[i]); - for (const auto& out : res) { - arg_names.push_back(out.name); - } - } - return arg_names; - } - - GenerateBodyOutput GenerateOpCall(const CallNode* call) { - const auto* op_node = call->op.as(); - ICHECK(op_node) << "Expect OpNode, but got " << call->op->GetTypeKey(); - - using ArgFunType = std::function(const CallNode*)>; - static const std::map> op_map = { - {"nn.conv2d", {"dnnl_conv2d", Conv2d}}, {"nn.dense", {"dnnl_dense", Dense}}, - {"nn.relu", {"dnnl_relu", Relu}}, {"nn.batch_norm", {"dnnl_bn", BatchNorm}}, - {"add", {"dnnl_binary_op", Add}}, {"multiply", {"dnnl_binary_op", Multiply}}, - }; - - const auto op_name = GetRef(op_node)->name; - const auto iter = op_map.find(op_name); - if (iter != op_map.end()) { - return GenerateBody(call, iter->second.first, iter->second.second(call)); - } - - LOG(FATAL) << "Unsupported op: " << AsText(call->op, false); - } - - GenerateBodyOutput GenerateCompositeFunctionCall(const FunctionNode* callee, - const CallNode* caller) { - const auto pattern_name = callee->GetAttr(attr::kComposite); - ICHECK(pattern_name.defined()) << "Only functions with composite attribute supported"; - - if (pattern_name == "dnnl.conv2d_bias_relu") { - const auto* conv_call = - GetRootCall(callee->body.as(), 2, {"nn.conv2d", "add", "nn.relu"}); - return GenerateBody(conv_call, "dnnl_fused_conv2d_bias_relu", GetArgumentNames(caller), - Conv2d(conv_call)); - } else if (pattern_name == "dnnl.conv2d_relu") { - const auto* conv_call = GetRootCall(callee->body.as(), 1, - (const std::vector){"nn.conv2d", "nn.relu"}); - return GenerateBody(conv_call, "dnnl_fused_conv2d_relu", GetArgumentNames(caller), - Conv2d(conv_call)); - } - - LOG(FATAL) << "Unknown composite function:" << pattern_name; - } - - GenerateBodyOutput GenerateBody(const CallNode* root_call, const std::string& func_name, - const std::vector& attribute_args) { - return GenerateBody(root_call, func_name, GetArgumentNames(root_call), attribute_args); - } - - GenerateBodyOutput GenerateBody(const CallNode* root_call, const std::string& func_name, - const std::vector& func_args, - const std::vector& attribute_args) { - // Make function call with input buffers when visiting arguments - ICHECK_GT(func_args.size(), 0); - std::ostringstream decl_stream; - decl_stream << "(" << func_args[0]; - for (size_t i = 1; i < func_args.size(); ++i) { - decl_stream << ", " << func_args[i]; - } - - // Analyze the output buffers - std::vector out_types; - if (root_call->checked_type()->IsInstance()) { - auto type_node = root_call->checked_type().as(); - for (auto field : type_node->fields) { - ICHECK(field->IsInstance()); - out_types.push_back(field); - } - } else if (root_call->checked_type()->IsInstance()) { - ICHECK(root_call->checked_type()->IsInstance()); - out_types.push_back(root_call->checked_type()); - } else { - LOG(FATAL) << "Unrecognized type node: " << AsText(root_call->checked_type(), false); - } - - GenerateBodyOutput ret; - for (const auto& out_type : out_types) { - this->PrintIndents(); - const std::string out = "buf_" + std::to_string(buf_idx_++); - const auto out_size = GetShape1DSize(out_type); - decl_stream << ", " << out; - - Output output; - output.name = out; - output.size = out_size; - output.dtype = GetDtypeString(out_type.as()); - output.need_copy = true; - ret.buffers.push_back("float* " + out + " = (float*)std::malloc(4 * " + - std::to_string(out_size) + ");"); - ret.outputs.push_back(output); - } - - // Attach attribute arguments - for (size_t i = 0; i < attribute_args.size(); ++i) { - decl_stream << ", " << attribute_args[i]; - } - decl_stream << ");"; - ret.decl = func_name + decl_stream.str(); - return ret; - } - - /*! \brief The id of the external dnnl ext_func. */ - std::string ext_func_id_{""}; - /*! - * \brief The index to track the output buffer. Each kernel will redirect the - * output to a buffer that may be consumed by other kernels. - */ - int buf_idx_{0}; - /*! \brief The index of global constants. */ - int const_idx_{0}; - /*! \brief The arguments used by a wrapped function that calls DNNL kernels. */ - Array ext_func_args_; - /*! \brief Statement of the function that will be compiled using DNNL kernels. */ - std::vector ext_func_body_; - /*! \brief The array declared to store the constant values. */ - std::string const_array_name_; - /*! \brief The declaration of intermeidate buffers. */ - std::vector buf_decl_; - /*! \brief The variable name to constant mapping. */ - Array const_vars_; - - friend class DNNLModuleCodegen; -}; - -/*! - * \brief The DNNL codegen helper to generate wrapepr function calls of DNNL - * libraries. The code is a CSourceModule that can be compiled separately and - * linked together with a DSOModule. - */ -class DNNLModuleCodegen : public CSourceModuleCodegenBase { - public: - // Create a corresponding DNNL function for the given relay Function. - std::pair> GenDNNLFunc(const Function& func) { - ICHECK(func.defined()) << "Input error: expect a Relay function."; - - // Record the external symbol for runtime lookup. - auto sid = GetExtSymbol(func); - - CodegenDNNL builder(sid); - auto out = builder.VisitExpr(func->body); - code_stream_ << builder.JIT(out); - - return {sid, builder.const_vars_}; - } - - /*! - * \brief The overridden function that will create a CSourceModule. In order - * to compile the generated C source code, users need to specify the paths to - * some libraries, including some TVM required and dnnl specific ones. To make - * linking simpiler, the DNNL kernels are wrapped in a TVM compatible manner - * and live under tvm/src/runtime/contrib/dnnl folder. - * - * \param ref An object ref that could be either a Relay function or module. - * - * \return The runtime module that contains C source code. - */ - runtime::Module CreateCSourceModule(const ObjectRef& ref) override { - // Create headers - code_stream_ << "#include \n"; - code_stream_ << "#include \n"; - code_stream_ << "#include \n"; - code_stream_ << "#include \n"; - code_stream_ << "#include \n"; - code_stream_ << "#include \n"; - code_stream_ << "#include \n"; - // dnnl_kernel file is saved under src/runtime/contrib/dnnl so that we don't - // expose it to ordinary users. To make export_library use it, users need to - // pass -I${PATH_TO_TVM}/src/runtime/contrib - code_stream_ << "#include \n"; - code_stream_ << "using namespace tvm::runtime;\n"; - code_stream_ << "using namespace tvm::runtime::contrib;\n"; - code_stream_ << "\n"; - - ICHECK(ref->IsInstance()); - auto res = GenDNNLFunc(Downcast(ref)); - std::string code = code_stream_.str(); - String sym = std::get<0>(res); - Array variables = std::get<1>(res); - - // Create a CSource module - const auto* pf = runtime::Registry::Get("runtime.CSourceModuleCreate"); - ICHECK(pf != nullptr) << "Cannot find csource module to create the external runtime module"; - // TODO(@manupa-arm): pass the function names to enable system-lib creation - return (*pf)(code, "c", Array{sym}, variables); - } - - private: - /*! - * \brief The code stream that prints the code that will be compiled using - * external codegen tools. - */ - std::ostringstream code_stream_; -}; - -#else // DNNL JSON runtime - -/*! \brief Serializer to DNNL JSON runtime module */ -class DNNLJSONSerializer : public backend::contrib::JSONSerializer { - using JSONGraphNode = tvm::runtime::json::JSONGraphNode; - using JSONGraphNodeEntry = tvm::runtime::json::JSONGraphNodeEntry; - - public: - DNNLJSONSerializer(const std::string& symbol, const Expr& expr) - : JSONSerializer("dnnl_" + symbol, expr) {} - - std::vector VisitExpr_(const CallNode* cn) override { - Expr expr = GetRef(cn); - std::string name; - tvm::Array args; - std::unordered_map extra_attrs; - - const CallNode* call = cn; - if (const auto* op_node = cn->op.as()) { - name = op_node->name; - args = cn->args; - } else if (const auto* fn = cn->op.as()) { - auto comp = fn->GetAttr(attr::kComposite); - ICHECK(comp.defined()) << "DNNL JSON runtime only supports composite functions."; - name = comp.value(); - - if (name.find("dnnl.deconv2d") != std::string::npos) { - call = GetRootCall(fn->body.as(), 10, "nn.conv2d_transpose"); - ICHECK(call->op.as()) << "Not op node"; - } else if (name.find("dnnl.deconv3d") != std::string::npos) { - call = GetRootCall(fn->body.as(), 10, "nn.conv3d_transpose"); - ICHECK(call->op.as()) << "Not op node"; - } else if (name.find("dnnl.conv1d") != std::string::npos) { - call = GetRootCall(fn->body.as(), 10, "nn.conv1d"); - ICHECK(call->op.as()) << "Not op node"; - } else if (name.find("dnnl.conv2d") != std::string::npos) { - call = GetRootCall(fn->body.as(), 10, "nn.conv2d"); - ICHECK(call->op.as()) << "Not op node"; - } else if (name.find("dnnl.conv3d") != std::string::npos) { - call = GetRootCall(fn->body.as(), 10, "nn.conv3d"); - ICHECK(call->op.as()) << "Not op node"; - } else if (name.find("dnnl.dense") != std::string::npos) { - call = GetRootCall(fn->body.as(), 10, "nn.dense"); - ICHECK(call->op.as()) << "Not op node"; - } else if (name.find("dnnl.qnn.conv2d") != std::string::npos || - name.find("dnnl.qnn.dense") != std::string::npos) { - std::vector args_loc; - call = ParseComposite(*fn, &extra_attrs, &args_loc); - args = BindToCallNodeArgs(args_loc, cn); - } else { - LOG(FATAL) << "Unrecognized DNNL pattern: " << name; - } - - if (args.empty()) { - args = cn->args; - } - } else { - LOG(FATAL) << "DNNL JSON runtime does not support calls to " << cn->op->GetTypeKey(); - } - - std::vector inputs; - for (const auto& arg : args) { - auto res = VisitExpr(arg); - inputs.insert(inputs.end(), res.begin(), res.end()); - } - auto node = std::make_shared(name, /* name_ */ - "kernel", /* op_type_ */ - inputs, 1 /* num_outputs_ */); - SetCallNodeAttribute(node, call); - // If has post-op `clip`. Assume the last op is clip, add clip's attrs to the pattern attrs. - if (name.find("_clip") != std::string::npos) { - auto clip_call = cn->op.as()->body.as(); - ICHECK(IsOp(clip_call, "clip")); - SetCallNodeAttribute(node, clip_call); - } - // For QNN. - for (const auto& kvp : extra_attrs) node->SetAttr(kvp.first, kvp.second); - - return AddNode(node, GetRef(cn)); - } -}; -#endif - -/*! - * \brief The external compiler/codegen tool. It takes a Relay expression/module and - * compile it into a runtime module. - */ -runtime::Module DNNLCompiler(const ObjectRef& ref) { -#ifdef USE_JSON_RUNTIME - ICHECK(ref->IsInstance()); - auto func = Downcast(ref); - auto func_name = GetExtSymbol(func); - DNNLJSONSerializer serializer(func_name, func); - serializer.serialize(); - std::string graph_json = serializer.GetJSON(); - - // Note that serializer.const_name_to_constant() is ignored. Instead the TECompiler invokes - // a callback which calls backend::UpdateConstants to capture the map before the function - // 'disappears' into lowered form, on the assumption the visit order and thus constant - // names match those generated by the JSONSerializer. - - const auto* pf = runtime::Registry::Get("runtime.DNNLJSONRuntimeCreate"); - ICHECK(pf != nullptr) << "Cannot find JSON runtime module to create"; - auto mod = (*pf)(func_name, graph_json, serializer.const_names()); - return mod; -#else - DNNLModuleCodegen dnnl; - return dnnl.CreateCSourceModule(ref); -#endif -} - -TVM_REGISTER_GLOBAL("relay.ext.dnnl").set_body_typed(DNNLCompiler); - -/*! - * \brief Constant Updater for DNNL JSON runtime - * - * Not all originally existing ConstantNode should be passed to JSON runtime. - * Some of them may be skipped or change ordering. So we have to apply the same traversing through - * the graph as DNNLJSONSerializer. - */ -struct DNNLConstantUpdater : public ConstantUpdater { - public: - DNNLConstantUpdater(const std::string& symbol, - std::unordered_map* params) - : ConstantUpdater("dnnl_" + symbol, params) {} - using ConstantUpdater::VisitExpr_; - - void VisitExpr_(const CallNode* cn) final { - this->VisitSpan(cn->span); - - if (const auto* fn = cn->op.as()) { - std::vector args_loc; - std::unordered_map attrs; - auto root_cn = ParseComposite(*fn, &attrs, &args_loc); - - auto args = root_cn ? BindToCallNodeArgs(args_loc, cn) : cn->args; - - // Customized visit order of args - for (const auto& arg : args) { - this->VisitExpr(arg); - } - } else { - // Original visit order of args - for (auto arg : cn->args) { - this->VisitExpr(arg); - } - } - } -}; - -/*! - * \brief The external compiler/codegen tool. It takes a Relay expression/module and - * produce collection of required constant NDArrays. - */ -Map DNNLConstantUpdaterFunc(Expr expr, std::string symbol) { - // Visit all suitable constant nodes - std::unordered_map res; - DNNLConstantUpdater const_updater(symbol, &res); - const_updater(expr); - - // Convert to tvm::Map - Map ret; - for (const auto& kvp : res) ret.Set(kvp.first, kvp.second); - return ret; -} - -TVM_REGISTER_GLOBAL("relay.ext.dnnl.constant_updater").set_body_typed(DNNLConstantUpdaterFunc); - -} // namespace contrib -} // namespace relay -} // namespace tvm diff --git a/src/relay/backend/contrib/dnnl/comp_op_matcher.h b/src/relay/backend/contrib/dnnl/comp_op_matcher.h deleted file mode 100644 index 364cc6e377ca..000000000000 --- a/src/relay/backend/contrib/dnnl/comp_op_matcher.h +++ /dev/null @@ -1,245 +0,0 @@ -/* - * 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 src/relay/backend/contrib/dnnl/comp_op_matcher.h - * \brief Implement matcher based function to parse complex composite nodes. - */ - -#ifndef TVM_RELAY_BACKEND_CONTRIB_DNNL_COMP_OP_MATCHER_H_ -#define TVM_RELAY_BACKEND_CONTRIB_DNNL_COMP_OP_MATCHER_H_ - -#include - -#include -#include -#include - -#include "../../../ir/dataflow_matcher_impl.h" - -/*! - * \brief Converter value to dmlc attr acceptable format - * - * \tparam T type of value (auto deduction) - * \param val value to convert - * \return resulting dmlc object - */ -template ::value, bool> = true> -dmlc::any dmlc_attr(const T& val) { - std::vector attr; - attr.emplace_back(std::vector{std::to_string(val)}); - return dmlc::any{attr}; -} - -template ::value, bool> = true> -dmlc::any dmlc_attr(const T& val) { - std::vector attr; - attr.emplace_back(std::vector{val}); - return dmlc::any{attr}; -} - -template >::value, bool> = true> -dmlc::any dmlc_attr(const T& val) { - std::vector attr; - attr.emplace_back(val); - return dmlc::any{attr}; -} - -/*! \brief Constructor of const scalar expression with defined type */ -tvm::relay::Expr constant(float val) { - auto value = tvm::runtime::NDArray::Empty({}, tvm::DataType::Float(32), {kDLCPU, 0}); - value.CopyFromBytes(&val, sizeof(val)); - auto res = tvm::relay::Constant(value); - tvm::relay::transform::InferTypeLocal(res); - return res; -} - -/*! - * \brief Simple helper to accumulate composite function arguments and corresponding attributes - * with indexes of them. - */ -class ArgPacker { - public: - ArgPacker(std::unordered_map* attrs, std::vector* args) - : attrs_(attrs), args_(args) {} - - int Put(const tvm::relay::Expr& arg, std::string tag_name = "") { - if (!arg.defined()) return -1; - int idx = args_->size(); - args_->push_back(arg); - if (!tag_name.empty()) { - attrs_->operator[](tag_name) = dmlc_attr(idx); - } - return idx; - } - - private: - std::unordered_map* attrs_; - std::vector* args_; -}; - -const tvm::relay::CallNode* ParseQnnConvComp(const tvm::relay::FunctionNode& comp_fn, - std::unordered_map* ext_attrs, - std::vector* args) { - using namespace tvm::relay; - - // Pattern - auto src = IsWildcard(); - auto wgh = IsWildcard(); - auto sum_src = IsWildcard(); - auto bias = IsConstant(); - - auto o_scl = IsConstant(); - auto act_scl = IsConstant(); - auto sum_scl = IsConstant(); - auto dst_zp = IsConstant(); - - DFPattern cnv; - DFPattern pat; - - cnv = IsOp("qnn.conv2d")({src, wgh, IsConstant(), IsConstant(), IsConstant(), IsConstant()}); - pat = IsOp("cast")({cnv}); - pat = IsOp("add")({pat, bias}) || pat; - pat = IsOp("multiply")({pat, o_scl}); - pat = IsOp("clip")({pat}); - pat = IsOp("multiply")({pat, act_scl}) || pat; - pat = IsOp("add")({pat, sum_scl * IsOp("cast")({sum_src})}) || pat; - pat = IsOp("add")({pat, dst_zp}) || pat; - pat = IsOp("cast")({pat}); - - // Check pattern match - auto indexed_body = CreateIndexedGraph(comp_fn.body); - DFPatternMatcher matcher(indexed_body.get()); - auto res = matcher.Match(pat, comp_fn.body); - ICHECK(res) << "Mismatch of DNNL partitioner and codegen logic"; - - // Handle arguments in deterministic order - auto map = matcher.GetMemo(); - auto find = [&map](const DFPattern& pat) -> tvm::relay::Expr { - if (map.count(pat)) return map.at(pat)[0]; - return {}; - }; - - ArgPacker arg_holder(ext_attrs, args); - arg_holder.Put(find(src)); - arg_holder.Put(find(wgh)); - arg_holder.Put(find(bias), "bias_idx"); - arg_holder.Put(find(sum_src), "sum_idx"); - arg_holder.Put(find(o_scl), "o_scl_idx"); - arg_holder.Put(find(act_scl), "act_scl_idx"); - arg_holder.Put(find(sum_scl), "sum_scl_idx"); - arg_holder.Put(find(dst_zp), "dst_zp_idx"); - - // Activation. Default clip to simulate relu via uint8 cast - std::vector clip_attr{"clip"}; - auto act_scl_val = map.count(act_scl) ? find(act_scl) : constant(1.0); - clip_attr.push_back(std::to_string(arg_holder.Put(act_scl_val))); // act_scale - clip_attr.push_back(std::to_string(arg_holder.Put(constant(0.0)))); // alpha - clip_attr.push_back(std::to_string(arg_holder.Put(constant(255.0)))); // beta - (*ext_attrs)["activation"] = dmlc_attr(clip_attr); - - return map.at(cnv)[0].as(); -} - -const tvm::relay::CallNode* ParseQnnDenseComp(const tvm::relay::FunctionNode& comp_fn, - std::unordered_map* ext_attrs, - std::vector* args) { - using namespace tvm::relay; - - // Pattern - auto src = IsWildcard(); - auto wgh = IsWildcard(); - auto sum_src = IsWildcard(); - auto bias = IsConstant(); - - auto o_scl = IsConstant(); - auto act_scl = IsConstant(); - auto sum_scl = IsConstant(); - auto dst_zp = IsConstant(); - - DFPattern dns, act, pat; - - dns = IsOp("qnn.dense")({src, wgh, IsConstant(), IsConstant(), IsConstant(), IsConstant()}); - pat = IsOp("cast")({dns}); - pat = IsOp("add")({pat, bias}) || pat; - pat = IsOp("multiply")({pat, o_scl}); - pat = IsOp("clip")({pat}); - pat = IsOp("multiply")({pat, act_scl}) || pat; - pat = IsOp("add")({pat, sum_scl * IsOp("cast")({sum_src})}) || pat; - pat = IsOp("add")({pat, dst_zp}) || pat; - pat = IsOp("cast")({pat}); - - // Check pattern match - auto indexed_body = CreateIndexedGraph(comp_fn.body); - DFPatternMatcher matcher(indexed_body.get()); - auto res = matcher.Match(pat, comp_fn.body); - ICHECK(res) << "Mismatch of DNNL partitioner and codegen logic"; - - // Handle arguments in deterministic order - auto memo = matcher.GetMemo(); - auto find = [&memo](const DFPattern& pat) -> tvm::relay::Expr { - if (memo.count(pat)) return memo.at(pat)[0]; - return {}; - }; - - ArgPacker arg_holder(ext_attrs, args); - arg_holder.Put(find(src)); - arg_holder.Put(find(wgh)); - arg_holder.Put(find(bias), "bias_idx"); - arg_holder.Put(find(sum_src), "sum_idx"); - arg_holder.Put(find(o_scl), "o_scl_idx"); - arg_holder.Put(find(act_scl), "act_scl_idx"); - arg_holder.Put(find(sum_scl), "sum_scl_idx"); - arg_holder.Put(find(dst_zp), "dst_zp_idx"); - - // Activation. Default clip to simulate relu via uint8 cast - std::vector clip_attr{"clip"}; - auto act_scl_val = memo.count(act_scl) ? find(act_scl) : constant(1.0); - clip_attr.push_back(std::to_string(arg_holder.Put(act_scl_val))); // act_scale - clip_attr.push_back(std::to_string(arg_holder.Put(constant(0.0)))); // alpha - clip_attr.push_back(std::to_string(arg_holder.Put(constant(255.0)))); // beta - (*ext_attrs)["activation"] = dmlc_attr(clip_attr); - - return memo.at(dns)[0].as(); -} - -/*! - * Parse composite function and return real args, additional attributes and root call node - * @param comp_fn composite function to parse - * @param ext_attrs attr collection with additional attributes - * @param args real arguments of node - * @return root call node - */ -const tvm::relay::CallNode* ParseComposite(const tvm::relay::FunctionNode& comp_fn, - std::unordered_map* ext_attrs, - std::vector* args) { - auto comp = comp_fn.GetAttr(tvm::relay::attr::kComposite); - ICHECK(comp.defined()) << "DNNL JSON runtime only supports composite functions."; - auto name = comp.value(); - - const tvm::relay::CallNode* res = nullptr; - if (name == "dnnl.qnn.conv2d") - res = ParseQnnConvComp(comp_fn, ext_attrs, args); - else if (name == "dnnl.qnn.dense") - res = ParseQnnDenseComp(comp_fn, ext_attrs, args); - return res; -} - -#endif // TVM_RELAY_BACKEND_CONTRIB_DNNL_COMP_OP_MATCHER_H_ diff --git a/src/relay/backend/contrib/dnnl/query_layout.cc b/src/relay/backend/contrib/dnnl/query_layout.cc deleted file mode 100755 index 2660481e00c2..000000000000 --- a/src/relay/backend/contrib/dnnl/query_layout.cc +++ /dev/null @@ -1,379 +0,0 @@ -/* - * 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 src/relay/backend/contrib/dnnl/query_layout.cc - * \brief layout auto-query func. - */ - -#include -#include -#include -#include -#include -#include - -#include -#include -#include - -#include "../../../../runtime/contrib/dnnl/dnnl_utils.h" -#include "../../../../runtime/regex.h" -#include "../../utils.h" -#include "dnnl.hpp" -namespace tvm { -namespace relay { -namespace contrib { - -using dim_t = dnnl_dim_t; -using dims_t = dnnl_dims_t; -using tvm::runtime::contrib::dtype_dl2dnnl; - -template -inline void array_set(T* arr, const U& val, size_t size) { - for (size_t i = 0; i < size; ++i) arr[i] = static_cast(val); -} - -template -inline void array_copy(T* dst, const T* src, size_t size) { - for (size_t i = 0; i < size; ++i) dst[i] = src[i]; -} - -template -inline void swap(T& t1, T& t2) { - T tmp(t1); - t1 = t2; - t2 = tmp; -} - -template -inline void simultaneous_sort(T* vals, T* vals_2nd_level, U* keys, size_t size, F comparator) { - if (size == 0) return; - - for (size_t i = 0; i < size - 1; ++i) { - bool swapped = false; - - for (size_t j = 0; j < size - i - 1; j++) { - auto res = comparator(vals[j], vals[j + 1]); - if (res == 0) res = comparator(vals_2nd_level[j], vals_2nd_level[j + 1]); - - if (res > 0) { - swap(vals[j], vals[j + 1]); - swap(vals_2nd_level[j], vals_2nd_level[j + 1]); - swap(keys[j], keys[j + 1]); - swapped = true; - } - } - - if (swapped == false) break; - } -} - -void compute_blocks(dims_t blocks, const dnnl::memory::desc* md) { - using format_kind_t = dnnl_format_kind_t; - const format_kind_t blocked = dnnl_blocked; - if (!(md->data.format_kind == blocked)) { - array_set(blocks, 0, md->data.ndims); - return; - } - array_set(blocks, 1, md->data.ndims); - const auto& bd = md->data.format_desc.blocking; - for (int iblk = 0; iblk < bd.inner_nblks; ++iblk) - blocks[bd.inner_idxs[iblk]] *= bd.inner_blks[iblk]; -} - -inline bool has_runtime_strides(const dnnl::memory::desc* md) { - using format_kind_t = dnnl_format_kind_t; - const format_kind_t blocked = dnnl_blocked; - if (!(md->data.format_kind == blocked)) return false; - for (int d = 0; d < md->data.ndims; ++d) - if (md->data.format_desc.blocking.strides[d] == DNNL_RUNTIME_DIM_VAL) return true; - return false; -} - -std::string md2fmt_tag_str(const dnnl::memory::desc* md) { - const auto& blk = md->data.format_desc.blocking; - - dims_t blocks = {0}; - compute_blocks(blocks, md); - - char dim_chars[DNNL_MAX_NDIMS + 1]; - - dims_t ou_blocks = {0}; - array_copy(ou_blocks, md->data.padded_dims, md->data.ndims); - - bool plain = true; - for (int d = 0; d < md->data.ndims; ++d) { - dim_chars[d] = (blocks[d] == 1 ? 'a' : 'A') + static_cast(d); - if (blocks[d] != 1) plain = false; - ou_blocks[d] /= blocks[d]; - } - - // Can't report meaningful tag for runtime dimensions. - if (has_runtime_strides(md)) return "*"; - - dims_t strides; - array_copy(strides, blk.strides, md->data.ndims); - - simultaneous_sort(strides, ou_blocks, dim_chars, md->data.ndims, - [](dim_t a, dim_t b) { return b - a; }); - - dim_chars[md->data.ndims] = '\0'; - - std::string s(dim_chars); - - if (!plain) { - for (int iblk = 0; iblk < blk.inner_nblks; ++iblk) { - char c = ('a' + static_cast(blk.inner_idxs[iblk])); - s += (std::to_string(blk.inner_blks[iblk]) + c); - } - } - return s; -} - -dnnl::memory::dims str2dims(const std::string& str_shape, bool dilates = false, - std::string interval = ",") { - // Split strings - std::vector str_dims; - size_t pos = 0, start = 0; - while ((pos = str_shape.find(interval, start)) != std::string::npos) { - std::string str_dim = str_shape.substr(start, pos - start); - if (pos > start) str_dims.push_back(str_dim); - start = pos + interval.size(); - } - if (str_shape.size() > start) { - str_dims.push_back(str_shape.substr(start)); - } - // transfer string to dims - dnnl::memory::dims out_dims; - if (dilates) { - std::transform(str_dims.begin(), str_dims.end(), std::back_inserter(out_dims), - [](const std::string& str) { return std::stoi(str) - 1; }); - } else { - std::transform(str_dims.begin(), str_dims.end(), std::back_inserter(out_dims), - [](const std::string& str) { return std::stoi(str); }); - } - return out_dims; -} - -void check_shapes(const std::vector shapes) { - std::string valid_pat("(\\d*)(,(\\d*))*"); - bool checked = tvm::runtime::regex_match(shapes[0], valid_pat); - for (size_t i = 1; i < shapes.size() - 1; i++) { - checked &= tvm::runtime::regex_match(shapes[i], valid_pat); - } - checked &= tvm::runtime::regex_match(shapes[shapes.size() - 1], "\\d*"); - if (!checked) { - LOG(FATAL) << "Invalid input args for query dnnl optimal layout."; - } -} - -void check_layout(bool var, bool ref) { - if (var != ref) { - LOG(FATAL) << "Invalid input layout for query dnnl optimal layout."; - } -} - -std::string get_optimal_layout_for_conv(std::string data_layout, std::string kernel_layout, - std::string weight_shape, std::string out_shape, - std::string paddings, std::string strides, - std::string dilates, std::string G, std::string dtype) { - check_layout(tvm::runtime::regex_match(data_layout, "NC(D?)(H?)W"), true); - check_layout(tvm::runtime::regex_match(kernel_layout, "(G?)OI(D?)(H?)W"), true); - check_shapes({weight_shape, out_shape, paddings, strides, dilates, G}); - - dnnl::engine eng(dnnl::engine::kind::cpu, 0); - dnnl::stream s(eng); - using tag = dnnl::memory::format_tag; - - dnnl::memory::dim groups = std::stoi(G); - dnnl::memory::dims weight_dims_ = str2dims(weight_shape); - dnnl::memory::dims weight_dims = weight_dims_; - - if (groups > 1) { - if (weight_dims_.size() == 5) { - weight_dims = {groups * weight_dims_[1], groups * weight_dims_[2], weight_dims_[3], - weight_dims_[4]}; - } else { - weight_dims[1] = weight_dims[1] * groups; - } - } - - dnnl::memory::dims out_dims = str2dims(out_shape); - dnnl::memory::dims padding_dims = str2dims(paddings); - dnnl::memory::dims padding_dims_l(padding_dims.begin(), - padding_dims.begin() + padding_dims.size() / 2); - dnnl::memory::dims padding_dims_r(padding_dims.end() - padding_dims.size() / 2, - padding_dims.end()); - dnnl::memory::dims strides_dims = str2dims(strides); - dnnl::memory::dims dilates_dims = str2dims(dilates, true); - - dnnl::memory::dims input_dims = out_dims; - input_dims[1] = weight_dims[1]; - for (size_t i = 2; i < out_dims.size(); i++) { - dnnl::memory::dim K = weight_dims[i]; - dnnl::memory::dim S = strides_dims[i - 2]; - dnnl::memory::dim D = dilates_dims[i - 2]; - dnnl::memory::dim PL = padding_dims_l[i - 2]; - dnnl::memory::dim PR = padding_dims_r[i - 2]; - dnnl::memory::dim DK = 1 + (K - 1) * (D + 1); - input_dims[i] = out_dims[i] * S - PL - PR + DK - 1; - } - - dnnl::memory::dims conv_src_dims = input_dims; - dnnl::memory::dims conv_weights_dims = weight_dims; - if (groups > 1) { - conv_weights_dims = {groups, out_dims[1] / groups, input_dims[1] / groups}; - conv_weights_dims.insert(conv_weights_dims.end(), weight_dims.begin() + 2, weight_dims.end()); - } - - dnnl::memory::dims conv_dst_dims = out_dims; - dnnl::memory::dims conv_strides = strides_dims; - dnnl::memory::dims conv_dilates = dilates_dims; - dnnl::memory::dims conv_padding_l = padding_dims_l; - dnnl::memory::dims conv_padding_r = padding_dims_r; - - auto dnnl_dtype = dtype_dl2dnnl(tvm::runtime::String2DLDataType(dtype)); - auto conv_src_md = dnnl::memory::desc({conv_src_dims}, dnnl_dtype, tag::any); - auto conv_weights_md = dnnl::memory::desc({conv_weights_dims}, dnnl_dtype, tag::any); - auto conv_dst_md = dnnl::memory::desc({conv_dst_dims}, dnnl_dtype, tag::any); - - auto conv_desc = dnnl::convolution_forward::desc( - dnnl::prop_kind::forward_inference, dnnl::algorithm::convolution_direct, conv_src_md, - conv_weights_md, conv_dst_md, conv_strides, conv_dilates, conv_padding_l, conv_padding_r); - - auto conv_prim_desc = dnnl::convolution_forward::primitive_desc(conv_desc, eng); - - auto src_format = conv_prim_desc.src_desc(); - auto weights_format = conv_prim_desc.weights_desc(); - auto dst_format = conv_prim_desc.dst_desc(); - std::string src_df, weight_df, dst_df; - - src_df = md2fmt_tag_str(&src_format); - weight_df = md2fmt_tag_str(&weights_format); - dst_df = md2fmt_tag_str(&dst_format); - std::string res = src_df + "," + weight_df + "," + dst_df; - return res; -} - -std::string get_optimal_layout_for_conv_transpose(std::string data_layout, - std::string kernel_layout, - std::string weight_shape, std::string out_shape, - std::string paddings, std::string output_paddings, - std::string strides, std::string dilates, - std::string G, std::string dtype) { - check_layout(tvm::runtime::regex_match(data_layout, "NC(D?)(H?)W"), true); - check_layout(tvm::runtime::regex_match(kernel_layout, "(G?)((IO)|(OI))(D?)(H?)W"), true); - check_shapes({weight_shape, out_shape, paddings, output_paddings, strides, dilates, G}); - - dnnl::engine eng(dnnl::engine::kind::cpu, 0); - dnnl::stream s(eng); - using tag = dnnl::memory::format_tag; - - dnnl::memory::dim groups = std::stoi(G); - dnnl::memory::dims weight_dims_ = str2dims(weight_shape); - dnnl::memory::dims weight_dims = weight_dims_; - if (groups > 1) { - if (weight_dims_.size() == 5) { - weight_dims = {groups * weight_dims_[1], groups * weight_dims_[2], weight_dims_[3], - weight_dims_[4]}; - } else { - weight_dims[1] = weight_dims[1] * groups; - } - } - dnnl::memory::dims out_dims = str2dims(out_shape); - dnnl::memory::dims padding_dims = str2dims(paddings); - dnnl::memory::dims padding_dims_l(padding_dims.begin(), - padding_dims.begin() + padding_dims.size() / 2); - dnnl::memory::dims padding_dims_r(padding_dims.end() - padding_dims.size() / 2, - padding_dims.end()); - dnnl::memory::dims output_padding_dims = str2dims(output_paddings); - dnnl::memory::dims strides_dims = str2dims(strides); - dnnl::memory::dims dilates_dims = str2dims(dilates, true); - - dnnl::memory::dims input_dims = out_dims; - if (out_dims[1] == weight_dims[0]) { - input_dims[1] = weight_dims[1]; - } else { - input_dims[1] = weight_dims[0]; - std::swap(weight_dims[0], weight_dims[1]); - } - for (size_t i = 2; i < out_dims.size(); i++) { - dnnl::memory::dim K = weight_dims[i]; - dnnl::memory::dim S = strides_dims[i - 2]; - dnnl::memory::dim D = dilates_dims[i - 2]; - dnnl::memory::dim PL = padding_dims_l[i - 2]; - dnnl::memory::dim PR = padding_dims_r[i - 2]; - dnnl::memory::dim OP = output_padding_dims[i - 2]; - dnnl::memory::dim DK = 1 + (K - 1) * (D + 1); - input_dims[i] = (out_dims[i] - DK + PL + PR - OP) / S + 1; - } - - dnnl::memory::dims deconv_src_dims = input_dims; - dnnl::memory::dims deconv_weights_dims = weight_dims; - if (groups > 1) { - deconv_weights_dims = {groups, out_dims[1] / groups, input_dims[1] / groups}; - deconv_weights_dims.insert(deconv_weights_dims.end(), weight_dims.begin() + 2, - weight_dims.end()); - } - dnnl::memory::dims deconv_dst_dims = out_dims; - dnnl::memory::dims deconv_strides = strides_dims; - dnnl::memory::dims deconv_dilates = dilates_dims; - dnnl::memory::dims deconv_padding_l = padding_dims_l; - dnnl::memory::dims deconv_padding_r = padding_dims_r; - - auto dnnl_dtype = dtype_dl2dnnl(tvm::runtime::String2DLDataType(dtype)); - auto deconv_src_md = dnnl::memory::desc({deconv_src_dims}, dnnl_dtype, tag::any); - auto deconv_weights_md = dnnl::memory::desc({deconv_weights_dims}, dnnl_dtype, tag::any); - auto deconv_dst_md = dnnl::memory::desc({deconv_dst_dims}, dnnl_dtype, tag::any); - - auto deconv_desc = dnnl::deconvolution_forward::desc( - dnnl::prop_kind::forward_inference, dnnl::algorithm::deconvolution_direct, deconv_src_md, - deconv_weights_md, deconv_dst_md, deconv_strides, deconv_dilates, deconv_padding_l, - deconv_padding_r); - - auto deconv_prim_desc = dnnl::deconvolution_forward::primitive_desc(deconv_desc, eng); - - auto src_format = deconv_prim_desc.src_desc(); - auto weights_format = deconv_prim_desc.weights_desc(); - auto dst_format = deconv_prim_desc.dst_desc(); - std::string src_df, weight_df, dst_df; - - src_df = md2fmt_tag_str(&src_format); - weight_df = md2fmt_tag_str(&weights_format); - dst_df = md2fmt_tag_str(&dst_format); - std::string res = src_df + "," + weight_df + "," + dst_df; - return res; -} - -TVM_REGISTER_GLOBAL("relay.ir.get_optimal_layout_for_conv") - .set_body([](TVMArgs args, TVMRetValue* rv) { - *rv = get_optimal_layout_for_conv(args[0], args[1], args[2], args[3], args[4], args[5], - args[6], args[7], args[8]); - }); - -TVM_REGISTER_GLOBAL("relay.ir.get_optimal_layout_for_conv_transpose") - .set_body([](TVMArgs args, TVMRetValue* rv) { - *rv = get_optimal_layout_for_conv_transpose(args[0], args[1], args[2], args[3], args[4], - args[5], args[6], args[7], args[8], args[9]); - }); - -} // namespace contrib -} // namespace relay -} // namespace tvm diff --git a/src/relay/backend/contrib/example_target_hooks/relay_to_tir.cc b/src/relay/backend/contrib/example_target_hooks/relay_to_tir.cc deleted file mode 100644 index 2b037181653c..000000000000 --- a/src/relay/backend/contrib/example_target_hooks/relay_to_tir.cc +++ /dev/null @@ -1,275 +0,0 @@ - -/* - * 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 "../../../op/call/call.h" -#include "tvm/tir/function.h" - -namespace tvm { -namespace relay { -namespace contrib { -namespace example_target_hooks { - -namespace { - -/*! - * \brief An example mutator for a "RelayToTIR" custom pass. Replaces every call to a Relay - * Function with "external_symbol" attribute of "replace_add_with_subtract" with a call to a - * TIR PrimFunc implementing subtraction. - * - * Illustrates six aspects a custom 'lowering' style pass may need to account for: - * - Lowerable functions can appear inline as call ops, bound to let-bound variables, or as - * global functions. - * - Let-bound lowerable functions should be inlined on-the-fly since after processing the - * let-binding is no longer required. - * - There may be multiple calls to the same lowerable function. All calls need to be - * rewritten, even though the function itself need be rewritten only once. - * - GlobalVars must be shared between all calls and the new definition itself. - * - Calls to lowered functions must use the "call_lowered" calling convention. - * - The Target::Current() may hold an instance of the TargetKind from which the custom Pass - * was extracted. - * - * Though not illustrated here, it is also valid for a "RelayToTIR" custom pass to add - * runtime::Modules to the output IRModule's "external_mods" attribute. In this case the - * IRModule must be left with an 'extern' Function definition with the matching "external_symbol" - * name. - */ -class ConvertAddToSubtract : public MixedModeMutator { - public: - explicit ConvertAddToSubtract(IRModule ir_module, Target host_target) - : ir_module_(ir_module), - host_target_(host_target), - custom_target_(Target("example_target_hook")) {} - - IRModule Mutate() { - GlobalVar main_global_var = ir_module_->GetGlobalVar("main"); - Function main = Downcast(ir_module_->Lookup(main_global_var)); - Function mutated_main = WithFields(main, main->params, VisitExpr(main->body)); - - ir_module_->Update(main_global_var, mutated_main); - - return ir_module_; - } - - private: - tir::BufferLoad LoadIndex(const tir::Buffer& buffer, const PrimExpr& index) { - return tir::BufferLoad(buffer, {index}); - } - - GlobalVar ReplaceAddWithSubtractPrimFunc(const Function& func) { - auto func_name = func->GetAttr(::tvm::attr::kGlobalSymbol); - ICHECK(func_name.defined()); - - // -------------------------------------------------------------------------------------------- - // Cases: - // - Inline function: - // - First encounter: create global var, rewrite to PrimFunc, add binding, replace call. - // - Thereafter (via object sharing): discover global var already in module, replace call - // - Global function: - // - Assume func_name == global_var->name_hint - // - First encounter: create global var, rewrite to PrimFunc, update binding, replace call - // - Thereafter (via global var): discover global var already in module, replace call - // -------------------------------------------------------------------------------------------- - - // If necessary, introduce a new global var to map the function to and copy the source type - // over for InferType. - GlobalVar global_var; - bool need_rewriting; - if (ir_module_->ContainGlobalVar(func_name.value())) { - global_var = ir_module_->GetGlobalVar(func_name.value()); - // Only rewrite to a PrimFunc if the global definition is still a Relay function. - need_rewriting = ir_module_->Lookup(global_var)->IsInstance(); - } else { - global_var = GlobalVar(func_name.value()); - global_var->checked_type_ = func->checked_type(); - need_rewriting = true; - } - - // For illustration only, check if the current target matches the example_target_hook kind, - // and if so extract the example attribute value. - int64_t example_attribute_value = 0; - Optional opt_current_target = Target::Current(); - if (opt_current_target.defined() && - opt_current_target.value()->kind->name == "example_target_hook") { - example_attribute_value = - opt_current_target.value()->GetAttr("example_attribute").value()->value; - } - - if (need_rewriting) { - // The called function is still in Relay form. Convert to TIR. - tir::Buffer x_buffer = tir::decl_buffer({8}, DataType::Float(32), "x"); - tir::Buffer y_buffer = tir::decl_buffer({8}, DataType::Float(32), "y"); - tir::Buffer out_buffer = tir::decl_buffer({8}, DataType::Float(32)); - - tir::Var x_var("x", DataType::Handle()); - tir::Var y_var("y", DataType::Handle()); - tir::Var out_var("out", DataType::Handle()); - - Map dict_attrs; - dict_attrs.Set("global_symbol", global_var->name_hint); - dict_attrs.Set("tir.noalias", Bool(true)); - - te::Var index("index", DataType::Int(32)); - tir::Sub indexed_sub = tir::Sub(LoadIndex(x_buffer, index), LoadIndex(y_buffer, index)); - if (example_attribute_value > 0) { - // For illustration only, fold the example attribute into the result. - indexed_sub = tir::Sub(indexed_sub, FloatImm(DataType::Float(32), - static_cast(example_attribute_value))); - } - - tir::Stmt math_body = tir::BufferStore(out_buffer, indexed_sub, {index}); - tir::Stmt math_loop = tir::For(index, 0, 8, tir::ForKind::kSerial, math_body); - - Map buffer_map = { - {x_var, x_buffer}, - {y_var, y_buffer}, - {out_var, out_buffer}, - }; - - tir::PrimFunc replacement_func = tir::PrimFunc({x_var, y_var, out_var}, math_loop, VoidType(), - buffer_map, DictAttrs(dict_attrs)); - - // Switch to TIRToRuntime hook for testing - Bool tir_to_runtime = func->GetAttr("tir_to_runtime").value_or(Bool(false)); - if (tir_to_runtime) { - replacement_func = WithAttr(replacement_func, ::tvm::attr::kTarget, custom_target_); - } else { - replacement_func = WithAttr(replacement_func, ::tvm::attr::kTarget, host_target_); - } - - ir_module_->Update(global_var, replacement_func); // Will Add if global_var is new. - } - - return global_var; - } - - using MixedModeMutator::VisitExpr_; - - Expr VisitExpr_(const LetNode* op) final { - auto pre_visit = [this](const LetNode* op) { - Expr var = this->VisitExpr(op->var); - Expr value = this->VisitExpr(op->value); - - if (AsLowerableFunction(value)) { - // Inline on-the-fly if the let-bound value is lowerable. - this->memo_[var] = value; - } - }; - auto post_visit = [this](const LetNode* op) { - // Rely on the Memoizer to cache pre-visit values - Expr value = this->VisitExpr(op->value); - Expr body = this->VisitExpr(op->body); - auto expr = GetRef(op); - - if (AsLowerableFunction(value)) { - // The let binding is no longer needed since inlined on-the-fly above. - this->memo_[expr] = this->VisitExpr(op->body); - } else { - Var var = Downcast(this->VisitExpr(op->var)); - if (var.same_as(op->var) && value.same_as(op->value) && body.same_as(op->body)) { - this->memo_[expr] = expr; - } else { - this->memo_[expr] = Let(var, value, body); - } - } - }; - ExpandANormalForm(op, pre_visit, post_visit); - return memo_[GetRef(op)]; - } - - const FunctionNode* AsLowerableFunction(const Expr& expr) { - if (const auto* function_node = expr.as()) { - auto func_name = function_node->GetAttr(::tvm::attr::kGlobalSymbol); - if (!func_name.defined()) { - return nullptr; - } - if (func_name != "replace_add_with_subtract") { - return nullptr; - } - return function_node; - } else if (auto global_var_node = expr.as()) { - return AsLowerableFunction(ir_module_->Lookup(global_var_node.value())); - } else { - return nullptr; - } - } - - const GlobalVarNode* AsAlreadyLoweredFunction(const Expr& expr) { - if (auto opt = expr.as()) { - auto global_var = opt.value(); - if (ir_module_->Lookup(global_var).as()) { - return global_var.get(); - } - } - return nullptr; - } - - Expr Rewrite_(const CallNode* pre, const Expr& post) override { - if (const auto* call = post.as()) { - GlobalVar new_op; - if (const auto* function_node = AsLowerableFunction(call->op)) { - // Add or replace the function with a PrimFunc. - new_op = ReplaceAddWithSubtractPrimFunc(GetRef(function_node)); - } else if (const auto* global_var_node = AsAlreadyLoweredFunction(call->op)) { - // The function has already been rewritten, so we just need to update the call. - new_op = GetRef(global_var_node); - } - if (new_op.defined()) { - // Since we are replacing the Relay function with a call to a TIR function, we must use - // the call_lowered op. - CallLoweredAttrs attrs; - attrs.metadata.Set("relay_attrs", call->attrs); - ICHECK(call->type_args.empty()) << "lowered functions cannot be polymorphic"; - return CallLowered(std::move(new_op), call->args, std::move(attrs), call->span); - } - } - - return post; - } - - public: - IRModule ir_module_; - Target host_target_; - Target custom_target_; -}; - -} // namespace - -transform::Pass RelayToTIR() { - runtime::TypedPackedFunc pass_func = - [=](IRModule ir_module, transform::PassContext pass_context) { - ConvertAddToSubtract relay_to_tir(std::move(ir_module), Target("c")); - return relay_to_tir.Mutate(); - }; - return tvm::transform::CreateModulePass(pass_func, 0, "RelayToTIR", {}); -} - -} // namespace example_target_hooks -} // namespace contrib -} // namespace relay -} // namespace tvm diff --git a/src/relay/backend/contrib/example_target_hooks/target.cc b/src/relay/backend/contrib/example_target_hooks/target.cc deleted file mode 100644 index de9c81a2706e..000000000000 --- a/src/relay/backend/contrib/example_target_hooks/target.cc +++ /dev/null @@ -1,43 +0,0 @@ -/* - * 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 - -namespace tvm { - -using FTVMTIRToRuntime = tvm::runtime::TypedPackedFunc; - -namespace relay { -namespace contrib { -namespace example_target_hooks { -tvm::transform::Pass RelayToTIR(); -runtime::Module TIRToRuntime(IRModule mod, Target target); -} // namespace example_target_hooks -} // namespace contrib -} // namespace relay - -TVM_REGISTER_TARGET_KIND("example_target_hook", kDLCPU) - .set_attr("use_device_api", Bool(true)) - .set_attr(attr::kRelayToTIR, - relay::contrib::example_target_hooks::RelayToTIR()) - .set_attr("TIRToRuntime", relay::contrib::example_target_hooks::TIRToRuntime) - .add_attr_option("example_attribute", Integer(0)); - -} // namespace tvm diff --git a/src/relay/backend/contrib/example_target_hooks/tir_to_runtime.cc b/src/relay/backend/contrib/example_target_hooks/tir_to_runtime.cc deleted file mode 100644 index 6f09e0a0c3f0..000000000000 --- a/src/relay/backend/contrib/example_target_hooks/tir_to_runtime.cc +++ /dev/null @@ -1,82 +0,0 @@ -/* - * 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 "../../../../target/source/codegen_c_host.h" - -namespace tvm { -namespace relay { -namespace contrib { -namespace example_target_hooks { - -using namespace tir; - -class CodeGenExampleTargetHook : public codegen::CodeGenCHost { - public: - using codegen::CodeGenCHost::VisitExpr_; - - /*! - * \brief Emit code that changes adds to multiplies for testing - */ - void VisitExpr_(const SubNode* op, std::ostream& os) final { - os << '('; - PrintExpr(op->a, os); - os << " * "; - PrintExpr(op->b, os); - os << ')'; - } -}; - -runtime::Module TIRToRuntime(IRModule mod, Target target) { - bool output_ssa = false; - bool emit_asserts = false; - bool emit_fwd_func_decl = false; - CodeGenExampleTargetHook codegen; - - std::unordered_set devices; - codegen.Init(output_ssa, emit_asserts, emit_fwd_func_decl, target->str(), devices); - - Map functions; - for (auto [gvar, base_func] : mod->functions) { - auto prim_func = Downcast(base_func); - functions.Set(gvar, prim_func); - } - - for (auto [gvar, prim_func] : functions) { - codegen.DeclareFunction(gvar, prim_func); - } - for (auto [gvar, prim_func] : functions) { - codegen.AddFunction(gvar, prim_func, emit_fwd_func_decl); - } - - std::string code = codegen.Finish(); - - Array function_names; - for (auto [gvar, prim_func] : functions) { - function_names.push_back(codegen.GetFunctionName(gvar)); - } - - return codegen::CSourceModuleCreate(code, "c", function_names); -} - -} // namespace example_target_hooks -} // namespace contrib -} // namespace relay -} // namespace tvm diff --git a/src/relay/backend/contrib/libtorch/libtorch_codegen.cc b/src/relay/backend/contrib/libtorch/libtorch_codegen.cc deleted file mode 100644 index 29fee504349c..000000000000 --- a/src/relay/backend/contrib/libtorch/libtorch_codegen.cc +++ /dev/null @@ -1,141 +0,0 @@ -/* - * 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 src/relay/backend/contrib/libtorch/codegen.cc - * \brief Implementation of libtorch codegen. - */ - -// clang-format off -#include -#include -#include -#include -#include -#include -#include -#include -#include -#include - -#include -#include -#include - -#include "../../utils.h" - -#include -#include -#include -#include -// clang-format on - -namespace tvm { -namespace relay { -namespace contrib { - -using namespace backend; - -/*! \brief Attributes of a TorchFunction node */ -struct TorchFunctionAttrs : public tvm::AttrsNode { - std::string serialized_function; - int64_t len; - - TVM_DECLARE_ATTRS(TorchFunctionAttrs, "relay.attrs.TorchFunctionAttrs") { - TVM_ATTR_FIELD(serialized_function).set_default("").describe("Function from fn.save(...)"); - TVM_ATTR_FIELD(len).set_default(-1).describe("Function from fn.save(...)"); - } -}; - -TVM_REGISTER_NODE_TYPE(TorchFunctionAttrs); - -bool TorchOpRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - const auto* sfattrs = attrs.as(); - std::stringstream str(sfattrs->serialized_function); - torch::jit::Module mod = torch::jit::load(str); - - std::vector inputs; - for (int i = 0; i < num_inputs; i++) { - auto* ty = types[i].as(); - ICHECK(ty) << "only accept tensors as inputs"; - std::vector shape; - for (const auto& s : ty->shape) { - auto* si = s.as(); - if (!si) { - return false; - } - shape.push_back(si->value); - } - auto torchScalarType = at::toScalarType(ty->dtype); - - inputs.emplace_back(torch::zeros(shape, at::TensorOptions().dtype(torchScalarType))); - } - auto res = mod.forward(inputs); - auto res_t = res.toTensor(); - ICHECK((int)types.size() == num_inputs + 1) << "only single output supported"; - Array res_sizes; - for (int d = 0; d < res_t.dim(); d++) { - res_sizes.push_back(IntImm(DataType::Int(32), res_t.size(d))); - } - reporter->Assign(types[num_inputs], TensorType(res_sizes, DataType(at::getDLDataType(res_t)))); - return true; -} - -RELAY_REGISTER_OP("torch_op") - .set_support_level(99) - .add_type_rel("TorchOpRel", TorchOpRel) - .set_attrs_type(); - -Expr MakeTorchOp(Array args, const std::string& serialized_function) { - static const Op& op = Op::Get("torch_op"); - auto attrs = make_object(); - attrs->serialized_function = serialized_function; - attrs->len = serialized_function.size(); - return Call(op, args, Attrs(attrs), {}); -} - -TVM_REGISTER_GLOBAL("relay.op._make.torchop").set_body_typed(MakeTorchOp); - -/*! - * \brief The external compiler/codegen tool. It takes a Relay expression/module and - * compile it into a runtime module. - */ -runtime::Module TorchCompiler(const ObjectRef& ref) { - ICHECK(ref->IsInstance()) << "The input ref is expected to be a Relay function."; - Function func = Downcast(ref); - std::string func_name = backend::GetExtSymbol(func); - - ICHECK(func.defined()) << "Input error: expect a Relay function."; - const auto* call = func->body.as(); - ICHECK(call) << "Expected call node\n"; - const auto* op_node = call->op.as(); - ICHECK(op_node) << "Expect OpNode, but got " << call->op->GetTypeKey(); - const auto op_name = GetRef(op_node)->name; - ICHECK(op_name == "torch_op") << "Unsupported op: " << AsText(call->op, false) << "\n"; - - const auto* attrs = call->attrs.as(); - return tvm::runtime::contrib::TorchRuntimeCreate(func_name, attrs->serialized_function); -} - -TVM_REGISTER_GLOBAL("relay.ext.torch").set_body_typed(TorchCompiler); - -} // namespace contrib -} // namespace relay -} // namespace tvm diff --git a/src/relay/backend/contrib/mrvl/codegen.cc b/src/relay/backend/contrib/mrvl/codegen.cc deleted file mode 100644 index 70c980e429b4..000000000000 --- a/src/relay/backend/contrib/mrvl/codegen.cc +++ /dev/null @@ -1,1524 +0,0 @@ -/* - * 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 src/relay/backend/contrib/mrvl/codegen.cc - * \brief Marvell MLIP specific API - */ - -#include -#include -#include -#include - -#include -#include -#include -#include -#include -#include -#include -#include - -#include "../../../../runtime/contrib/json/json_node.h" -#include "../../../../support/base64.h" -#include "../../../qnn/utils.h" -#include "../../utils.h" -#include "../codegen_json/codegen_json.h" - -namespace tvm { -namespace relay { - -namespace contrib { -namespace mrvl { - -using namespace backend; - -struct const_struct { - std::string name; - std::string shape; - std::string dtype; - std::string min; - std::string max; - std::string data_base64; -}; - -using JSONGraphNode = tvm::runtime::json::JSONGraphNode; -using JSONGraphNodeEntry = tvm::runtime::json::JSONGraphNodeEntry; - -/*! - * \brief Generates an MrvlModule from a relay expression. This "compilation" - * does not require Mrvl driver since the actual conversion using Mrvl APIs is - * deferred until creation of the runtime. This step simply serializes the - * relay program into a JSON string. - */ -class MrvlJSONSerializer : public backend::contrib::JSONSerializer { - public: - /*! - * \brief Constructor - * - * \param symbol The symbol that represents the graph being converted. - * \param expr The Relay expression to be converted to the JSON form. - */ - MrvlJSONSerializer(const std::string& symbol, const Expr& expr) : JSONSerializer(symbol, expr) { - layer_name_ = symbol; - } - - /*! \brief Return the required params. */ - Array GetParams() const { - tvm::runtime::Array base_params = JSONSerializer::const_names(); - Array mrvl_params; - for (size_t idx = 0; idx < base_params.size(); idx++) { - mrvl_params.push_back(base_params[idx]); - } - for (size_t idx = 0; idx < batch_norm_params_.size(); idx++) { - mrvl_params.push_back(batch_norm_params_[idx]); - } - return mrvl_params; - } - - template - std::string FloatToString(T val, size_t precision = 17) { - // Method to serialize floating point values (double, float) - // to a string with required precision. - std::ostringstream s; - s.precision(precision); - s << val; - return s.str(); - } - - /*! \brief Return the Const Json Strings. */ - std::string GetConstJSONString() { - std::string json_string; - auto names = const_names(); - auto map = const_name_to_constant(); - std::vector const_info_vec; - for (auto name_const : names) { - const_struct a; - std::string const_string = name_const; - auto arr = map[const_string]; - a.name = const_string; - a.dtype = "float" + std::to_string(static_cast(arr->dtype.bits)); - std::string shape; - shape += "[ "; - - int ndim = arr->ndim; - if (ndim == 1 || ndim == 3) { - shape += "1, "; - } - int tot_dim = 1; - for (int i = 0; i < ndim; i++) { - tot_dim *= arr->shape[i]; - shape += std::to_string(arr->shape[i]); - if (i != ndim - 1) shape += ", "; - } - shape += " ]"; - a.shape = shape; - int size = (arr->dtype.bits + 7) / 8; - int num_bytes = tot_dim * size; - std::string blob; - dmlc::MemoryStringStream mstrm(&blob); - support::Base64OutStream b64strm(&mstrm); - b64strm.Write(arr->data, num_bytes); - b64strm.Finish(); - a.data_base64 = blob; - // Populate min and max - float min_val = std::numeric_limits::infinity(); - float max_val = -min_val; - for (int i = 0; i < tot_dim; i++) { - auto val = static_cast(arr->data)[i]; - if (val > max_val) max_val = val; - if (val < min_val) min_val = val; - } - - a.min = FloatToString(min_val); - a.max = FloatToString(max_val); - - const_info_vec.push_back(a); - } - - json_string += "{\n"; - for (unsigned int i = 0; i < const_info_vec.size(); i++) { - auto a = const_info_vec[i]; - json_string += "\t\"" + a.name + "\": {\n"; - json_string += "\t\"shape\": " + a.shape + ",\n"; - json_string += "\t\"dtype\": \"" + a.dtype + "\"" + ",\n"; - json_string += "\t\"min\": \"" + a.min + "\"" + ",\n"; - json_string += "\t\"max\": \"" + a.max + "\"" + ",\n"; - json_string += "\t\"data_base64\": \"" + a.data_base64 + "\"\n"; - if (i == const_info_vec.size() - 1) { - json_string += "\t}\n"; - } else { - json_string += "\t},\n"; - } - } - json_string += "}\n"; - return json_string; - } - - protected: - /*! - * \brief A series of operators that form a composite - * convolution. Supports both nn.conv2d and qnn.conv2d. - */ - struct CompositeConvNode { - const CallNode* pad = nullptr; - const CallNode* conv = nullptr; - const CallNode* add = nullptr; - const CallNode* batch_norm = nullptr; - const CallNode* activation = nullptr; - }; - - /*! - * \brief A series of operators that form a composite - * sum. - */ - struct CompositeSumNode { - const CallNode* add = nullptr; - const CallNode* activation = nullptr; - }; - - /*! - * \brief A series of operators that form a composite - * maxpool or avgpool. Supports both nn.max_pool2d and qnn.conv2d. - */ - struct CompositePoolNode { - const CallNode* pad = nullptr; - const CallNode* pool = nullptr; - }; - - /*! - * \brief A series of operators that form a composite - * concat. - */ - struct CompositeConcatNode { - const CallNode* concat = nullptr; - }; - - /*! - * \brief A series of operators that form a reshape node. - */ - struct CompositeReshapeNode { - const CallNode* reshape = nullptr; - }; - - /*! - * \brief A series of operators that form a batch flatten node. - */ - struct CompositeBatchFlattenNode { - const CallNode* batch_flatten = nullptr; - }; - - /*! - * \brief A series of operators that form a Squeeze node. - */ - struct CompositeSqueezeNode { - const CallNode* squeeze = nullptr; - }; - - /*! - * \brief A series of operators that form a composite - * fc layer. Supports both nn.fc_ni2no and qnn.fc_ni2no. - */ - struct CompositeFcNode { - const CallNode* transform = nullptr; - const CallNode* flatten = nullptr; - const CallNode* fc = nullptr; - const CallNode* add = nullptr; - const CallNode* activation = nullptr; - }; - - /*! - * \brief Visit call nodes and generate appropriate JSON node. - * - * \param cn The current call node. - * \return A list of graph entry nodes. - */ - std::vector VisitExpr_(const CallNode* cn) override { - const auto* op_node = cn->op.as(); - if (op_node) { - // handle certain op node types specially - String op_name = op_node->name; - bool handle_by_mrvl = (op_name == "layout_transform" || op_name == "transpose"); - if (!handle_by_mrvl) { - return JSONSerializer::VisitExpr_(cn); - } - - // setup json attributes and then add the Mrvl Layer to JSON files - std::shared_ptr json_kernel_node; - json_kernel_node = CreateMrvlLayer4OpNode(cn); - return AddNode(json_kernel_node, GetRef(cn)); - } - - // handle only mrvl composite functions - if (!cn->op.as()) { - LOG(FATAL) << "Mrvl JSON runtime does not support calls to " << cn->op->GetTypeKey(); - } - auto fn = cn->op.as(); - auto comp = fn->GetAttr(attr::kComposite); - ICHECK(comp.defined()) << "Marvell-Compiler-ERROR-Internal::Illegal Mrvl composite function."; - const std::string name = comp.value(); - std::shared_ptr json_kernel_node; - if (name == "mrvl.conv2d_nhwc2nhwc") { - json_kernel_node = CreateCompositeMrvlConv2DLayer(cn); - } else if (name == "mrvl.fc_ni2no") { - json_kernel_node = CreateCompositeMrvlFcLayer(cn); - } else if (name == "mrvl.maxpool2d_nhwc2nhwc") { - json_kernel_node = CreateCompositeMrvlMaxpool2DLayer(cn); - } else if (name == "mrvl.avgpool2d_nhwc2nhwc") { - json_kernel_node = CreateCompositeMrvlAvgpool2DLayer(cn); - } else if (name == "mrvl.globalavgpool2d_nhwc2nhwc") { - json_kernel_node = CreateCompositeMrvlGlobalAvgpool2DLayer(cn); - } else if (name == "mrvl.globalmaxpool2d_nhwc2nhwc") { - json_kernel_node = CreateCompositeMrvlGlobalMaxpool2DLayer(cn); - } else if (name == "mrvl.sum") { - json_kernel_node = CreateCompositeMrvlSumLayer(cn); - } else if (name == "mrvl.concat") { - json_kernel_node = CreateMrvlConcatLayer(cn); - } else if (name == "mrvl.reshape") { - json_kernel_node = CreateMrvlReshapeLayer(cn); - } else if (name == "mrvl.batch_flatten") { - json_kernel_node = CreateMrvlBatchFlattenLayer(cn); - } else if (name == "mrvl.squeeze") { - json_kernel_node = CreateMrvlSqueezeLayer(cn); - } else { - LOG(FATAL) << "Unrecognized Mrvl pattern: " << name; - } - // calling codegen_json.h::AddNode() - return AddNode(json_kernel_node, GetRef(cn)); - } - - private: - /*! \brief The symbol that represents the layer json graph. */ - std::string layer_name_; - Array batch_norm_params_; - int node_idx_{0}; - int const_suffix_{0}; - - void resizeInputOutputLayoutTo4dim(std::shared_ptr json_node, const CallNode* cn, - std::string node_name) { - const uint64_t new_layout_size = 4; - std::string data_layout = "NHWC"; - std::string out_layout = "NHWC"; - - auto num_inputs = GetInputNum(cn); - auto num_outputs = GetOutputNum(cn); - uint64_t max_old_input_layout_size = 0; - // Inputs - if (num_inputs > 1) { - for (uint64_t in_idx = 0; in_idx < num_inputs; in_idx++) { - std::vector layout; - GetInputTensorShapeViaArgN(cn, &layout, in_idx); - uint64_t old_layout_size = layout.size(); - max_old_input_layout_size = std::max(old_layout_size, max_old_input_layout_size); - ICHECK(old_layout_size <= 4) << "Marvell-Compiler-ERROR-Internal::" << node_name - << " with input tensor shape > 4 is not supported yet."; - layout.resize(new_layout_size, 1); - - if (!cn->args[in_idx].as()) { - JsonNodeSetVecAttr(json_node, "data_layout_shape_" + std::to_string(in_idx), layout); - if (in_idx == 0) { - JsonNodeSetVecAttr(json_node, "data_layout_shape", layout); - } - } - } - for (uint64_t in_idx = 0; in_idx < num_inputs; in_idx++) { - std::vector layout; - GetInputTensorShapeViaArgN(cn, &layout, in_idx); - uint64_t old_layout_size = layout.size(); - ICHECK(old_layout_size <= 4) << "Marvell-Compiler-ERROR-Internal::" << node_name - << " with input tensor shape > 4 is not supported yet."; - layout.resize(max_old_input_layout_size, 1); - std::rotate(layout.begin(), layout.end() - (max_old_input_layout_size - old_layout_size), - layout.end()); - layout.resize(new_layout_size, 1); - if (cn->args[in_idx].as()) { - std::vector const_name = {layer_name_ + "_const_" + - std::to_string(const_suffix_++)}; - JsonNodeSetAttr(json_node, "input_const_name", const_name); - JsonNodeSetVecAttr(json_node, "input_const_shape", layout); - } - } - } else { - std::vector layout; - GetInputTensorShapeViaArgN(cn, &layout, 0); - layout.resize(new_layout_size, 1); - JsonNodeSetVecAttr(json_node, "data_layout_shape", layout); - } - // Outputs - if (num_outputs > 1) { - std::vector> layout; - GetOutputTensorShapes(cn, &layout); - for (size_t out_idx = 0; out_idx < num_outputs; out_idx++) { - ICHECK(layout.at(out_idx).size() <= 4) - << "Marvell-Compiler-ERROR-Internal::" << node_name - << " with output tensor shape > 4 is not supported yet."; - layout.at(out_idx).resize(new_layout_size, 1); - JsonNodeSetVecAttr(json_node, "out_layout_shape_" + std::to_string(out_idx), - layout.at(out_idx)); - if (out_idx == 0) { - JsonNodeSetVecAttr(json_node, "out_layout_shape", layout.at(out_idx)); - } - } - } else { - std::vector layout; - GetOutputTensorShape(cn, &layout); - layout.resize(new_layout_size, 1); - JsonNodeSetVecAttr(json_node, "out_layout_shape", layout); - } - - std::vector layout_format_vec = {data_layout}; - JsonNodeSetAttr(json_node, "data_layout", layout_format_vec); - JsonNodeSetAttr(json_node, "out_layout", layout_format_vec); - } - - /*! - * \brief Extract convolution nodes from a composite function. - * - * \param call The call node of the composite function. - * \return Extracted composite convolution nodes. - */ - CompositeConvNode UnpackCompositeConvolution(const CallNode* call) { - CompositeConvNode nodes{}; - const auto* fn = call->op.as(); - ICHECK(fn) << "Marvell-Compiler-ERROR-Internal::Downcast to FunctionNode failed."; - // - conv2d + [ bias_add ] + [ batch_norm + tuple.getitem(0) ] + [ relu ] - // Traverse composite convolution function from child to parent - const TupleGetItemNode* tuple_get_item_node = nullptr; - const CallNode* current_call = fn->body.as(); - if (current_call) { - if (backend::IsOp(current_call, "nn.relu")) { - nodes.activation = current_call; - if (current_call->args[0].as()) { - tuple_get_item_node = current_call->args[0].as(); - } else { - current_call = current_call->args[0].as(); - } - } else { - ICHECK(current_call) << "Marvell-Compiler-ERROR-Internal::Downcast to CallNode failed."; - } - } else { - tuple_get_item_node = fn->body.as(); - } - - if (tuple_get_item_node != nullptr) { - ICHECK(tuple_get_item_node->index == 0) - << "Marvell-Compiler-ERROR-Internal::(index == 0) failed for the TupleGetItem node."; - current_call = tuple_get_item_node->tuple.as(); - - ICHECK(backend::IsOp(current_call, "nn.batch_norm")) - << "Marvell-Compiler-ERROR-Internal::nn.batch_norm Op missing."; - nodes.batch_norm = current_call; - current_call = nodes.batch_norm->args[0].as(); - } - - ICHECK(current_call) << "Marvell-Compiler-ERROR-Internal::Downcast to CallNode failed."; - if (backend::IsOp(current_call, "add")) { - nodes.add = current_call; - current_call = current_call->args[0].as(); - } - - ICHECK(backend::IsOp(current_call, "nn.conv2d")) - << "Marvell-Compiler-ERROR-Internal::nn.conv2d Op missing."; - nodes.conv = current_call; - current_call = current_call->args[0].as(); - - if (current_call && backend::IsOp(current_call, "nn.pad")) { - nodes.pad = current_call; - } - return nodes; - } - - /*! - * \brief Extract sum nodes from a composite function. - * - * \param call The call node of the composite function. - * \return Extracted composite sum nodes. - */ - CompositeSumNode UnpackCompositeSum(const CallNode* call) { - CompositeSumNode nodes{}; - const auto* fn = call->op.as(); - ICHECK(fn) << "Marvell-Compiler-ERROR-Internal::Downcast to FunctionNode failed."; - - const auto* current_call = fn->body.as(); - if (backend::IsOp(current_call, "nn.relu")) { - nodes.activation = current_call; - current_call = current_call->args[0].as(); - } - ICHECK(backend::IsOp(current_call, "add")) - << "Marvell-Compiler-ERROR-Internal::add Op missing."; - nodes.add = current_call; - - return nodes; - } - - /*! - * \brief Extract Concat nodes from a composite function. - * - * \param call The call node of the composite function. - * \return Extracted composite Concat nodes. - */ - CompositeConcatNode UnpackCompositeConcat(const CallNode* call) { - CompositeConcatNode nodes{}; - const auto* fn = call->op.as(); - ICHECK(fn) << "Marvell-Compiler-ERROR-Internal::Downcast to FunctionNode failed."; - - const auto* current_call = fn->body.as(); - - ICHECK(backend::IsOp(current_call, "concatenate")) - << "Marvell-Compiler-ERROR-Internal::concatenate Op missing."; - nodes.concat = current_call; - - return nodes; - } - - /*! - * \brief Extract Reshape nodes from a composite function. - * - * \param call The call node of the composite function. - * \return Extracted composite Reshape nodes. - */ - CompositeReshapeNode UnpackCompositeReshape(const CallNode* call) { - CompositeReshapeNode nodes{}; - const auto* fn = call->op.as(); - ICHECK(fn) << "Marvell-Compiler-ERROR-Internal::Downcast to FunctionNode failed."; - const auto* current_call = fn->body.as(); - ICHECK(backend::IsOp(current_call, "reshape")) - << "Marvell-Compiler-ERROR-Internal::reshape missing."; - nodes.reshape = current_call; - return nodes; - } - - /*! - * \brief Extract Batch flatten nodes from a composite function. - * - * \param call The call node of the composite function. - * \return Extracted composite batch flatten nodes. - */ - CompositeBatchFlattenNode UnpackCompositeBatchFlatten(const CallNode* call) { - CompositeBatchFlattenNode nodes{}; - const auto* fn = call->op.as(); - ICHECK(fn) << "Marvell-Compiler-ERROR-Internal::Downcast to FunctionNode failed."; - const auto* current_call = fn->body.as(); - ICHECK(backend::IsOp(current_call, "nn.batch_flatten")) - << "Marvell-Compiler-ERROR-Internal::batch_flatten missing."; - nodes.batch_flatten = current_call; - return nodes; - } - - /*! - * \brief Extract squeeze nodes from a composite function. - * \param call The call node of the composite function. - * \return Extracted composite squeeze nodes. - */ - CompositeSqueezeNode UnpackCompositeSqueeze(const CallNode* call) { - CompositeSqueezeNode nodes{}; - const auto* fn = call->op.as(); - ICHECK(fn) << "Marvell-Compiler-ERROR-Internal::Downcast to FunctionNode failed."; - const auto* current_call = fn->body.as(); - ICHECK(backend::IsOp(current_call, "squeeze")) - << "Marvell-Compiler-ERROR-Internal::squeeze missing."; - nodes.squeeze = current_call; - return nodes; - } - - /*! - * \brief Extract maxpool nodes from a composite function. - * - * \param call The call node of the composite function. - * \return Extracted composite maxpool nodes. - */ - CompositePoolNode UnpackCompositePool(const CallNode* call, const std::string& mrvlLayerName) { - CompositePoolNode nodes{}; - const auto* fn = call->op.as(); - ICHECK(fn) << "Marvell-Compiler-ERROR-Internal::Downcast to FunctionNode failed."; - - // Traverse composite maxpool function from child to parent - const auto* current_call = fn->body.as(); - - if (mrvlLayerName == "Maxpool2D") { - ICHECK(backend::IsOp(current_call, "nn.max_pool2d")) - << "Marvell-Compiler-ERROR-Internal::nn.max_pool2d Op missing."; - } else if (mrvlLayerName == "Avgpool2D") { - ICHECK(mrvlLayerName == "Avgpool2D") - << "Marvell-Compiler-ERROR-Internal::nn.avg_pool2d Op missing."; - ICHECK(backend::IsOp(current_call, "nn.avg_pool2d")) - << "Marvell-Compiler-ERROR-Internal::nn.avg_pool2d Op missing."; - } else if (mrvlLayerName == "GlobalMaxpool2D") { - ICHECK(mrvlLayerName == "GlobalMaxpool2D") - << "Marvell-Compiler-ERROR-Internal::nn.global_max_pool2d Op missing."; - ICHECK(backend::IsOp(current_call, "nn.global_max_pool2d")) - << "Marvell-Compiler-ERROR-Internal::nn.global_max_pool2d Op missing."; - } else { - ICHECK(mrvlLayerName == "GlobalAvgpool2D") - << "Marvell-Compiler-ERROR-Internal::nn.global_avg_pool2d Op missing."; - ICHECK(backend::IsOp(current_call, "nn.global_avg_pool2d")) - << "Marvell-Compiler-ERROR-Internal::nn.global_avg_pool2d Op missing."; - } - nodes.pool = current_call; - current_call = current_call->args[0].as(); - if (current_call && backend::IsOp(current_call, "nn.pad")) { - nodes.pad = current_call; - } - - return nodes; - } - - /*! - * \brief Extract fc nodes from a composite function. - * - * \param call The call node of the composite function. - * \return Extracted composite fc nodes. - */ - CompositeFcNode UnpackCompositeFc(const CallNode* call) { - CompositeFcNode nodes{}; - const auto* fn = call->op.as(); - ICHECK(fn) << "Marvell-Compiler-ERROR-Internal::Downcast to FunctionNode failed."; - const auto* current_call = fn->body.as(); - - // Traverse composite fc function from child to parent - if (backend::IsOp(current_call, "nn.batch_flatten")) { - current_call = current_call->args[0].as(); - } - if (backend::IsOp(current_call, "nn.relu")) { - nodes.activation = current_call; - current_call = current_call->args[0].as(); - } - if (backend::IsOp(current_call, "add")) { - nodes.add = current_call; - current_call = current_call->args[0].as(); - } - ICHECK(backend::IsOp(current_call, "nn.dense")) - << "Marvell-Compiler-ERROR-Internal::nn.dense Op missing."; - nodes.fc = current_call; - current_call = current_call->args[0].as(); - if (current_call) { - if (backend::IsOp(current_call, "reshape") | - backend::IsOp(current_call, "nn.batch_flatten")) { - nodes.flatten = current_call; - current_call = current_call->args[0].as(); - ICHECK(backend::IsOp(current_call, "layout_transform")) - << "Marvell-Compiler-ERROR-Internal::layout_transform Op missing."; - nodes.transform = current_call; - } - } - - return nodes; - } - - void JsonNodeSetAttr(std::shared_ptr json_node, const std::string& key, - const std::vector& string_vec) { - std::vector json_attr; - json_attr.emplace_back(string_vec); - json_node->SetAttr(key, json_attr); - } - - void JsonNodeSetVecAttr(std::shared_ptr json_node, const std::string& key, - const std::vector& tvec) { - size_t tvec_size = tvec.size(); - std::vector tvec_str; - if (tvec_size == 4) { - tvec_str = {std::to_string(tvec[0]), std::to_string(tvec[1]), std::to_string(tvec[2]), - std::to_string(tvec[3])}; - } else if (tvec_size == 3) { - tvec_str = {std::to_string(tvec[0]), std::to_string(tvec[1]), std::to_string(tvec[2])}; - } else if (tvec_size == 2) { - tvec_str = {std::to_string(tvec[0]), std::to_string(tvec[1])}; - } else { - tvec_str = {std::to_string(tvec[0])}; - } - std::vector json_attr; - json_attr.emplace_back(tvec_str); - json_node->SetAttr(key, json_attr); - } - - void SetMrvlLayerBatchnormAttrs(std::shared_ptr json_node, - const CallNode* cn_batchnorm) { - if (cn_batchnorm == nullptr) return; - - SetCallNodeAttribute(json_node, cn_batchnorm); - - std::vector gamma_const_name; - std::vector beta_const_name; - std::vector mean_const_name; - std::vector var_const_name; - std::string batch_norm_layout = "-O"; - - gamma_const_name = {layer_name_ + "_const_" + std::to_string(const_suffix_++)}; - beta_const_name = {layer_name_ + "_const_" + std::to_string(const_suffix_++)}; - mean_const_name = {layer_name_ + "_const_" + std::to_string(const_suffix_++)}; - var_const_name = {layer_name_ + "_const_" + std::to_string(const_suffix_++)}; - - JsonNodeSetAttr(json_node, "gamma_const_name", gamma_const_name); - JsonNodeSetAttr(json_node, "beta_const_name", beta_const_name); - JsonNodeSetAttr(json_node, "mean_const_name", mean_const_name); - JsonNodeSetAttr(json_node, "var_const_name", var_const_name); - JsonNodeSetAttr(json_node, "gamma_layout", {batch_norm_layout}); - JsonNodeSetAttr(json_node, "beta_layout", {batch_norm_layout}); - JsonNodeSetAttr(json_node, "mean_layout", {batch_norm_layout}); - JsonNodeSetAttr(json_node, "var_layout", {batch_norm_layout}); - } - - void SetMrvlLayerPadAttrs(std::shared_ptr json_node, const CallNode* cn_pad) { - if (cn_pad == nullptr) return; - - const auto* pad_attr = cn_pad->attrs.as(); - ICHECK(pad_attr) << "Marvell-Compiler-ERROR-Internal::Downcast to PadAttrs failed."; - ICHECK(cn_pad->args[1].as() == 0) - << "Marvell-Compiler-ERROR-Internal::padded value is non-zero."; - ICHECK(pad_attr->pad_mode == "constant") - << "Marvell-Compiler-ERROR-Internal::unsupported padding mode."; - - auto p = pad_attr->pad_width; - // Convert to TVM layout for now, conversion to Mrvl layout takes place in runtime. - // Standard pad layout for TVM: top, left, bottom, right. - std::vector padding = {std::to_string(p[1][0].as()->value), - std::to_string(p[2][0].as()->value), - std::to_string(p[1][1].as()->value), - std::to_string(p[2][1].as()->value)}; - - JsonNodeSetAttr(json_node, "padding", {padding}); - } - - void SetMrvlLayerCommonAttrs(std::shared_ptr json_node, const CallNode* cn, - const std::string& func_name, const std::string& mrvlLayerName, - const std::string& data_layout, const std::string& kernel_layout, - const std::string& out_layout) { - JsonNodeSetAttr(json_node, "layer_name", {mrvlLayerName}); - JsonNodeSetAttr(json_node, "func_node_name", {func_name}); - std::vector data_layout_vec; - - auto num_inputs = GetInputNum(cn); - auto num_outputs = GetOutputNum(cn); - auto counter = num_inputs; - for (size_t i = 0; i < counter; i++) { - if (cn->args[i].as()) num_inputs--; - } - - std::vector tuple_idx_vec; - int tuple_idx = -1; - if (num_inputs > 1) { - for (size_t in_idx = 0; in_idx < num_inputs; in_idx++) { - std::vector data_layout_vec_n; - tuple_idx = GetInputTensorShapeViaArgN(cn, &data_layout_vec_n, in_idx); - std::string attr_name = "data_layout_shape_" + std::to_string(in_idx); - JsonNodeSetVecAttr(json_node, attr_name, data_layout_vec_n); - tuple_idx_vec.push_back(tuple_idx); - if (in_idx == 0) { - JsonNodeSetVecAttr(json_node, "data_layout_shape", data_layout_vec_n); - } - } - } else { - tuple_idx = GetInputTensorShapeViaArgN(cn, &data_layout_vec, 0); - JsonNodeSetVecAttr(json_node, "data_layout_shape", data_layout_vec); - tuple_idx_vec.push_back(tuple_idx); - } - JsonNodeSetVecAttr(json_node, "from_tuple_idx", tuple_idx_vec); - - if (data_layout != "") { - std::vector data_layout_format_vec = {data_layout}; - JsonNodeSetAttr(json_node, "data_layout", data_layout_format_vec); - } - - std::vector out_layout_vec; - if (num_outputs > 1) { - std::vector> output_layout_vec_vec; - GetOutputTensorShapes(cn, &output_layout_vec_vec); - for (size_t out_idx = 0; out_idx < num_outputs; out_idx++) { - std::string attr_name = "out_layout_shape_" + std::to_string(out_idx); - JsonNodeSetVecAttr(json_node, attr_name, output_layout_vec_vec.at(out_idx)); - } - // For compatibility with backend - JsonNodeSetVecAttr(json_node, "out_layout_shape", output_layout_vec_vec.at(0)); - } else { - GetOutputTensorShape(cn, &out_layout_vec); - JsonNodeSetVecAttr(json_node, "out_layout_shape", out_layout_vec); - } - - if (kernel_layout != "") { - std::vector kernel_layout_format_vec = {kernel_layout}; - JsonNodeSetAttr(json_node, "kernel_layout", kernel_layout_format_vec); - } - if (out_layout != "") { - std::vector out_layout_format_vec = {out_layout}; - JsonNodeSetAttr(json_node, "out_layout", out_layout_format_vec); - } - - // setup n<#>_ as GUI node name ("func_name") in nodes JSON file - std::string node_id_func_name = ""; - node_id_func_name = "n" + std::to_string(node_idx_++) + "_" + mrvlLayerName; - - // - add posfix layout(s) if applicable - if ((data_layout != "") && (out_layout != "")) { - node_id_func_name += "_" + data_layout; - if (data_layout != out_layout) { - node_id_func_name += "2" + out_layout; - } - } - - JsonNodeSetAttr(json_node, "func_name", {node_id_func_name}); - - const auto* fn = cn->op.as(); - if (fn != nullptr) { - ICHECK(fn->IsInstance()) - << "Marvell-Compiler-ERROR-Internal::Downcast to FunctionNode failed."; - auto composite = fn->GetAttr(attr::kComposite); - ICHECK(composite.defined()) - << "Marvell-Compiler-ERROR-Internal::Illegal Mrvl composite function."; - std::string composite_name = composite.value(); - JsonNodeSetAttr(json_node, "composite_name", {composite_name}); - } - } - - void GetInputTensorShapeFromTuple(const CallNode* call_node_ptr, size_t index, - std::vector* tensor_shape) { - ICHECK(!call_node_ptr->args.empty()); - const TensorTypeNode* tensor_type = nullptr; - if (call_node_ptr->args[0].as()) { - const auto* arg0 = call_node_ptr->args[0].as(); - tensor_type = arg0->checked_type_.as(); - } else if (call_node_ptr->args[0].as()) { - const auto* arg0 = call_node_ptr->args[0].as(); - ICHECK((arg0 != nullptr) && arg0->IsInstance()) - << "Marvell-Compiler-ERROR-Internal::Downcast to VarNode failed."; - tensor_type = arg0->checked_type_.as(); - const TupleTypeNode* tuple_type = arg0->checked_type_.as(); - if (tuple_type) { - tensor_type = tuple_type->fields[index].as(); - } - } else { - LOG(INFO) << "TVM Mrvl runtime does not support calls to " - << call_node_ptr->args[0]->GetTypeKey(); - } - - ICHECK((tensor_type != nullptr) && tensor_type->IsInstance()) - << "Marvell-Compiler-ERROR-Internal::Downcast to TensorTypeNode failed."; - for (IndexExpr dim_val : tensor_type->shape) { - tensor_shape->push_back(*(tir::as_const_int(dim_val))); - } - } - - size_t GetInputNum(const CallNode* call_node_ptr) { - size_t num_inputs = call_node_ptr->args.size(); - ICHECK(!call_node_ptr->args.empty()); - const TupleGetItemNode* tuple_get_item_node = call_node_ptr->args[0].as(); - const TensorTypeNode* tensor_type = nullptr; - if (tuple_get_item_node) { - tensor_type = tuple_get_item_node->checked_type().as(); - } else if (call_node_ptr->args[0].as()) { - num_inputs = call_node_ptr->args.size(); - } else if (call_node_ptr->args[0].as()) { - const auto* arg_0 = call_node_ptr->args[0].as(); - ICHECK((arg_0 != nullptr) && arg_0->IsInstance()) - << "Marvell-Compiler-ERROR-Internal::Downcast to VarNode failed."; - tensor_type = arg_0->checked_type_.as(); - if (tensor_type == nullptr) { - const TupleTypeNode* tuple_type = arg_0->checked_type_.as(); - if (tuple_type) { - num_inputs = tuple_type->fields.size(); - } - } - } else { - LOG(INFO) << "TVM Mrvl runtime does not support calls to " - << call_node_ptr->args[0]->GetTypeKey(); - } - return num_inputs; - } - - size_t GetOutputNum(const CallNode* call_node_ptr) { - ICHECK(call_node_ptr != nullptr); - const TupleTypeNode* tuple_type = call_node_ptr->checked_type_.as(); - if (tuple_type) { - return tuple_type->fields.size(); - } - // If output isn't a tuple, there is a single output - return 1; - } - - void GetInputTensorShapeViaArg(const CallNode* call_node_ptr, std::vector* tensor_shape, - int* tuple_index, size_t n) { - *tuple_index = -1; - ICHECK(!call_node_ptr->args.empty()); - const TensorTypeNode* tensor_type = nullptr; - const TupleGetItemNode* tuple_get_item_node = call_node_ptr->args[n].as(); - if (tuple_get_item_node) { - *tuple_index = tuple_get_item_node->index; - tensor_type = tuple_get_item_node->checked_type().as(); - } else if (call_node_ptr->args[n].as()) { - const auto* arg_n = call_node_ptr->args[n].as(); - tensor_type = arg_n->checked_type().as(); - } else if (call_node_ptr->args[n].as()) { - const auto* arg_n = call_node_ptr->args[n].as(); - ICHECK((arg_n != nullptr) && arg_n->IsInstance()) - << "Marvell-Compiler-ERROR-Internal::Downcast to VarNode failed."; - tensor_type = arg_n->checked_type().as(); - if (tensor_type == nullptr) { - const TupleTypeNode* tuple_type = arg_n->checked_type().as(); - if (tuple_type) { - tensor_type = tuple_type->fields[n].as(); - } - } else if (call_node_ptr->args[n].as()) { - const auto* arg_n = call_node_ptr->args[n].as(); - ICHECK((arg_n != nullptr) && arg_n->IsInstance()) - << "Marvell-Compiler-ERROR-Internal::Downcast to ConstantNode failed."; - tensor_type = arg_n->checked_type().as(); - if (tensor_type == nullptr) { - const TupleTypeNode* tuple_type = arg_n->checked_type().as(); - if (tuple_type) { - tensor_type = tuple_type->fields[n].as(); - } - } - } - } else { - LOG(INFO) << "TVM Mrvl runtime does not support calls to " - << call_node_ptr->args[n]->GetTypeKey(); - } - - ICHECK((tensor_type != nullptr) && tensor_type->IsInstance()) - << "Marvell-Compiler-ERROR-Internal::Downcast to TensorTypeNode failed."; - // use only data types supported by json.h (e.g., int or int64_t or size_t) - for (IndexExpr dim_val : tensor_type->shape) { - tensor_shape->push_back(*(tir::as_const_int(dim_val))); - } - } - - int GetInputTensorShapeViaArgN(const CallNode* call_node_ptr, std::vector* tensor_shape, - int64_t n = 0) { - int tuple_idx = -1; - GetInputTensorShapeViaArg(call_node_ptr, tensor_shape, &tuple_idx, n); - return tuple_idx; - } - - void GetTensorShape(const VarNode* var_node_ptr, std::vector* tensor_shape) { - ICHECK((var_node_ptr != nullptr) && var_node_ptr->IsInstance()) - << "Marvell-Compiler-ERROR-Internal::Downcast to VarNode failed."; - const TensorTypeNode* tensor_type = var_node_ptr->checked_type_.as(); - ICHECK((tensor_type != nullptr) && tensor_type->IsInstance()) - << "Marvell-Compiler-ERROR-Internal::Downcast to TensorTypeNode failed."; - // use only data types supported by json.h (e.g., int or int64_t or size_t) - for (IndexExpr dim_val : tensor_type->shape) { - tensor_shape->push_back(*(tir::as_const_int(dim_val))); - } - } - - void GetOutputTensorShape(const CallNode* call_node_ptr, std::vector* tensor_shape) { - ICHECK(call_node_ptr != nullptr); - const TensorTypeNode* tensor_type = call_node_ptr->checked_type_.as(); - ICHECK((tensor_type != nullptr) && tensor_type->IsInstance()) - << "Marvell-Compiler-ERROR-Internal::Downcast to TensorTypeNode failed."; - for (IndexExpr dim_val : tensor_type->shape) { - tensor_shape->push_back(*(tir::as_const_int(dim_val))); - } - } - - void GetOutputTensorShapes(const CallNode* call_node_ptr, - std::vector>* tensor_shapes) { - ICHECK(call_node_ptr != nullptr); - - const TupleTypeNode* tuple_type = call_node_ptr->checked_type_.as(); - ICHECK((tuple_type != nullptr) && tuple_type->IsInstance()) - << "Marvell-Compiler-ERROR-Internal::Downcast to TupleTypeNode failed."; - for (auto field : tuple_type->fields) { - const TensorTypeNode* tensor_type = field.as(); - ICHECK((tensor_type != nullptr) && tensor_type->IsInstance()) - << "Marvell-Compiler-ERROR-Internal::Downcast to TensorTypeNode failed."; - // use only data types supported by json.h (e.g., int or int64_t or size_t) - std::vector tensor_shape; - for (IndexExpr dim_val : tensor_type->shape) { - tensor_shape.push_back(*(tir::as_const_int(dim_val))); - } - tensor_shapes->push_back(tensor_shape); - } - } - - /*! - * \brief Create a JSON representation of a composite convolution. - * - * \param cn The call to be represented. - * \return A JSON representation of a specific operator. - */ - std::shared_ptr CreateCompositeMrvlConv2DLayer(const CallNode* cn) { - CompositeConvNode nodes = UnpackCompositeConvolution(cn); - const auto* conv_attrs = nodes.conv->attrs.as(); - ICHECK(conv_attrs) << "Marvell-Compiler-ERROR-Internal::Downcast to Conv2DAttrs failed."; - - std::string name; - std::string mrvlLayerName = ""; - std::string data_layout; - std::string kernel_layout; - std::string out_layout; - std::vector inputs; - - // data input tensor - inputs.push_back(VisitExpr(cn->args[0])[0]); - // weight tensor - inputs.push_back(VisitExpr(nodes.conv->args[1])[0]); - if (nodes.add) { - // bias tensor - inputs.push_back(VisitExpr(nodes.add->args[1])[0]); - } - if (nodes.batch_norm) { - // get gamma, beta, mean, and var of batch-norm - for (size_t const_idx = 0; const_idx <= 3; const_idx++) { - size_t arg_idx = const_idx + 1; - ICHECK(nodes.batch_norm->args[arg_idx].as()) - << "Marvell-Compiler-ERROR-Internal::Downcast to ConstantNode failed."; - auto n = nodes.batch_norm->args[arg_idx]; - auto it = memo_.find(n); - if (it != memo_.end()) { - memo_.erase(n); - } - inputs.push_back(VisitExpr(n)[0]); - } - } - - // Distinguish between normal and depth-wise convolution - data_layout = conv_attrs->data_layout; - kernel_layout = conv_attrs->kernel_layout; - out_layout = conv_attrs->out_layout; - int groups = conv_attrs->groups; - if ((groups != 1) && conv_attrs->channels.defined() && - tvm::tir::ExprDeepEqual()(conv_attrs->channels, conv_attrs->groups)) { - name = "nn.dw_conv2d_nhwc2nhwc"; - mrvlLayerName = "Conv2D"; - if (conv_attrs->groups == 1) { - ICHECK(kernel_layout == "IHWO") - << "Marvell-Compiler-ERROR-Internal::" - << "Kernel layout must be IHWO, has the module been pre-processed correctly?"; - } - } else { - name = "nn.conv2d_nhwc2nhwc"; - mrvlLayerName = "Conv2D"; - ICHECK(data_layout == "NHWC") - << "Marvell-Compiler-ERROR-Internal::" - << "Data layout must be NHWC, has the module been pre-processed correctly?"; - ICHECK(kernel_layout == "OHWI") - << "Marvell-Compiler-ERROR-Internal::" - << "Kernel layout must be OHWI, has the module been pre-processed correctly?"; - ICHECK(out_layout == "NHWC") - << "Marvell-Compiler-ERROR-Internal::" - << "Out layout must be NHWC, has the module been pre-processed correctly?"; - } - - // add json node attributes - auto json_node = std::make_shared(name, "kernel", inputs, 1); - SetCallNodeAttribute(json_node, nodes.conv); - std::vector kernel_const_name = {layer_name_ + "_const_" + - std::to_string(const_suffix_++)}; - JsonNodeSetAttr(json_node, "kernel_const_name", kernel_const_name); - - if (nodes.add) { - SetCallNodeAttribute(json_node, nodes.add); - std::vector bias_const_name = {layer_name_ + "_const_" + - std::to_string(const_suffix_++)}; - JsonNodeSetAttr(json_node, "bias_const_name", bias_const_name); - JsonNodeSetAttr(json_node, "bias_layout", {"---O"}); - } - if (nodes.pad) SetMrvlLayerPadAttrs(json_node, nodes.pad); - if (nodes.batch_norm) SetMrvlLayerBatchnormAttrs(json_node, nodes.batch_norm); - if (nodes.activation) JsonNodeSetAttr(json_node, "activation_type", {"relu"}); - SetMrvlLayerCommonAttrs(json_node, cn, layer_name_, mrvlLayerName, data_layout, "", out_layout); - return json_node; - } - - /*! - * \brief Create a JSON representation of a composite sum. - * - * \param cn The call to be represented. - * \return A JSON representation of a specific operator. - */ - std::shared_ptr CreateCompositeMrvlSumLayer(const CallNode* cn) { - CompositeSumNode nodes = UnpackCompositeSum(cn); - ICHECK(nodes.add != nullptr) - << "Marvell-Compiler-ERROR-Internal::attribute add can't be nullptr"; - - std::string mrvlLayerName = "Sum2D"; - std::string name = "sum"; - std::string data_layout; - std::string out_layout; - std::vector layout_vec; - std::vector inputs; - - for (auto arg : cn->args) { - inputs.push_back(VisitExpr(arg)[0]); - } - - // add json node attributes - auto json_node = std::make_shared(name, "kernel", inputs, 1); - SetCallNodeAttribute(json_node, nodes.add); - if (nodes.activation) JsonNodeSetAttr(json_node, "activation_type", {"relu"}); - SetMrvlLayerCommonAttrs(json_node, cn, layer_name_, mrvlLayerName, data_layout, "", out_layout); - resizeInputOutputLayoutTo4dim(json_node, cn, "Sum"); - return json_node; - } - - /*! - * \brief Create a JSON representation of a composite reshape. - * - * \param cn The call to be represented. - * \return A JSON representation of a specific operator. - */ - std::shared_ptr CreateMrvlReshapeLayer(const CallNode* cn) { - CompositeReshapeNode nodes = UnpackCompositeReshape(cn); - - std::string name = "reshape"; - std::string data_layout; - std::string out_layout; - std::vector layout_vec; - std::vector inputs; - - inputs.push_back(VisitExpr(cn->args[0])[0]); - GetInputTensorShapeViaArgN(nodes.reshape, &layout_vec); - ICHECK(layout_vec.size() == 2 || layout_vec.size() == 4) - << "Marvell-Compiler-ERROR-Internal::" - << "Reshape with input tensor dim != 2 or != 4 is not supported yet."; - if (layout_vec.size() == 4) { - data_layout = "NHWC"; - } else { - data_layout = "NC"; - } - layout_vec.clear(); - GetOutputTensorShape(cn, &layout_vec); - ICHECK(layout_vec.size() == 2 || layout_vec.size() == 4) - << "Marvell-Compiler-ERROR-Internal::" - << "Reshape with output tensor dim != 2 or !=4 is not supported yet."; - if (layout_vec.size() == 4) { - out_layout = "NHWC"; - } else { - out_layout = "NC"; - } - - auto json_node = std::make_shared(name, "kernel", inputs, 1); - SetMrvlLayerCommonAttrs(json_node, cn, layer_name_, name, data_layout, - "" /* no kernel_layout */, out_layout); - return json_node; - } - - /*! - * \brief Create a JSON representation of a composite batch flatten. - * - * \param cn The call to be represented. - * \return A JSON representation of a specific operator. - */ - std::shared_ptr CreateMrvlBatchFlattenLayer(const CallNode* cn) { - CompositeBatchFlattenNode nodes = UnpackCompositeBatchFlatten(cn); - - std::string name = "nn.batch_flatten"; - std::string data_layout; - std::string out_layout = "NC"; - std::vector layout_vec; - std::vector inputs; - - inputs.push_back(VisitExpr(cn->args[0])[0]); - GetInputTensorShapeViaArgN(nodes.batch_flatten, &layout_vec); - ICHECK(layout_vec.size() == 2 || layout_vec.size() == 4) - << "Marvell-Compiler-ERROR-Internal::" - << "nn.batch_flatten with input tensor dim != 2 or != 4 is not supported yet."; - if (layout_vec.size() == 4) { - data_layout = "NHWC"; - } else { - data_layout = "NC"; - } - layout_vec.clear(); - GetOutputTensorShape(cn, &layout_vec); - ICHECK(layout_vec.size() == 2) - << "Marvell-Compiler-ERROR-Internal::" - << "nn.batch_flatten with output tensor dim != 2 is not supported yet."; - - auto json_node = std::make_shared(name, "kernel", inputs, 1); - SetMrvlLayerCommonAttrs(json_node, cn, layer_name_, name, data_layout, - "" /* no kernel_layout */, out_layout); - return json_node; - } - - /*! - * \brief Create a JSON representation of a composite Squeeze. - * - * \param cn The call to be represented. - * \return A JSON representation of a specific operator. - */ - std::shared_ptr CreateMrvlSqueezeLayer(const CallNode* cn) { - CompositeSqueezeNode nodes = UnpackCompositeSqueeze(cn); - std::vector inputs; - std::string name = "squeeze"; - inputs.push_back(VisitExpr(cn->args[0])[0]); - std::vector layout_vec; - GetInputTensorShapeViaArgN(nodes.squeeze, &layout_vec); - std::string data_layout; - if (layout_vec.size() == 4) { - data_layout = "NHWC"; - } else { - data_layout = "NC"; - } - layout_vec.clear(); - std::string out_layout = "NC"; - auto json_node = std::make_shared(name, "kernel", inputs, 1); - SetMrvlLayerCommonAttrs(json_node, cn, layer_name_, name, data_layout, - "" /* no kernel_layout */, out_layout); - return json_node; - } - - /*! - * \brief Create a JSON representation of a composite concat. - * - * \param cn The call to be represented. - * \return A JSON representation of a specific operator. - */ - std::shared_ptr CreateMrvlConcatLayer(const CallNode* cn) { - CompositeConcatNode nodes = UnpackCompositeConcat(cn); - ICHECK(nodes.concat != nullptr) - << "Marvell-Compiler-ERROR-Internal::attribute concat can't be nullptr"; - - std::string mrvlLayerName = "Concat"; - std::string name = "concat"; - std::string data_layout; - std::string out_layout; - std::vector inputs; - - for (auto arg : cn->args) { - inputs.push_back(VisitExpr(arg)[0]); - } - - std::vector layout_vec; - GetInputTensorShapeViaArgN(cn, &layout_vec); - if (layout_vec.size() == 4) { - data_layout = "NHWC"; - out_layout = "NHWC"; - } else if (layout_vec.size() == 2) { - data_layout = "NC"; - out_layout = "NC"; - } - - auto json_node = std::make_shared(name, "kernel", inputs, 1); - SetCallNodeAttribute(json_node, nodes.concat); - SetMrvlLayerCommonAttrs(json_node, cn, layer_name_, mrvlLayerName, data_layout, "", out_layout); - - return json_node; - } - - /*! - * \brief Create a JSON representation of a composite fc (fully-connected) operator. - * - * \param cn The call to be represented. - * \return A JSON representation of a specific operator. - */ - std::shared_ptr CreateCompositeMrvlFcLayer(const CallNode* cn) { - CompositeFcNode nodes = UnpackCompositeFc(cn); - - std::string name = "nn.fc_ni2no"; - std::string mrvlLayerName = "FC"; - std::string data_layout = "NC"; - std::string kernel_layout = "OI"; - std::string out_layout = "NC"; - std::string bias_layout = "-O"; - std::vector inputs; - - inputs.push_back(VisitExpr(cn->args[0])[0]); - inputs.push_back(VisitExpr(nodes.fc->args[1])[0]); - if (nodes.add) { - inputs.push_back(VisitExpr(nodes.add->args[1])[0]); - } - - auto json_node = std::make_shared(name, "kernel", inputs, 1); - std::vector kernel_const_name = {layer_name_ + "_const_" + - std::to_string(const_suffix_++)}; - JsonNodeSetAttr(json_node, "kernel_const_name", kernel_const_name); - SetCallNodeAttribute(json_node, nodes.fc); - if (nodes.add) { - SetCallNodeAttribute(json_node, nodes.add); - std::vector bias_const_name = {layer_name_ + "_const_" + - std::to_string(const_suffix_++)}; - JsonNodeSetAttr(json_node, "bias_const_name", bias_const_name); - JsonNodeSetAttr(json_node, "bias_layout", {bias_layout}); - } - if (nodes.activation) JsonNodeSetAttr(json_node, "activation_type", {"relu"}); - if (nodes.transform && nodes.flatten) { - JsonNodeSetAttr(json_node, "weights_need_transform", {"yes"}); - data_layout = "NHWC"; - } - SetMrvlLayerCommonAttrs(json_node, cn, layer_name_, mrvlLayerName, data_layout, kernel_layout, - out_layout); - return json_node; - } - - /*! - * \brief Create a JSON representation of a composite (global) maxpooling operator. - * - * \param cn The call to be represented. - * \return A JSON representation of a specific operator. - */ - std::shared_ptr CreateCompositeMrvlMaxpool2DLayer(const CallNode* cn) { - std::string mrvlLayerName = "Maxpool2D"; - CompositePoolNode nodes = UnpackCompositePool(cn, mrvlLayerName); - const auto* maxpool_attr = nodes.pool->attrs.as(); - std::string name = "nn.maxpool2d_nhwc2nhwc"; - std::string data_layout = maxpool_attr->layout; - std::string out_layout = maxpool_attr->layout; - std::vector inputs; - - ICHECK(maxpool_attr) << "Marvell-Compiler-ERROR-Internal::Downcast to MaxPool2DAttrs failed."; - ICHECK(maxpool_attr->layout == "NHWC") - << "Marvell-Compiler-ERROR-Internal::" - << "Layout must be NHWC, has the module been pre-processed correctly?"; - - inputs.push_back(VisitExpr(cn->args[0])[0]); - auto json_node = std::make_shared(name, "kernel", inputs, 1); - SetCallNodeAttribute(json_node, nodes.pool); - auto pool_attrs = nodes.pool->attrs.as(); - std::vector kernel_layout_vec; - kernel_layout_vec.push_back(*(tir::as_const_int(pool_attrs->pool_size[0]))); - kernel_layout_vec.push_back(*(tir::as_const_int(pool_attrs->pool_size[1]))); - JsonNodeSetVecAttr(json_node, "kernel_layout_shape", kernel_layout_vec); - if (nodes.pad) SetMrvlLayerPadAttrs(json_node, nodes.pad); - SetMrvlLayerCommonAttrs(json_node, cn, layer_name_, mrvlLayerName, data_layout, "HW", - out_layout); - return json_node; - } - - /*! - * \brief Create a JSON representation of a composite (global) avgpooling operator. - * - * \param cn The call to be represented. - * \return A JSON representation of a specific operator. - */ - std::shared_ptr CreateCompositeMrvlAvgpool2DLayer(const CallNode* cn) { - std::string mrvlLayerName = "Avgpool2D"; - CompositePoolNode nodes = UnpackCompositePool(cn, mrvlLayerName); - const auto* avgpool_attr = nodes.pool->attrs.as(); - std::string name = "nn.avgpool2d_nhwc2nhwc"; - std::string data_layout = avgpool_attr->layout; - std::string out_layout = avgpool_attr->layout; - std::vector inputs; - - ICHECK(avgpool_attr) << "Marvell-Compiler-ERROR-Internal::Downcast to AvgPool2DAttrs failed."; - ICHECK(avgpool_attr->layout == "NHWC") - << "Marvell-Compiler-ERROR-Internal::" - << "Layout must be NHWC, has the module been pre-processed correctly?"; - - inputs.push_back(VisitExpr(cn->args[0])[0]); - auto json_node = std::make_shared(name, "kernel", inputs, 1); - SetCallNodeAttribute(json_node, nodes.pool); - auto pool_attrs = nodes.pool->attrs.as(); - std::vector kernel_layout_vec; - kernel_layout_vec.push_back(*(tir::as_const_int(pool_attrs->pool_size[0]))); - kernel_layout_vec.push_back(*(tir::as_const_int(pool_attrs->pool_size[1]))); - JsonNodeSetVecAttr(json_node, "kernel_layout_shape", kernel_layout_vec); - if (nodes.pad) SetMrvlLayerPadAttrs(json_node, nodes.pad); - SetMrvlLayerCommonAttrs(json_node, cn, layer_name_, mrvlLayerName, data_layout, "HW", - out_layout); - return json_node; - } - - /*! - * \brief Create a JSON representation of a composite globalavgpooling operator. - * - * \param cn The call to be represented. - * \return A JSON representation of a specific operator. - */ - std::shared_ptr CreateCompositeMrvlGlobalAvgpool2DLayer(const CallNode* cn) { - std::string mrvlLayerName = "GlobalAvgpool2D"; - CompositePoolNode nodes = UnpackCompositePool(cn, mrvlLayerName); - const auto* globalavgpool_attr = nodes.pool->attrs.as(); - std::string name = "nn.globalavgpool2d_nhwc2nhwc"; - std::string data_layout = globalavgpool_attr->layout; - std::string out_layout = globalavgpool_attr->layout; - std::vector inputs; - - ICHECK(globalavgpool_attr) - << "Marvell-Compiler-ERROR-Internal::Downcast to GlobalPool2DAttrs failed."; - ICHECK(globalavgpool_attr->layout == "NHWC") - << "Marvell-Compiler-ERROR-Internal::" - << "Layout must be NHWC, has the module been pre-processed correctly?"; - - inputs.push_back(VisitExpr(cn->args[0])[0]); - std::vector kernel_layout_vec; - std::vector data_layout_vec; - GetInputTensorShapeViaArgN(cn, &data_layout_vec); - ICHECK(data_layout_vec.size() == 4); - kernel_layout_vec.push_back(data_layout_vec[1]); - kernel_layout_vec.push_back(data_layout_vec[2]); - auto json_node = std::make_shared(name, "kernel", inputs, 1); - SetCallNodeAttribute(json_node, nodes.pool); - JsonNodeSetVecAttr(json_node, "kernel_layout_shape", kernel_layout_vec); - if (nodes.pad) SetMrvlLayerPadAttrs(json_node, nodes.pad); - - SetMrvlLayerCommonAttrs(json_node, cn, layer_name_, mrvlLayerName, data_layout, "HW", - out_layout); - return json_node; - } - - /*! - * \brief Create a JSON representation of a composite globalmaxpooling operator. - * - * A composite function is only created when using the uint8 datatype for these operators. - * - * \param cn The call to be represented. - * \return A JSON representation of a specific operator. - */ - std::shared_ptr CreateCompositeMrvlGlobalMaxpool2DLayer(const CallNode* cn) { - std::string mrvlLayerName = "GlobalMaxpool2D"; - std::string name = "nn.globalmaxpool2d_nhwc2nhwc"; - CompositePoolNode nodes = UnpackCompositePool(cn, mrvlLayerName); - - const auto* globalmaxpool_attr = nodes.pool->attrs.as(); - ICHECK(globalmaxpool_attr) - << "Marvell-Compiler-ERROR-Internal::Downcast to GlobalPool2DAttrs failed."; - ICHECK(globalmaxpool_attr->layout == "NHWC") - << "Marvell-Compiler-ERROR-Internal::" - << "Layout must be NHWC, has the module been pre-processed correctly?"; - - std::string data_layout = globalmaxpool_attr->layout; - std::string out_layout = globalmaxpool_attr->layout; - std::vector inputs; - std::vector kernel_layout_vec; - std::vector data_layout_vec; - GetInputTensorShapeViaArgN(cn, &data_layout_vec); - ICHECK(data_layout_vec.size() == 4); - kernel_layout_vec.push_back(data_layout_vec[1]); - kernel_layout_vec.push_back(data_layout_vec[2]); - inputs.push_back(VisitExpr(cn->args[0])[0]); - - // op_type_ is "kernel" - auto json_node = std::make_shared(name, "kernel", inputs, 1); - SetCallNodeAttribute(json_node, nodes.pool); - JsonNodeSetVecAttr(json_node, "kernel_layout_shape", kernel_layout_vec); - if (nodes.pad) SetMrvlLayerPadAttrs(json_node, nodes.pad); - - SetMrvlLayerCommonAttrs(json_node, cn, layer_name_, mrvlLayerName, data_layout, "HW", - out_layout); - return json_node; - } - - /*! - * \brief Create a JSON representation of an OpNode layer. - * - * \param cn The call to be represented. - * \return A JSON representation of a specific operator. - */ - std::shared_ptr CreateMrvlLayer4OpNode(const CallNode* cn) { - const auto* op_node = cn->op.as(); - ICHECK(op_node) << "Marvell-Compiler-ERROR-Internal::Downcast to OpNode failed."; - String op_name = op_node->name; - - std::string name = op_name; - std::string mrvlLayerName = op_name; - std::string data_layout = ""; - std::string out_layout = ""; - std::vector inputs; - inputs.push_back(VisitExpr(cn->args[0])[0]); - // op_type_ is "kernel" - auto json_node = std::make_shared(name, "kernel", inputs, 1); - if (op_name == "transpose") { - SetCallNodeAttribute(json_node, cn); - } else if (op_name == "layout_transform") { - SetCallNodeAttribute(json_node, cn); - auto layout_transform_attr = cn->attrs.as(); - data_layout = layout_transform_attr->src_layout; - out_layout = layout_transform_attr->dst_layout; - } else { - LOG(FATAL) << "Can't handle this OpNode: " << AsText(GetRef(cn), false); - } - SetMrvlLayerCommonAttrs(json_node, cn, layer_name_, mrvlLayerName, data_layout, - "" /* no kernel_layout */, out_layout); - return json_node; - } -}; - -std::vector split(const std::string& s, char delim) { - std::vector result; - std::stringstream ss(s); - std::string item; - while (getline(ss, item, delim)) { - result.push_back(item); - } - return result; -} - -/*! - * \brief Generate compiled model binary and then return a runtime module for Mrvl. - * - * \note This consists of a series of IR functions, which each represents - * a full Mrvl subgraph/region (in tvmc mode) or one fused Mrvl backend layer - * macro function (in dbg mode), that they can be computed on Mrvl accelerator. - * - * \param ref The ext_func Relay expression/module to be executed using extern ops. - * \return A runtime module. - */ -runtime::Module MrvlCompiler(const ObjectRef& ref) { - ICHECK(ref->IsInstance()) - << "Marvell-Compiler-ERROR-Internal::Downcast to FunctionNode failed."; - - Function func = Downcast(ref); - std::string func_name = backend::GetExtSymbol(func); - const std::string mrvl_run_mode = func->GetAttr("mode").value(); - runtime::Module runtime_lib; - - // Extract attributes from the frontend to be passed to the runtime - const std::string compiler_opt = func->GetAttr("compiler_opts_string").value(); - MrvlJSONSerializer serializer(func_name, func); - serializer.serialize(); - std::string graph_json = serializer.GetJSON(); - - // Collect Nodes.json and Const.json - const auto* get_json = runtime::Registry::Get("tvm.mrvl.GetNodesJSONString"); - std::string nodes_json_string = (*get_json)(graph_json); - auto consts_json_string = serializer.GetConstJSONString(); - - // Rename constants to a form acceptable by backend - const auto* modifyConsts = runtime::Registry::Get("tvm.mrvl.ModifyConstNames"); - std::string modified_json = (*modifyConsts)(nodes_json_string, consts_json_string); - auto json_vec = split(modified_json, '|'); - - // Extract attributes from the nodes_json by key-value lookup using Python API - // These are passed to hardware runtime module for initialization - const tvm::runtime::PackedFunc* json_lookup; - json_lookup = runtime::Registry::Get("tvm.mrvl.find_value_in_KV_pair"); - const std::string string_inp = (*json_lookup)(nodes_json_string, "num_subgraph_inputs"); - const int num_inputs = std::stoi(string_inp); - const std::string string_out = (*json_lookup)(nodes_json_string, "num_subgraph_outputs"); - const int num_outputs = std::stoi(string_out); - const std::string string_bsize = (*json_lookup)(nodes_json_string, "batch_size"); - const int batch_size = std::stoi(string_bsize); - - // Invoke Marvell Backend compiler to generate binary for sub graph - const auto* compile = runtime::Registry::Get("tvm.mrvl.CompileModel"); - std::string bin = (*compile)(func_name, json_vec[0], json_vec[1], compiler_opt); - - if (mrvl_run_mode == "sim") { - const auto* pf = runtime::Registry::Get("runtime.mrvl_runtime_create"); - ICHECK(pf != nullptr) << "Cannot find software simulator runtime module to create"; - runtime_lib = (*pf)(func_name, json_vec[0], bin); - } else if (mrvl_run_mode == "hw") { - const auto* pf = runtime::Registry::Get("runtime.mrvl_hw_runtime_create"); - ICHECK(pf != nullptr) << "Cannot find hardware runtime module to create"; - runtime_lib = (*pf)(func_name, json_vec[0], bin, num_inputs, num_outputs, batch_size); - } else { - ICHECK(0) << "Unrecognized Marvell Run Mode! " << mrvl_run_mode; - } - - return runtime_lib; -} - -TVM_REGISTER_GLOBAL("relay.ext.mrvl").set_body_typed(MrvlCompiler); - -} // namespace mrvl -} // namespace contrib - -} // namespace relay -} // namespace tvm diff --git a/src/relay/backend/contrib/mrvl/compiler_attr.cc b/src/relay/backend/contrib/mrvl/compiler_attr.cc deleted file mode 100644 index 86cb04ab3936..000000000000 --- a/src/relay/backend/contrib/mrvl/compiler_attr.cc +++ /dev/null @@ -1,68 +0,0 @@ -/* - * 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 src/relay/backend/contrib/mrvl/compiler_attr.cc - * \brief Marvell MLIP specific attributes - */ - -#include -#include -#include - -namespace tvm { -namespace relay { -namespace contrib { -namespace mrvl { - -/*! \brief Attributes to store the compiler options for Mrvl MLIP */ -struct MrvlCompilerConfigNode : public tvm::AttrsNode { - String mcpu; - IntImm num_tiles; - String mattr; - - TVM_DECLARE_ATTRS(MrvlCompilerConfigNode, "ext.attrs.MrvlCompilerConfigNode") { - TVM_ATTR_FIELD(mcpu) - .describe( - "The CPU class of Marvell(R) ML Inference Processor;" - "possible values = {cn10ka, cnf10kb}") - .set_default("cn10ka"); - TVM_ATTR_FIELD(num_tiles) - .describe("Maximum number of tiles that may be used, possible values = {1,2,4,8}") - .set_default(IntImm(DataType::Int(64), 8)); - TVM_ATTR_FIELD(mattr) - .describe("Attributes for MLIP; possible values = {quantize,wb_pin_ocm}") - .set_default(""); - } -}; - -class MrvlCompilerConfig : public Attrs { - public: - TVM_DEFINE_NOTNULLABLE_OBJECT_REF_METHODS(MrvlCompilerConfig, Attrs, MrvlCompilerConfigNode); -}; - -TVM_REGISTER_NODE_TYPE(MrvlCompilerConfigNode); -TVM_REGISTER_PASS_CONFIG_OPTION("relay.ext.mrvl.options", MrvlCompilerConfig); - -TVM_REGISTER_TARGET_KIND("mrvl", kDLCPU); - -} // namespace mrvl -} // namespace contrib -} // namespace relay -} // namespace tvm diff --git a/src/relay/backend/contrib/tensorrt/codegen.cc b/src/relay/backend/contrib/tensorrt/codegen.cc deleted file mode 100644 index 1dd5e3a4d772..000000000000 --- a/src/relay/backend/contrib/tensorrt/codegen.cc +++ /dev/null @@ -1,416 +0,0 @@ -/* - * 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 src/relay/backend/contrib/tensorrt/codegen.cc - * \brief Implementation of the TensorRT JSON serializer. - */ -#include -#include -#include - -#include -#include -#include - -#include "../../../transforms/compiler_function_utils.h" -#include "../../utils.h" -#include "../codegen_json/codegen_json.h" - -#if TVM_GRAPH_EXECUTOR_TENSORRT -#include "NvInfer.h" -#endif - -namespace tvm { -namespace relay { -namespace contrib { -namespace tensorrt { - -/*! - * \brief Check whether TensorRT graph executor is enabled. - * \return True if enabled, False if not. - */ -inline constexpr bool IsRuntimeEnabled() { -#if TVM_GRAPH_EXECUTOR_TENSORRT - return true; -#else - return false; -#endif // TVM_GRAPH_EXECUTOR_TENSORRT -} - -TVM_REGISTER_GLOBAL("relay.ext.tensorrt.is_runtime_enabled").set_body_typed(IsRuntimeEnabled); - -/*! - * \brief Get TensorRT version that TVM is built against. - * \return Array of three integers for major, minor, and patch, or empty array if TensorRT graph - * runtime is not enabled. - */ -Array GetVersion() { -#if TVM_GRAPH_EXECUTOR_TENSORRT - return {Integer(NV_TENSORRT_MAJOR), Integer(NV_TENSORRT_MINOR), Integer(NV_TENSORRT_PATCH)}; -#else - return {}; -#endif // TVM_GRAPH_EXECUTOR_TENSORRT -} - -TVM_REGISTER_GLOBAL("relay.ext.tensorrt.get_version").set_body_typed(GetVersion); - -/*! - * \brief Returns the "tensorrt" Target instance to use for compilation. - */ -Target GetTensorRTTarget() { - Target target = Target::Current(/*allow_not_defined=*/true); - if (!target.defined() || target->kind->name != "tensorrt") { - // Since we allow partition_for_tensorrt to use the default "tensorrt" target, we should - // similarly allow the custom pass to execute without a specific "tensorrt" target in scope. - target = Target("tensorrt"); - } - return target; -} - -using JSONGraphNode = tvm::runtime::json::JSONGraphNode; -using JSONGraphNodeEntry = tvm::runtime::json::JSONGraphNodeEntry; -using JSONGraphObjectPtr = backend::contrib::JSONGraphObjectPtr; -using OpAttrExtractor = backend::contrib::OpAttrExtractor; -using JSONSerializer = backend::contrib::JSONSerializer; - -class TensorRTJSONSerializer; - -/*! - * \brief Collect the constants and attributes from all operator calls in the body - * of a "Composite" function. - */ -class CollectFromCompositeFunctionBody : public ExprVisitor { - public: - explicit CollectFromCompositeFunctionBody(TensorRTJSONSerializer* serializer) - : serializer_(serializer), node_(std::make_shared()) {} - - // We'll need to implement these out-of-band since they use the serializer. - void VisitExpr_(const ConstantNode* constant_node) final; - void VisitExpr_(const CallNode* call_node) final; - - void SetPadNodeAttribute(const CallNode* call_node) { - const auto* pad_attr = call_node->attrs.as(); - ICHECK(pad_attr); - auto p = pad_attr->pad_width; - const int dim_h = (p.size() == 5) ? 3 : 2; - const int dim_w = (p.size() == 5) ? 4 : 3; - std::vector padding = {std::to_string(p[dim_h][0].as()->value), - std::to_string(p[dim_w][0].as()->value), - std::to_string(p[dim_h][1].as()->value), - std::to_string(p[dim_w][1].as()->value)}; - std::vector padding_attr; - padding_attr.emplace_back(padding); - node_->SetAttr("padding", padding_attr); - } - - void SetStridedSliceNodeAttribute(const CallNode* call_node) { - const auto* attrs = call_node->attrs.as(); - ICHECK(attrs && attrs->begin && attrs->end && attrs->strides) - << "StridedSlice must have static begin, end, and strides."; - const bool default_strides = - !attrs->strides.value().defined() || attrs->strides.value().size() == 0; - auto ishape = backend::GetShape(call_node->args[0]->checked_type()); - - auto process_slice_index = [](Integer x, int default_value, int dim_value) { - if (!x.defined()) return default_value; - int value = x.as()->value; - if (value < 0) value += dim_value; - return value; - }; - - std::vector start, size, strides; - for (size_t i = 0; i < attrs->begin.value().size(); ++i) { - const int begin_value = process_slice_index(attrs->begin.value()[i], 0, ishape[i]); - ICHECK_GE(begin_value, 0); - start.push_back(std::to_string(begin_value)); - const int stride_value = (default_strides || i >= attrs->strides.value().size() || - !attrs->strides.value()[i].defined()) - ? 1 - : attrs->strides.value()[i].as()->value; - ICHECK_GT(stride_value, 0); - strides.push_back(std::to_string(stride_value)); - int size_value; - if (attrs->slice_mode == "end") { - const int end_value = process_slice_index(attrs->end.value()[i], ishape[i], ishape[i]); - size_value = (end_value - begin_value + stride_value - 1) / stride_value; - } else if (attrs->slice_mode == "size") { - // with slice_mode = "size", attrs->end_value mean the size of the slice - int end_value = attrs->end.value()[i].as()->value; - size_value = (end_value == -1) ? ishape[i] - begin_value : end_value; - } else { - LOG(FATAL) << "Unexpected slice_mode " << attrs->slice_mode << ", expected end or size"; - throw; - } - ICHECK_GT(size_value, 0); - size.push_back(std::to_string(size_value)); - } - std::vector start_attr, size_attr, strides_attr; - start_attr.emplace_back(start); - size_attr.emplace_back(size); - strides_attr.emplace_back(strides); - node_->SetAttr("start", start_attr); - node_->SetAttr("size", size_attr); - node_->SetAttr("strides", strides_attr); - } - - void SetSplitNodeAttribute(const CallNode* call_node) { - const auto* split_attr = call_node->attrs.as(); - ICHECK(split_attr); - - std::vector indices_or_sections; - std::vector mode; - std::vector axis = {std::to_string(split_attr->axis)}; - if (const auto* sections = split_attr->indices_or_sections.as()) { - mode.emplace_back("sections"); - indices_or_sections.emplace_back(std::to_string(sections->value)); - } else { - mode.emplace_back("indices"); - auto indices = Downcast>(split_attr->indices_or_sections); - for (const auto& i : indices) { - indices_or_sections.emplace_back(std::to_string(i->value)); - } - } - - std::vector indices_or_sections_attr; - std::vector mode_attr; - std::vector axis_attr; - indices_or_sections_attr.emplace_back(indices_or_sections); - mode_attr.emplace_back(mode); - axis_attr.emplace_back(axis); - node_->SetAttr("indices_or_sections", indices_or_sections_attr); - node_->SetAttr("mode", mode_attr); - node_->SetAttr("axis", axis_attr); - } - - void SetGenericAttributes(const CallNode* call_node) { - OpAttrExtractor extractor(node_); - const Object* attr_obj = call_node->attrs.get(); - extractor.Extract(const_cast(attr_obj)); - } - - /*! \brief The parent serializer for the overall TensorRT partition. */ - TensorRTJSONSerializer* serializer_; - /*! \brief Accumulated translated arguments. */ - std::vector args_; - /*! - * \brief Temporary node into which we'll accumulate attributes. Ideally this would be the - * final JSONGraphNode however we don't yet know how many inputs that will have. - */ - JSONGraphObjectPtr node_; -}; - -/*! - * \brief Generates an TensorRTModule from a relay expression by serializing the expression to a - * json representation. TensorRT is not required here because use of TensorRT APIs is deferred until - * runtime. - */ -class TensorRTJSONSerializer : public JSONSerializer { - public: - TensorRTJSONSerializer(Target target, const std::string& symbol, const Expr& expr) - : JSONSerializer(symbol, expr), target_(std::move(target)) {} - - private: - using JSONSerializer::VisitExpr_; - - std::vector VisitExpr_(const CallNode* call_node) final { - // The call must be to an inline "Composite" function - const auto* function_node = call_node->op.as(); - ICHECK(function_node != nullptr); - auto opt_composite = function_node->GetAttr(attr::kComposite); - ICHECK(opt_composite.defined()); - std::string name = opt_composite.value(); - - // Collect the constants and attributes of all operator calls inside the composite body. - CollectFromCompositeFunctionBody collector(this); - collector.VisitExpr(function_node->body); - - // Capture the args to the "Composite" function as inputs for this node. - std::vector inputs; - for (const auto& arg : call_node->args) { - auto res = VisitExpr(arg); - inputs.insert(inputs.end(), res.begin(), res.end()); - } - - // Capture constants from the composite function body as additional inputs for this node. - for (const auto& node : collector.args_) { - inputs.emplace_back(node); - } - - // Create the final node. - auto node = std::make_shared(name, - /*op_type=*/"kernel", inputs, - /*num_output=*/1); - - // Transfer attributes from the collector's node to the final node. - node->CaptureAttrs(*collector.node_); - - // Capture global settings on the JSON node. - // TODO(mbs): Why on every call? - SaveGlobalAttributes(node.get()); - - VLOG(1) << name << " has " << node->GetInputs().size() << " inputs"; - - return AddNode(node, GetRef(call_node)); - } - - static void SetAttr(JSONGraphNode* node, const std::string& key, - std::vector values) { - node->SetAttr(key, std::vector({std::move(values)})); - } - - /*! \brief Capture the compilation options as attributes on \p node. */ - void SaveGlobalAttributes(JSONGraphNode* node) { - { - // cf logic in tensorrt.py::get_tensorrt_version. - // First check for version in target. - Array target_attr = target_->GetAttr>("tensorrt_version").value(); - if (target_attr.empty()) { - // Next, ask runtime for its version. - target_attr = GetVersion(); - } - if (target_attr.empty()) { - // Finally, use default. - target_attr = {6, 0, 1}; - } - ICHECK_EQ(target_attr.size(), 3); - SetAttr(node, "tensorrt_version", - {std::to_string(target_attr[0]->value), std::to_string(target_attr[1]->value), - std::to_string(target_attr[2]->value)}); - } - - { - Bool target_attr = target_->GetAttr("use_implicit_batch").value(); - SetAttr(node, "use_implicit_batch", {std::to_string(target_attr->value)}); - } - - { - Integer target_attr = target_->GetAttr("max_workspace_size").value(); - SetAttr(node, "max_workspace_size", {std::to_string(target_attr->value)}); - } - - { - Bool target_attr = target_->GetAttr("use_fp16").value(); - SetAttr(node, "use_fp16", {std::to_string(target_attr->value)}); - } - - { - Bool target_attr = target_->GetAttr("use_uint8").value(); - SetAttr(node, "use_uint8", {std::to_string(target_attr->value)}); - } - } - - /*! \brief The "tensorrt" Target guiding compilation. */ - Target target_; -}; - -void CollectFromCompositeFunctionBody::VisitExpr_(const ConstantNode* constant_node) { - for (const auto& entry : serializer_->VisitExpr(GetRef(constant_node))) { - args_.emplace_back(entry); - } -} - -void CollectFromCompositeFunctionBody::VisitExpr_(const CallNode* call_node) { - const auto* op_node = call_node->op.as(); - ICHECK(op_node != nullptr); - std::string name = op_node->name; - if (name == "nn.pad") { - SetPadNodeAttribute(call_node); - } else if (name == "strided_slice") { - SetStridedSliceNodeAttribute(call_node); - } else if (name == "split") { - SetSplitNodeAttribute(call_node); - } else { - SetGenericAttributes(call_node); - } - ExprVisitor::VisitExpr_(call_node); -} - -/*! - * \brief The main TensorRT compiler. - * - * TODO(mbs): Currently we create a \p TensorRTRuntimeModule for every function with - * Compiler="tensorrt" (ie for each partition). Since the TensorRT engine is only designed to - * handle a single entry point this is mostly sensible, however there are probably opportunities - * for more sharing between functions. However, note this means each call to a TensorRT-compiled - * function will require a linear scan of imported runtime modules to find the matching - * TensorRTRuntimeModule implementing it. - */ -tvm::transform::Pass CompileForTensorRTImpl() { - auto pass_func = [](IRModule mod, const tvm::transform::PassContext& pass_ctx) { - VLOG(1) << "CompileForTensorRT input:" << std::endl << PrettyPrint(mod); - Target target = GetTensorRTTarget(); - - const auto* pf = runtime::Registry::Get("runtime.tensorrt_runtime_create"); - ICHECK(pf != nullptr) << "Cannot find TensorRT runtime module create function."; - - // The accumulated external runtime modules. - Array external_mods = - mod->GetAttr>(tvm::attr::kExternalMods).value_or({}); - // The accumulated constant bindings. - Map const_name_to_constant = - mod->GetAttr>(tvm::attr::kConstNameToConstant).value_or({}); - - for (const auto& kv : mod->functions) { - if (const auto* function_node = kv.second.as()) { - if (function_node->HasNonzeroAttr(attr::kPrimitive)) { - Optional opt_compiler = function_node->GetAttr(attr::kCompiler); - if (opt_compiler && opt_compiler.value() == "tensorrt") { - // Serialize the function to JSON. - TensorRTJSONSerializer serializer(target, kv.first->name_hint, - GetRef(function_node)); - serializer.serialize(); - std::string graph_json = serializer.GetJSON(); - VLOG(1) << "TensorRT JSON for '" << kv.first->name_hint << "':" << std::endl - << graph_json; - - // Remember all the constant bindings. - for (const auto& kv2 : serializer.const_name_to_constant()) { - ICHECK_EQ(const_name_to_constant.count(kv2.first), 0); - VLOG(1) << "binding constant '" << kv2.first << "' for function '" - << kv.first->name_hint << "'"; - const_name_to_constant.Set(kv2.first, kv2.second); - } - - // Create the actual runtime module. - runtime::Module runtime_mod = - (*pf)(kv.first->name_hint, graph_json, serializer.const_names()); - - // Remember the runtime module. - external_mods.push_back(runtime_mod); - } - } - } - } - return WithAttrs(mod, {{tvm::attr::kExternalMods, external_mods}, - {tvm::attr::kConstNameToConstant, const_name_to_constant}}); - }; - return tvm::transform::CreateModulePass(pass_func, 0, "CompileForTensorRT", {}); -} - -tvm::transform::Pass CompileForTensorRT() { - return transform::Sequential( - {transform::OutlineCompilerFunctionsWithExistingGlobalSymbols("tensorrt"), - CompileForTensorRTImpl(), transform::MarkCompilerFunctionsAsExtern("tensorrt")}); -} - -} // namespace tensorrt -} // namespace contrib -} // namespace relay -} // namespace tvm diff --git a/src/relay/backend/contrib/tensorrt/codegen.h b/src/relay/backend/contrib/tensorrt/codegen.h deleted file mode 100644 index 813a8663756d..000000000000 --- a/src/relay/backend/contrib/tensorrt/codegen.h +++ /dev/null @@ -1,47 +0,0 @@ -/* - * 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 src/relay/backend/contrib/tensorrt/codegen.h - * \brief The 'custom' compilation pass for TensorRT (invoked by the RelayToTIRTargetHook pass). - */ - -#ifndef TVM_RELAY_BACKEND_CONTRIB_TENSORRT_CODEGEN_H_ -#define TVM_RELAY_BACKEND_CONTRIB_TENSORRT_CODEGEN_H_ - -#include - -namespace tvm { -namespace relay { -namespace contrib { -namespace tensorrt { - -/*! - * \brief Returns the pass which replaces all calls to "Primitive" functions with a "Compiler" - * attribute of "tensorrt" with calls to an extern which is implemented by a \p TensorRTRuntime - * runtime module added to the IRModule's "external_mods" attribute. - */ -transform::Pass CompileForTensorRT(); - -} // namespace tensorrt -} // namespace contrib -} // namespace relay -} // namespace tvm - -#endif // TVM_RELAY_BACKEND_CONTRIB_TENSORRT_CODEGEN_H_ diff --git a/src/relay/backend/contrib/tensorrt/target.cc b/src/relay/backend/contrib/tensorrt/target.cc deleted file mode 100644 index a62dc25e329c..000000000000 --- a/src/relay/backend/contrib/tensorrt/target.cc +++ /dev/null @@ -1,69 +0,0 @@ -/* - * 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 src/relay/backend/contrib/tensorrt/target.cc - * \brief Registers the "tensorrt" external codegen TargetKind. - */ - -#include - -#include "./codegen.h" - -namespace tvm { -namespace relay { -namespace contrib { -namespace tensorrt { - -/*! - * \brief This external codegen target can offload compilation to the TensorRT compiler. - * - Patterns: python/tvm/relay/op/contrib/tensorrt.py - * - Custom compiler: src/relay/backend/contrib/tensorrt/codegen.cc - * - Runtime: src/runtime/contrib/tensorrt/... - */ -TVM_REGISTER_TARGET_KIND("tensorrt", kDLCUDA) - .set_attr(tvm::attr::kIsExternalCodegen, runtime::Bool(true)) - .set_attr("RelayToTIR", CompileForTensorRT()) - // A array of three integers given the major, minor, and patch numbers for the supported - // TensorRT compiler version. If empty will be auto-detected from linked library. Default empty. - .add_attr_option>("tensorrt_version", Array()) - // If true, the first tensor dimension for most operators is allowed to be Any and - // TensorRT will assume it represents a batch dimension only known at inference time. - // Fewer Relay operators are supported in implicit batch mode. Default true. - .add_attr_option("use_implicit_batch", runtime::Bool(true)) - // If true, excludes sub-graphs which do not have multiply-accumulate operations, even though - // TensorRT supports them. ad. This is a simple heuristic to optimize the partitioning between - // TensorRT and TVM. Not required if using Collage for partitioning. Defalut false. - .add_attr_option("remove_no_mac_subgraphs", runtime::Bool(false)) - // How many bytes of workspace size to allow each subgraph to use for TensorRT engine creation. - // Default 1G. - .add_attr_option("max_workspace_size", runtime::Int(1 << 30)) - // If true, allows TensorRT to automatically convert float32 operations to float16. Must also be - // enabled if any float16 operations are in the model. Note that TensorRT may still choose a - // higher-precision kernel if it results in overall lower runtime, or if no low-precision - // implementation exists. Default false. - .add_attr_option("use_fp16", runtime::Bool(false)) - // If true, allows TensorRT to automatically convert float32 operations to uint8 - // (aka quantized). Default false. - .add_attr_option("use_uint8", runtime::Bool(false)); - -} // namespace tensorrt -} // namespace contrib -} // namespace relay -} // namespace tvm diff --git a/src/relay/backend/contrib/uma/relay_to_tir.cc b/src/relay/backend/contrib/uma/relay_to_tir.cc deleted file mode 100644 index ca3ae0ebec6b..000000000000 --- a/src/relay/backend/contrib/uma/relay_to_tir.cc +++ /dev/null @@ -1,175 +0,0 @@ -/* - * 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 relay/backend/contrib/uma/codegen.cc - * - * \brief this file contains the target hooks for the Universal Modular Accelerator Interface (UMA). - */ - -#include -#include -#include -#include -#include -#include -#include -#include - -#include -#include -#include -#include - -namespace tvm { -namespace relay { -namespace contrib { -namespace uma { - -// TODO(@mjklaiber, @manupa-arm, @areusch) move this to include -/*! - * \brief This mutator outlines functions that are marked with a named - * "Compiler" attribute. Functions that do not match this condition remain - * unaltered. - */ -class OutlineCompilerFunctionsMutator : public MixedModeMutator { - public: - explicit OutlineCompilerFunctionsMutator(const IRModule& mod, const std::string& compiler_name) - : mod_(mod), compiler_name_(compiler_name) {} - - Expr VisitExpr_(const LetNode* op) final { - auto pre_visit = [this](const LetNode* op) { - Expr var = this->VisitExpr(op->var); - Expr value = this->VisitExpr(op->value); - - // Outlineable function no longer needs let binding - if (this->CanOutlineExpr(value)) { - this->memo_[var] = value; - } - }; - auto post_visit = [this](const LetNode* op) { - // Rely on the Memoizer to cache pre-visit values - Expr value = this->VisitExpr(op->value); - Expr body = this->VisitExpr(op->body); - auto expr = GetRef(op); - - // Drop the let binding - if (this->CanOutlineExpr(value)) { - this->memo_[expr] = this->VisitExpr(op->body); - } else { - Var var = Downcast(this->VisitExpr(op->var)); - if (var.same_as(op->var) && value.same_as(op->value) && body.same_as(op->body)) { - this->memo_[expr] = expr; - } else { - this->memo_[expr] = Let(var, value, body); - } - } - }; - ExpandANormalForm(op, pre_visit, post_visit); - return memo_[GetRef(op)]; - } - - Expr Rewrite_(const CallNode* pre, const Expr& post) override { - Call call = Downcast(post); - if (CanOutlineExpr(call->op)) { - Function func = Downcast(call->op); - auto gv_name = func->GetAttr("global_symbol").value_or(""); - ICHECK_NE(gv_name, "") - << "Function to be outlined must have global_symbol attribute, but didn't."; - GlobalVar gv(gv_name); - if (func->checked_type_.defined()) { - gv->checked_type_ = func->checked_type(); - } - mod_->Update(gv, func); - return Call(gv, call->args, call->attrs, call->type_args); - } - return post; - } - - private: - /*! - * \brief Check if the expr is a function and has the same - * compiler name as compiler_name_. - * - * \param expr The input expr. - * \return True if is outlineable else False. - */ - bool CanOutlineExpr(const Expr& expr) { - if (!expr->IsInstance()) { - return false; - } - Function func = Downcast(expr); - auto compiler = func->GetAttr(attr::kCompiler); - if (!compiler.defined()) { - return false; - } - if (compiler != compiler_name_) { - return false; - } - return true; - } - - /*! \brief The module that the pass will run on. */ - IRModule mod_; - /*! \brief The name of the compiler to enable outlining on external functions for. */ - std::string compiler_name_; -}; - -/*! - * \brief A pass to outline compiler specific functions. - */ -tvm::transform::Pass OutlineCompilerFunctions(const std::string& compiler_name) { - runtime::TypedPackedFunc pass_func = - [=](IRModule mod, transform::PassContext ctx) { - GlobalVar gv = mod->GetGlobalVar("main"); - Function main_func = Downcast(mod->Lookup("main")); - auto new_main_body = - OutlineCompilerFunctionsMutator(mod, compiler_name).VisitExpr(main_func->body); - if (!new_main_body.same_as(main_func->body)) { - Function new_main_func = WithFields(main_func, main_func->params, new_main_body); - mod->Update(gv, new_main_func); - } - return mod; - }; - return tvm::transform::CreateModulePass(pass_func, 0, - "relay.backend.contrib.uma.OutlineCompilerFunctions", {}); -} - -TVM_REGISTER_GLOBAL("relay.ext.uma.OutlineCompilerFunctions") - .set_body_typed(OutlineCompilerFunctions); - -/*! - * \brief This pass will lower UMA functions in a Relay module to scheduled TIR prim functions. - */ -tvm::transform::Pass RelayToTIR(String target_name) { - runtime::TypedPackedFunc pass_func = - [=](IRModule ir_module, transform::PassContext pass_context) { - auto relay_to_tir_pf = - tvm::runtime::Registry::Get("relay.ext.uma." + target_name + ".relay_to_tir"); - ICHECK(relay_to_tir_pf); - ir_module = (*relay_to_tir_pf)(ir_module); - return ir_module; - }; - return tvm::transform::CreateModulePass(pass_func, 0, "relay.contrib.uma.RelayToTIR", {}); -} - -} // namespace uma -} // namespace contrib -} // namespace relay -} // namespace tvm diff --git a/src/relay/backend/contrib/uma/targets.cc b/src/relay/backend/contrib/uma/targets.cc deleted file mode 100644 index 0499c0bba198..000000000000 --- a/src/relay/backend/contrib/uma/targets.cc +++ /dev/null @@ -1,89 +0,0 @@ -/* - * 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 relay/backend/contrib/uma/targets.cc - * - * \brief this file contains the targets for the Universal Modular Accelerator Interface (UMA). - */ - -#include -#include - -namespace tvm { - -using FTVMTIRToRuntime = tvm::runtime::TypedPackedFunc; - -namespace relay { -namespace contrib { -namespace uma { -transform::Pass RelayToTIR(String target_name); -runtime::Module TIRToRuntime(IRModule mod, Target target); -} // namespace uma -} // namespace contrib -} // namespace relay - -TVM_REGISTER_GLOBAL("relay.backend.contrib.uma.RegisterTarget") - .set_body_typed([](String target_name, Map attr_options) -> bool { - // create only new target and init only once - for (const String registered_target_name : TargetKindRegEntry::ListTargetKinds()) { - if (registered_target_name == target_name) { - LOG(FATAL) << "TVM UMA Error: Target is already registered: " << target_name; - } - } - - auto target_kind = - TargetKindRegEntry::RegisterOrGet(target_name) - .set_name() - .set_default_device_type(kDLCPU) - .add_attr_option>("keys") - .add_attr_option("tag") - .add_attr_option("device") - .add_attr_option("model") - .add_attr_option>("libs") - .add_attr_option("host") - .add_attr_option("from_device") - .set_attr( - attr::kRelayToTIR, relay::contrib::uma::RelayToTIR(target_name)) - .set_attr("TIRToRuntime", relay::contrib::uma::TIRToRuntime); - - // target kind attrs inventory - auto kind = TargetKind::Get(target_name).value(); - auto list_attrs = TargetKindRegEntry::ListTargetKindOptions(kind); - - for (auto& attr_option : attr_options) { - auto option_name = attr_option.first; - auto default_value = attr_option.second; - if (list_attrs.find(option_name) != list_attrs.end()) { - LOG(FATAL) << "TVM UMA Error: Attribute is already registered: " << option_name; - } - if (default_value->IsInstance()) { - target_kind.add_attr_option(option_name, Downcast(default_value)); - } else if (default_value->IsInstance()) { - target_kind.add_attr_option(option_name, - Downcast(default_value)); - } else { - LOG(FATAL) << "TypeError: Only String, Integer, or Bool are supported. " - << "Given attribute option type: " << attr_option.second->GetTypeKey(); - } - } - return true; - }); - -} // namespace tvm diff --git a/src/relay/backend/contrib/uma/tir_to_runtime.cc b/src/relay/backend/contrib/uma/tir_to_runtime.cc deleted file mode 100644 index 487e247f5d38..000000000000 --- a/src/relay/backend/contrib/uma/tir_to_runtime.cc +++ /dev/null @@ -1,89 +0,0 @@ -/* - * 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 "../../../../runtime/file_utils.h" -#include "../../../../target/source/codegen_c.h" -#include "../../../../target/source/codegen_c_host.h" - -namespace tvm { -using namespace tir; -namespace relay { -namespace contrib { -namespace uma { - -class UMACodegen : public codegen::CodeGenCHost { - public: - explicit UMACodegen(String target_str) : target_str_(target_str) {} - - void Init(bool output_ssa, bool emit_asserts, bool emit_fwd_func_decl) { - auto includes_pf = - tvm::runtime::Registry::Get("relay.ext.uma.codegen_c_includes_" + target_str_); - if (includes_pf) { - String includes = (*includes_pf)(); - decl_stream << includes; - } - std::unordered_set devices; - devices.insert(target_str_); - CodeGenCHost::Init(output_ssa, emit_asserts, emit_fwd_func_decl, target_str_, devices); - } - - private: - String target_str_; -}; - -runtime::Module TIRToRuntime(IRModule mod, Target target) { - bool output_ssa = false; - bool emit_asserts = false; - bool emit_fwd_func_decl = true; - UMACodegen codegen(target->kind->name); - codegen.Init(output_ssa, emit_asserts, emit_fwd_func_decl); - - Map functions; - for (auto [gvar, base_func] : mod->functions) { - auto prim_func = Downcast(base_func); - functions.Set(gvar, prim_func); - } - - for (auto [gvar, prim_func] : functions) { - codegen.DeclareFunction(gvar, prim_func); - } - for (auto [gvar, prim_func] : functions) { - codegen.AddFunction(gvar, prim_func, emit_fwd_func_decl); - } - - std::string code = codegen.Finish(); - - Array function_names; - for (auto [gvar, prim_func] : functions) { - function_names.push_back(codegen.GetFunctionName(gvar)); - } - - return codegen::CSourceModuleCreate(code, "c", function_names); -} - -} // namespace uma -}; // namespace contrib -} // namespace relay -} // namespace tvm diff --git a/src/relay/backend/contrib/verilator/codegen.cc b/src/relay/backend/contrib/verilator/codegen.cc deleted file mode 100644 index 2e6fb1326314..000000000000 --- a/src/relay/backend/contrib/verilator/codegen.cc +++ /dev/null @@ -1,146 +0,0 @@ -/* - * 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 src/relay/backend/contrib/verilator/codegen.cc - * \brief Implementation of Verilator codegen APIs. - */ - -#include -#include -#include -#include -#include -#include - -#include -#include -#include - -#include "../../../../runtime/contrib/json/json_node.h" -#include "../../../../runtime/contrib/verilator/verilator_runtime.h" -#include "../../utils.h" -#include "../codegen_json/codegen_json.h" - -namespace tvm { -namespace relay { -namespace contrib { - -using namespace backend; - -/*! \brief Verilator JSON serializer */ -class VerilatorJSONSerializer : public backend::contrib::JSONSerializer { - using JSONGraphNode = tvm::runtime::json::JSONGraphNode; - using JSONGraphNodeEntry = tvm::runtime::json::JSONGraphNodeEntry; - - public: - VerilatorJSONSerializer(const std::string& symbol, const Expr& expr) - : JSONSerializer(symbol, expr) {} - - std::vector VisitExpr_(const CallNode* cn) override { - Expr expr = GetRef(cn); - std::string name; - const CallNode* call = cn; - if (const auto* op_node = cn->op.as()) { - name = op_node->name; - } else { - LOG(FATAL) << "Verilator JSON runtime does not support calls to " << cn->op->GetTypeKey(); - } - - std::vector inputs; - for (const auto& arg : cn->args) { - auto res = VisitExpr(arg); - inputs.insert(inputs.end(), res.begin(), res.end()); - } - auto node = std::make_shared(name, /* name_ */ - "kernel", /* op_type_ */ - inputs, 1 /* num_outputs_ */); - SetCallNodeAttribute(node, call); - return AddNode(node, GetRef(cn)); - } -}; - -/*! \brief Attributes to store options for Verilator */ -struct VerilatorOptionsNode : public tvm::AttrsNode { - String lib_path; - int reset_cycles; - bool profiler_enable; - int profiler_cycle_counter_id; - - TVM_DECLARE_ATTRS(VerilatorOptionsNode, "ext.attrs.VerilatorOptionsNode") { - TVM_ATTR_FIELD(lib_path).describe("the design library path").set_default("libverilator.so"); - TVM_ATTR_FIELD(reset_cycles).describe("the number of reset cycles").set_default(1); - TVM_ATTR_FIELD(profiler_enable).describe("enable profiler").set_default(false); - TVM_ATTR_FIELD(profiler_cycle_counter_id).describe("profiler cycle counter id").set_default(0); - } -}; - -class VerilatorOptions : public Attrs { - public: - TVM_DEFINE_NOTNULLABLE_OBJECT_REF_METHODS(VerilatorOptions, Attrs, VerilatorOptionsNode); -}; - -TVM_REGISTER_NODE_TYPE(VerilatorOptionsNode); -TVM_REGISTER_PASS_CONFIG_OPTION("relay.ext.verilator.options", VerilatorOptions); - -/*! - * \brief The Verilator codegen tool. It takes a Relay expression/module and - * compile it into a Verilator runtime module. - */ -runtime::Module VerilatorBackend(const ObjectRef& ref) { - VLOG(0) << "compiling for verilator runtime"; - CHECK(ref->IsInstance()); - auto func = Downcast(ref); - auto func_name = GetExtSymbol(func); - VerilatorJSONSerializer serializer(func_name, func); - serializer.serialize(); - std::string graph_json = serializer.GetJSON(); - - // Note that serializer.const_name_to_constant() is ignored. Instead the TECompiler invokes - // a callback which calls backend::UpdateConstants to capture the map before the function - // 'disappears' into lowered form, on the assumption the visit order and thus constant - // names match those generated by the JSONSerializer. - - // Create runtime object - auto n = make_object(func_name, graph_json, - serializer.const_names()); - - // Get Verilator compiler options - auto ctx = transform::PassContext::Current(); - auto cfg = ctx->GetConfig("relay.ext.verilator.options"); - if (!cfg.defined()) { - cfg = AttrsWithDefaultValues(); - } - - n->SetLibrary(cfg.value()->lib_path); - n->SetResetCycles(cfg.value()->reset_cycles); - - if (cfg.value()->profiler_enable) { - n->EnableProfiler(); - n->SetProfilerCycleCounterId(cfg.value()->profiler_cycle_counter_id); - } - - return runtime::Module(n); -} - -TVM_REGISTER_GLOBAL("relay.ext.verilator").set_body_typed(VerilatorBackend); - -} // namespace contrib -} // namespace relay -} // namespace tvm diff --git a/src/relay/backend/contrib/vitis_ai/config_vitis_ai.cc b/src/relay/backend/contrib/vitis_ai/config_vitis_ai.cc deleted file mode 100644 index a5d3879553ba..000000000000 --- a/src/relay/backend/contrib/vitis_ai/config_vitis_ai.cc +++ /dev/null @@ -1,82 +0,0 @@ -/* - * 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 src/relay/backend/contrib/vitis_ai/config_vitis_ai.cc - * \brief Register Vitis-AI codegen options. Main codegen is implemented in python. - */ - -#include - -namespace tvm { -namespace relay { -namespace contrib { -namespace vitis_ai { - -/*! \brief Attributes to store the compiler options for Vitis AI */ -struct VitisAICompilerConfigNode : public tvm::AttrsNode { - String dpu; - String build_dir; - String work_dir; - String export_runtime_module; - String load_runtime_module; - TVM_DECLARE_ATTRS(VitisAICompilerConfigNode, "ext.attrs.VitisAICompilerConfigNode") { - TVM_ATTR_FIELD(dpu).describe("Vitis AI DPU identifier").set_default(""); - TVM_ATTR_FIELD(build_dir) - .describe("Build directory to be used (optional, debug)") - .set_default(""); - TVM_ATTR_FIELD(work_dir) - .describe("Work directory to be used (optional, debug)") - .set_default(""); - TVM_ATTR_FIELD(export_runtime_module) - .describe("Export the Vitis AI runtime module to this file") - .set_default(""); - TVM_ATTR_FIELD(load_runtime_module) - .describe("Load the Vitis AI runtime module to this file") - .set_default(""); - } -}; - -class VitisAICompilerConfig : public Attrs { - public: - TVM_DEFINE_NOTNULLABLE_OBJECT_REF_METHODS(VitisAICompilerConfig, Attrs, - VitisAICompilerConfigNode); -}; - -TVM_REGISTER_NODE_TYPE(VitisAICompilerConfigNode); -TVM_REGISTER_PASS_CONFIG_OPTION("relay.ext.vitis_ai.options", VitisAICompilerConfig); -TVM_REGISTER_GLOBAL("relay.ext.vitis_ai.available") - .set_body([](tvm::TVMArgs args, tvm::TVMRetValue* rv) { *rv = true; }); - -// Following config options are here for backward compatibility (deprecated API's) -/*! \brief The target Vitis-AI accelerator device */ -TVM_REGISTER_PASS_CONFIG_OPTION("relay.ext.vitis_ai.options.target", String); -/*! \brief (Optional config) The build directory to be used by Vitis-AI */ -TVM_REGISTER_PASS_CONFIG_OPTION("relay.ext.vitis_ai.options.build_dir", String); -/*! \brief (Optional config) The work directory to be used by Vitis-AI */ -TVM_REGISTER_PASS_CONFIG_OPTION("relay.ext.vitis_ai.options.work_dir", String); -/*! \brief (Optional config) Export PyXIR runtime module to disk during serialization if provided */ -TVM_REGISTER_PASS_CONFIG_OPTION("relay.ext.vitis_ai.options.export_runtime_module", String); -/*! \brief (Optional config) Load PyXIR runtime module from disk */ -TVM_REGISTER_PASS_CONFIG_OPTION("relay.ext.vitis_ai.options.load_runtime_module", String); - -} // namespace vitis_ai -} // namespace contrib -} // namespace relay -} // namespace tvm diff --git a/src/relay/backend/executor.cc b/src/relay/backend/executor.cc deleted file mode 100644 index 66feac4699e6..000000000000 --- a/src/relay/backend/executor.cc +++ /dev/null @@ -1,112 +0,0 @@ -/* - * 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 src/relay/backend/executor.cc - * \brief Executor Registry - */ - -#include - -#include "../../node/attr_registry.h" -namespace tvm { -namespace relay { - -TVM_REGISTER_NODE_TYPE(ExecutorNode); - -TVM_STATIC_IR_FUNCTOR(ReprPrinter, vtable) - .set_dispatch([](const ObjectRef& obj, ReprPrinter* p) { - const Executor& executor = Downcast(obj); - p->stream << executor->name; - p->stream << executor->attrs; - }); - -/********** Registry-related code **********/ - -using ExecutorRegistry = AttrRegistry; - -Executor Executor::Create(String name, Map attrs) { - const ExecutorRegEntry* reg = ExecutorRegistry::Global()->Get(name); - if (reg == nullptr) { - throw Error("Executor \"" + name + "\" is not defined"); - } - - for (const auto& kv : attrs) { - if (!reg->key2vtype_.count(kv.first)) { - throw Error("Attribute \"" + kv.first + "\" is not available on this Executor"); - } - std::string expected_type = reg->key2vtype_.at(kv.first).type_key; - std::string actual_type = kv.second->GetTypeKey(); - if (expected_type != actual_type) { - throw Error("Attribute \"" + kv.first + "\" should have type \"" + expected_type + - "\" but instead found \"" + actual_type + "\""); - } - } - - for (const auto& kv : reg->key2default_) { - if (!attrs.count(kv.first)) { - attrs.Set(kv.first, kv.second); - } - } - - return Executor(name, DictAttrs(attrs)); -} - -Array Executor::ListExecutors() { return ExecutorRegistry::Global()->ListAllNames(); } - -Map Executor::ListExecutorOptions(const String& name) { - Map options; - const ExecutorRegEntry* reg = ExecutorRegistry::Global()->Get(name); - if (reg == nullptr) { - throw Error("Executor \"" + name + "\" is not defined"); - } - for (const auto& kv : reg->key2vtype_) { - options.Set(kv.first, kv.second.type_key); - } - return options; -} - -ExecutorRegEntry& ExecutorRegEntry::RegisterOrGet(const String& name) { - return ExecutorRegistry::Global()->RegisterOrGet(name); -} - -/********** Register Executors and options **********/ - -TVM_REGISTER_EXECUTOR("aot") - .add_attr_option("link-params", runtime::Bool(true)) - .add_attr_option("unpacked-api") - .add_attr_option("interface-api") - .add_attr_option("workspace-byte-alignment") - .add_attr_option("constant-byte-alignment"); - -TVM_REGISTER_EXECUTOR("graph").add_attr_option("link-params", runtime::Bool(false)); - -/********** Registry **********/ - -TVM_REGISTER_GLOBAL("relay.backend.CreateExecutor").set_body_typed(Executor::Create); -TVM_REGISTER_GLOBAL("relay.backend.GetExecutorAttrs").set_body_typed([](const Executor& executor) { - return executor->attrs->dict; -}); - -TVM_REGISTER_GLOBAL("relay.backend.ListExecutors").set_body_typed(Executor::ListExecutors); -TVM_REGISTER_GLOBAL("relay.backend.ListExecutorOptions") - .set_body_typed(Executor::ListExecutorOptions); - -} // namespace relay -} // namespace tvm diff --git a/src/relay/backend/graph_executor_codegen.cc b/src/relay/backend/graph_executor_codegen.cc deleted file mode 100644 index 734b3d6e4360..000000000000 --- a/src/relay/backend/graph_executor_codegen.cc +++ /dev/null @@ -1,787 +0,0 @@ -/* - * 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 relay/backend/graph_codegen.cc - * \brief Graph executor codegen - */ - -#include -#include -#include -#include -#include -#include -#include -#include -#include -#include - -#include -#include -#include - -#include "../op/annotation/annotation.h" -#include "../op/call/call.h" -#include "../op/memory/device_copy.h" -#include "../transforms/device_aware_visitors.h" -#include "./te_compiler.h" -#include "./utils.h" - -namespace tvm { -namespace relay { - -// TODO(@jroesch, @csullivan): declare directly elsewhere -backend::StaticMemoryPlan GraphPlanMemory(const Function& func); - -namespace backend { - -class GraphNode; -class GraphInputNode; -class GraphOpNode; - -using IntegerArray = Array; -using ShapeVector = std::vector>; -using GraphAttrs = std::unordered_map; -using GraphObjectPtr = std::shared_ptr; -using GraphInputObjectPtr = std::shared_ptr; -using GraphOpObjectPtr = std::shared_ptr; - -/*! \brief Node types */ -enum GraphNodeType { - kGraphNop, - kGraphInputNode, - kGraphOpNode, -}; - -class GraphNodeRef { - public: - GraphNodeRef() {} - GraphNodeRef(int ident, int index, int version = 0) - : ident_(ident), index_(index), version_(version) {} - - inline void Save(dmlc::JSONWriter* writer) const { - writer->BeginArray(); - writer->WriteArrayItem(ident_); - writer->WriteArrayItem(index_); - writer->WriteArrayItem(version_); - writer->EndArray(); - } - - inline void Load(dmlc::JSONReader* reader) { LOG(FATAL) << "Not implemented."; } - - protected: - int ident_; - int index_{0}; - int version_{0}; -}; - -/*! \brief Base Node class */ -class GraphNode { - public: - GraphNode() {} - virtual void Save(dmlc::JSONWriter* writer) const {} - virtual void Load(dmlc::JSONReader* reader) {} - virtual GraphNodeType Type() const { return kGraphNop; } - virtual ~GraphNode() {} - - public: - int num_outputs_{1}; - std::string name_; - GraphAttrs attrs_; -}; - -/*! \brief Input Node */ -class GraphInputNode : public GraphNode { - public: - GraphInputNode() {} - GraphInputNode(const std::string& name, const GraphAttrs& attrs) { - name_ = name; - attrs_ = attrs; - } - - GraphNodeType Type() const override { return kGraphInputNode; } - - void Save(dmlc::JSONWriter* writer) const override { - const std::string op_name{"null"}; - writer->BeginObject(); - writer->WriteObjectKeyValue("op", op_name); - writer->WriteObjectKeyValue("name", this->name_); - writer->WriteObjectKeyValue("inputs", std::list()); - writer->EndObject(); - } - static std::shared_ptr make_node_ptr(const std::string& name, - const GraphAttrs& attrs) { - auto ptr = std::make_shared(name, attrs); - return std::dynamic_pointer_cast(ptr); - } -}; - -/*! \brief Op Node */ -class GraphOpNode : public GraphNode { - public: - GraphOpNode() {} - GraphOpNode(const std::string& name, const GraphAttrs& nd_attrs, const std::string& op_name, - const std::vector& inputs, const GraphAttrs& attrs, - size_t num_outputs = 1) { - name_ = name; - attrs_ = nd_attrs; - op_name_ = op_name; - inputs_ = inputs; - op_attrs_ = attrs; - num_outputs_ = num_outputs; - op_attrs_["func_name"] = op_name_; - op_attrs_["flatten_data"] = std::string("0"); - op_attrs_["num_inputs"] = std::to_string(inputs_.size()); - op_attrs_["num_outputs"] = std::to_string(num_outputs_); - } - - GraphNodeType Type() const override { return kGraphOpNode; } - - void Save(dmlc::JSONWriter* writer) const override { - GraphAttrs attrs = op_attrs_; - attrs["func_name"] = this->op_name_; - attrs["flatten_data"] = std::string("0"); - attrs["num_inputs"] = std::to_string(this->inputs_.size()); - attrs["num_outputs"] = std::to_string(this->num_outputs_); - writer->BeginObject(); - writer->WriteObjectKeyValue("op", op_type_name_); - writer->WriteObjectKeyValue("name", name_); - writer->WriteObjectKeyValue("attrs", attrs); - writer->WriteObjectKeyValue("inputs", this->inputs_); - writer->EndObject(); - } - static std::shared_ptr make_node_ptr(const std::string& name, - const GraphAttrs& nd_attrs, - const std::string& op_name, - const std::vector& inputs, - const GraphAttrs& attrs, size_t num_outputs = 1) { - auto ptr = std::make_shared(name, nd_attrs, op_name, inputs, attrs, num_outputs); - return std::dynamic_pointer_cast(ptr); - } - - public: - std::string op_name_; - std::vector inputs_; - GraphAttrs op_attrs_; - - private: - const std::string op_type_name_{"tvm_op"}; -}; - -/*! \brief Code generator for the graph executor, produces a module containing the graph JSON, - * module, and parameters. - */ -class GraphExecutorCodegen : public backend::MemoizedExprTranslator> { - public: - GraphExecutorCodegen(runtime::Module* mod, const Array& targets) - : mod_(mod), config_(transform::PassContext::Current(), targets) {} - - StorageInfo GetStorageInfo(const Expr& e) { - size_t count = memory_plan_->expr_to_storage_info.count(e); - ICHECK_GT(count, 0) << "Expr is not existing in storage plan"; - auto storage_info = memory_plan_->expr_to_storage_info[e]; - return storage_info; - } - - LoweredOutput Codegen(IRModule mod, relay::Function func, String mod_name) { - mod_name_ = mod_name; - VLOG_CONTEXT << "GraphExecutorCodegen"; - VLOG(1) << "compiling:" << std::endl << PrettyPrint(func); - - // TODO(mbs): Why plan memory and update workspace sizes before lowering? - memory_plan_ = GraphPlanMemory(func); - - backend::FunctionInfo func_info; - - if (memory_plan_.defined()) { - // TODO(@electriclilies, @jroesch): remove UpdateMainWorkspaceSize - func_info = - relay::tec::UpdateMainWorkspaceSize(mod, config_, memory_plan_->expr_to_storage_info); - mod = WithAttr(mod, "main_func_info", func_info); - } - - IRModule lowered_mod = tec::LowerTE(mod_name_, config_, [this](BaseFunc func) { - // We need to maintain the constant map for external - // functions so we pass this processing function which - // allows us to process each function as we lower it. - if (func->GetAttr(attr::kCompiler).defined()) { - UpdateConstants(func, ¶ms_); - } - - // TODO(@areusch, @jroesch): We should refactor this to - // execute as a further pass, instead writing data to the - // lowering process directly. - tec::UpdateFunctionMetadata(func, this->function_metadata_); - })(mod); - - Optional main_func_info = - lowered_mod->GetAttr("main_func_info"); - - function_metadata_.Set(runtime::symbol::tvm_module_main, main_func_info.value()); - - Function lowered_main_func = Downcast(lowered_mod->Lookup("main")); - - // Now that we have lowered all operators to TIR code, we can proceed with compilation. - // - // We need to unfortunately re-plan as the previous results have been invalidated by lowering - // we will fix this in future refactors. - memory_plan_ = GraphPlanMemory(lowered_main_func); - - // The graph planner also can not handle planning calls to global variables to we must remap - - // First we convert all the parameters into input nodes. - for (auto param : lowered_main_func->params) { - auto node_ptr = GraphInputNode::make_node_ptr(param->name_hint(), GraphAttrs()); - var_map_[param.get()] = AddNode(node_ptr, param); - } - - heads_ = VisitExpr(lowered_main_func->body); - std::ostringstream os; - - dmlc::JSONWriter writer(&os); - GetJSON(&writer); - LoweredOutput ret; - ret.graph_json = os.str(); - - // Collect any runtime modules generated by external codegen. - ret.external_mods = - lowered_mod->GetAttr>(tvm::attr::kExternalMods).value_or({}); - - // Collect any constants extracted by external codegen. - ret.params = std::unordered_map(); - Map const_name_to_constant = - lowered_mod->GetAttr>(tvm::attr::kConstNameToConstant) - .value_or({}); - for (const auto& kv : const_name_to_constant) { - VLOG(1) << "constant '" << kv.first << "' contributed by external codegen"; - ICHECK(ret.params.emplace(kv.first, kv.second).second); - } - - // Collect any constants extracted during lowering. - for (const auto& kv : params_) { - VLOG(1) << "constant '" << kv.first << "' contributed by TECompiler"; - ICHECK(ret.params.emplace(kv.first, kv.second).second); - } - - ret.function_metadata = std::move(function_metadata_); - - // This is the point where we separate the functions in the module by target - ret.lowered_funcs = tec::GetPerTargetModules(lowered_mod); - ret.metadata = - ExecutorCodegenMetadata({} /* inputs */, {} /* input_tensor_types */, {} /* outputs */, - {} /* output_tensor_types */, {} /* pools */, {} /* devices */, - runtime::kTvmExecutorGraph /* executor */, mod_name_ /* mod_name */, - "packed" /* interface_api */, Bool(false) /* unpacked_api */); - return ret; - } - - protected: - /*! - * \brief Add node to graph - * - * \param node - * \param expr - * \return std::vector<_NodeRef> - */ - std::vector AddNode(GraphObjectPtr node, Expr expr) { - auto checked_type = expr->checked_type(); - - auto storage_info = GetStorageInfo(expr); - // storage - std::vector storage_ids; - for (auto v : storage_info->storage_ids) { - storage_ids.push_back(v); - } - node->attrs_["storage_id"] = std::move(storage_ids); - // type - std::vector device_types; - for (const auto& virtual_device : storage_info->virtual_devices) { - // TODO(mbs): Keeping only the device type. - ICHECK_GT(virtual_device->device_type(), 0); - device_types.push_back(virtual_device->device_type()); - } - size_t num_unknown_devices = std::count(device_types.begin(), device_types.end(), 0); - if (num_unknown_devices != 0 && num_unknown_devices != device_types.size()) { - LOG(FATAL) << "The graph contains not annotated nodes for " - << "heterogeneous execution. All nodes must be " - << "annotated."; - } - if (num_unknown_devices == 0) { - node->attrs_["device_index"] = device_types; - } - // storage scope - std::vector storage_scope; - for (const auto& virtual_device : storage_info->virtual_devices) { - storage_scope.push_back(std::string(virtual_device->memory_scope)); - } - node->attrs_["storage_scope"] = std::move(storage_scope); - auto node_id = nodes_.size(); - nodes_.push_back(node); - // Tuple return value, flatten as tuple - if (const auto* tuple_type = checked_type.as()) { - std::vector ret; - ShapeVector shape; - std::vector dtype; - for (size_t i = 0; i < tuple_type->fields.size(); ++i) { - if (const auto* typ = tuple_type->fields[i].as()) { - ret.push_back(GraphNodeRef(node_id, i)); - shape.emplace_back(ShapeToJSON(typ->shape)); - dtype.emplace_back(DType2String(typ->dtype)); - } else { - LOG(FATAL) << "type " << checked_type->GetTypeKey() << " not supported"; - } - } - ICHECK_EQ(node->Type(), kGraphOpNode); - auto op_nd = std::dynamic_pointer_cast(node); - op_nd->attrs_["shape"] = shape; - op_nd->attrs_["dtype"] = dtype; - op_nd->num_outputs_ = tuple_type->fields.size(); - return ret; - } - // Normal tensor return type - if (const auto* tensor_type = checked_type.as()) { - ShapeVector shape; - std::vector dtype; - shape.emplace_back(ShapeToJSON(tensor_type->shape)); - dtype.emplace_back(DType2String(tensor_type->dtype)); - node->attrs_["shape"] = shape; - node->attrs_["dtype"] = dtype; - } else { - LOG(FATAL) << "type " << checked_type->GetTypeKey() << " not supported"; - } - return {GraphNodeRef(node_id, 0)}; - } - - std::vector VisitExpr_(const VarNode* op) override { - Expr expr = GetRef(op); - return var_map_[expr.get()]; - } - - std::vector VisitExpr_(const ConstantNode* op) override { - Expr expr = GetRef(op); - size_t index = params_.size(); - std::string name = "p" + std::to_string(index); - auto node = GraphInputNode::make_node_ptr(name, GraphAttrs()); - auto to_return = AddNode(node, expr); - CHECK_EQ(to_return.size(), 1) << "Expected exactly 1 parameter node created"; - param_storage_ids_[name] = GetStorageInfo(expr)->storage_ids[0]; - params_[name] = op->data; - return to_return; - } - - std::vector VisitExpr_(const TupleNode* op) override { - std::vector fields; - for (auto field : op->fields) { - auto ref_vec = VisitExpr(field); - for (auto ref : ref_vec) { - fields.push_back(ref); - } - } - return fields; - } - - bool ShareSameStorage(const Expr& lhs, const Expr& rhs) { - StorageInfo lit = GetStorageInfo(lhs); - StorageInfo rit = GetStorageInfo(rhs); - int64_t lhs_storage_id = lit->storage_ids[0]; - int64_t rhs_storage_id = rit->storage_ids[0]; - return lhs_storage_id == rhs_storage_id; - } - - std::vector GraphAddCallNode(const CallNode* call_node, GraphAttrs attrs) { - Call call = GetRef(call_node); - std::vector inputs; - std::string func_name; - - DeviceCopyProps device_copy_props = GetDeviceCopyProps(call_node); - CallLoweredProps call_lowered_props = GetCallLoweredProps(call_node); - if (device_copy_props.body.defined()) { - // The graph executor expects to see a normal call to the undefined @__copy function. - // The source and destination device annotations are no longer needed since they have - // been captured in the StorageInfos for both input and output. - // TODO(mbs): device_copy cleanup - func_name = "__copy"; - for (const auto& n : VisitExpr(device_copy_props.body)) { - inputs.push_back(n); - } - } else if (call_lowered_props.lowered_func.defined()) { - // Extract function and arguments from the call_lowered op - - func_name = call_lowered_props.lowered_func->name_hint; - - for (const Expr& arg : call_lowered_props.arguments) { - for (auto n : VisitExpr(arg)) { - inputs.push_back(n); - } - } - if (call_lowered_props.attrs.metadata.count("relay_attrs")) { - if (auto relay_attrs = - call_lowered_props.attrs.metadata["relay_attrs"].as()) { - for (auto p : relay_attrs->dict) { - if (p.second.as()) { - attrs[p.first] = std::string(Downcast(p.second)); - } - } - } - } - // TODO(mbs): "reshape" cleanup. - if (IsReshapeOnly(call_lowered_props) && - ShareSameStorage(GetRef(call_node), call_lowered_props.arguments[0])) { - auto node = GraphOpNode::make_node_ptr("reshape_nop", GraphAttrs(), "__nop", inputs, attrs); - return AddNode(node, call); - } - } else if (!call_node->attrs.defined()) { // Call is an extern function - const auto* func = call_node->op.as(); - ICHECK(func) << "Expected the operator to be a global var, but got " - << call_node->op->GetTypeKey(); // getting a relay fn here, not sure why. - func_name = func->name_hint; - - for (const Expr& arg : call_node->args) { - for (auto n : VisitExpr(arg)) { - inputs.push_back(n); - } - } - } else { - LOG(FATAL) << "Non-primitive-call nodes should have been transformed away.\n" - << "The graph executor code generator expects all calls to be call_lowered, " - << "but found: " << std::endl - << PrettyPrint(call); - } - - // Compute the operator name, because we used the get unique name when generating the kernel. - auto op_name = name_supply_->FreshName(func_name); - auto node = GraphOpNode::make_node_ptr(op_name, GraphAttrs(), func_name, inputs, attrs); - return AddNode(node, call); - } - - std::vector VisitExpr_(const CallNode* call_node) override { - relay::Call call = GetRef(call_node); - OnDeviceProps props = GetOnDeviceProps(call_node); - if (props.body.defined()) { - // See through "on_device" calls. - return VisitExpr(props.body); - } - return GraphAddCallNode(call_node, GraphAttrs()); - } - - std::vector VisitExpr_(const LetNode* op) override { - ICHECK_EQ(var_map_.count(op->var.get()), 0); - var_map_[op->var.get()] = VisitExpr(op->value); - return VisitExpr(op->body); - } - std::vector VisitExpr_(const TupleGetItemNode* op) override { - auto vtuple = VisitExpr(op->tuple); - return {vtuple[op->index]}; - } - - std::vector VisitExpr_(const OpNode* op) override { - LOG(FATAL) << "All OpNodes should have been expanded"; - } - std::vector VisitExpr_(const GlobalVarNode* op) override { - LOG(FATAL) << "All GlobalVarNodes should be removed before graph executor's Codegen is called"; - } - std::vector VisitExpr_(const IfNode* op) override { - LOG(FATAL) << "Graph executor does not support control flow (found IfNode)"; - } - std::vector VisitExpr_(const FunctionNode* op) override { - ICHECK(op->GetAttr(attr::kCompiler).defined()) - << "Only functions supported by custom codegen"; - return {}; - } - std::vector VisitExpr_(const RefCreateNode* op) override { - LOG(FATAL) << "Graph executor does not support references (found RefCreateNode)"; - } - std::vector VisitExpr_(const RefReadNode* op) override { - LOG(FATAL) << "Graph executor does not support references (found RefReadNode)"; - } - std::vector VisitExpr_(const RefWriteNode* op) override { - LOG(FATAL) << "Graph executor does not support references (found RefWriteNode)"; - } - std::vector VisitExpr_(const ConstructorNode* op) override { - LOG(FATAL) << "Graph executor does not support ADTs (found ConstructorNode)"; - } - std::vector VisitExpr_(const MatchNode* op) override { - LOG(FATAL) << "Graph executor does not support matching (found MatchNode)"; - } - /*! - * \brief Generate Graph JSON - * - * \param writer json writer - */ - void GetJSON(dmlc::JSONWriter* writer) { - std::vector arg_nodes; - for (size_t i = 0; i < nodes_.size(); ++i) { - auto node = nodes_[i]; - if (node->Type() == kGraphInputNode) { - arg_nodes.push_back(i); - } - } - size_t num_entry = 0; - ShapeVector shapes; - std::vector storage_ids; - std::vector storage_scopes; - std::vector device_types; - std::vector dltypes; - std::vector node_row_ptr{0}; - for (auto node : nodes_) { - const auto& shape_vec = dmlc::get(node->attrs_["shape"]); - const auto& storage_id = dmlc::get>(node->attrs_["storage_id"]); - const auto& storage_scope = - dmlc::get>(node->attrs_["storage_scope"]); - const auto& dtype_vec = dmlc::get>(node->attrs_["dtype"]); - - ICHECK_EQ(node->num_outputs_, shape_vec.size()); - num_entry += node->num_outputs_; - - shapes.insert(shapes.end(), shape_vec.begin(), shape_vec.end()); - dltypes.insert(dltypes.end(), dtype_vec.begin(), dtype_vec.end()); - storage_ids.insert(storage_ids.end(), storage_id.begin(), storage_id.end()); - storage_scopes.insert(storage_scopes.end(), storage_scope.begin(), storage_scope.end()); - if (node->attrs_.count("device_index")) { - const auto& dev_types = dmlc::get>(node->attrs_["device_index"]); - device_types.insert(device_types.end(), dev_types.begin(), dev_types.end()); - } - node_row_ptr.push_back(num_entry); - } - - // verification if storage_scope contains any non global memory scope - // in other case it's better not to write scopes to the JSON at all - bool global_only_scope = true; - for (const auto& ss : storage_scopes) { - if (!(ss.empty() || ss == "global")) { - global_only_scope = false; - } - } - if (global_only_scope) { - storage_scopes.clear(); - } - writer->BeginObject(); - writer->WriteObjectKeyValue("nodes", nodes_); - writer->WriteObjectKeyValue("arg_nodes", arg_nodes); - writer->WriteObjectKeyValue("heads", heads_); - std::unordered_map> attrs; - attrs["shape"].emplace_back(std::string("list_shape")); - attrs["shape"].emplace_back(shapes); - attrs["storage_id"].emplace_back(std::string("list_int")); - attrs["storage_id"].emplace_back(storage_ids); - if (device_types.size()) { - attrs["device_index"].emplace_back(std::string("list_int")); - attrs["device_index"].emplace_back(device_types); - } - if (storage_scopes.size()) { - attrs["storage_scope"].emplace_back(std::string("list_str")); - attrs["storage_scope"].emplace_back(storage_scopes); - } - attrs["dltype"].emplace_back(std::string("list_str")); - attrs["dltype"].emplace_back(dltypes); - writer->WriteObjectKeyValue("attrs", attrs); - writer->WriteObjectKeyValue("node_row_ptr", node_row_ptr); - writer->EndObject(); - } - - protected: - /*! \brief nodes */ - std::vector nodes_; - /*! \brief output of graph */ - std::vector heads_; - /*! \brief mod */ - runtime::Module* mod_; - /*! \brief variable map */ - std::unordered_map> var_map_; - /*! \brief Available targets */ - CompilationConfig config_; - /*! - * \brief parameters (i.e. ConstantNodes found in the graph). - * These are take as inputs to the GraphExecutor. - * Maps param name to a pair of storage_id and NDArray. At runtime, the storage_id can be - * used to lookup the parameter. - */ - std::unordered_map params_; - std::unordered_map param_storage_ids_; - /*! \brief plan memory of device result */ - StaticMemoryPlan memory_plan_; - /*! \brief the module name we use to mangle the function names */ - String mod_name_; - /*! \brief function metadata */ - Map function_metadata_; - /*! \brief NameSupply */ - NameSupply name_supply_; -}; - -class GraphExecutorCodegenModule : public runtime::ModuleNode { - public: - GraphExecutorCodegenModule() {} - virtual PackedFunc GetFunction(const String& name, const ObjectPtr& sptr_to_self) { - if (name == "init") { - return PackedFunc([sptr_to_self, this](TVMArgs args, TVMRetValue* rv) { - ICHECK_EQ(args.num_args, 2) << "The expected of arguments are: " - << "runtime::Module mod and Array targets"; - void* mod = args[0]; - Array targets = args[1]; - codegen_ = std::make_shared(reinterpret_cast(mod), - std::move(targets)); - }); - } else if (name == "codegen") { - return PackedFunc([sptr_to_self, this](TVMArgs args, TVMRetValue* rv) { - IRModule mod = args[0]; - Function func = args[1]; - String mod_name = args[2]; - this->output_ = this->codegen_->Codegen(mod, func, mod_name); - }); - } else if (name == "get_graph_json") { - return PackedFunc( - [sptr_to_self, this](TVMArgs args, TVMRetValue* rv) { *rv = this->output_.graph_json; }); - } else if (name == "list_params_name") { - return PackedFunc([sptr_to_self, this](TVMArgs args, TVMRetValue* rv) { - Array ret; - for (const auto& kv : this->output_.params) { - ret.push_back(kv.first); - } - *rv = ret; - }); - } else if (name == "get_param_by_name") { - return PackedFunc([sptr_to_self, this](TVMArgs args, TVMRetValue* rv) { - String key = args[0]; - auto it = this->output_.params.find(key); - CHECK(it != this->output_.params.end()) << "no such parameter " << key; - *rv = (*it).second; - }); - } else if (name == "get_irmodule") { - return PackedFunc([sptr_to_self, this](TVMArgs args, TVMRetValue* rv) { - *rv = this->output_.lowered_funcs; - }); - } else if (name == "get_external_modules") { - return PackedFunc([sptr_to_self, this](TVMArgs args, TVMRetValue* rv) { - *rv = this->output_.external_mods; - }); - } else if (name == "get_devices") { - return PackedFunc([sptr_to_self](TVMArgs args, TVMRetValue* rv) { *rv = Array(); }); - } else if (name == "get_executor_codegen_metadata") { - return PackedFunc( - [sptr_to_self, this](TVMArgs args, TVMRetValue* rv) { *rv = this->output_.metadata; }); - } else if (name == "get_function_metadata") { - return PackedFunc([sptr_to_self, this](TVMArgs args, TVMRetValue* rv) { - *rv = this->output_.function_metadata; - }); - } else { - return PackedFunc([](TVMArgs args, TVMRetValue* rv) {}); - } - } - - const char* type_key() const final { return "RelayGraphExecutorCodegenModule"; } - - /*! \brief Get the property of the runtime module .*/ - int GetPropertyMask() const final { return runtime::ModulePropertyMask::kRunnable; } - - private: - std::shared_ptr codegen_; - LoweredOutput output_; -}; - -runtime::Module CreateGraphCodegenMod() { - auto ptr = make_object(); - return runtime::Module(ptr); -} - -TVM_REGISTER_GLOBAL("relay.build_module._GraphExecutorCodegen") - .set_body([](TVMArgs args, TVMRetValue* rv) { *rv = CreateGraphCodegenMod(); }); - -} // namespace backend -} // namespace relay -} // namespace tvm - -namespace dmlc { -namespace json { -// JSON utils -template -inline bool SameType(const dmlc::any& data) { - return std::type_index(data.type()) == std::type_index(typeid(T)); -} - -template <> -struct Handler> { - inline static void Write(dmlc::JSONWriter* writer, - const std::shared_ptr& data) { - data->Save(writer); - } - inline static void Read(dmlc::JSONReader* reader, - std::shared_ptr* data) { - LOG(FATAL) << "Not implemented."; - } -}; -template <> -struct Handler> { - inline static void Write(dmlc::JSONWriter* writer, - const std::unordered_map& data) { - writer->BeginObject(); - for (const auto& kv : data) { - auto k = kv.first; - const dmlc::any& v = kv.second; - if (SameType(v)) { - writer->WriteObjectKeyValue(k, dmlc::get(v)); - } else if (SameType(v)) { - writer->WriteObjectKeyValue(k, dmlc::get(v)); - } else if (SameType>(v)) { - writer->WriteObjectKeyValue(k, dmlc::get>(v)); - } else if (SameType>>(v)) { - writer->WriteObjectKeyValue(k, dmlc::get>>(v)); - } else if (SameType>(v)) { - writer->WriteObjectKeyValue(k, dmlc::get>(v)); - } else if (SameType>(v)) { - writer->WriteObjectKeyValue(k, dmlc::get>(v)); - } else { - LOG(FATAL) << "Not supported"; - } - } - writer->EndObject(); - } - inline static void Read(dmlc::JSONReader* reader, - std::unordered_map* data) { - LOG(FATAL) << "Not implemented."; - } -}; - -template <> -struct Handler> { - inline static void Write(dmlc::JSONWriter* writer, const std::vector& data) { - writer->BeginArray(); - for (const auto& v : data) { - if (SameType(v)) { - writer->WriteArrayItem(dmlc::get(v)); - } else if (SameType(v)) { - writer->WriteArrayItem(dmlc::get(v)); - } else if (SameType>(v)) { - writer->WriteArrayItem(dmlc::get>(v)); - } else if (SameType>>(v)) { - writer->WriteArrayItem(dmlc::get>>(v)); - } else if (SameType>(v)) { - writer->WriteArrayItem(dmlc::get>(v)); - } else { - LOG(FATAL) << "Not supported"; - } - } - writer->EndArray(); - } - inline static void Read(dmlc::JSONReader* reader, std::vector* data) { - LOG(FATAL) << "Not implemented."; - } -}; -} // namespace json -} // namespace dmlc diff --git a/src/relay/backend/graph_plan_memory.cc b/src/relay/backend/graph_plan_memory.cc deleted file mode 100644 index 33b3adea5f2f..000000000000 --- a/src/relay/backend/graph_plan_memory.cc +++ /dev/null @@ -1,416 +0,0 @@ -/* - * 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 relay/backend/graph_plan_memory.cc - * \brief Memory index assignment pass for executing - * the program in the graph executor. - */ -#include -#include -#include -#include -#include -#include -#include -#include - -#include "../../runtime/texture.h" -#include "../../support/arena.h" -#include "../op/annotation/annotation.h" -#include "../op/call/call.h" -#include "../op/memory/memory.h" -#include "../transforms/device_aware_visitors.h" -#include "./token_allocator.h" -#include "./utils.h" - -namespace tvm { -namespace relay { - -using TargetsMap = Map; -using Texture2DShape = runtime::Texture2DShape; -constexpr auto Is2DStorage = runtime::IsTextureStorage; - -using backend::StaticMemoryPlan; -using backend::StorageInfo; -using IntegerArray = Array; - -class StorageAllocaBaseVisitor : public transform::DeviceAwareExprVisitor { - public: - StorageAllocaBaseVisitor() : transform::DeviceAwareExprVisitor(Optional()) {} - - // run the visitor on a global function. - void Run(const Function& func) { VisitExpr(func); } - - using transform::DeviceAwareExprVisitor::VisitExpr_; - - void VisitExpr_(const ConstantNode* op) final { this->CreateToken(op, false); } - - void VisitExpr_(const VarNode* op) final { - // Do nothing. - } - - void DeviceAwareVisitExpr_(const FunctionNode* func_node) final { - if (function_nesting() > 1) { - // do not recurse into sub functions. - return; - } - if (func_node->HasNonzeroAttr(attr::kPrimitive)) { - // No storage needed for primitive functions. - return; - } - for (const auto& param : func_node->params) { - CreateToken(param.get(), /*can_realloc=*/false); - } - // Process the function body, and make sure all result tokens are considered 'alive'. - for (StorageToken* tok : GetToken(func_node->body)) { - tok->ref_counter += 1; - } - } - - void VisitExpr_(const GlobalVarNode* op) final { - // Do nothing. - } - - void VisitExpr_(const OpNode* op) final { - // Do nothing. - } - - void VisitExpr_(const TupleNode* op) final { - std::vector fields; - for (Expr field : op->fields) { - auto tokens = GetToken(field); - fields.insert(fields.end(), tokens.begin(), tokens.end()); - } - token_map_[op] = fields; - } - - void VisitExpr_(const TupleGetItemNode* op) final { - const auto& tok = GetToken(op->tuple); - ICHECK_LT(static_cast(op->index), tok.size()); - token_map_[op] = {tok[op->index]}; - } - - void VisitExpr_(const IfNode* op) final { LOG(FATAL) << "if is not supported."; } - - void PreVisitLetBinding_(const Var& var, const Expr& value) final { - token_map_[var.get()] = GetToken(value); - } - - void PostVisitLet_(const LetNode* let_node) final { - token_map_[let_node] = GetToken(let_node->body); - } - - protected: - /*! \brief internal token map */ - std::unordered_map> token_map_; - /*! \brief empty token map */ - const std::vector no_tokens_; - - /*! - * \brief Get the necessary token. - * \param expr The expression. - * \return The corresponding token. - */ - const std::vector& GetToken(const Expr& expr) { - this->VisitExpr(expr); - // See through on_device calls. - Expr real_expr = IgnoreOnDevice(expr); - - // Functions don't require data storage, represented by the empty token - if (real_expr->checked_type().as()) { - return no_tokens_; - } - this->VisitExpr(real_expr); - auto it = token_map_.find(real_expr.get()); - ICHECK(it != token_map_.end()) << "Expression not found in storage map:" << std::endl - << PrettyPrint(real_expr); - return it->second; - } - - /*! - * \brief Allocates (or reuses if \p can_realloc is true) a storage token for holding - * the result of evaluating \p op. - */ - void CreateToken(const ExprNode* expr_node, bool can_realloc) { - return CreateTokenOnDevice(expr_node, GetVirtualDevice(GetRef(expr_node)), can_realloc); - } - - /*! - * \brief Allocates (or reuses if \p can_realloc is true) a storage token for holding - * the result of evaluating \p op on \p device_type. - */ - virtual void CreateTokenOnDevice(const ExprNode* op, const VirtualDevice& virtual_device, - bool can_realloc) = 0; -}; - -/*! \brief Associate storage with every expression without any concern for sharing. */ -class StorageAllocaInit : protected StorageAllocaBaseVisitor { - public: - explicit StorageAllocaInit(support::Arena* arena) : arena_(arena) {} - - /*! \return The internal token map */ - std::unordered_map> GetInitTokenMap( - const Function& func) { - this->Run(func); - return std::move(token_map_); - } - - protected: - using StorageAllocaBaseVisitor::VisitExpr_; - - void CreateTokenOnDevice(const ExprNode* op, const VirtualDevice& virtual_device, - bool can_realloc) override { - ICHECK(!token_map_.count(op)); - std::vector tokens; - for (const auto& ttype : FlattenTupleType(op->checked_type())) { - auto* token = arena_->make(); - token->ttype = ttype; - token->virtual_device = virtual_device; - tokens.push_back(token); - } - token_map_[op] = tokens; - } - - using StorageAllocaBaseVisitor::DeviceAwareVisitExpr_; - - void DeviceAwareVisitExpr_(const CallNode* call_node) final { - // create token for the call node. - CreateToken(call_node, true); - - // for each input, visit argument token. - for (Expr arg : call_node->args) { - for (StorageToken* tok : GetToken(arg)) { - tok->ref_counter += 1; - } - } - } - - private: - // allocator - support::Arena* arena_; - Map> node_storage_map_; -}; - -/*! \brief Associate storage with every expression, reusing storage where possible. */ -class StorageAllocator : public StorageAllocaBaseVisitor { - public: - StorageAllocator() = default; - - /*! - * \return total number of bytes allocated - */ - size_t TotalAllocBytes() const { - size_t total = 0; - for (const auto* p : data_) { - total += p->max_bytes; - } - return total; - } - - // Run storage allocation for a function. - StaticMemoryPlan Plan(const Function& func) { - VLOG_CONTEXT << "StorageAllocator"; - VLOG(1) << "planning:" << std::endl << PrettyPrint(func); - prototype_ = StorageAllocaInit(&arena_).GetInitTokenMap(func); - // Backup the virtual devices as token reuse might lost the original memory scope - std::unordered_map> virtual_device_map_; - for (const auto& kv : prototype_) { - std::vector virtual_devices; - virtual_devices.reserve(kv.second.size()); - for (StorageToken* tok : kv.second) { - virtual_devices.push_back(tok->virtual_device); - } - virtual_device_map_.insert({kv.first, virtual_devices}); - } - this->Run(func); - - // The value of smap contains two integer arrays where the first array - // contains the planned storage ids and the second holds the device types. - Map smap; - int num_annotated_nodes = 0; - int num_nodes = 0; - - for (const auto& kv : token_map_) { - std::vector storage_ids; - storage_ids.reserve(kv.second.size()); - std::vector virtual_devices; - virtual_devices.reserve(kv.second.size()); - std::vector sid_sizes_byte; - sid_sizes_byte.reserve(kv.second.size()); - - for (StorageToken* tok : kv.second) { - VLOG(1) << "token: " << tok->ToString(); - if (tok->is_valid()) { - num_annotated_nodes++; - } - num_nodes++; - storage_ids.push_back(tok->storage_id); - sid_sizes_byte.push_back(allocator_.GetMemorySize(tok)); - } - ICHECK(kv.second.size() == virtual_device_map_[kv.first].size()) - << "Mismatch of tokens and virtual devices"; - for (auto vdev : virtual_device_map_[kv.first]) { - virtual_devices.push_back(vdev); - } - auto storage_info = backend::StorageInfo(std::move(storage_ids), std::move(virtual_devices), - std::move(sid_sizes_byte)); - smap.Set(GetRef(kv.first), storage_info); - } - // Either all or none of the nodes should be annotated. - VLOG(1) << "num annotated nodes / num_nodes: " << num_annotated_nodes << " / " << num_nodes - << std::endl; - if (num_annotated_nodes != 0 && num_annotated_nodes != num_nodes) { - LOG(FATAL) << num_annotated_nodes << " out of " << num_nodes - << "expressions are assigned with virtual device types. Either all " - "or none of the expressions are expected to be annotated."; - } - return backend::StaticMemoryPlan(smap); - } - - protected: - // override create token by getting token as prototype requirements. - void CreateTokenOnDevice(const ExprNode* op, const VirtualDevice& virtual_device, - bool can_realloc) final { - ICHECK(!token_map_.count(op)); - auto it = prototype_.find(op); - ICHECK(it != prototype_.end()); - std::vector tokens; - - for (StorageToken* tok : it->second) { - ICHECK(tok->virtual_device == virtual_device); - if (can_realloc) { - tokens.push_back(allocator_.Request(tok)); - } else { - // Allocate a new token, - StorageToken* allocated_tok = allocator_.Alloc(tok); - allocated_tok->virtual_device = tok->virtual_device; - // ensure it never get de-allocated. - allocated_tok->ref_counter += 1; - tokens.push_back(allocated_tok); - } - } - token_map_[op] = tokens; - } - - // Mark op to reuse the input_token - // tie the two memories together - void ReuseInputToken(const ExprNode* op, StorageToken* input_token) { - ICHECK(!token_map_.count(op)); - auto it = prototype_.find(op); - ICHECK(it != prototype_.end()); - ICHECK_EQ(it->second.size(), 1U); - StorageToken* prototype = it->second[0]; - // add the reference counter of the output - // so the input token can only be deleted after references - // to both are expired - input_token->ref_counter += prototype->ref_counter; - // reuse the input token - token_map_[op] = {input_token}; - } - - using StorageAllocaBaseVisitor::DeviceAwareVisitExpr_; - - // The call map - void DeviceAwareVisitExpr_(const CallNode* call_node) final { - std::vector args; - // for each input, visit argument token. - - for (const Expr& arg : call_node->args) { - // Note: GetToken skips GlobalVars and handles tuples properly, so we don't need to treat - // call_lowered specially. - for (StorageToken* tok : GetToken(arg)) { - args.push_back(tok); - } - } - - // Under the flat-memory setting. - // we can force aliasing the input and output of reshape - // to make it an nop. Note that this is not true - // for non-flat memory case. Given the current graph plan memory - // only works for flat memory case, we will go with this choice - // - // TODO(tvm-team) Update checks of flat memory enablement when we support - // opaque-nd memory planning to skip this path. - // TODO(mbs): "reshape" cleanup. - CallLoweredProps call_lowered_props = GetCallLoweredProps(call_node); - if (call_lowered_props.lowered_func.defined() && IsReshapeOnly(call_lowered_props)) { - ICHECK_EQ(call_lowered_props.arguments.size(), 1U); - ReuseInputToken(call_node, args[0]); - } else { - // create token for the call node. - CreateToken(call_node, true); - } - - // check if there is orphaned output that can be released immediately. - for (StorageToken* tok : token_map_.at(call_node)) { - allocator_.CheckForRelease(tok); - } - for (StorageToken* tok : args) { - tok->ref_counter -= 1; - allocator_.CheckForRelease(tok); - } - } - - class TokenAllocator { - public: - StorageToken* Alloc(StorageToken* proto) { return token_mixed_.Alloc(proto, storage_ids_++); } - StorageToken* Request(StorageToken* proto) { - StorageToken* token = token_mixed_.Request(proto); - return token ? token : this->Alloc(proto); - } - void CheckForRelease(StorageToken* tok) { return token_mixed_.CheckForRelease(tok); } - - size_t GetMemorySize(StorageToken* tok) { - // TODO(amalyshe): figure out who requries sizes and for what - // size in case of texture is not enough - we can return any value if it - // assumed to be used for memory allocatoion or we can return real size - // if it is just for information - return token_mixed_.GetMemorySize(tok); - } - static bool Is2DStorage(StorageToken* tok) { - return relay::Is2DStorage(tok->virtual_device->memory_scope); - } - - private: - int64_t storage_ids_{0}; - TokenAllocatorMixed token_mixed_; - }; - - private: - // allocator - support::Arena arena_; - // scale used for rough match - // size_t match_range_{16}; - // free list of storage entry - std::multimap free_; - // all the storage resources available - std::vector data_; - /*! \brief internal prototype token map */ - std::unordered_map> prototype_; - /*! \brief token allocator for optimizing 1d and 2d token alloc requests */ - TokenAllocator allocator_; -}; - -StaticMemoryPlan GraphPlanMemory(const Function& func) { return StorageAllocator().Plan(func); } - -TVM_REGISTER_GLOBAL("relay.backend.GraphPlanMemory").set_body_typed(GraphPlanMemory); - -} // namespace relay -} // namespace tvm diff --git a/src/relay/backend/interpreter.cc b/src/relay/backend/interpreter.cc deleted file mode 100644 index 865b5616edab..000000000000 --- a/src/relay/backend/interpreter.cc +++ /dev/null @@ -1,1124 +0,0 @@ -/* - * 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 src/relay/interpreter.cc - * \brief An interpreter for the Relay IR. - */ - -#include -#include -#include -#include -#include -#include -#include -#include -#include -#include -#include -#include -#include -#include -#include - -#include "../op/annotation/annotation.h" -#include "../op/call/call.h" -#include "../op/memory/device_copy.h" -#include "../transforms/pass_utils.h" -#include "te_compiler.h" - -namespace tvm { -namespace relay { - -using runtime::ADT; -using runtime::ADTObj; -using runtime::NDArray; -using runtime::TVMArgsSetter; -using runtime::operator<<; - -namespace { -// TODO(mbs): Centralize. -struct PairHash { - template - std::size_t operator()(const std::pair& k) const { - return dmlc::HashCombine(std::hash()(k.first), std::hash()(k.second)); - } - template - std::size_t operator()(const std::pair& k) const { - return dmlc::HashCombine(ObjectHash()(k.first), std::hash()(k.second)); - } -}; - -// Analogue of FlattenTupleType for runtime ADT vs NDArray values. -// TODO(mbs): Hoist somewhere sensible, maybe op/memory.h? -void FlattenADTAux(const ObjectRef& object_ref, std::vector* out) { - if (auto ndarray = object_ref.as()) { - out->push_back(ndarray.value()); - } else if (const ADTObj* adt = object_ref.as()) { - for (size_t i = 0; i < adt->size; ++i) { - FlattenADTAux((*adt)[i], out); - } - } else { - LOG(FATAL) << "unsupported " << object_ref; - } -} - -std::vector FlattenADT(const ObjectRef& object_ref) { - std::vector out; - FlattenADTAux(object_ref, &out); - return out; -} - -std::vector FlattenADTs(const std::vector& object_refs) { - std::vector out; - for (const auto& object_ref : object_refs) { - FlattenADTAux(object_ref, &out); - } - return out; -} - -// Analogue of ToTupleType for runtime ADT vs NDArray values. -// TODO(mbs): Hoist somewhere sensible, maybe op/memory.h? -void ToADTOrNDArrayAux(const Type& type, const std::vector& nd_arrays, int* index, - std::vector* out) { - if (type.as()) { - out->push_back(nd_arrays[*index]); - *index += 1; - } else if (const TupleTypeNode* ttn = type.as()) { - std::vector tuple_out; - for (size_t i = 0; i < ttn->fields.size(); i++) { - ToADTOrNDArrayAux(ttn->fields[i], nd_arrays, index, &tuple_out); - } - out->push_back(ADT::Tuple(tuple_out)); - } else { - LOG(FATAL) << "unsupported " << type; - } -} - -ObjectRef ToADTOrNDArray(const Type& type, const std::vector& nd_arrays) { - if (type.as() && nd_arrays.size() == 1) { - return nd_arrays[0]; - } else { - std::vector out; - int index = 0; - ToADTOrNDArrayAux(type, nd_arrays, &index, &out); - return out[0]; - } -} - -} // namespace - -InterpreterClosure::InterpreterClosure(Map env, Function func) { - ObjectPtr n = make_object(); - n->env = std::move(env); - n->func = std::move(func); - data_ = std::move(n); -} - -TVM_STATIC_IR_FUNCTOR(ReprPrinter, vtable) - .set_dispatch([](const ObjectRef& ref, ReprPrinter* p) { - auto* node = static_cast(ref.get()); - p->stream << "InterpreterClosureNode(" << node->func << ", " << node->env << ")"; - }); - -inline const PackedFunc& GetPackedFunc(const std::string& name) { - const PackedFunc* pf = runtime::Registry::Get(name); - ICHECK(pf != nullptr) << "Cannot find function " << name << " in registry"; - return *pf; -} - -// TODO(@jroesch): this doesn't support mutual letrec -/* Object Implementation */ -RecClosure::RecClosure(InterpreterClosure clos, Var bind) { - ObjectPtr n = make_object(); - n->clos = std::move(clos); - n->bind = std::move(bind); - data_ = std::move(n); -} - -TVM_STATIC_IR_FUNCTOR(ReprPrinter, vtable) - .set_dispatch([](const ObjectRef& ref, ReprPrinter* p) { - auto* node = static_cast(ref.get()); - p->stream << "RecClosureObj(" << node->clos << ")"; - }); - -RefValue::RefValue(ObjectRef value) { - ObjectPtr n = make_object(); - n->value = value; - data_ = std::move(n); -} - -TVM_REGISTER_GLOBAL("relay._make.RefValue").set_body_typed([](ObjectRef value) { - return RefValue(value); -}); - -TVM_REGISTER_NODE_TYPE(RefValueObj); - -TVM_STATIC_IR_FUNCTOR(ReprPrinter, vtable) - .set_dispatch([](const ObjectRef& ref, ReprPrinter* p) { - auto* node = static_cast(ref.get()); - p->stream << "RefValueObj(" << node->value << ")"; - }); - -ConstructorValue::ConstructorValue(int32_t tag, Array fields, Constructor constructor) { - ObjectPtr n = make_object(); - n->tag = tag; - n->fields = fields; - n->constructor = constructor; - data_ = std::move(n); -} - -TVM_REGISTER_GLOBAL("relay._make.ConstructorValue") - .set_body_typed([](int32_t tag, Array fields, Constructor constructor) { - return ConstructorValue(tag, fields, constructor); - }); - -TVM_REGISTER_NODE_TYPE(ConstructorValueObj); - -TVM_STATIC_IR_FUNCTOR(ReprPrinter, vtable) - .set_dispatch([](const ObjectRef& ref, ReprPrinter* p) { - auto* node = static_cast(ref.get()); - p->stream << "ConstructorValueObj(" << node->tag << "," << node->fields << ")"; - }); - -/*! - * \brief A stack frame in the Relay interpreter. - * - * Contains a mapping from relay::Var to relay::ObjectRef. - */ -struct Frame { - /*! \brief The set of local variables and arguments for the frame. */ - Map locals; - - explicit Frame(Map locals) : locals(locals) {} -}; - -/*! - * \brief The call stack in the Relay interpreter. - * - * Contains a stack of frames; each corresponding to - * a function call. - */ -struct Stack { - /*! \brief The stack frames. */ - std::vector frames; - Stack() : frames() { frames.push_back(Frame({})); } - - Frame& current_frame() { return frames.back(); } - - ObjectRef Lookup(const Var& local) { - for (auto frame = frames.rbegin(); frame != frames.rend(); frame++) { - auto elem = frame->locals.find(local); - if (elem != frame->locals.end()) { - return (*elem).second; - } - } - - LOG(FATAL) << "could not find variable binding for " << local - << "address= " << local.operator->(); - return ObjectRef(); - } - /*! - * A wrapper around Frame to add RAII semantics to pushing and popping - * stack frames. - */ - struct LocalFrame { - Stack& st; - explicit LocalFrame(Stack& st, const Frame& fr) : st(st) { st.frames.push_back(fr); } - ~LocalFrame() { st.frames.pop_back(); } - }; -}; - -/*! \brief A representation of the interpreter state which can be passed back to Python. */ -class InterpreterState; - -/*! \brief A container capturing the state of the interpreter. */ -class InterpreterStateObj : public Object { - public: - using Frame = Map; - using Stack = Array; - - /*! \brief The current expression under evaluation. */ - Expr current_expr; - - /*! \brief The call stack of the interpreter. */ - Stack stack; - - void VisitAttrs(AttrVisitor* v) { - v->Visit("current_expr", ¤t_expr); - v->Visit("stack", &stack); - } - - static constexpr const char* _type_key = "relay.InterpreterState"; - TVM_DECLARE_FINAL_OBJECT_INFO(InterpreterStateObj, Object); -}; - -class InterpreterState : public ObjectRef { - public: - using Frame = Map; - using Stack = Array; - - InterpreterState(Expr current_expr, Stack stack); - - TVM_DEFINE_OBJECT_REF_METHODS(InterpreterState, ObjectRef, InterpreterStateObj); -}; - -InterpreterState::InterpreterState(Expr current_expr, InterpreterState::Stack stack) { - ObjectPtr n = make_object(); - n->current_expr = std::move(current_expr); - n->stack = std::move(stack); - data_ = std::move(n); -} - -// NOTE: the current interpreter assumes A-normal form. -// which is better for execution. -// -// It will run duplicated computations when taking program that -// contains DAG in dataflow-form. -// -// Conversion to ANF is recommended before running the interpretation. -class Interpreter : public ExprFunctor, - PatternFunctor { - public: - Interpreter(IRModule unified_mod, CompilationConfig config, Device device) - : unified_mod_(unified_mod), - config_(std::move(config)), - device_(device), - debug_op_(Op::Get("debug")) {} - - template - T WithFrame(const Frame& fr, const std::function& f) { - Stack::LocalFrame lf(stack_, fr); - return f(); - } - - void extend(const Var& id, ObjectRef v) { stack_.current_frame().locals.Set(id, v); } - - ObjectRef Lookup(const Var& local) { return stack_.Lookup(local); } - - ObjectRef Eval(const Expr& expr) { return VisitExpr(expr); } - - ObjectRef VisitExpr_(const VarNode* var_node) final { return Lookup(GetRef(var_node)); } - - ObjectRef VisitExpr_(const GlobalVarNode* op) final { - return Eval(unified_mod_->Lookup(GetRef(op))); - } - - ObjectRef VisitExpr_(const OpNode* id) override { - // TODO(@jroesch): Eta-expand and return in this case. - LOG(FATAL) << "internal error, need to wrap intrinsic into call synthetic call node " - << "in this case, eta expand"; - return ObjectRef(); - } - - ObjectRef VisitExpr_(const ConstantNode* op) final { return op->data.CopyTo(device_); } - - ObjectRef VisitExpr_(const TupleNode* op) final { - std::vector values; - - for (const auto& field : op->fields) { - ObjectRef field_value = Eval(field); - values.push_back(field_value); - } - - return ADT::Tuple(values); - } - - ObjectRef MakeClosure(const Function& func, Var letrec_name = Var()) { - Map captured_mod; - Array free_vars = FreeVars(func); - - for (const auto& var : free_vars) { - // Evaluate the free var (which could be a function call) if it hasn't - // shown up in a letting binding that has invoked the function. - if (letrec_name.defined() && letrec_name == var) { - continue; - } - - captured_mod.Set(var, Eval(var)); - } - - // We must use mutation here to build a self referential closure. - InterpreterClosure closure(captured_mod, func); - if (letrec_name.defined()) { - return RecClosure(closure, letrec_name); - } - return std::move(closure); - } - - ObjectRef VisitExpr_(const FunctionNode* func_node) final { - auto func = GetRef(func_node); - return MakeClosure(func); - } - - /*! - * \brief Returns the packed function implementing the TIR function bound to \p tir_fn_var. - * - * \param tir_fn_var Global var for the already lowered TIR function. - * \param all_tir_fn_vars Global vars for all lowered TIR functions the above - * may reference, plus \p tir_fn_var itself. - * \param target Target for which the TIR function should be compiled. For primitives this - * will be the interpreter's target_. However for shape functions this will be the generic - * 'cpu' target, since shape functions are always executed on the host cpu. - */ - PackedFunc TIRToPackedFunc(const GlobalVar& tir_fn_var, const Array& all_tir_fn_vars, - Target target) { - std::pair packed_func_key(target, tir_fn_var->name_hint); - auto packed_itr = compiled_packed_funcs_.find(packed_func_key); - if (packed_itr != compiled_packed_funcs_.end()) { - // Already compiled. - return packed_itr->second; - } - - // Project out just the function(s) we need. - IRModule lowered_projected_mod; - Map per_target_module = tec::GetPerTargetModules(unified_mod_); - std::unordered_map - per_target_module_std_map = backend::TargetModuleMapToTargetStrModuleMap(per_target_module); - auto mod_itr = per_target_module_std_map.find(target); - ICHECK(mod_itr != per_target_module_std_map.end()) - << "No target module for target " << target->ToDebugString(); - const IRModule& target_module = (*mod_itr).second; - for (const auto& var : all_tir_fn_vars) { - ICHECK(target_module->ContainGlobalVar(var->name_hint)) - << "No global var for '" << var->name_hint << "' in module for target " - << target->ToDebugString(); - lowered_projected_mod->Add(var, target_module->Lookup(var->name_hint)); - } - - // Compile (aka 'build') the projected module into a runtime module of packed functions. - runtime::Module runtime_module; - if (const auto* f = runtime::Registry::Get("relay.backend.build")) { - // TODO(mbs): Cleanup hooks. - runtime_module = (*f)(lowered_projected_mod, target); - } else { - runtime_module = build(lowered_projected_mod, target, /*target_host=*/Target(nullptr)); - } - - // Extract all the packed functions. - for (const auto& var : all_tir_fn_vars) { - PackedFunc packed_func = runtime_module.GetFunction(var->name_hint); - ICHECK(packed_func != nullptr) - << "No packed function for global var '" << var->name_hint - << "' in compiled module for target " << target->ToDebugString(); - compiled_packed_funcs_.emplace(std::make_pair(target, var->name_hint), packed_func); - } - - // Return just what we need for this call. - packed_itr = compiled_packed_funcs_.find(packed_func_key); - ICHECK(packed_itr != compiled_packed_funcs_.end()) << " " << tir_fn_var->name_hint; - ICHECK_NOTNULL(packed_itr->second); - return packed_itr->second; - } - - /*! - * \brief Call the dynamic shape function bound to \p prim_shape_fn_var passing the - * shapes of args, and return the resulting shapes. - * - * \param prim_shape_fn_var Global var bound to lowered shape function. - * \param all_prim_shape_fn_vars All the global vars needed to build the above, including - * the shape function itself. - * \param prim_shape_fn_states For each primitive arg, indicate whether the primitive shape - * function requires the shape of the argument and/or the actual argument tensor. - * \param num_shape_inputs The number of inputs, after accounting for both shapes vs data - * inputs and unfolding of tuple types. - * \param num_shape_outputs The number of outputs, after accounting for flattening of - * tuple types. - * \param args Arguments to the primitive this shape function is for. - * \return Expected shapes of the underlying primitive's flattened outputs. - */ - Array ComputeDynamicShape(const GlobalVar& prim_shape_fn_var, - const Array& all_prim_shape_fn_vars, - const Array& prim_shape_fn_states, - size_t num_shape_inputs, size_t num_shape_outputs, - Target prim_shape_target, const std::vector& args) { - VLOG_CONTEXT << "ComputeDynamicShape"; - ICHECK(prim_shape_fn_var.defined()); - ICHECK(prim_shape_fn_var->checked_type().defined()); - VLOG(1) << "prim_shape_fn_var:" << std::endl << PrettyPrint(prim_shape_fn_var); - ICHECK(prim_shape_fn_states.defined()); - for (size_t i = 0; i < prim_shape_fn_states.size(); ++i) { - VLOG(1) << "prim_shape_fn_states[" << i << "]: " << prim_shape_fn_states[i]; - } - VLOG(1) << "num_shape_inputs: " << num_shape_inputs; - VLOG(1) << "num_shape_outputs: " << num_shape_outputs; - VLOG(1) << "args.size(): " << args.size(); - VLOG(1) << "prim_shape_target: " << prim_shape_target->ToDebugString(); - - // The function type is that of the shape function rather than the original primitive the shape - // function is for. - const auto* func_type_node = prim_shape_fn_var->checked_type().as(); - ICHECK(func_type_node); - // The shape function states are w.r.t. the original primitive's arguments in - // non-flattened form. - // TODO(mbs): Clean this up so we don't mix flattened vs original conventions. - ICHECK_EQ(args.size(), prim_shape_fn_states.size()); - - // num_shape_inputs will account for which primitive function arguments are dynamic, - // whether the shape and or data needs to be passed, and flattening of tuples. - // Similarly, num_shape_outputs will account for flattening of tuples. - - // TODO(mbs): Take this from the host_virtual_device. - Device shape_device; - shape_device.device_type = static_cast(prim_shape_target->GetTargetDeviceType()); - shape_device.device_id = 0; - - // 'Compile' the TIR shape function to appropriate callable form. - PackedFunc packed_shape_func = - TIRToPackedFunc(prim_shape_fn_var, all_prim_shape_fn_vars, prim_shape_target); - - size_t arity = num_shape_inputs + num_shape_outputs; - std::vector values(arity); - std::vector codes(arity); - TVMArgsSetter setter(values.data(), codes.data()); - std::vector inputs(num_shape_inputs); - std::vector outputs(num_shape_outputs); - - // Collect the shapes and/or data needed by the shape function from - // the primitive's arguments. - size_t arg_counter = 0; - for (size_t i = 0; i < args.size(); ++i) { - // TODO(mbs): The same need data/need shape arg state applies to everything in the - // flattened form of this arg. Does that match what lowering actually does? - int64_t state = prim_shape_fn_states[i]->value; - for (const auto& nd_array : FlattenADT(args[i])) { - if (state & tec::kNeedInputData) { - auto arr = nd_array.CopyTo(shape_device); - inputs[arg_counter] = arr; - setter(arg_counter, arr); - ++arg_counter; - } - if (state & tec::kNeedInputShape) { - int64_t ndim = nd_array.Shape().size(); - NDArray shape_arr; - if (ndim == 0) { - shape_arr = NDArray::Empty({}, DataType::Int(64), shape_device); - } else { - shape_arr = NDArray::Empty({ndim}, DataType::Int(64), shape_device); - int64_t* data = reinterpret_cast(shape_arr->data); - for (auto j = 0; j < ndim; ++j) { - data[j] = nd_array.Shape()[j]; - } - } - inputs[arg_counter] = shape_arr; - setter(arg_counter, shape_arr); - ++arg_counter; - } - } - } - ICHECK_EQ(arg_counter, num_shape_inputs) << "Shape function input sizes mismatch"; - - // Prepare NDArrays to hold the output shapes. - size_t out_cnt = 0; - for (const auto& ttype : FlattenTupleType(func_type_node->ret_type)) { - ICHECK(out_cnt < num_shape_outputs); - std::vector concrete_shape; - for (const auto& dim : ttype->shape) { - const auto* ivalue = tir::as_const_int(dim); - ICHECK(ivalue) << "expected concrete dimensions"; - concrete_shape.push_back(ivalue[0]); - } - auto arr = NDArray::Empty(concrete_shape, ttype->dtype, shape_device); - outputs[out_cnt] = arr; - setter(arg_counter + out_cnt, arr); - ++out_cnt; - } - ICHECK_EQ(out_cnt, num_shape_outputs) << "Shape function output sizes mismatch"; - - // Call the dynamic shape function. - TVMRetValue rv; // ignored - packed_shape_func.CallPacked(TVMArgs(values.data(), codes.data(), arity), &rv); - - // Convert result tensors back to shapes. - Array out_shapes; - for (auto out_tensor : outputs) { - int64_t* shape_data = reinterpret_cast(out_tensor->data); - Shape out_shape; - for (int i = 0; i < out_tensor->shape[0]; ++i) { - out_shape.push_back(Integer(shape_data[i])); - } - out_shapes.push_back(out_shape); - } - return out_shapes; - } - - /*! - * \brief Call primitive op bound to \p prim_fn_var with \p args. If necessary, evaluate dynamic - * shape function bound to \p prim_shape_fn_var to calculate shapes of result tensors. - * - * @param prim_fn_var Global bound to lowered primitive. - * @param all_prim_fn_vars All globals references by lowered primitive, plus prim_fn_var itself. - * @param prim_shape_fn_var Global bound to lowered shape function for primitive, if needed. - * @param all_prim_shape_fn_vars All globals references by lowered shape function, plus - * prim_shape_fn_var itself. - * @param prim_shape_fn_states Records whether shape and/or data is needed by the dynamic - * shape function (if any) for each (flattened) argument. - * @param num_shape_inputs Number of arguments to the dynamic shape function (if any). - * @param num_shape_outputs Number of outputs from the dynamic shape function (if any). - * @param args Already evaluated arguments to primitive. - * @return Result of primitive. - */ - ObjectRef InvokePrimitiveOp(const GlobalVar& prim_fn_var, const Array all_prim_fn_vars, - Target prim_target, const GlobalVar& prim_shape_fn_var, - const Array& all_prim_shape_fn_vars, - const Array& prim_shape_fn_states, size_t num_shape_inputs, - size_t num_shape_outputs, Target prim_shape_target, - const std::vector& args) { - ICHECK(prim_fn_var->checked_type().defined()); - const FuncTypeNode* ftn = prim_fn_var->checked_type().as(); - ICHECK(ftn); - - // 'Compile' the TIR primitive to appropriate callable form (on the desired target). - PackedFunc packed_func = TIRToPackedFunc(prim_fn_var, all_prim_fn_vars, prim_target); - - // Argument tuples are flattened. - std::vector arg_nd_arrays = FlattenADTs(args); - const size_t num_inputs = arg_nd_arrays.size(); - // num_inputs should equal size(concat(map(FlattenTupleType, function arg types))) - - // TVM's primitive calling convention is for the final arguments to be for output - // buffers. We must allocate space for those buffers based on the return type. - std::vector result_tensor_types = FlattenTupleType(ftn->ret_type); - const size_t arg_len = num_inputs + result_tensor_types.size(); - - std::vector values(arg_len); - std::vector codes(arg_len); - TVMArgsSetter setter(values.data(), codes.data()); - - // Marshall the call's arguments in flattened form. - int arg_counter = 0; - for (const auto& nd_array : arg_nd_arrays) { - setter(arg_counter++, nd_array); - Device arg_dev = nd_array->device; - ICHECK(arg_dev.device_type == device_.device_type && arg_dev.device_id == device_.device_id) - << "Interpreter expect device to be " << device_ << ", but got " << arg_dev; - } - - // If necessary, retrieve concrete shapes for outputs from shape function rather - // than relying on TensorType shapes. - Array runtime_shapes; - bool is_dyn = IsDynamic(ftn->ret_type); - if (is_dyn) { - ICHECK(prim_shape_fn_var.defined()); - ICHECK(prim_shape_fn_states.defined()); - runtime_shapes = - ComputeDynamicShape(prim_shape_fn_var, all_prim_shape_fn_vars, prim_shape_fn_states, - num_shape_inputs, num_shape_outputs, prim_shape_target, args); - ICHECK_EQ(runtime_shapes.size(), result_tensor_types.size()); - } - - // Prepare the result tensors for the call. - TVMRetValue rv; // ignored - std::vector result_nd_arrays; - for (size_t i = 0; i < result_tensor_types.size(); ++i) { - const auto& ttype = result_tensor_types[i]; - const Shape& shape = is_dyn ? runtime_shapes[i] : ttype->shape; - // Allocate output tensor of appropriate shape. - std::vector concrete_shape; - for (const auto& dim : shape) { - const auto* ivalue = tir::as_const_int(dim); - ICHECK(ivalue) << "expected concrete dimensions"; - concrete_shape.push_back(ivalue[0]); - } - NDArray nd_array = NDArray::Empty(concrete_shape, ttype->dtype, device_); - setter(num_inputs + i, nd_array); - result_nd_arrays.emplace_back(nd_array); - } - - // Call the primitive. - packed_func.CallPacked(TVMArgs(values.data(), codes.data(), static_cast(arg_len)), &rv); - - // Unflatten the results. - return ToADTOrNDArray(ftn->ret_type, result_nd_arrays); - } - - /*! - * \brief Invoke \p closure with \p args. If \p bind is defined then this is a recursive - * closure and \p bind should refer to itself. - */ - ObjectRef Invoke(const InterpreterClosure& closure, const Array& args, - const Var& bind = Var()) { - // Get a reference to the function inside the closure. - Function func = closure->func; - ICHECK_EQ(func->params.size(), args.size()); - - if (func->HasNonzeroAttr(attr::kPrimitive)) { - if (const CallNode* call_node = closure->func->body.as()) { - if (call_node->op == debug_op_) { - // Special case: Calling the debug tracing function. - auto dattrs = call_node->attrs.as(); - auto interp_state = get_state(call_node->args[0]); - - if (dattrs->debug_func.defined()) { - dattrs->debug_func(interp_state); - } else { - RELAY_DEBUG_INTERP(interp_state); - } - - return args[0]; - } - } - } - - ICHECK(!func->HasNonzeroAttr(attr::kPrimitive)) - << "Calls to primitive functions should have been removed by lowering"; - - // Allocate a frame with the parameters and free variables. - Map locals; - for (size_t i = 0; i < func->params.size(); i++) { - ICHECK_EQ(locals.count(func->params[i]), 0); - locals.Set(func->params[i], args[i]); - } - - // Add the var to value mappings from the Closure's environment. - for (auto it = closure->env.begin(); it != closure->env.end(); ++it) { - ICHECK_EQ(locals.count((*it).first), 0); - locals.Set((*it).first, (*it).second); - } - - if (bind.defined()) { - locals.Set(bind, RecClosure(closure, bind)); - } - - return WithFrame(Frame(locals), [&]() { return Eval(func->body); }); - } - - ObjectRef VisitExpr_(const CallNode* call_node) final { - DeviceCopyProps device_copy_props = GetDeviceCopyProps(call_node); - CallLoweredProps call_lowered_props = GetCallLoweredProps(call_node); - - if (device_copy_props.body.defined()) { - // TODO(mbs): device_copy cleanup - LOG(FATAL) << "The interpreter does not support device_copy"; - } else if (call_lowered_props.lowered_func.defined()) { - // Special case: Call a lowered TIR function. - - // Evaluate only function args - std::vector args; - for (auto arg : call_lowered_props.arguments) { - args.push_back(Eval(arg)); - } - - // TODO(mbs): Make calling convention first-class in Relay. - Array all_prim_fn_vars; - if (call_lowered_props.attrs.metadata.count("all_prim_fn_vars")) { - all_prim_fn_vars = - Downcast>(call_lowered_props.attrs.metadata.at("all_prim_fn_vars")); - } - GlobalVar prim_shape_fn_var; - if (call_lowered_props.attrs.metadata.count("prim_shape_fn_var")) { - prim_shape_fn_var = - Downcast(call_lowered_props.attrs.metadata.at("prim_shape_fn_var")); - } - Array all_prim_shape_fn_vars; - if (call_lowered_props.attrs.metadata.count("all_prim_shape_fn_vars")) { - all_prim_shape_fn_vars = Downcast>( - call_lowered_props.attrs.metadata.at("all_prim_shape_fn_vars")); - } - Array prim_shape_fn_states; - if (call_lowered_props.attrs.metadata.count("prim_shape_fn_states")) { - prim_shape_fn_states = - Downcast>(call_lowered_props.attrs.metadata.at("prim_shape_fn_states")); - } - - size_t num_shape_inputs = 0; - if (call_lowered_props.attrs.metadata.count("prim_shape_fn_num_inputs")) { - num_shape_inputs = static_cast( - Downcast(call_lowered_props.attrs.metadata.at("prim_shape_fn_num_inputs")) - ->value); - } - size_t num_shape_outputs = 0; - if (call_lowered_props.attrs.metadata.count("prim_shape_fn_num_outputs")) { - num_shape_outputs = static_cast( - Downcast(call_lowered_props.attrs.metadata.at("prim_shape_fn_num_outputs")) - ->value); - } - ICHECK(config_->optional_homogeneous_target.defined()); - return InvokePrimitiveOp(call_lowered_props.lowered_func, all_prim_fn_vars, - config_->optional_homogeneous_target, prim_shape_fn_var, - all_prim_shape_fn_vars, prim_shape_fn_states, num_shape_inputs, - num_shape_outputs, config_->host_virtual_device->target, args); - } else { // All other calls - // Evaluate all arguments - std::vector args; - for (auto arg : call_node->args) { - args.push_back(Eval(arg)); - } - - if (call_node->op == OnDeviceOp()) { - // Special case: The call 'on_device(expr)' denotes that expr should be executed on - // a particular device. We can ignore this during interpretation. - ICHECK_EQ(call_node->args.size(), 1UL); - return args[0]; - } - if (const ConstructorNode* con = call_node->op.as()) { - // Special case: ADT constructor - - return ConstructorValue(con->tag, args, GetRef(con)); - } - - if (const OpNode* op_node = call_node->op.as()) { - // Except for call_lowered and on_device, we should not find calls to operators after - // running fusion and lowering. - LOG(FATAL) << "found " << op_node->name - << "; operators should have been removed by previous passes; try " - "fusing and lowering"; - } - - // Now we just evaluate and expect to find a closure. - // TODO(@electriclilies): How should call_lowered behave with closures? - ObjectRef fn_val = Eval(call_node->op); - if (auto closure = fn_val.as()) { - return Invoke(closure.value(), args); - } else if (const RecClosureObj* closure_node = fn_val.as()) { - return Invoke(closure_node->clos, args, closure_node->bind); - } else { - LOG(FATAL) << "internal error: type error, expected function value in the call " - << "position"; - return ObjectRef(); - } - } - } - - ObjectRef VisitExpr_(const LetNode* let) final { - if (auto func = let->value.as()) { - auto clo = MakeClosure(func.value(), let->var); - this->extend(let->var, clo); - } else { - auto value = Eval(let->value); - this->extend(let->var, value); - } - - return Eval(let->body); - } - - ObjectRef VisitExpr_(const TupleGetItemNode* op) final { - ObjectRef val = Eval(op->tuple); - const auto* adt_obj = val.as(); - ICHECK(adt_obj) << "internal error: when evaluating TupleGetItem expected an ADT value"; - auto adt = GetRef(adt_obj); - ICHECK_LT(static_cast(op->index), adt.size()) << "internal error: index out of bounds"; - return adt[op->index]; - } - - ObjectRef VisitExpr_(const IfNode* op) final { - ObjectRef v = Eval(op->cond); - if (v->IsInstance()) { - auto nd_array = Downcast(v); - Device cpu_dev; - cpu_dev.device_type = kDLCPU; - cpu_dev.device_id = 0; - NDArray cpu_array = nd_array.CopyTo(cpu_dev); - ICHECK_EQ(DataType(cpu_array->dtype), DataType::Bool()); - // TODO(@jroesch, @MK): Refactor code into helper from DCE. - if (reinterpret_cast(cpu_array->data)[0]) { - return Eval(op->true_branch); - } else { - return Eval(op->false_branch); - } - } else { - LOG(FATAL) << "type error, type system should have caught this"; - } - } - - ObjectRef VisitExpr_(const RefWriteNode* op) final { - ObjectRef r = Eval(op->ref); - if (const RefValueObj* rv = r.as()) { - rv->value = Eval(op->value); - return ADT::Tuple(std::vector()); - } else { - LOG(FATAL) << "type error, type system should have caught this"; - } - } - - ObjectRef VisitExpr_(const RefCreateNode* op) final { return RefValue(Eval(op->value)); } - - ObjectRef VisitExpr_(const RefReadNode* op) final { - ObjectRef r = Eval(op->ref); - if (const RefValueObj* rv = r.as()) { - return rv->value; - } else { - LOG(FATAL) << "type error, type system should have caught this"; - } - } - - ObjectRef VisitExpr_(const MatchNode* op) final { - ObjectRef v = Eval(op->data); - for (const Clause& c : op->clauses) { - if (VisitPattern(c->lhs, v)) { - return VisitExpr(c->rhs); - } - } - LOG(FATAL) << "did not find any match"; - } - - bool VisitPattern_(const PatternConstructorNode* op, const ObjectRef& v) final { - const ConstructorValueObj* cvn = v.as(); - ICHECK(cvn) << "need to be a constructor for match"; - ICHECK_NE(op->constructor->tag, -1); - ICHECK_NE(cvn->tag, -1); - if (op->constructor->tag == cvn->tag) { - ICHECK_EQ(op->patterns.size(), cvn->fields.size()); - for (size_t i = 0; i < op->patterns.size(); ++i) { - if (!VisitPattern(op->patterns[i], cvn->fields[i])) { - return false; - } - } - return true; - } - return false; - } - - bool VisitPattern_(const PatternTupleNode* op, const ObjectRef& v) final { - auto adt = Downcast(v); - ICHECK_EQ(op->patterns.size(), adt.size()); - for (size_t i = 0; i < op->patterns.size(); ++i) { - if (!VisitPattern(op->patterns[i], adt[i])) { - return false; - } - } - return true; - } - - bool VisitPattern_(const PatternWildcardNode* op, const ObjectRef& v) final { return true; } - - bool VisitPattern_(const PatternVarNode* op, const ObjectRef& v) final { - extend(op->var, v); - return true; - } - - InterpreterState get_state(Expr e = Expr()) const { - InterpreterStateObj::Stack stack; - for (auto fr : this->stack_.frames) { - InterpreterStateObj::Frame frame = fr.locals; - stack.push_back(frame); - } - auto state = InterpreterState(e, stack); - return state; - } - - private: - // Unified module. Functions are annotated with their target. - // All expressions are eval'ed w.r.t. the definitions in this module. - // This module contains functions that used to be in main_module and the per_target_module (TIR - // functions) in one module. - IRModule unified_mod_; - // Cached packed functions for the primitives and shape functions, keyed by target and - // global var name. - std::unordered_map, PackedFunc, PairHash> compiled_packed_funcs_; - /*! \brief Compilation config describing the available targets. */ - CompilationConfig config_; - // Unique device on which primitives (but not shape functions) will be executed. - // (For simplicity we only run the interpreter on a single device.) - Device device_; - // Call stack. - Stack stack_; - // The distinguished 'debug' operator, which is handled specially. - const Op& debug_op_; -}; - -/*! - * Lowers all calls to primitives in \p mod appropriate for \p config. Returns the - * rewritten \p mod and target-specific modules containing bindings for all TIR primitive - * functions needed by the rewritten module. - */ -IRModule Prepare(IRModule mod, const CompilationConfig& config) { - // Run minimal transforms on module to establish invariants needed by interpreter. - transform::Sequential seq( - {transform::SimplifyInference(), qnn::transform::Legalize(), - // Figure out which devices should be used to execute. - // TODO(mbs): Should ignore all existing annotations when constant folding - transform::PlanDevices(config), - // FuseOps will mark wrapped calls to prim-ops with the 'Primitive' - // attribute. - transform::FuseOps(/*fuse_opt_level=*/0), - // Use ANF to reduce number of cases to handle. - transform::ToANormalForm(), - // eta expand to support constructors in argument position. - transform::EtaExpand( - /*expand_constructor=*/true, /*expand_global_var=*/false), - transform::InferType(), tec::LowerTE(/*module_name=*/"intrp", config)}); - - transform::PassContext pass_ctx = transform::PassContext::Current(); - With ctx(pass_ctx); - mod = seq(mod); - - return mod; -} - -/*! \brief Check if an expression could be changed by \p Prepare. - * - * If not we can evaluate it directly and don't need to bind it into a fresh module. - */ -class NeedsPreparationVisitor : public ExprVisitor { - public: - bool needs_preparation = false; - - private: - void VisitExpr_(const VarNode* vn) override { - // Could be prim. - needs_preparation = true; - } - // ConstantNode ok - // GlobalVarNode ok - void VisitExpr_(const OpNode* op) override { - // Could be prim. - needs_preparation = true; - } - // TupleNode recurse - void VisitExpr_(const FunctionNode* op) override { - // Could be prim. - needs_preparation = true; - } - // CallNode recurse - void VisitExpr_(const LetNode* ln) override { - // May bind prim. - needs_preparation = true; - } - // IfNode recurse - // TupleGetItemNode recurse - // RefCreateNode recurse - // RefReadNode recurse - // RefWriteNode recurse - // ConstructorNode ok - void VisitExpr_(const MatchNode* op) override { - // Needs eta-expansion. - needs_preparation = true; - } -}; - -TypedPackedFunc)> EvalFunction(IRModule mod, Expr expr, Device device, - Target target) { - VLOG_CONTEXT << "EvalFunction"; - VLOG(1) << "evaling module:" << std::endl - << PrettyPrint(mod) << "and expression:" << std::endl - << PrettyPrint(expr); - - ICHECK_EQ(device.device_type, target->GetTargetDeviceType()); - Array raw_targets = {target}; - CompilationConfig config(transform::PassContext::Current(), raw_targets); - - // - // Step 1: Prepare mod. - // - - // If expr is simple enough we can avoid binding it into the module and - // just eval it directly. - NeedsPreparationVisitor visitor; - visitor.VisitExpr(expr); - - Expr expr_to_eval; - IRModule mod_with_expr; // default empty - if (visitor.needs_preparation) { - GlobalVar main; - // Bind expr to a new zero-argument function so it can be prepared along with the module - // (if any). - std::pair mod_and_global; - if (mod.defined()) { - // TODO(mbs): Type inference currently assumes all global functions in modules have - // known result types, and so each global function has it's body types inferred independently - // and in arbitrary order. However, the interpreter may be called with an expression relative - // to a 'main' which has no result type annotation, and that expressions will be bound into a - // fresh global below. Type inference then fails since 'main' has unknown type. We should - // allow inference on mutually recursive global functions. To workaround, infer the type - // of mod now. Obviously that won't work if 'main' itself calls other global functions of - // partial type, but it at least maintains legacy behavior. - transform::PassContext pass_ctx = transform::PassContext::Current(); - With ctx(pass_ctx); - mod = transform::InferType()(mod); - mod_and_global = - IRModule::FromExprInContext(expr, mod->functions, mod->type_definitions, mod->Imports()); - } else { - mod_and_global = IRModule::FromExprInContext(expr); - } - mod_with_expr = mod_and_global.first; - expr_to_eval = mod_and_global.second; - } else { - if (mod.defined()) { - mod_with_expr = mod; - } - // Prepare won't change expr, so we don't need to worry about binding it into a module - // and can just eval it directly. - expr_to_eval = expr; - } - IRModule lowered_mod = Prepare(mod_with_expr, config); - - std::shared_ptr intrp = std::make_shared(lowered_mod, config, device); - - // - // Step 2: Evaluate target function to a closure. - // - ObjectRef object_ref = intrp->Eval(expr_to_eval); - if (auto opt = object_ref.as()) { - InterpreterClosure closure = opt.value(); - ICHECK(closure->func.defined()); - - return TypedPackedFunc)>([intrp, closure](Array args) { - VLOG_CONTEXT << "EvalFunction::Apply"; - VLOG(1) << "evaling closure with " << args.size() << " arguments"; - // - // Step 3: Apply closure to arguments. - // - ICHECK_NOTNULL(intrp); - ICHECK(closure.defined()); - ICHECK(closure->func.defined()); - Array evaled_args; - for (auto arg : args) { - NeedsPreparationVisitor visitor; - visitor.VisitExpr(arg); - ICHECK(!visitor.needs_preparation) - << "attempting to apply closure to expression which needs preparation: " - << PrettyPrint(arg); - evaled_args.push_back(intrp->Eval(arg)); - } - return intrp->Invoke(closure, evaled_args); - }); - } else { - LOG(FATAL) << "expecting expression to have function type and evaluate to a closure"; - } -} - -ObjectRef Eval(Expr expr, Map type_definitions, - std::unordered_set import_set, Device device, Target target, - Map attrs) { - ICHECK_EQ(device.device_type, target->GetTargetDeviceType()); - Array raw_targets = {target}; - CompilationConfig config(transform::PassContext::Current(), raw_targets); - - std::pair mod_and_global = - IRModule::FromExprInContext(expr, /*global_funcs=*/{}, type_definitions, import_set); - - IRModule mod = Prepare(WithAttrs(mod_and_global.first, {attrs}), config); - - Interpreter intrp(mod, config, device); - Expr expr_to_eval = mod->GetGlobalVar(mod_and_global.second->name_hint); - if (expr.as() == nullptr) { - // TODO(mbs): IRModule::FromExpr will implicitly close over the free vars of expr - // unless it is a function, so we must reverse that in the expression to eval. - // This should done more systematically. - expr_to_eval = Call(expr_to_eval, {}); - } - return intrp.Eval(expr_to_eval); -} - -TVM_REGISTER_GLOBAL("relay.backend.EvalFunction").set_body_typed(EvalFunction); - -} // namespace relay -} // namespace tvm diff --git a/src/relay/backend/liveness_analysis.cc b/src/relay/backend/liveness_analysis.cc deleted file mode 100644 index 52db9e6a4c23..000000000000 --- a/src/relay/backend/liveness_analysis.cc +++ /dev/null @@ -1,232 +0,0 @@ -/* - * 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 src/relay/backend/liveness_analysis.cc - * \brief Analysis that collects the live variables before and after each node. - * NOTE: the input IR should be in ANF. - */ - -#include "./liveness_analysis.h" - -#include -#include -#include -#include - -namespace tvm { -namespace relay { -namespace transform { - -using support::Arena; -using VarSet = std::unordered_set; - -ControlFlowGraph ControlFlowGraph::Create(Arena* arena, const Expr& body) { - return Creator().Create(arena, body); -} - -ControlFlowGraph ControlFlowGraph::Creator::Create(Arena* arena, const Expr& body) { - arena_ = arena; - cfg_.entry = BasicBlock::Make(arena); - VisitExpr(body, cfg_.entry); - return std::move(cfg_); -} - -void ControlFlowGraph::Creator::Succ(BasicBlockPtr from, BasicBlockPtr to) { - from->succ.push_back(to); - to->pred.push_back(from); -} - -void ControlFlowGraph::Creator::VisitExpr_(const FunctionNode* f, BasicBlockPtr parent) { - ICHECK(!in_func_) << "nested functions not supported by CFG analysis"; - in_func_ = true; - - // Unwrap the nested function and proceed normally. - if (f->HasNonzeroAttr(attr::kClosure)) { - ICHECK(f->body.as()); - return VisitExpr(Downcast(f->body)->body, parent); - } - - return VisitExpr(f->body, parent); -} - -void ControlFlowGraph::Creator::VisitExpr_(const LetNode* let_node, BasicBlockPtr parent) { - Expr expr = GetRef(let_node); - - while (const LetNode* inner_let_node = expr.as()) { - NodePtr curr_node = Node::Make(arena_, parent, expr); - - ICHECK(!cfg_.let_map.count(expr)); - cfg_.let_map[expr] = curr_node; - cfg_.reverse_post_order.push_back(curr_node); - - // The basic block ends upon reaching control flow, with successor blocks corresponding to the - // control flow branch exprs (true/false in If, and one for each clause in Match). - if (const IfNode* ite = AsIgnoringOnDevice(inner_let_node->value)) { - // Create the basic blocks for each branch and mark them as successors to the current block. - BasicBlockPtr t_block = BasicBlock::Make(arena_); - BasicBlockPtr f_block = BasicBlock::Make(arena_); - Succ(parent, t_block); - Succ(parent, f_block); - - VisitExpr(ite->true_branch, t_block); - VisitExpr(ite->false_branch, f_block); - - // All subsequent bindings (and/or the body expr) will be in a new basic block. - BasicBlockPtr next = BasicBlock::Make(arena_); - Succ(t_block, next); - Succ(f_block, next); - parent = next; - } else if (const MatchNode* match = AsIgnoringOnDevice(inner_let_node->value)) { - // Same as above but one for each pattern. - std::vector clause_blocks; - BasicBlockPtr next = BasicBlock::Make(arena_); - for (const Clause& clause : match->clauses) { - BasicBlockPtr clause_block = BasicBlock::Make(arena_); - Succ(parent, clause_block); - Succ(clause_block, next); - VisitExpr(clause->rhs, clause_block); - } - parent = next; - } - - expr = inner_let_node->body; - } - - VisitExpr(expr, parent); -} - -void ControlFlowGraph::Creator::VisitExpr_(const IfNode* if_node, BasicBlockPtr parent) { - // TODO(@altanh): is there a way of making this work? - LOG(FATAL) << "If expressions should be bound to variables."; -} - -void ControlFlowGraph::Creator::VisitExpr_(const MatchNode* match_node, BasicBlockPtr parent) { - // TODO(@altanh): same as If - LOG(FATAL) << "Match expressions should be bound to variables."; -} - -VarSet VarUseCollector::VisitExpr_(const VarNode* var_node) { return {GetRef(var_node)}; } - -VarSet VarUseCollector::VisitExpr_(const CallNode* call_node) { - VarSet use = VisitExpr(call_node->op); - for (const Expr& arg : call_node->args) { - VarSet arg_use = VisitExpr(arg); - use.insert(arg_use.begin(), arg_use.end()); - } - return use; -} - -VarSet VarUseCollector::VisitExpr_(const TupleNode* tuple_node) { - VarSet use; - for (const Expr& field : tuple_node->fields) { - VarSet field_use = VisitExpr(field); - use.insert(field_use.begin(), field_use.end()); - } - return use; -} - -VarSet VarUseCollector::VisitExpr_(const TupleGetItemNode* get_node) { - return VisitExpr(get_node->tuple); -} - -VarSet VarUseCollector::VisitExpr_(const IfNode* if_node) { return VisitExpr(if_node->cond); } - -VarSet VarUseCollector::VisitExpr_(const MatchNode* match_node) { - return VisitExpr(match_node->data); -} - -UseDefAnalysis UseDefAnalysis::Analyze(const CFG& cfg) { - UseDefAnalysis a; - - // One pass is sufficient. - for (auto it = cfg.reverse_post_order.begin(); it != cfg.reverse_post_order.end(); ++it) { - const CFG::NodePtr& node = *it; - if (const LetNode* let_node = AsIgnoringOnDevice(node->expr)) { - a.use[node] = a.use_collector.VisitExpr(let_node->value); - a.def[node] = let_node->var; - } else { - a.use[node] = a.use_collector.VisitExpr(node->expr); - a.def[node] = Var(); - } - } - - return a; -} - -bool SetEqual(const VarSet& a, const VarSet& b) { - if (a.size() != b.size()) { - return false; - } - for (auto& xa : a) { - if (!b.count(xa)) { - return false; - } - } - return true; -} - -LivenessAnalysis LivenessAnalysis::Analyze(const ControlFlowGraph& cfg, - const UseDefAnalysis& use_def) { - LivenessAnalysis a; - std::list worklist; - - // Initialize worklist to post-order traversal for quick convergence. - worklist.insert(worklist.end(), cfg.reverse_post_order.rbegin(), cfg.reverse_post_order.rend()); - - // See https://lambda.uta.edu/cse5317/notes/node40.html for an overview of the algorithm. - auto visitor = [&](const CFG::NodePtr n) { - VarSet old_in_n = a.live_in[n]; - VarSet old_out_n = a.live_out[n]; - - a.live_in[n] = use_def.use.at(n); - for (const Var& v : a.live_out[n]) { - if (!v.same_as(use_def.def.at(n))) { - a.live_in[n].insert(v); - } - } - - a.live_out[n] = VarSet(); - for (const CFG::NodePtr& s : n->GetSucc()) { - a.live_out[n].insert(a.live_in[s].begin(), a.live_in[s].end()); - } - - if (SetEqual(old_in_n, a.live_in[n]) && SetEqual(old_out_n, a.live_out[n])) { - // No need to update the worklist. - } else { - // Add predecessor nodes back to worklist (no need to add successors, since each node's - // in/out sets are not dependent on its predecessors). - for (const CFG::NodePtr& p : n->GetPred()) { - worklist.push_back(p); - } - } - }; - - while (!worklist.empty()) { - const CFG::NodePtr n = worklist.front(); - worklist.pop_front(); - visitor(n); - } - - return a; -} - -} // namespace transform -} // namespace relay -} // namespace tvm diff --git a/src/relay/backend/liveness_analysis.h b/src/relay/backend/liveness_analysis.h deleted file mode 100644 index 4e9514056b86..000000000000 --- a/src/relay/backend/liveness_analysis.h +++ /dev/null @@ -1,270 +0,0 @@ -/* - * 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 src/relay/backend/liveness_analysis.h - * \brief Analysis that collects the live variables before and after each node. - * NOTE: the input IR should be in ANF. - */ - -#ifndef TVM_RELAY_BACKEND_LIVENESS_ANALYSIS_H_ -#define TVM_RELAY_BACKEND_LIVENESS_ANALYSIS_H_ - -#include - -#include -#include -#include - -#include "../../support/arena.h" -#include "../op/memory/device_copy.h" -#include "../transforms/device_aware_visitors.h" -#include "../transforms/let_list.h" - -namespace tvm { -namespace relay { -namespace transform { - -using support::Arena; -using VarSet = std::unordered_set; - -// TODO(@altanh, @mbs, @mbrookhart): we should do a survey of all "*-flow graphs" in the codebase -// to see what can be deduplicated. - -// TODO(@altanh): support Relay Refs once/if they are supported by the VM. - -/*! - * \brief A representation of an input expression (typically a Function) as a directed graph of - * basic blocks, with edges between basic blocks corresponding to control flow branching. - */ -class ControlFlowGraph { - public: - struct Node; - struct BasicBlock; - - using NodePtr = Node*; - using BasicBlockPtr = BasicBlock*; - - /*! - * \brief A chunk of IR that does not have any control flow branching. At this stage in the IR, - * basic blocks correspond to: - * (1) a sequence of nested Let expressions, where each node in the block corresponds to a - * binding and the last node is either the (non-Let) body or a binding that branches - * (e.g. "let %x = if (%c) { true_block } else { false_block }"). - * (2) an atomic expression representing the target expression of a control flow branch, e.g. - * %v and %u in "let %x = if (%c) { %v } else { %u }". - */ - struct BasicBlock { - // The nodes of the basic block. - std::vector nodes; - // The predecessor basic blocks. - std::vector pred; - // The successor basic blocks. - std::vector succ; - - static BasicBlockPtr Make(support::Arena* arena) { return arena->make(); } - }; - - /*! - * \brief Roughly corresponds to a "statement" in the IR, such as an individual binding in a - * basic block or the "return value" of a block. Each node maps to a single corresponding expr in - * the IR, but the converse is not true (e.g. in the case of variables). - */ - struct Node { - /*! \brief The basic block this node belongs to. */ - BasicBlockPtr parent; - /*! \brief The index into the parent basic block where this node is. */ - size_t index; - /*! \brief The expr this node corresponds to. */ - Expr expr; - - /*! \brief Returns whether or not this node is the first one in the parent basic block. */ - bool IsFirst() const { return index == 0; } - - /*! \brief Returns whether or not this node is the last one in the parent basic block. */ - bool IsLast() const { return index == parent->nodes.size() - 1; } - - /*! \brief Returns the predecessor nodes of this node. */ - std::vector GetPred() const { - std::vector pred; - if (IsFirst()) { - for (const BasicBlockPtr& pred_block : parent->pred) { - pred.push_back(pred_block->nodes.back()); - } - } else { - pred.push_back(parent->nodes[index - 1]); - } - return pred; - } - - /*! \brief Returns the successor nodes of this node. */ - std::vector GetSucc() const { - std::vector succ; - if (IsLast()) { - for (const BasicBlockPtr& succ_block : parent->succ) { - succ.push_back(succ_block->nodes.front()); - } - } else { - succ.push_back(parent->nodes[index + 1]); - } - return succ; - } - - /*! \brief Creates a node with the given expr and appends it to the parent basic block. */ - static NodePtr Make(Arena* arena, BasicBlockPtr parent, Expr expr) { - NodePtr n = arena->make(); - n->parent = parent; - n->expr = expr; - n->index = parent->nodes.size(); - parent->nodes.push_back(n); - return n; - } - }; - - /*! \brief The basic block where control flow begins. */ - BasicBlockPtr entry; - - /*! - * \brief Mapping from Let expressions to their corresponding nodes. Note that Let expressions - * are never shared in ANF (unlike vars), so this is an injection. - */ - std::unordered_map let_map; - - /*! \brief The nodes of the CFG in reverse post order. */ - std::vector reverse_post_order; - - /*! \brief Creates and returns the CFG of the given expression. */ - static ControlFlowGraph Create(Arena* arena, const Expr& body); - - private: - class Creator; -}; - -/*! \brief Helper class for building CFGs. */ -class ControlFlowGraph::Creator : private ExprFunctor { - public: - Creator() {} - - ControlFlowGraph Create(Arena* arena, const Expr& body); - - private: - /*! \brief The arena allocator. */ - Arena* arena_; - - /*! \brief The CFG being built. */ - ControlFlowGraph cfg_; - /*! - * \brief Whether or not we are in a function. CFGs do not support nested functions so this is - * used to error out in such a case. - */ - bool in_func_ = false; - - /*! - * \brief Link \p to as a successor block to \p from. - */ - void Succ(BasicBlockPtr from, BasicBlockPtr to); - -#define DEFAULT_CFG(OP) \ - void VisitExpr_(const OP* op, BasicBlockPtr parent) final { \ - NodePtr n = Node::Make(arena_, parent, GetRef(op)); \ - cfg_.reverse_post_order.push_back(n); \ - } - - void VisitExpr_(const FunctionNode* f, BasicBlockPtr parent) final; - void VisitExpr_(const LetNode* let_node, BasicBlockPtr parent) final; - void VisitExpr_(const IfNode* if_node, BasicBlockPtr parent); - void VisitExpr_(const MatchNode* match_node, BasicBlockPtr parent); - - DEFAULT_CFG(VarNode); - DEFAULT_CFG(GlobalVarNode); - DEFAULT_CFG(ConstantNode); - DEFAULT_CFG(CallNode); - DEFAULT_CFG(OpNode); - DEFAULT_CFG(TupleNode); - DEFAULT_CFG(TupleGetItemNode); -}; - -/*! - * \brief Helper class for collecting the variables used/read by an expression. NOTE: for If exprs, - * only the condition is included (not the branches). Similarly, for Match exprs only the value - * being deconstructed is included. - */ -class VarUseCollector : public ExprFunctor { - public: - VarSet VisitExpr_(const VarNode* var_node); - VarSet VisitExpr_(const CallNode* call_node); - VarSet VisitExpr_(const TupleNode* tuple_node); - VarSet VisitExpr_(const TupleGetItemNode* get_node); - VarSet VisitExpr_(const IfNode* if_node); - VarSet VisitExpr_(const MatchNode* match_node); - - VarSet VisitExpr_(const ConstructorNode* cons_node) { return {}; } - VarSet VisitExpr_(const GlobalVarNode* gvar_node) { return {}; } - VarSet VisitExpr_(const ConstantNode* const_node) { return {}; } - VarSet VisitExpr_(const OpNode* op_node) { return {}; } - VarSet VisitExpr_(const FunctionNode* func_node) { return {}; } -}; - -/*! - * \brief Analysis that collects the variables used and defined at each node. - */ -struct UseDefAnalysis { - using CFG = ControlFlowGraph; - - /*! \brief Mapping of node -> variables used/read by node. */ - std::unordered_map use; - - /*! \brief Mapping of node -> variable defined/written by node. */ - std::unordered_map def; - - VarUseCollector use_collector; - - static UseDefAnalysis Analyze(const CFG& cfg); -}; - -/*! \brief Returns whether \p a and \p b are the same set of vars. */ -bool SetEqual(const VarSet& a, const VarSet& b); - -/*! - * \brief Analysis that collects the live variables before and after each node. - */ -struct LivenessAnalysis { - using CFG = ControlFlowGraph; - - /*! \brief Mapping of node -> set of variables live before node. */ - std::unordered_map live_in; - - /*! \brief Mapping of node -> set of variables live after node. */ - std::unordered_map live_out; - - /*! - * \brief Analyze the input \p cfg (using info from \p use_def). - * - * \param cfg The input control flow graph. - * \param use_def Use-def analysis of \p cfg. - * \return LivenessAnalysis - */ - static LivenessAnalysis Analyze(const ControlFlowGraph& cfg, const UseDefAnalysis& use_def); -}; - -} // namespace transform -} // namespace relay -} // namespace tvm - -#endif // TVM_RELAY_BACKEND_LIVENESS_ANALYSIS_H_ diff --git a/src/relay/backend/name_transforms.cc b/src/relay/backend/name_transforms.cc deleted file mode 100644 index a527d38fb84e..000000000000 --- a/src/relay/backend/name_transforms.cc +++ /dev/null @@ -1,127 +0,0 @@ -/* - * 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 "name_transforms.h" - -#include -#include - -#include -#include - -namespace tvm { -namespace relay { -namespace backend { - -std::string ToCamel(const std::string& original_name) { - std::string camel_name; - camel_name.reserve(original_name.size()); - - bool new_block = true; - for (const char& symbol : original_name) { - if (std::isalpha(symbol)) { - if (new_block) { - camel_name.push_back(std::toupper(symbol)); - new_block = false; - } else { - camel_name.push_back(std::tolower(symbol)); - } - } else if (symbol == '_') { - new_block = true; - } - } - return camel_name; -} - -std::string ToCFunctionStyle(const std::string& original_name) { - ICHECK(!original_name.empty()) << "Function name is empty"; - ICHECK_EQ(original_name.find("TVM"), 0) << "Function not TVM prefixed"; - - int tvm_prefix_length = 3; - std::string function_prefix("TVM"); - - return function_prefix + ToCamel(original_name.substr(tvm_prefix_length)); -} - -std::string ToCVariableStyle(const std::string& original_name) { - ICHECK(!original_name.empty()) << "Variable name is empty"; - ICHECK_EQ(original_name.find("TVM"), 0) << "Variable not TVM prefixed"; - - std::string variable_name; - variable_name.resize(original_name.size()); - - std::transform(original_name.begin(), original_name.end(), variable_name.begin(), ::tolower); - return variable_name; -} - -std::string ToCConstantStyle(const std::string& original_name) { - ICHECK_EQ(original_name.find("TVM"), 0) << "Constant not TVM prefixed"; - std::string constant_name = ToCVariableStyle(original_name); - - std::transform(constant_name.begin(), constant_name.end(), constant_name.begin(), ::toupper); - return constant_name; -} - -std::string ToRustStructStyle(const std::string& original_name) { - ICHECK(!original_name.empty()) << "Struct name is empty"; - return ToCamel(original_name); -} - -std::string ToRustMacroStyle(const std::string& original_name) { - ICHECK(!original_name.empty()) << "Macro name is empty"; - - std::string macro_name; - macro_name.resize(original_name.size()); - - std::transform(original_name.begin(), original_name.end(), macro_name.begin(), ::tolower); - return macro_name; -} - -std::string ToRustConstantStyle(const std::string& original_name) { - ICHECK(!original_name.empty()) << "Constant name is empty"; - std::string constant_name; - constant_name.resize(original_name.size()); - - std::transform(original_name.begin(), original_name.end(), constant_name.begin(), ::toupper); - return constant_name; -} - -std::string CombineNames(const Array& names) { - std::stringstream combine_stream; - ICHECK(!names.empty()) << "Name segments empty"; - - for (const String& name : names) { - ICHECK(!name.empty()) << "Name segment is empty"; - combine_stream << name << "_"; - } - - std::string combined_name = combine_stream.str(); - combined_name.pop_back(); - return combined_name; -} - -TVM_REGISTER_GLOBAL("relay.backend.ToCFunctionStyle").set_body_typed(ToCFunctionStyle); -TVM_REGISTER_GLOBAL("relay.backend.ToCVariableStyle").set_body_typed(ToCVariableStyle); -TVM_REGISTER_GLOBAL("relay.backend.ToCConstantStyle").set_body_typed(ToCConstantStyle); -TVM_REGISTER_GLOBAL("relay.backend.PrefixName").set_body_typed(PrefixName); -TVM_REGISTER_GLOBAL("relay.backend.PrefixGeneratedName").set_body_typed(PrefixGeneratedName); - -} // namespace backend -} // namespace relay -} // namespace tvm diff --git a/src/relay/backend/name_transforms.h b/src/relay/backend/name_transforms.h deleted file mode 100644 index fab518debc63..000000000000 --- a/src/relay/backend/name_transforms.h +++ /dev/null @@ -1,133 +0,0 @@ -/* - * 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 relay/backend/name_transforms.h - * \brief Transformations which are applied on names to generate appropriately named compiler - * artifacts - * - * Example: - * ToCFunctionStyle(PrefixName(CombineNames({"Device", "target", "Invoke"}))) - * // TVMDeviceTargetInvoke - * - * ToCFunctionStyle(PrefixGeneratedName(CombineNames({"model", "Run"}))) - * // TVMGenModelRun - * - * ToCVariableStyle(PrefixName(CombineNames({"Device", "target", "t"}))) - * // tvm_device_target_t - * - * ToCVariableStyle(PrefixGeneratedName(CombineNames({"model", "Devices"}))) - * // tvmgen_model_devices - * - * ToCConstantStyle(PrefixGeneratedName(CombineNames({"model", "Devices"}))) - * // TVMGEN_MODEL_DEVICES - * - */ - -#include -#include -#include - -#include -#include -#include - -#ifndef TVM_RELAY_BACKEND_NAME_TRANSFORMS_H_ -#define TVM_RELAY_BACKEND_NAME_TRANSFORMS_H_ - -namespace tvm { -namespace relay { -namespace backend { - -/*! - * \brief Transform a name to the C variable style assuming it is - * appropriately constructed using the prefixing functions - * \param original_name Original name - * \return Transformed function in the C function style - */ -std::string ToCFunctionStyle(const std::string& original_name); - -/*! - * \brief Transform a name to the C variable style assuming it is - * appropriately constructed using the prefixing functions - * \param name Original name - * \return Transformed function in the C variable style - */ -std::string ToCVariableStyle(const std::string& original_name); - -/*! - * \brief Transform a name to the C constant style assuming it is - * appropriately constructed using the prefixing functions - * \param name Original name - * \return Transformed function in the C constant style - */ -std::string ToCConstantStyle(const std::string& original_name); - -/*! - * \brief Transform a name to the Rust struct style assuming it is - * appropriately constructed using the combining functions - * \param name Original name - * \return Transformed function in the Rust struct style - */ -std::string ToRustStructStyle(const std::string& original_name); - -/*! - * \brief Transform a name to the Rust macro style assuming it is - * appropriately constructed using the combining functions - * \param name Original name - * \return Transformed function in the Rust macro style - */ -std::string ToRustMacroStyle(const std::string& original_name); - -/*! - * \brief Transform a name to the Rust constant style assuming it is - * appropriately constructed using the combining functions - * \param name Original name - * \return Transformed function in the Rust constant style - */ -std::string ToRustConstantStyle(const std::string& original_name); - -/*! - * \brief Combine names together for use as a generated name - * \param names Vector of strings to combine - * \return Combined together names - */ -std::string CombineNames(const Array& names); - -/*! - * \brief Apply TVM-specific prefix to a name - * \param names Vector of names to combine to form a combined name - * \return Name with prefix applied or prefix-only if no name passed - */ -inline std::string PrefixName(const Array& names) { return "TVM_" + CombineNames(names); } - -/*! - * \brief Apply generated TVM-specific prefix to a name - * \param names Vector of names to combine to form a combined name - * \return Name with prefix applied or prefix-only if no name passed - */ -inline std::string PrefixGeneratedName(const Array& names) { - return "TVMGen_" + CombineNames(names); -} - -} // namespace backend -} // namespace relay -} // namespace tvm - -#endif // TVM_RELAY_BACKEND_NAME_TRANSFORMS_H_ diff --git a/src/relay/backend/param_dict.cc b/src/relay/backend/param_dict.cc deleted file mode 100644 index bb0fad9142c1..000000000000 --- a/src/relay/backend/param_dict.cc +++ /dev/null @@ -1,54 +0,0 @@ -/* - * 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 param_dict.cc - * \brief Implementation and registration of parameter dictionary - * serializing/deserializing functions. - */ -#include "param_dict.h" - -#include -#include - -#include -#include -#include - -#include "../../runtime/file_utils.h" - -namespace tvm { -namespace relay { - -using namespace runtime; - -TVM_REGISTER_GLOBAL("tvm.relay._save_param_dict") - .set_body_typed([](const Map& params) { - std::string s = ::tvm::runtime::SaveParams(params); - // copy return array so it is owned by the ret value - TVMRetValue rv; - rv = TVMByteArray{s.data(), s.size()}; - return rv; - }); -TVM_REGISTER_GLOBAL("tvm.relay._load_param_dict").set_body_typed([](const String& s) { - return ::tvm::runtime::LoadParams(s); -}); - -} // namespace relay -} // namespace tvm diff --git a/src/relay/backend/param_dict.h b/src/relay/backend/param_dict.h deleted file mode 100644 index 96e17a9da07b..000000000000 --- a/src/relay/backend/param_dict.h +++ /dev/null @@ -1,38 +0,0 @@ -/* - * 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 param_dict.h - * \brief Definitions for serializing and deserializing parameter dictionaries. - */ -#ifndef TVM_RELAY_BACKEND_PARAM_DICT_H_ -#define TVM_RELAY_BACKEND_PARAM_DICT_H_ - -#include -#include -#include -#include - -#include - -namespace tvm { -namespace relay {} // namespace relay -} // namespace tvm - -#endif // TVM_RELAY_BACKEND_PARAM_DICT_H_ diff --git a/src/relay/backend/runtime.cc b/src/relay/backend/runtime.cc deleted file mode 100644 index 0534298ea44d..000000000000 --- a/src/relay/backend/runtime.cc +++ /dev/null @@ -1,106 +0,0 @@ -/* - * 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 src/relay/backend/runtime.cc - * \brief Runtime Registry - */ - -#include - -#include "../../node/attr_registry.h" - -namespace tvm { -namespace relay { - -TVM_REGISTER_NODE_TYPE(RuntimeNode); - -TVM_STATIC_IR_FUNCTOR(ReprPrinter, vtable) - .set_dispatch([](const ObjectRef& obj, ReprPrinter* p) { - const Runtime& runtime = Downcast(obj); - p->stream << runtime->name; - }); - -/********** Registry-related code **********/ - -using RuntimeRegistry = AttrRegistry; - -Runtime Runtime::Create(String name, Map attrs) { - const RuntimeRegEntry* reg = RuntimeRegistry::Global()->Get(name); - if (reg == nullptr) { - throw Error("Runtime \"" + name + "\" is not defined"); - } - - for (const auto& kv : attrs) { - if (!reg->key2vtype_.count(kv.first)) { - throw Error("Attribute \"" + kv.first + "\" is not available on this Runtime"); - } - std::string expected_type = reg->key2vtype_.at(kv.first).type_key; - std::string actual_type = kv.second->GetTypeKey(); - if (expected_type != actual_type) { - throw Error("Attribute \"" + kv.first + "\" should have type \"" + expected_type + - "\" but instead found \"" + actual_type + "\""); - } - } - - for (const auto& kv : reg->key2default_) { - if (!attrs.count(kv.first)) { - attrs.Set(kv.first, kv.second); - } - } - - return Runtime(name, DictAttrs(attrs)); -} - -Array Runtime::ListRuntimes() { return RuntimeRegistry::Global()->ListAllNames(); } - -Map Runtime::ListRuntimeOptions(const String& name) { - Map options; - const RuntimeRegEntry* reg = RuntimeRegistry::Global()->Get(name); - if (reg == nullptr) { - throw Error("Runtime \"" + name + "\" is not defined"); - } - for (const auto& kv : reg->key2vtype_) { - options.Set(kv.first, kv.second.type_key); - } - return options; -} - -RuntimeRegEntry& RuntimeRegEntry::RegisterOrGet(const String& name) { - return RuntimeRegistry::Global()->RegisterOrGet(name); -} - -/********** Register Runtimes and options **********/ - -TVM_REGISTER_RUNTIME(kTvmRuntimeCrt).add_attr_option("system-lib"); - -TVM_REGISTER_RUNTIME(kTvmRuntimeCpp).add_attr_option("system-lib"); - -/********** Registry **********/ - -TVM_REGISTER_GLOBAL("relay.backend.CreateRuntime").set_body_typed(Runtime::Create); -TVM_REGISTER_GLOBAL("relay.backend.GetRuntimeAttrs").set_body_typed([](const Runtime& runtime) { - return runtime->attrs->dict; -}); - -TVM_REGISTER_GLOBAL("relay.backend.ListRuntimes").set_body_typed(Runtime::ListRuntimes); -TVM_REGISTER_GLOBAL("relay.backend.ListRuntimeOptions").set_body_typed(Runtime::ListRuntimeOptions); - -} // namespace relay -} // namespace tvm diff --git a/src/relay/backend/task_extraction.cc b/src/relay/backend/task_extraction.cc deleted file mode 100644 index 6ac7a99d3509..000000000000 --- a/src/relay/backend/task_extraction.cc +++ /dev/null @@ -1,143 +0,0 @@ -/* - * 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 "../../meta_schedule/module_equality.h" -#include "../../te/operation/create_primfunc.h" -#include "./te_compiler_cache.h" -#include "./utils.h" - -namespace tvm { -namespace relay { -namespace backend { - -class OpCounter : public ExprVisitor { - public: - static size_t GetOpCount(relay::Function func) { - OpCounter counter; - counter(func->body); - return counter.count; - } - - private: - void VisitExpr_(const CallNode* call) final { - if (call->op->IsInstance()) { - ++count; - } - ExprVisitor::VisitExpr_(call); - } - - size_t count{0}; -}; - -Array ExtractTask(IRModule mod, Target target, - Map params, - String mod_eq_name) { - using meta_schedule::ExtractedTask; - using meta_schedule::ModuleEqual; - using meta_schedule::ModuleHash; - backend::BindParamsInModule(mod, params); - // is_vm=true for backward compatibility - Array pass_seqs = relay::backend::GetPassPrefix(/*is_homogenous=*/true, /*is_vm=*/true); - pass_seqs.push_back(transform::FuseOps()); - - mod = transform::Sequential(pass_seqs)(std::move(mod)); - - std::vector tasks; - - auto mod_eq = meta_schedule::ModuleEquality::Create(mod_eq_name); - - std::unordered_map cache( - /*bucket_count*/ 0, ModuleHash(*mod_eq), ModuleEqual(*mod_eq)); - - std::vector> lower_results; - - NameSupply constant_name_supply; - - PostOrderVisit(mod->Lookup("main"), [&](const Expr& exp) { - if (exp->IsInstance()) { - Function relay_func = Downcast(exp); - if (!relay_func->HasNonzeroAttr(attr::kPrimitive)) { - return; - } - - auto [f, fused_name] = tec::LowerToPrimFunc(relay_func, target, constant_name_supply); - if (f) { - IRModule tir_mod = PrimFuncToIRModule(f.value()); - lower_results.push_back(std::make_tuple(fused_name, relay_func, tir_mod)); - } - } - }); - - std::vector indices(lower_results.size()); - std::iota(indices.begin(), indices.end(), 0); - - if (mod_eq_name == "anchor-block") { - std::vector op_counts(lower_results.size()); - for (size_t i = 0; i < op_counts.size(); ++i) { - op_counts[i] = OpCounter::GetOpCount(std::get<1>(lower_results[i])); - } - - // When anchor-block based equality is used, tuning tasks "nn_conv2d_add_nn_relu" and - // "nn_conv2d_add_add_nn_relu", for example, can be identified as equal. Thus, one of - // them will be filtered by the cache below. - // - // To make sure that we tune "nn_conv2d_add_nn_relu" and not "nn_conv2d_add_add_nn_relu", - // we sort the TE lowering results based on the number of relay ops. This way, - // "nn_conv2d_add_nn_relu" will be added to the cache first, and "nn_conv2d_add_add_nn_relu" - // will be filtered. - std::sort(indices.begin(), indices.end(), - [&op_counts](int i1, int i2) { return op_counts[i1] < op_counts[i2]; }); - } - - for (auto i : indices) { - const auto& [fused_name, relay_func, tir_mod] = lower_results[i]; - auto it = cache.find(tir_mod); - if (it != cache.end()) { - it->second->weight += 1; - continue; - } - // Note that the cache is key-ed on the tir mod, rather than the relay mod - IRModule relay_mod({{GlobalVar(fused_name), relay_func}}); - ExtractedTask task(fused_name, relay_mod, target, {tir_mod}, 1); - tasks.push_back(task); - cache.emplace(tir_mod, task); - } - - // Tasks are extracted via post order visit, return the reversed list. - std::reverse(tasks.begin(), tasks.end()); - NameSupply name_supply; - for (ExtractedTask task : tasks) { - task->task_name = name_supply->FreshName(task->task_name); - } - return tasks; -} - -TVM_REGISTER_GLOBAL("relay.backend.MetaScheduleExtractTask").set_body_typed(ExtractTask); - -} // namespace backend -} // namespace relay -} // namespace tvm diff --git a/src/relay/backend/te_compiler.cc b/src/relay/backend/te_compiler.cc deleted file mode 100644 index eab4837ba882..000000000000 --- a/src/relay/backend/te_compiler.cc +++ /dev/null @@ -1,1322 +0,0 @@ -/* - * 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 relay/backend/te_compiler.cc - * \brief Manages the transition from Relay "Primitive" \p Functions to TIR \p PrimFuncs. Also - * handles invocation of external codegen. - * - * \p LowerTEPass handles the following (as a monolithic blob of code): - * - * - Most importantly, any function with the "Primitive" attribute is first converted to TE by - * \p LowerToTECompute (see te_compiler_cache.cc) using each operator's 'compute' function. - * The TE is then 'scheduled' to TIR using the 'anchor' operator's 'schedule' function. Both - * of those functions come from the \p OpStrategy returned by the Python - * 'relay.backend.lower_call' function (see te_compiler.py). - * The TIR is packed as a \p PrimFunc and introduced as a new global function. Calls to the - * original "Primitive" function are then rewritten to the form: - * \code - * call_lowered(@new_global, (... original args...), attributes) - * \endcode - * - * - The above "Primitive" function can appear: - * - As a global function - * - As a let-bound function - * - As an inline function, ie the 'op' of calls. - * In all three cases it is possible for the same "Primitive" function to be called multiple - * times, and that sharing must be respected. - * - * - "Primitive" functions must have a "global_symbol" attribute matching their desired or - * existing global name. Care is taken to ensure GlobalVars with the same name are shared. - * - * - It is possible for multiple structurally equal "Primitive" functions to appear in the same - * \p IRModule. Only one implementation should be generated, and all calls should share that - * implementation. - * - * - When later converting to DPS (see memory_alloc.cc) we must handle functions who's result - * tensor shapes depend at runtime on the input tensor shapes and/or data. - * - That dependency is first described in TE form (see \p MakeShapeFunc in - * te_compiler_cache.cc), then scheduled to yield a 'dynamic shape function' \p PrimFunc. - * This relies on each operator's "FShapeFunc" and "TShapeDataDependent" attributes. - * Since shapes are rank-1 tensors everything can be reflected back down into the regular - * TE/TIR forms. - * - Then the call_lowered attributes must record everything about the dynamic shape function - * later needed by memory_alloc.cc. We call this 'cross linking' the call with the shape - * function. - * - * - Two external codegen mechanisms are supported, both triggered by "Primitive" functions which - * also have a "Compiler" attribute bound to $compiler: - * - Function-at-a-time (old style): The primitive function is passed to the function - * registered as 'relay.ext.$compiler'. The function returns a runtime::Module which - * should return true for \p ImplementsFunction for the function's global name. That - * module is added to the IRModule's "external_mods" attributes. - * - IRModule-at-a-item (new style): The \p RelayToTIRTargetHook sub-pass looks for - * $compiler names which correspond to TargetKind names with a \p RelayToTIR attribute. - * The \p Pass bound to that attribute is run, and each such 'custom' pass can do what - * it likes, including replacing Functions with PrimFuncs, or adding new runtime::Modules - * to the IRModule's "external_mods" attribute. - * - * - Calls to functions added by external codegen are also rewritten to call_lowered form, and - * may also require cross-linking to dynamic shape functions. However, since the functions - * are/will be implemented by a runtime::Module all the Relay type information is no longer - * available. So the Relay definitions for these "Primitive" "Compiler" functions are retained - * in the \p IRModule, but marked with the "Extern" attribute to signal the function is now - * just for carrying metadata. - * - * - Some operators are handled specially: - * - 'reshape', since it's a no-op on the underlying tensor buffer, and this is handled by - * condition tests in many passes. - * - 'debug', since it's intercepted differently depending on runtimes. - * - * TODO(mbs): This desperately deserves a refactor to separate all these concerns. See Relax. - */ - -#include "./te_compiler.h" - -#include -#include -#include -#include -#include -#include -#include -#include -#include -#include -#include -#include -#include -#include -#include -#include -#include - -#include -#include -#include -#include -#include -#include -#include - -#include "../op/annotation/annotation.h" -#include "../op/call/call.h" -#include "../op/memory/device_copy.h" -#include "../transforms/device_aware_visitors.h" -#include "./te_compiler_cache.h" -#include "./utils.h" - -namespace tvm { -namespace relay { -// TODO(@jroesch, @csullivan): declare directly elsewhere -backend::StaticMemoryPlan GraphPlanMemory(const Function& func); - -namespace tec { - -using namespace tvm::relay::transform; - -TVM_REGISTER_OBJECT_TYPE(TECompilerNode); - -class TECompilerImpl : public TECompilerNode { - public: - explicit TECompilerImpl(Optional opt_mod, Optional opt_mod_name) - : global_var_supply_(GlobalVarSupply(NameSupply(opt_mod_name.value_or("")))) { - // Make sure we don't collide with any existing globals in the module. - if (opt_mod) { - for (const auto& kv : opt_mod.value()->functions) { - global_var_supply_->name_supply_->ReserveName(kv.first->name_hint, false); - } - } - } - - // Lower the function. - CachedFunc Lower(const CCacheKey& key) { - return LowerInternal(key, global_var_supply_)->cached_func; - } - - // TODO(gigiblender): Only to be called by the global TE compiler. - // Remove this when the global TE compiler is removed. - CachedFunc Lower(const CCacheKey& key, const String mod_name) { - global_var_supply_->name_supply_->prefix_ = mod_name; - return LowerInternal(key, global_var_supply_)->cached_func; - } - - // For now, build one module per function. - PackedFunc JIT(const CCacheKey& key) final { - CCacheValue value = LowerInternal(key, GlobalVarSupply()); - if (value->packed_func != nullptr) { - return value->packed_func; - } - auto m = build(value->cached_func->funcs, key->target, Target(nullptr)); - value->packed_func = m.GetFunction(value->cached_func->prim_fn_var->name_hint); - return value->packed_func; - } - - CachedFunc LowerShapeFunc(const CCacheKey& key) final { - return LowerShapeFuncInternal(key)->cached_func; - } - - IRModule GetLoweredFunctions() { - VLOG(1) << "GetLoweredFunctions"; - IRModule mod; - // Extract lowered functions from the cache - for (const auto& it : cache_) { - auto source_func = it.first; - auto lowered_func = it.second; - - IRModule lowered_mod = lowered_func->cached_func->funcs; - - // Annotate functions with their target and put them in the return module - for (const auto& kv : lowered_mod->functions) { - const GlobalVar& var = kv.first; - const BaseFunc& func = kv.second; - - // Only add functions that are not external functions - if (!func->GetAttr(attr::kCompiler).defined()) { - ICHECK(func->IsInstance()) - << "Expected all functions that are not external to be PrimFuncs, but found:" - << std::endl - << PrettyPrint(func); - const tir::PrimFunc& prim_func = Downcast(func); - mod->Update(var, WithAttr(prim_func, tvm::attr::kTarget, source_func->target)); - } - } - } - - // Extract lowered dynamic shape functions from the shape cache - for (const auto& it : shape_func_cache_) { - auto source_func = it.first; - auto lowered_func = it.second; - auto target = source_func->target; - IRModule lowered_mod = lowered_func->cached_func->funcs; - - // Annotate functions with their target and put them in the return module - for (auto kv : lowered_mod->functions) { - const GlobalVar& var = kv.first; - const BaseFunc& func = kv.second; - const tir::PrimFunc& prim_func = Downcast(func); - mod->Update(var, WithAttr(prim_func, tvm::attr::kTarget, source_func->target)); - } - } - - return mod; - } - - void AddExterns(IRModule module) { - // Everything tagged with "Compiler" has been compiled, so remove those definitions. - std::vector to_be_deleted; - for (const auto& kv : module->functions) { - if (kv.second->GetAttr(attr::kCompiler).defined()) { - to_be_deleted.push_back(kv.first); - } - } - for (const auto& global_var : to_be_deleted) { - VLOG(1) << "Removing definition for external codegened '" << global_var->name_hint << "'"; - module->Remove(global_var); - } - // HOWEVER we still need a Relay definition to go with those now external functions, so - // retrieve them from the cache and mark them with "ExternalSymbol". - for (const auto& kv1 : cache_) { - auto src_func = kv1.first->source_func; - ICHECK(src_func.defined()); - if (src_func->GetAttr(attr::kCompiler).defined()) { - for (const auto& kv2 : kv1.second->cached_func->funcs->functions) { - if (const auto* function_node = kv2.second.as()) { - // Abandon the existing function annotations. - - // Unfortunately, Optional() is indistinguishable from - // NullValue(), and DictAttrs() is nullptr, so to erase the attributes, we - // need pass in DictAttrs()), which is a DictAttrs containing no - // attributes. - Function function = - WithFields(GetRef(function_node), function_node->params, - function_node->body, function_node->ret_type, function_node->type_params, - /* erase attributes */ DictAttrs(Map())); - // Mark function as 'extern'. - function = WithAttr(std::move(function), attr::kExtern, Integer(1)); - module->Add(kv2.first, function); - } - } - } - } - } - - Array LowerExternalFunctions() { - Array ret; - std::vector cached_ext_funcs; - - for (const auto& it : cache_) { - auto src_func = it.first->source_func; - ICHECK(src_func.defined()); - Optional opt_compiler = src_func->GetAttr(attr::kCompiler); - if (opt_compiler.defined()) { - Optional opt_symbol_name = src_func->GetAttr(tvm::attr::kGlobalSymbol); - ICHECK(opt_symbol_name.defined()) << "No external symbol is set for:" << std::endl - << PrettyPrint(src_func); - VLOG(1) << "using external codegen '" << opt_compiler.value() << "' for name '" - << opt_symbol_name.value() << "' and function:" << std::endl - << PrettyPrint(src_func); - cached_ext_funcs.push_back(it.first); - - std::string ext_name = "relay.ext." + opt_compiler.value(); - auto pf = tvm::runtime::Registry::Get(ext_name); - ICHECK(pf) << "Failed to find the external codegen tool for " << ext_name; - // No need to keep compiler attribute at this point, functions have been - // extracted for specific codegen. - src_func = WithAttr(std::move(src_func), attr::kCompiler, NullValue()); - VLOG_CONTEXT << opt_compiler.value(); - With with_target(it.first->target); - runtime::Module ext_mod = (*pf)(src_func); - if (ext_mod.defined()) { - // TODO(mbs): Can this be an ICHECKs? - if (!ext_mod->ImplementsFunction(opt_symbol_name.value())) { - VLOG(1) << "Note that the external codegen for '" << opt_compiler.value() - << "' returned a runtime module which does not appear to implement '" - << opt_symbol_name.value() << "'"; - } - ret.push_back(ext_mod); - } else { - // It is valid for the external codegen function to return null: - // - Unit tests can use it. - // - The true compilation may have already been handled by a RelayToTIR custom pass - // on the Target's kind. The original Relay functions will be left in place so - // that we can capture that their function names are now externally defined. - VLOG(1) << "Note that no external runtime module was generated by external codegen '" - << opt_compiler.value() << "'"; - } - } - } - - // No need to cache external functions as we collected them all to create - // external runtime modules. - for (const auto& it : cached_ext_funcs) { - cache_.erase(it); - } - return ret; - } - - Map GetDeviceContexts() { return device_contexts_; } - void SetDeviceContexts(const Map& device_contexts) { - device_contexts_ = device_contexts; - } - - void Clear() final { cache_.clear(); } - - // List all items in the cache. - Array ListItems() { - std::lock_guard lock(mutex_); - Array items; - for (auto& kv : cache_) { - items.push_back(kv.first); - items.push_back(kv.second); - } - return items; - } - - /*! - * \brief Get the cache key of the function that is being lowered currently - * \return the cache key - */ - CCacheKey GetCurrentCCacheKey() { return cur_ccache_key_; } - - private: - // implement lowered func - CCacheValue LowerInternal(const CCacheKey& key, GlobalVarSupply global_var_supply) { - VLOG(1) << "lowering:" << std::endl - << PrettyPrint(key->source_func) << std::endl - << "for target:" << std::endl - << key->target->ToDebugString(); - std::lock_guard lock(mutex_); - CCacheValue value; - auto it = cache_.find(key); - if (it != cache_.end()) { - VLOG(1) << "already lowered to name:" << std::endl - << PrettyPrint(it->second->cached_func->prim_fn_var); - it->second->use_count += 1; - if (it->second->cached_func.defined()) return it->second; - value = it->second; - } else { - value = CCacheValue(make_object()); - value->use_count = 1; - cache_[key] = value; - } - cur_ccache_key_ = key; - - Optional opt_compiler = key->source_func->GetAttr(attr::kCompiler); - if (opt_compiler.defined()) { - // Don't compile now since we don't have anywhere to put the resulting runtime module. - // Instead place the original definition in the cache and wait for LowerExternalFunctions. - IRModule ir_module({}, {}); - Optional opt_global_symbol = - key->source_func->GetAttr(tvm::attr::kGlobalSymbol); - ICHECK(opt_global_symbol.defined()) << "External function has not been attached a name yet."; - // Note that the source_func may already be bound to a global function in the module - // we are compiling, in which case we should not attempt to make its name unique w.r.t. - // the module's globals. Furthermore, the external codegen tool must bind the compiled - // function to the "global_symbol" attribute on the source_func. So do not use GetUniqueName - // here. - auto global_var = global_var_supply->UniqueGlobalFor(opt_global_symbol.value(), false); - global_var->checked_type_ = key->source_func->checked_type(); - ir_module->Add(global_var, key->source_func); - value->cached_func = CachedFunc(key->target, global_var, {}, {}, te::Schedule{nullptr}, - tir::PrimFunc{nullptr}, {}, ir_module); - // Collect these here as it's removed in LowerExternalFunctions() - device_contexts_.Set(value->cached_func->prim_fn_var, opt_compiler.value()); - VLOG(1) << "preparing to use external codegen '" << opt_compiler.value() - << "' with name:" << std::endl - << PrettyPrint(value->cached_func->prim_fn_var) << std::endl - << "and definitions:" << std::endl - << PrettyPrint(value->cached_func->funcs); - return value; - } - - // Enforce use the target. - With target_scope(key->target); - - ICHECK(!value->cached_func.defined()); - value->cached_func = - PrimFuncFor(key->source_func, key->target, global_var_supply, constant_name_supply_); - - if (value->cached_func->prim_func.defined()) { - VLOG(1) << "Lowering PrimFunc"; - IRModule lowered = tvm::LowerPrimFunc(value->cached_func->prim_func.value(), - value->cached_func->prim_fn_var->name_hint, false); - ICHECK_EQ(lowered->functions.size(), 1); - for (const auto& kv : lowered->functions) { - value->cached_func->funcs->Add(value->cached_func->prim_fn_var, kv.second); - } - } else { - // NOTE: array will copy on write. - Array all_args = Array(value->cached_func->inputs); - for (te::Tensor arg : value->cached_func->outputs) { - all_args.push_back(arg); - } - Array all_consts; - for (auto kv : value->cached_func->constant_tensors) { - all_args.push_back(kv.second); - all_consts.push_back(kv.first->data); - } - // lower the function - std::unordered_map binds; - - // If we have memory scopes, need to create tir::Buffer knowing this info - size_t i = 0; // for corresponding from tensor array - for (Var param : key->source_func->params) { - if (!param->virtual_device()->memory_scope.empty()) { - for (const auto& ttype : FlattenTupleType(param->checked_type())) { - te::Tensor x_ref = value->cached_func->inputs[i]; - // verification if we have synced params and tensors - ICHECK(ttype->dtype == x_ref->dtype && ttype->shape.size() == x_ref->shape.size()) - << "function parameter does not correspond to prepared tensor"; - binds[x_ref] = - tir::BufferWithOffsetAlignment(x_ref->shape, x_ref->dtype, x_ref->op->name, -1, 0, - false, param->virtual_device()->memory_scope); - } - } - i++; - } - if (key->virtual_device != VirtualDevice::FullyUnconstrained() && - !key->virtual_device->memory_scope.empty() && - key->virtual_device->memory_scope != "global") { - ICHECK(value->cached_func->outputs.size() == 1) - << "Expect only one output for defined memory scope"; - te::Tensor x_ref = value->cached_func->outputs[0]; - binds[x_ref] = - tir::BufferWithOffsetAlignment(x_ref->shape, x_ref->dtype, x_ref->op->name, -1, 0, - false, key->virtual_device->memory_scope); - } - auto func_name = value->cached_func->prim_fn_var->name_hint; - VLOG(1) << "scheduling"; - IRModule scheduled_module = tvm::LowerSchedule(value->cached_func->schedule, all_args, - func_name, binds, global_var_supply); - scheduled_module->Update(tir::transform::BindParams(all_consts)(scheduled_module)); - for (const auto& kv : scheduled_module->functions) { - GlobalVar global_var = kv.first; - auto func = kv.second; - // Propagate the structural hash of the relay function to the tir - // function so associations can be made between the two. - Optional hash = key->source_func->attrs.GetAttr("hash"); - if (hash) { - func = WithAttrs(Downcast(func), {{String("hash"), hash.value()}}); - } - value->cached_func->funcs->Add(global_var, func); - } - ICHECK(value->cached_func->funcs->Lookup(value->cached_func->prim_fn_var) - .as()); - } - VLOG(1) << "lowered to name:" << std::endl - << PrettyPrint(value->cached_func->prim_fn_var) << std::endl - << "with definitions:" << std::endl - << PrettyPrint(value->cached_func->funcs); - - return value; - } - - // implement lowered shape func - CCacheValue LowerShapeFuncInternal(const CCacheKey& key) { - VLOG(1) << "lowering dynamic shape function for:" << std::endl - << PrettyPrint(key->source_func) << std::endl - << "for target:" << std::endl - << key->target->ToDebugString(); - std::lock_guard lock(mutex_); - CCacheValue value; - auto it = shape_func_cache_.find(key); - if (it != shape_func_cache_.end()) { - it->second->use_count += 1; - if (it->second->cached_func.defined()) return it->second; - value = it->second; - } else { - value = CCacheValue(make_object()); - value->use_count = 0; - shape_func_cache_[key] = value; - } - // Enforce use the target. - With target_scope(key->target); - - ICHECK(!value->cached_func.defined()); - - using tvm::transform::PassContext; - With fresh_pass_ctx_scope(PassContext::Create()); - value->cached_func = ShapeFuncFor(key->source_func, key->target, global_var_supply_); - - ICHECK( - value->cached_func->funcs->Lookup(value->cached_func->prim_fn_var).as()); - - VLOG(1) << "lowered to name:" << std::endl - << PrettyPrint(value->cached_func->prim_fn_var) << std::endl - << "with definitions:" << std::endl - << PrettyPrint(value->cached_func->funcs); - return value; - } - - Map GetOpWeights() const { - Map weights; - for (const auto& kv : cache_) { - auto value = kv.second; - auto name = value->cached_func->prim_fn_var->name_hint; - weights.Set(name, value->use_count); - } - return weights; - } - - // TODO(mbs): Hold the output module here and reduce the cache_ to just be from - // Function to GlobalVar. - - /*! \brief compiler cache lock*/ - std::mutex mutex_; - /*! \brief internal GlobalVarSupply to get unique GlobalVars */ - GlobalVarSupply global_var_supply_; - /*! \brief A NameSupply object for assigning unique names to constants, across different - * invocations of PrimFuncFor. */ - NameSupply constant_name_supply_; - /*! \brief internal compiler cache */ - std::unordered_map cache_; - /*! \brief internal compiler cache for shape funcs */ - std::unordered_map shape_func_cache_; - /*! \brief the cache key of the function that is being lowered currently*/ - CCacheKey cur_ccache_key_; - /*! \brief Map of GlobalVar to C Device API context names */ - Map device_contexts_; -}; - -TECompiler::TECompiler(Optional opt_mod, Optional mod_name) { - auto object = make_object(std::move(opt_mod), std::move(mod_name)); - data_ = object; -} - -/*! \brief The global TE compiler */ -// TODO(mbs): To be terminated with extreme prejudice. -TECompiler& TECompiler::Global() { - static TECompiler* inst = - new TECompiler(make_object(Optional(), Optional())); - return *inst; -} -TVM_REGISTER_PASS_CONFIG_OPTION("relay.backend.use_auto_scheduler", Bool); -TVM_REGISTER_PASS_CONFIG_OPTION("relay.backend.use_meta_schedule", Bool); -TVM_REGISTER_PASS_CONFIG_OPTION("relay.backend.use_meta_schedule_dispatch", Integer); -TVM_REGISTER_PASS_CONFIG_OPTION("relay.backend.tir_converter", String); - -TVM_REGISTER_GLOBAL("relay.backend._TECompilerGlobal").set_body_typed([]() { - return TECompiler::Global(); -}); - -TVM_REGISTER_GLOBAL("relay.backend._make_CCacheKey") - .set_body_typed([](Function source_func, Target target) { - return CCacheKey(source_func, target); - }); - -TVM_REGISTER_GLOBAL("relay.backend._make_LoweredOutput") - .set_body_typed([](tvm::Array outputs, OpImplementation impl) { - return LoweredOutput(outputs, impl); - }); - -TVM_REGISTER_GLOBAL("relay.backend._TECompilerClear").set_body_typed([](TECompiler self) { - self->Clear(); -}); - -TVM_REGISTER_GLOBAL("relay.backend._TECompilerLower") - .set_body_typed([](TECompiler self, CCacheKey key, const String mod_name) { - return self->Lower(key, mod_name); - }); - -TVM_REGISTER_GLOBAL("relay.backend._TECompilerJIT") - .set_body_typed([](TECompiler self, CCacheKey key) { return self->JIT(key); }); - -TVM_REGISTER_GLOBAL("relay.backend._TECompilerListItems").set_body_typed([](TECompiler self) { - TECompilerImpl* ptr = dynamic_cast(self.operator->()); - ICHECK(ptr != nullptr); - return ptr->ListItems(); -}); - -using AnalysisRemapping = std::unordered_map; - -/*! - * \brief Rewrites call expressions to Relay Functions marked as "primitive" - * to calls to the corresponding TIR PrimFunc for the appropriate target. - * - * \code - * %0 = fn(...) { prim_op(...) } OR let %p = fn(...) { prim_op(...) } - * ... %0(...) ... ... %p(...) ... - * ==> - * def @q(..., target=) { } - * ... @q(...) ... - * \endcode - * - * Requires FuseOps, ToANormalForm, EtaExpand and InferType to have run. - * - * FuseOps is needed to identify and lift all prim op calls: - * \code - * ... prim_op(...) ... - * ==> - * %0 = fn(...) { prim_op(...) } - * ... %0(...) ... - * \endcode - * - * ToANormalForm is needed so we only need to consider vars and function literals as the call - * target. - * - * EtaExpand is needed to ensures all calls to primitives are direct: - * \code - * let %p1 = fn(...) { prim_op1(...) } - * let %p2 = fn(...) { prim_op2(...) } - * let %p = if (...) { %p1 } else { %p2 } - * ... %p(...) ... - * ==> - * let %p1 = fn(...) { prim_op1(...) } - * let %p2 = fn(...) { prim_op2(...) } - * let %p = fn(...) { if (...) { %p1(...) } else { %p2(...) } } - * ... %p(...) ... - * \endcode - */ -class LowerTensorExprMutator : public DeviceAwareExprMutator { - public: - LowerTensorExprMutator(IRModule module, ProcessFn process_fn, CompilationConfig config, - TECompiler compiler) - : DeviceAwareExprMutator(module), - module_(std::move(module)), - process_fn_(std::move(process_fn)), - config_(std::move(config)), - compiler_(std::move(compiler)), - debug_op_(Op::Get("debug")) {} - - /*! - * \brief Returns the primitive function associated with \p expr, or nullptr if none. - */ - BaseFunc ResolveToPrimitive(const Expr& expr) { - // NOTE: We can't assume expr->checked_type_ is defined, so can't early exit for first-order - // expressions. - if (const auto* global_var_node = expr.as()) { - if (!module_->ContainGlobalVar(global_var_node->name_hint)) { - // TODO(mbs): extern function cleanup - // Assume the function is extern and thus no longer in the IRModule. - return {}; - } else { - BaseFunc base_func = module_->Lookup(GetRef(global_var_node)); - return ResolveToPrimitive(base_func); - } - } else if (auto prim_func = expr.as()) { - return prim_func.value(); - } else if (const auto* var_node = expr.as()) { - auto itr = primitive_functions_.find(var_node); - if (itr == primitive_functions_.end()) { - // Not bound to a primitive function. - return {}; - } else { - return itr->second; - } - } else if (const auto* function_node = expr.as()) { - if (function_node->HasNonzeroAttr(attr::kExtern)) { - // We have a regular call to an 'extern' function. The call itself needs to be rewritten - // to call_lowered form, and any required dynamic shape functions generated and - // cross-linked. - return GetRef(function_node); - } else if (function_node->HasNonzeroAttr(attr::kPrimitive)) { - if (const auto* call_node = function_node->body.as()) { - if (call_node->op == debug_op_) { - // Debug 'primitives' are not lowered. - return {}; - } - } - // We have a regular call to a 'primitive' function (possibly with a 'Compiler' attribute). - // We need to lower and rewrite the call. - return GetRef(function_node); - } else { - // Not marked as primitive during partitioning or TVM fusion. - return {}; - } - } else { - return {}; - } - } - - /*! - * \brief Returns a 'call_lowered' call to \p prim_fn_var with \p args and \p span with all the - * required attributes filled in. Generally \p prim_fn_var will correspond to the lowered or - * externally codegen-ed form of \p original_function, where \p lowered_functions binds all - * the required lowered functions. - * - * The call's attributes will capture: - * - Any attributes on the original_function. - * - All the lowered functions. - * TODO(mbs): Pretty sure that's no longer needed. - * - Details needed to cross-link the call to it's dynamic shape function, if any. - */ - Expr MakeLoweredCall(const BaseFunc& original_function, const GlobalVar& prim_fn_var, - Array args, Span span, const Target& target, - const Map& lowered_functions, - const te::Schedule& sch = {}) { - auto opt_compiler = original_function->GetAttr(attr::kCompiler); - - // Add some metadata on top of the *original function* and invoke the callback so it can - // be captured. - // TODO(@areusch, @jroesch): this metadata is for AOT, this should be our interface for AOT - Map prim_fns; - Array all_prim_fn_vars; - for (const auto& kv : lowered_functions) { - if (opt_compiler) { - // We expect the original function to have just the "Extern" attribute signaling the - // function (will be) compiled externally. - ICHECK(kv.second.as()) - << PrettyPrint(kv.first) << " must be bound to an (external) Function"; - } else { - // We expect one or more PrimFuncs, one of which corresponds to 'the' lowered primitive, - // and the rest are in support of that via tir::Calls. - ICHECK(kv.second.as()) - << PrettyPrint(kv.first) << " must be bound to a PrimFunc"; - prim_fns.Set(kv.first, Downcast(kv.second)); - all_prim_fn_vars.push_back(kv.first); - } - } - - // Alas, WithAttr cannot work with base classes. - if (auto opt = original_function.as()) { - auto func_with_metadata = opt.value(); - func_with_metadata = WithAttr(func_with_metadata, "prim_fn_var", prim_fn_var); - func_with_metadata = WithAttr(func_with_metadata, "prim_funcs", prim_fns); - func_with_metadata = WithAttr(func_with_metadata, tvm::attr::kTarget, target); - // Store generated Schedules of operator - if (sch.defined() && sch->keep_schedule_record) { - func_with_metadata = WithAttr(func_with_metadata, "schedule", sch); - } - this->process_fn_(func_with_metadata); - } else { - auto func_with_metadata = original_function.as().value(); - func_with_metadata = WithAttr(func_with_metadata, "prim_fn_var", prim_fn_var); - func_with_metadata = WithAttr(func_with_metadata, "prim_funcs", prim_fns); - func_with_metadata = WithAttr(func_with_metadata, tvm::attr::kTarget, target); - // Store generated Schedules of operator - if (sch.defined() && sch->keep_schedule_record) { - func_with_metadata = WithAttr(func_with_metadata, "schedule", sch); - } - this->process_fn_(func_with_metadata); - } - - // Now prepare the attributes of the call_lowered. - CallLoweredAttrs call_lowered_attrs; - - // TODO(mbs): "reshape" cleanup. - if (!opt_compiler && original_function->HasNonzeroAttr(attr::kReshapeOnly)) { - call_lowered_attrs.metadata.Set(attr::kReshapeOnly, tvm::Integer(1)); - } - - call_lowered_attrs.metadata.Set("relay_attrs", original_function->attrs); - call_lowered_attrs.metadata.Set("all_prim_fn_vars", all_prim_fn_vars); - - if (const auto* function_node = original_function.as()) { - if (IsDynamic(function_node->ret_type)) { - // Create a dynamic shape function to calculate the expected shape of the results of - // the lowered function. - // Shape function keys use the original function as their 'function', but the generic 'cpu' - // target as the target since all shape functions run on the host cpu irrespective of where - // the primitive runs. - CCacheKey shape_key(GetRef(function_node), config_->host_virtual_device->target); - CachedFunc lowered_shape_func = compiler_->LowerShapeFunc(shape_key); - - // Capture the shape function's global var and parameters 'states' in call - // annotations so calling convention can be recovered. - // TODO(mbs): Shape cleanup. - call_lowered_attrs.metadata.Set("prim_shape_fn_var", lowered_shape_func->prim_fn_var); - call_lowered_attrs.metadata.Set("prim_shape_fn_states", - lowered_shape_func->shape_func_param_states); - call_lowered_attrs.metadata.Set( - "prim_shape_fn_num_inputs", - Integer(static_cast(lowered_shape_func->inputs.size()))); - call_lowered_attrs.metadata.Set( - "prim_shape_fn_num_outputs", - Integer(static_cast(lowered_shape_func->outputs.size()))); - Array all_prim_shape_fn_vars; - for (const auto& kv : lowered_shape_func->funcs->functions) { - CHECK(kv.second.as()) << "must be a prim fn"; - all_prim_shape_fn_vars.push_back(kv.first); - } - call_lowered_attrs.metadata.Set("all_prim_shape_fn_vars", all_prim_shape_fn_vars); - } - } - - return CallLowered(prim_fn_var, std::move(args), std::move(call_lowered_attrs), - std::move(span)); - } - - std::pair PreVisitLetBinding_(const Var& var, const Expr& value) final { - Var new_var = Downcast(Mutate(var)); - Expr new_value = Mutate(value); - BaseFunc prim_func = ResolveToPrimitive(new_value); - - if (prim_func.defined()) { - // Remember let var is bound (possibly indirectly) to a primitive function. - primitive_functions_.emplace(var.get(), prim_func); - } - return {new_var, new_value}; - } - - Expr PostVisitLet_(const LetNode* pre_let_node, const LetNode* post_let_node) final { - BaseFunc prim_func = ResolveToPrimitive(post_let_node->value); - if (prim_func.defined()) { - // Leaving let var scope - primitive_functions_.erase(pre_let_node->var.get()); - // Drop the let node - return post_let_node->body; - } - return DeviceAwareExprMutator::PostVisitLet_(pre_let_node, post_let_node); - } - - Expr DeviceAwareVisitExpr_(const FunctionNode* function_node) override { - if (function_node->HasNonzeroAttr(attr::kPrimitive) || - function_node->HasNonzeroAttr(attr::kExtern)) { - // Nothing to lower inside primitive/external functions. - return GetRef(function_node); - } else { - return DeviceAwareExprMutator::DeviceAwareVisitExpr_(function_node); - } - } - - Expr DeviceAwareVisitExpr_(const CallNode* call_node) override { - // We can see six forms of calls: - // 1. A 'normal' Relay call to a Function with the "Primitive" attribute and not "Compiler" - // attribute. We will need to lower that to a global PrimFunc and rewrite the call to: - // call_lowered(@new_global, (arg1, ..., argn), ) - // If needed, the call needs to be cross-linked with any dynamic shape functions. - // (However, some primitives are special and handled separately.) - // 2. A 'normal' Relay call to a Function with the "Primitive" and "Compiler" attributes. We - // will need to invoke the "relay.ext." function to yield a runtime module, and - // rewrite the call to the same form as above. Dynamic shape function cross-linking may - // also be needed. - // 3. A 'normal' Relay call to a Function with the "Extern" attribute. This function has - // already been compiled by an external codegen and a definition for it exists in some - // runtime module. Again, we rewrite to call_lowered form, and cross-link with a dynamic - // shape function if needed. - // 4. A 'normal' Relay call to a PrimFunc which has already been supplied via a global - // definition. We rewrite those to use the call_lowered form, but otherwise nothing else - // needs to be done. - // 5. A 'call_lowered' call from an earlier invocation of this pass or otherwise deliberately - // inserted. It has all the required attributes, and any associated dynamic shape function - // has been generated and cross-linked. These calls are not changed. - // 6. A 'normal' Relay call to a Relay Function without any special attribute. These - // calls are not changed. - // - // Note that ResolveToPrimitive will yield non-null only for cases 1-4. - - // Prepare the arguments and op. - Array new_args; - for (const auto& arg : call_node->args) { - new_args.push_back(VisitExpr(arg)); - } - Expr new_op = VisitExpr(call_node->op); - - // Look for (possibly indirect) calls to primitives. - BaseFunc primitive_func = ResolveToPrimitive(call_node->op); - if (!primitive_func.defined()) { - // Cases 5 and 6: Leave as ordinary call. - if (auto function = call_node->op.as()) { - process_fn_(function.value()); - } - return WithFields(GetRef(call_node), std::move(new_op), std::move(new_args)); - } - - // Special case for case 1: device_copies are left as calls to primitive operators - // so that each backend can handle them directly. - // TODO(mbs): device_copy cleanup. Would be better for FuseOps to just leave device_copy alone. - if (const auto* function_node = primitive_func.as()) { - DeviceCopyProps device_copy_props = GetDeviceCopyProps(function_node->body); - if (device_copy_props.body.defined()) { - ICHECK_EQ(new_args.size(), 1); - return DeviceCopy(new_args[0], device_copy_props.src_virtual_device, - device_copy_props.dst_virtual_device); - } - } - - ICHECK(call_node->type_args.empty()) << "lowered functions cannot be polymorphic"; - - // Case 4: If the function has already been lowered we just need to update the call. - if (auto prim_func = primitive_func.as()) { - // Function should already be Target annotated by this point - // but the TE Compiler metadata is still needed for the callback - // TODO(Mousius) - Robustify this to not assume we're in the GlobalVar for Target Hooks - Optional opt_target = primitive_func->GetAttr(tvm::attr::kTarget); - ICHECK(opt_target.defined()); - auto prim_fn_var = Downcast(call_node->op); - Map prim_fns = {{prim_fn_var, prim_func.value()}}; - return MakeLoweredCall(primitive_func, prim_fn_var, std::move(new_args), call_node->span, - opt_target.value(), prim_fns); - } - - // Determine the target for lowering or external codegen. - Target target; - Optional opt_compiler = primitive_func->GetAttr(attr::kCompiler); - if (opt_compiler.defined()) { - // This function needs to be compiled with external codegen. - Optional opt_target = config_->FindPrimitiveTargetForKind(opt_compiler.value()); - if (opt_target.defined()) { - // The target is what's supplied by the compilation config for kind matching the - // "Compiler" name. - target = opt_target.value(); - } else { - // Legacy fallback. - target = Target("ext_dev"); - } - } else { - // The target corresponding to the call_node expression's annotation. - VirtualDevice virtual_device = GetVirtualDevice(GetRef(call_node)); - ICHECK(!virtual_device->IsFullyUnconstrained()) << PrettyPrint(GetRef(call_node)); - target = virtual_device->target; - ICHECK(target.defined()); - } - - if (primitive_func->HasNonzeroAttr(attr::kExtern)) { - // Case 3: Function has already been compiled. - GlobalVar prim_fn_var = Downcast(call_node->op); - return MakeLoweredCall(primitive_func, prim_fn_var, std::move(new_args), call_node->span, - target, /*lowered_functions=*/{}); - } else { - // Cases 1 and 2: lower the primitive function for the desired target, possibly using external - // codegen. - CCacheKey key(Downcast(primitive_func), target, - GetVirtualDevice(GetRef(call_node))); - CachedFunc cfunc = compiler_->Lower(key); - ICHECK(cfunc.defined()); - return MakeLoweredCall(primitive_func, cfunc->prim_fn_var, std::move(new_args), - call_node->span, target, cfunc->funcs->functions, cfunc->schedule); - } - } - - IRModule module_; - ProcessFn process_fn_; - /*! \brief All available targets. */ - CompilationConfig config_; - // Map from in-scope let-bound variables to Functions known to be primitive, or PrimFuncs which - // have already been lowered. We'll rewrite these to the fresh global vars bound to the lowered - // primitive function as we go. Those vars will be bound in the target device-type specific - // module we'll ultimately emit for each required device-type. Note that a primitive may be - // lowered for multiple device types, each which will be assigned a fresh var. - std::unordered_map primitive_functions_; - TECompiler compiler_; - // Cache ops that need to be frequently used later to reduce lookup overhead. - const Op& debug_op_; -}; - -Pass LowerTensorExpr(TECompiler compiler, ProcessFn process_fn, CompilationConfig config) { - runtime::TypedPackedFunc pass_func = - [=](Function func, IRModule module, PassContext ctx) { - LowerTensorExprMutator lower_te(module, process_fn, config, compiler); - return Downcast(lower_te.Mutate(func)); - }; - return CreateFunctionPass(pass_func, 0, "LowerTensorExpr", {}); -} - -backend::FunctionInfo UpdateMainWorkspaceSize(const IRModule& mod, const CompilationConfig& config, - Map storage_info_map) { - Function func = Downcast(mod->Lookup("main")); - - VLOG_CONTEXT << "UpdateMainWorkspaceSize"; - VLOG(1) << "calculating FunctionInfo for main:" << std::endl << PrettyPrint(func); - - // This is a Map> - // TODO(mbs): Collapsing VirtualDevices to just device type. - std::unordered_map, backend::EnumClassHash> - sid_workspace; - // This is a Map - std::unordered_map device_io; - // This is a Map - std::unordered_map device_consts; - - // Initialize the mapping from all storage identifiers to workspace sizes, - // the amount of device io, and the device constants. - for (const auto& kv : storage_info_map) { - const backend::StorageInfo& storage_info = kv.second; - const std::vector& storage_ids = storage_info->storage_ids; - const std::vector& virtual_devices = storage_info->virtual_devices; - CHECK_EQ(storage_ids.size(), virtual_devices.size()); - for (uint32_t i = 0; i < virtual_devices.size(); i++) { - DLDeviceType device_type = virtual_devices[i]->device_type(); - sid_workspace[device_type][storage_ids[i]] = 0; - device_io[device_type] = 0; - device_consts[device_type] = 0; - } - } - - // Iterate the storage map to compute all the tensor sizes in the program. - // There are 3 cases in this code: - // - // First we need to compute the sizes of all - // inline constants. - // - // Second we compute the size of any bound variable as these are input and output - // sizes of the program. - // - // Finally for all other expressions we check which storage identifier they have - // been assigned and we compute the maximal size of the storage, as tensors can - // share storage with other tensors which are the same size or larger. - // - // In this final case there is only one allocation for all tensors which share storage - // which will be the maximal size of all tensors which were assigned to it. - for (const auto& kv : storage_info_map) { - const Expr& expr = kv.first; - const backend::StorageInfo& storage_info = kv.second; - int64_t size_bytes = backend::CalculateRelayExprSizeBytes(expr->checked_type()); - VLOG(1) << "expression:" << std::endl - << PrettyPrint(expr) << std::endl - << "of type:" << std::endl - << PrettyPrint(expr->checked_type()) << std::endl - << "has size " << size_bytes << " and storage info:" << std::endl - << storage_info; - const std::vector& storage_ids = storage_info->storage_ids; - const std::vector& virtual_devices = storage_info->virtual_devices; - - if (expr->IsInstance()) { - for (const auto& virtual_device : virtual_devices) { - DLDeviceType device_type = virtual_device->device_type(); - ICHECK_EQ(device_consts.count(device_type), 1); - device_consts[device_type] += size_bytes; - } - } else if (expr->IsInstance() || expr.same_as(func->body)) { - CHECK(size_bytes == 0 || virtual_devices.size() >= 1) << "must be at least one device"; - for (const auto& virtual_device : virtual_devices) { - DLDeviceType device_type = virtual_device->device_type(); - device_io[device_type] += size_bytes; - } - } else { - // TODO(@electriclilies): This code is never being called which means sid_workspace is not - // updated.. This means that storage info is probably not being created correctly. Or is not - // equivalent to what was here previously - for (uint32_t i = 0; i < storage_ids.size(); i++) { - // Here we record the largest size of the tensor - // that share the same storage id, because storage_id will - // be shared between multiple tensors that are not live simultaneously. - DLDeviceType device_type = virtual_devices[i]->device_type(); - if (size_bytes > sid_workspace[device_type][storage_ids[i]]) { - sid_workspace[device_type][storage_ids[i]] = size_bytes; - } - } - } - } - - // This is a Map - std::unordered_map device_workspace; - // Once we know the sizes of sids, we need to accumulate per device - for (const auto& dev_sid_size : sid_workspace) { - auto dev = dev_sid_size.first; - device_workspace[dev] = 0; - for (const auto& sid_size : dev_sid_size.second) { - device_workspace[dev] += sid_size.second; - } - } - - Map workspace_sizes; - Map io_sizes; - Map constant_sizes; - Map tir_primfuncs; - Map relay_primfuncs; - - // Initialize all target workspaces to zero - for (const auto& target : config->primitive_targets) { - workspace_sizes.Set(target, 0); - } - - for (const auto& dev_and_size : device_workspace) { - Target target = config->FindPrimitiveTargetForDeviceOrFail(dev_and_size.first); - workspace_sizes.Set(target, dev_and_size.second); - relay_primfuncs.Set(target, func); - } - for (const auto& dev_and_size : device_io) { - Target target = config->FindPrimitiveTargetForDeviceOrFail(dev_and_size.first); - io_sizes.Set(target, dev_and_size.second); - } - - for (const auto& dev_and_size : device_consts) { - Target target = config->FindPrimitiveTargetForDeviceOrFail(dev_and_size.first); - ICHECK_EQ(constant_sizes.count(target), 0); - constant_sizes.Set(target, dev_and_size.second); - } - - backend::FunctionInfo func_info(std::move(workspace_sizes), std::move(io_sizes), - std::move(constant_sizes), std::move(tir_primfuncs), - std::move(relay_primfuncs)); - VLOG(1) << "func_info: " << func_info; - return std::move(func_info); -} - -/*! - * \brief A function to create the function metadata for an input function (ie calculate buffer - * input/output sizes) - * \param func The function to calculate function metadata for - * \param function_metadata The map that stores all the function metadatas - */ -void UpdateFunctionMetadata(BaseFunc func, - Map& function_metadata, // NOLINT(*) - Integer workspace_byte_alignment) { - VLOG_CONTEXT << "UpdateFunctionMetadata"; - VLOG(1) << "updating function metadata for:" << std::endl << PrettyPrint(func); - // Originally UpdateFunctionMetadata took in CCachedFunc and looped through all the funcs stored - // there Now the goal is to take only one func because process_fn should be controlling the - // iteration However, to do the workspace calculations we need the primfuncs. So process_fn - // needs to either access the cached funcs or be directly passed primfuncs This is bad and - // ideally we don't want process_fn to look at primfuncs There's also the question now of what - // the function metadatas are and how they are used if we can do something else to replicate the - // behavior of the function metadatas that might be good (ie annotating functions or something). - Map workspace_sizes; - Map io_sizes; - Map constant_sizes; - Map tir_primfuncs; - Map relay_primfuncs; - - Optional> prim_fns = - func->GetAttr>("prim_funcs"); - CHECK(prim_fns) << "primitive functions not set on Relay function by TECompiler."; - - Optional prim_fn_var = func->GetAttr("prim_fn_var"); - CHECK(prim_fn_var) << "prim_fn_var must be set on Relay functions by TECompiler."; - - Optional relay_target = func->GetAttr(tvm::attr::kTarget); - CHECK(relay_target) << "target must be set on Relay functions by the TECompiler."; - - for (const auto& kv : prim_fns.value()) { - auto prim_fn = Downcast(kv.second); - CHECK(prim_fn.defined()) << "the primitive function must be defined"; - - Integer workspace_size = CalculateWorkspaceBytes(prim_fn, workspace_byte_alignment); - - // Workspace sizes - Target prim_fn_target; - if (prim_fn->attrs->dict.count(tvm::attr::kTarget)) { - prim_fn_target = Downcast(prim_fn->attrs->dict[tvm::attr::kTarget]); - } else { - prim_fn_target = relay_target.value(); - } - - workspace_sizes.Set(prim_fn_target, workspace_size); - - // Calculating size for I/O - // TODO(mbs): See also the other three utils for calculating tensor bytesize. - for (auto const& param : prim_fn->params) { - bool not_a_buffer = prim_fn->buffer_map.count(param) == 0; - if (not_a_buffer) { - io_sizes.Set(prim_fn_target, 0); - continue; - } - - auto p_shape = prim_fn->buffer_map[param]->shape; - int num_of_elements = 1; - for (const auto& dim_index_expr : p_shape) { - if (dim_index_expr->IsInstance()) { - num_of_elements *= dim_index_expr.as()->value; - } else { - // If shape is dynamic, we cannot calculate workspace in compile time. - num_of_elements = 0; - } - } - int element_size = prim_fn->buffer_map[param]->dtype.bytes(); - io_sizes.Set(prim_fn_target, element_size * num_of_elements); - } - - constant_sizes.Set(prim_fn_target, 0); - tir_primfuncs.Set(prim_fn_target, prim_fn); - if (func->IsInstance()) { - relay_primfuncs.Set(prim_fn_target, Downcast(func)); - } - } - - backend::FunctionInfo fi = backend::FunctionInfo( - std::move(workspace_sizes), std::move(io_sizes), std::move(constant_sizes), - std::move(tir_primfuncs), std::move(relay_primfuncs)); - - VLOG(1) << "FunctionInfo: " << PrettyPrint(prim_fn_var.value()) << " = " << PrettyPrint(fi); - - // The primitive function name here corresponds to the string we will use to generate - // this Relay function at the low level. - function_metadata.Set(prim_fn_var.value()->name_hint, fi); -} - -/*! \brief Main lowering driving. */ -IRModule LowerTE(const IRModule& module, const String& module_name, ProcessFn process_fn, - CompilationConfig config) { - TECompiler compiler(module, module_name); - - // TODO(mbs): This is all unnecessarily convoluted. Better would be to accumulate the rewritten - // module as we go (including rewritten Functions, lowered primitives, and runtime modules - // generated by external toolchains), and use a pair of maps over vars and global vars - // to global vars to remember which functions have already been lowered. - - // Lower all the callees in module: - // - Functions tagged with "Compiler" are unchanged (checked by CreateFunctionPass) - // - Functions tagged with "Primitive" are unchanged (checked by LowerTensorExprMutator) - // - Called functions tagged with "Compiler" are copied into the compiler cache with a fresh - // GlobalVar, and calls updated (sticking with regular Relay Call). - // - Calls to functions tagged with "Primitive" are compiled to PrimFuncs, and calls updated - // (using call_lowered convention). - IRModule updated_module = - LowerTensorExpr(compiler, std::move(process_fn), std::move(config))(module); - - // The Functions tagged with "Compiler" are now residing in the cache ready to be - // compiled by LowerExternalFunctions. However we still need a record of them in the - // IRModule so that the various executors can see which function names need to be - // retrieved. They may, however, have been renamed. - compiler->AddExterns(updated_module); - - // Add the lowered functions. - IRModule lowered_module = compiler->GetLoweredFunctions(); - VLOG(1) << "capturing " << lowered_module->functions.size() << " new lowered functions"; - for (const auto& kv : lowered_module->functions) { - if (updated_module->ContainGlobalVar(kv.first->name_hint)) { - LOG(FATAL) << "duplicate bindings for '" << kv.first->name_hint - << "'. Existing is:" << std::endl - << PrettyPrint(updated_module->Lookup(kv.first->name_hint)) << std::endl - << "while new is:" << std::endl - << PrettyPrint(kv.second); - } - updated_module->Add(kv.first, kv.second); - } - - // Invoke external codegen for all Functions in the cache tagged with "Compiler", and - // annotate the module with the resulting runtime modules. - // TODO(mbs): runtime modules should be first class rather than attributes. - Array external_mods = - module->GetAttr>(tvm::attr::kExternalMods).value_or({}); - Array new_external_mods = compiler->LowerExternalFunctions(); - VLOG(1) << "capturing " << external_mods.size() << " existing and " << new_external_mods.size() - << " new external modules"; - for (const auto& mod : new_external_mods) { - external_mods.push_back(mod); // copy-on-write. - } - - // Annotate the module with C Device API context mapping (this is until we have Targets - // annotated for the C Device API) - // TODO(Mousius) - Remove "device_contexts" as soon as we have the graph annotated properly with - // Targets - Map device_contexts = - module->GetAttr>("device_contexts", Map()).value(); - Map new_device_contexts = compiler->GetDeviceContexts(); - VLOG(1) << "capturing " << device_contexts.size() << " existing and " - << new_device_contexts.size() << " new device contexts for external functions"; - for (const auto& kv : new_device_contexts) { - ICHECK_EQ(device_contexts.count(kv.first), 0); - device_contexts.Set(kv.first, kv.second); // copy-on-write. - } - - updated_module = WithAttrs(updated_module, {{tvm::attr::kExternalMods, std::move(external_mods)}, - {"device_contexts", std::move(device_contexts)}}); - - if (backend::IsAutoSchedulerEnabled()) { - // Capture all the 'operator weights', ie usage counts for each PrimFunc. - Map op_weights = - module->GetAttr>("op_weights", Map()).value(); - Map new_op_weights = compiler->GetOpWeights(); - VLOG(1) << "capturing " << op_weights.size() << " existing and " << new_op_weights.size() - << " new operator weights for PrimFuncs"; - for (const auto& kv : new_op_weights) { - ICHECK_EQ(op_weights.count(kv.first), 0); - op_weights.Set(kv.first, kv.second); // copy-on-write. - } - updated_module = WithAttr(updated_module, "op_weights", std::move(op_weights)); - } - - return updated_module; -} - -Map GetPerTargetModules(IRModule mod) { - std::unordered_map - per_target_modules; - for (const auto& kv : mod->functions) { - const GlobalVar& var = kv.first; - const BaseFunc& func = kv.second; - if (func->IsInstance()) { - // Extract target - Optional target = func->GetAttr(tvm::attr::kTarget); - ICHECK(target) << "Target should be set at this point"; - - // Put the function in per_target_modules - if (!per_target_modules.count(target.value())) { - // Initialize the IRModule for this target with the attributes from the input IRModule - IRModule target_module = IRModule({}, {}, {}, {}, mod->attrs); - // Add the function to the IRModule - target_module->Add(var, func); - per_target_modules[target.value()] = target_module; - } else { - // The IRModule for this target is initialized, so just add the function. - IRModule target_module = per_target_modules.at(target.value()); - target_module->Add(var, func); - } - } else if (!func->IsInstance()) { - LOG(FATAL) - << "The function types in the IRModule should be RelayFunction or PrimFunc, but got " - << func->GetTypeKey(); - } - } - return per_target_modules; -} - -Pass LowerTE(String module_name, CompilationConfig complilation_config, ProcessFn process_fn) { - runtime::TypedPackedFunc pass_func = [=](IRModule module, - PassContext ctx) { - return LowerTE(module, module_name, process_fn, complilation_config); - }; - - return tvm::transform::Sequential( - {tvm::relay::transform::RelayToTIRTargetHook(complilation_config), - tvm::transform::CreateModulePass(pass_func, 0, "LowerTE", {"InferType"}), InferType(), - tvm::tir::transform::ExtractPrimFuncConstants()}); -} - -TVM_REGISTER_GLOBAL("relay.tec.LowerTE") - .set_body_typed([](String module_name, CompilationConfig compilation_config) { - return LowerTE(std::move(module_name), std::move(compilation_config)); - }); - -} // namespace tec -} // namespace relay -} // namespace tvm diff --git a/src/relay/backend/te_compiler.h b/src/relay/backend/te_compiler.h deleted file mode 100644 index f2ba84014a09..000000000000 --- a/src/relay/backend/te_compiler.h +++ /dev/null @@ -1,197 +0,0 @@ -/* - * 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 relay/backend/te_compiler.h - * \brief Internal compilation layer which lowers Relay "primitive functions" to TIR PrimFns. - * - * - * This represents the new design of the Relay compilation flow and will replace the interface - * contained in compile_engine.h as we migrate towards a standard pass based lowering of - * Relay functions. - * - * This files provides an internal API which lowers Relay programs to components which - * can be combined with TVM produced kernels to compile an entire program. - * - * The result of lowering contains a combination of `runtime::Module`s produced by external - * compilers and a set of lowered PrimFns which can be code generated for targets. - */ -#ifndef TVM_RELAY_BACKEND_TE_COMPILER_H_ -#define TVM_RELAY_BACKEND_TE_COMPILER_H_ - -#include -#include -#include -#include -#include -#include -#include -#include -#include - -#include -#include -#include - -#include "../transforms/infer_layout_utils.h" -#include "../transforms/pass_utils.h" -#include "./te_compiler_cache.h" -#include "./utils.h" - -namespace tvm { -namespace relay { -namespace tec { - -using ProcessFn = std::function; - -/*! - * \brief A compiler which lowers primitive Relay functions to tensor expressions - * and schedules them into TIR functions. - */ -class TECompilerNode : public Object { - public: - /*! \brief destructor */ - virtual ~TECompilerNode() {} - /*! - * \brief Get lowered result. - * \param key The key to the cached function. - * \return The result. - */ - virtual CachedFunc Lower(const CCacheKey& key) = 0; - - /*! - * \brief Get lowered result. - * \param key The key to the cached function. - * \return The result. - */ - virtual CachedFunc Lower(const CCacheKey& key, const String mod_name) = 0; - - /* Return all functions which have been lowered by the compiler in an IRModule, annotated with - * their target. */ - virtual IRModule GetLoweredFunctions() = 0; - - /*! - * \brief Just in time compile to get a PackedFunc. - * \param key The key to the cached function. - * \return The result. - */ - virtual PackedFunc JIT(const CCacheKey& key) = 0; - /*! - * \brief Lower the shape function. - * \param key The key to the cached function. - * \return The result. - */ - virtual CachedFunc LowerShapeFunc(const CCacheKey& key) = 0; - /*! - * \brief Lower the external function using external codegen tools. - * \return The runtime modules for each needed external codegen tool. - */ - virtual tvm::Array LowerExternalFunctions() = 0; - - /*! - * \brief Update \p module to remove functions marked with the "Compiler" attribute and replace - * them with their 'external' representation using the "ExternalSymbol" attribute. - * - * TODO(mbs): This is a stepping stone while we migrate to a more official representation - * of 'external functions' in the IRModule and allow lowering to incrementally updatethe - * module stead of forcing everything via the cache. - * - */ - virtual void AddExterns(IRModule module) = 0; - - /*! - * \brief Get C Device API context mapping - * \return Map of GlobalVar to associated C Device API context name (either Target or kCompiler - * annotated) - */ - virtual Map GetDeviceContexts() = 0; - virtual void SetDeviceContexts(const Map& device_contexts) = 0; - - virtual Map GetOpWeights() const = 0; - - /*! \brief clear the cache. */ - virtual void Clear() = 0; - - void VisitAttrs(AttrVisitor*) {} - - static constexpr const char* _type_key = "relay.TECompiler"; - TVM_DECLARE_FINAL_OBJECT_INFO(TECompilerNode, Object); -}; - -/*! \brief cache entry used in compile engine */ -class TECompiler : public ObjectRef { - public: - explicit TECompiler(Optional opt_mod = {}, Optional mod_name = {}); - explicit TECompiler(ObjectPtr n) : ObjectRef(n) {} - TECompilerNode* operator->() { return static_cast(get_mutable()); } - using ContainerType = TECompilerNode; - TVM_DLL static TECompiler& Global(); -}; - -/*! - * \brief A function to create the function metadata for an input function (ie calculate buffer - * input/output sizes) - * \param func The function to calculate function metadata for - * \param function_metadata The map that stores all the function metadatas - * \param workspace_byte_alignment Byte alignment for allocations - */ -void UpdateFunctionMetadata(BaseFunc relay_func, - Map& function_metadata, // NOLINT(*) - Integer workspace_byte_alignment = 16); - -/*! - * \brief Update the "main" control function's metadata - * - * \param mod The module - * \param config All the available targets. - * \return function_infos Function info for each function in the module - */ -backend::FunctionInfo UpdateMainWorkspaceSize(const IRModule& mod, const CompilationConfig& config, - Map storage_info_map); - -/*! \brief Returns all the global \p PrimFunc functions in \p mod, but separated into an \p IRModule - * per \p Target. - * - * \param mod The IRModule to extract the per target module from - * \return The map from Target to IRModule - */ -Map GetPerTargetModules(IRModule mod); - -inline void DefaultProcessFn(BaseFunc) {} - -/*! - * \brief Pass to lower an IRModule's primitive functions to TIR. - * - * This is the "back half" of the Relay compiler which lowers "primitive functions" - * to TE expressions, schedules them, and emits PrimFuncs. - * - * \param module_name The name of this module, used as a prefix for generated globals. - * \param config All available targets. - * \param process_fn Callback allowing one-level up code generators to process - * each function that we lower (default is no-op). - * \returns The pass which lowers primitive functions to TIR - */ -transform::Pass LowerTE(String module_name, CompilationConfig config, - ProcessFn process_fn = DefaultProcessFn); - -} // namespace tec -} // namespace relay -} // namespace tvm - -#endif // TVM_RELAY_BACKEND_TE_COMPILER_H_ diff --git a/src/relay/backend/te_compiler_cache.cc b/src/relay/backend/te_compiler_cache.cc deleted file mode 100644 index 79a41ae050c6..000000000000 --- a/src/relay/backend/te_compiler_cache.cc +++ /dev/null @@ -1,1155 +0,0 @@ -/* - * 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 "./te_compiler_cache.h" - -#include -#include -#include -#include -#include -#include -#include -#include -#include -#include -#include -#include -#include -#include -#include -#include -#include -#include -#include -#include -#include -#include -#include -#include - -#include -#include -#include -#include -#include -#include -#include - -#include "../../te/operation/create_primfunc.h" -#include "../op/memory/memory.h" -#include "../src/meta_schedule/module_equality.h" -#include "../src/meta_schedule/trace_apply.h" -#include "../transforms/meta_schedule_layout_rewrite.h" -#include "utils.h" - -namespace tvm { -namespace relay { -namespace tec { - -TVM_REGISTER_NODE_TYPE(LoweredOutputNode); -TVM_REGISTER_NODE_TYPE(CachedFuncNode); -TVM_REGISTER_NODE_TYPE(CCacheKeyNode); -TVM_REGISTER_NODE_TYPE(CCacheValueNode); - -LoweredOutput::LoweredOutput(tvm::Array outputs, OpImplementation impl) { - auto n = make_object(); - n->outputs = std::move(outputs); - n->implementation = std::move(impl); - data_ = std::move(n); -} - -CCacheKey::CCacheKey(Function source_func, Target target, VirtualDevice vd) { - auto n = make_object(); - n->source_func = std::move(source_func); - n->target = std::move(target); - n->virtual_device = std::move(vd); - data_ = std::move(n); -} - -CachedFunc::CachedFunc(tvm::Target target, GlobalVar prim_fn_var, tvm::Array inputs, - tvm::Array outputs, te::Schedule schedule, - tir::PrimFunc prim_func, tvm::Array shape_func_param_states, - IRModule funcs, - std::unordered_map constant_tensors) { - auto n = make_object(); - n->target = target; - n->prim_fn_var = prim_fn_var; - n->inputs = inputs; - n->outputs = outputs; - n->schedule = schedule; - n->prim_func = prim_func; - n->shape_func_param_states = shape_func_param_states; - n->funcs = funcs; - n->constant_tensors = constant_tensors; - data_ = std::move(n); -} - -Array GetShape(const Array& shape) { - // for now, we always use int32 shape when possible - // even if the result of shape inference becomes int64. - Array res; - for (IndexExpr val : shape) { - const int64_t* pval = tir::as_const_int(val); - if (pval != nullptr) { -#ifndef TVM_INDEX_DEFAULT_I64 - ICHECK_LE(pval[0], std::numeric_limits::max()) - << "dimension must be less then int32_t's max value"; - ICHECK_GE(pval[0], std::numeric_limits::min()) - << "dimension must be less then int32_t's max value"; - res.push_back(IntImm(DataType::Int(32), *pval)); -#else - res.push_back(val); -#endif // TVM_INDEX_DEFAULT_I64 - } else if (val->IsInstance()) { - // currently all 'any' we meet in shape function are non-negative. - res.push_back(val.as()->ToSizeVar()); - } else { - res.push_back(val); - } - } - return res; -} - -// Helper class that is used during lowering to TE. -// It matches sequence of Ops and lower them into single TOPI operation. All supported patterns are -// enumerated in "supported_patterns_". -class QnnPatternMatcher { - public: - QnnPatternMatcher() - : qnn_conv2d_op_(Op::Get("qnn.conv2d")), - qnn_dense_op_(Op::Get("qnn.dense")), - qnn_dense_pack_op_(Op::Get("qnn.contrib_dense_pack")), - qnn_requantize_op_(Op::Get("qnn.requantize")), - bias_add_op_(Op::Get("add")) {} - - // Memoize visited operations - void Register(const CallNode* call_node) { - ICHECK(call_node->op.as()); - Op op = Downcast(call_node->op); - if (op == qnn_conv2d_op_) { - registered_ops_.push_front(P_QConv2d); - ICHECK(anchor_op_ == nullptr); - anchor_op_ = call_node; - } else if (op == qnn_requantize_op_) { - registered_ops_.push_front(P_QRequantize); - } else if (op == bias_add_op_) { - registered_ops_.push_front(P_BiasAdd); - } else if (op == qnn_dense_op_) { - registered_ops_.push_front(P_QDense); - ICHECK(anchor_op_ == nullptr); - anchor_op_ = call_node; - } else if (op == qnn_dense_pack_op_) { - registered_ops_.push_front(P_QDensePack); - ICHECK(anchor_op_ == nullptr); - anchor_op_ = call_node; - } else { - registered_ops_.push_front(P_Opaque); - } - } - - // Check whether given Op is a part of matched pattern. - bool find(const Op& op) { - if (registered_ops_.empty()) return false; - - if (op == qnn_conv2d_op_ || op == qnn_requantize_op_ || op == bias_add_op_ || - op == qnn_dense_op_ || op == qnn_dense_pack_op_) { - for (const auto& pat : supported_patterns_) { - auto it = - std::search(registered_ops_.begin(), registered_ops_.end(), pat.begin(), pat.end()); - if (it != registered_ops_.end()) return true; - } - } - return false; - } - - // returns whether given Op is last in the pattern sequence. - bool IsLeafOp(const Op& op) { return op == qnn_requantize_op_; } - - const CallNode* GetAnchorOp() { return anchor_op_; } - - void Clear() { registered_ops_.clear(); } - - private: - const Op& qnn_conv2d_op_; - const Op& qnn_dense_op_; - const Op& qnn_dense_pack_op_; - const Op& qnn_requantize_op_; - const Op& bias_add_op_; - - // Main (complicated) operation in the primitive (for example qnn.conv2d, qnn.dense etc.). - const CallNode* anchor_op_ = nullptr; - - enum POper { P_QConv2d, P_QDense, P_QDensePack, P_BiasAdd, P_QRequantize, P_Opaque }; - - std::deque registered_ops_; - - const std::vector> supported_patterns_ = { - {P_QDense, P_BiasAdd, P_QRequantize}, // qnn.dense -> bias_add -> qnn.requantize - {P_QDense, P_QRequantize}, // qnn.dense -> qnn.requantize - {P_QDensePack, P_BiasAdd, P_QRequantize}, // qnn.contrib_dense_pack -> bias -> qnn.requantize - {P_QDensePack, P_QRequantize}, // qnn.contrib_dense_pack -> qnn.requantize - {P_QConv2d, P_BiasAdd, P_QRequantize}, // qnn.conv2d -> bias_add -> qnn.requantize - {P_QConv2d, P_QRequantize} // qnn.conv2d -> qnn.requantize - }; -}; - -// Lowers Relay primitive Function to TE Compute -class LowerToTECompute : public backend::MemoizedExprTranslator> { - public: - LowerToTECompute(Target target, NameSupply constants_name_supply) - : target_(target), - device_copy_op_(Op::Get("device_copy")), - constants_name_supply_(constants_name_supply) {} - - Array Lower(const Function& relay_func) { - for (Var param : relay_func->params) { - Array inputs; - for (const auto& ttype : FlattenTupleType(param->checked_type())) { - auto name_hint = param->vid->name_hint; - tvm::te::Tensor tensor = tvm::te::placeholder( - GetShape(ttype->shape), ttype->dtype, (name_hint == "") ? "placeholder" : name_hint); - inputs.push_back(tensor); - fn_inputs_.push_back(tensor); - } - memo_[param] = inputs; - } - readable_name_stream_ << "fused"; - - Array outputs = this->VisitExpr(relay_func->body); - - candidate_name_ = readable_name_stream_.str(); - - constexpr static size_t kMaxFuncNameLength = 80; - // WARNING: Please make sure to also update TVM_CRT_MAX_STRLEN_FUNCTION_NAME - // whenever the value of kMaxFuncNameLength changes - if (candidate_name_.size() > kMaxFuncNameLength) { - std::stringstream truncated_name; - truncated_name << candidate_name_.substr(0, kMaxFuncNameLength); - truncated_name << "_" << std::hex << std::hash{}(candidate_name_) << "_"; - candidate_name_ = truncated_name.str(); - } - - return outputs; - } - - Array VisitExpr_(const VarNode* op) final { - LOG(FATAL) << "Unexpected free variable " << PrettyPrint(GetRef(op)); - } - - Array VisitExpr_(const ConstantNode* op) final { - using tir::make_const; - void* data = op->data->data; - DataType dtype = DataType(op->data->dtype); - if (op->is_scalar()) { - auto value = te::compute( - {}, - [&](const Array&) { - if (dtype == DataType::Int(16)) { - return make_const(dtype, static_cast(data)[0]); - } else if (dtype == DataType::Int(8)) { - return make_const(dtype, static_cast(data)[0]); - } else if (dtype == DataType::UInt(8) || dtype == DataType::Bool()) { - return make_const(dtype, static_cast(data)[0]); - } else if (dtype == DataType::Int(32)) { - return make_const(dtype, static_cast(data)[0]); - } else if (dtype == DataType::Int(64)) { - return make_const(dtype, static_cast(data)[0]); - } else if (dtype == DataType::Float(16)) { - return make_const(dtype, __gnu_h2f_ieee(static_cast(data)[0])); - } else if (dtype == DataType::Float(32)) { - return make_const(dtype, static_cast(data)[0]); - } else if (dtype == DataType::Float(64)) { - return make_const(dtype, static_cast(data)[0]); - } else { - LOG(FATAL) << dtype << " not handled"; - } - }, - "compile_engine_const", topi::kBroadcast); - scalars_.push_back(value->op); - return {value}; - } else { - const auto* ttype = op->checked_type().as(); - std::stringstream ss; - std::string s = readable_name_stream_.str(); - std::replace(s.begin(), s.end(), '.', '_'); - ss << constants_name_supply_->FreshName(s + "_constant"); - tvm::te::Tensor tensor = tvm::te::placeholder(GetShape(ttype->shape), ttype->dtype, ss.str()); - constant_tensors_[op] = tensor; - return {tensor}; - } - } - - Array VisitExpr_(const CallNode* call_node) final { - static auto flower_call = tvm::runtime::Registry::Get("relay.backend.lower_call"); - ICHECK(flower_call) << "relay.backend.lower_call is not registered."; - - pattern_matcher_.Register(call_node); - - Array inputs; - // int count_tuple = 0; - for (Expr arg : call_node->args) { - if (arg->checked_type().as()) { - // ++count_tuple; - } - for (te::Tensor tensor : VisitExpr(arg)) { - inputs.push_back(tensor); - } - } - - ICHECK(call_node->op.as()) << "Primitive function only allows call into primitive ops"; - Op op = Downcast(call_node->op); - - // TODO(mbs): device_copy cleanup - ICHECK_NE(op, device_copy_op_) << "device_copy cannot be lowered"; - - Array outputs; - - if (pattern_matcher_.find(op)) { - if (pattern_matcher_.IsLeafOp(op)) { - // Lower anchor op when pattern leaf op was reached - auto anchor_op = pattern_matcher_.GetAnchorOp(); - LoweredOutput lowered_out = - (*flower_call)(GetRef(anchor_op), inputs, target_, call_node->checked_type()); - outputs = lowered_out->outputs; - Op a_op = Downcast(anchor_op->op); - op_implementations_[a_op.operator->()] = lowered_out->implementation; - - pattern_matcher_.Clear(); - } else { - // Forward inputs as "outputs" for successor. - readable_name_stream_ << '_' << op->name; - return inputs; - } - } else { - LoweredOutput lowered_out = (*flower_call)(GetRef(call_node), inputs, target_); - outputs = lowered_out->outputs; - op_implementations_[op.operator->()] = lowered_out->implementation; - } - - if (outputs.size() != 1) { - const auto* tuple_type = call_node->checked_type().as(); - ICHECK(tuple_type) << "Expected output to be a tuple type " - << PrettyPrint(call_node->checked_type()); - - ICHECK_EQ(tuple_type->fields.size(), outputs.size()); - } - - readable_name_stream_ << '_' << op->name; - return outputs; - } - - Array VisitExpr_(const FunctionNode* op) final { - LOG(FATAL) << "Primitive Functions can not contain nested functions."; - } - - Array VisitExpr_(const LetNode* op) final { - Array val = VisitExpr(op->value); - ICHECK(!memo_.count(op->var)); - memo_[op->var] = val; - return VisitExpr(op->body); - } - - Array VisitExpr_(const TupleNode* op) final { - Array fields; - for (Expr field : op->fields) { - // TODO(mbs): Generalize to be equivalent to FlattenTupleType. - ICHECK(field->checked_type().as()) << "Only allow Tuple of Tensor"; - Array res = VisitExpr(field); - ICHECK_EQ(res.size(), 1); - fields.push_back(res[0]); - } - return fields; - } - - Array VisitExpr_(const TupleGetItemNode* op) final { - const auto* tuple_type = op->tuple->type_as(); - Array tuple = VisitExpr(op->tuple); - ICHECK_EQ(tuple_type->fields.size(), tuple.size()); - ICHECK_GE(op->index, 0); - ICHECK_LT(static_cast(op->index), tuple.size()); - return {tuple[op->index]}; - } - - public: - // Additional outputs - Array fn_inputs_; - Array scalars_; - std::unordered_map constant_tensors_; - std::unordered_map op_implementations_; - std::string candidate_name_; - - private: - QnnPatternMatcher pattern_matcher_; - - tvm::Target target_; - std::ostringstream readable_name_stream_; - // Cache device copy op for equivalence checking to reduce registry lookup - // overhead for each invocation of call node when retrieving schedules. - const Op& device_copy_op_; - // A NameSupply object passed from a caller, used to assign unique names to constants - // across different invocations of LowerToTECompute. - NameSupply constants_name_supply_; -}; - -using namespace tvm::tir; - -class LayoutFreeConstantCollector : public StmtVisitor { - public: - Array constants; - - private: - void VisitStmt_(const BlockNode* op) final { - StmtVisitor::VisitStmt_(op); - if (Optional ann = op->annotations.Get("layout_free_placeholders")) { - for (Buffer buffer : Downcast>(ann)) { - layout_free_buffer_vars_.insert(buffer->data.get()); - } - } - } - - void VisitStmt_(const AllocateConstNode* op) final { - StmtVisitor::VisitStmt_(op); - if (auto it = layout_free_buffer_vars_.find(op->buffer_var.get()); - it != layout_free_buffer_vars_.end()) { - constants.push_back(op->data.value()); - } - } - - std::unordered_set layout_free_buffer_vars_; -}; - -using NDArrayMap = - std::unordered_map; - -// Replace constants in AllocateConst nodes according to the given mapping -class AllocateConstReplaceConstant : public StmtExprMutator { - public: - explicit AllocateConstReplaceConstant(const NDArrayMap& constant_map) - : constant_map_(constant_map) {} - - static PrimFunc Rewrite(PrimFunc f, const NDArrayMap& constant_map) { - AllocateConstReplaceConstant rewriter(constant_map); - PrimFuncNode* n = f.CopyOnWrite(); - n->body = rewriter(std::move(n->body)); - return f; - } - - private: - Stmt VisitStmt_(const AllocateConstNode* op) final { - if (auto it = constant_map_.find(op->data.value()); it != constant_map_.end()) { - auto rewriten_constant = it->second; - Array rewritten_extents; - for (auto s : rewriten_constant.Shape()) { - rewritten_extents.push_back(PrimExpr(static_cast(s))); - } - return AllocateConst(op->buffer_var, op->dtype, rewritten_extents, rewriten_constant, - op->body, op->annotations, op->span); - } - return StmtExprMutator::VisitStmt_(op); - } - - NDArrayMap constant_map_; -}; - -// Construct a schedule for a given Relay primitive function and target. -class ScheduleBuilder : public ExprVisitor { - public: - explicit ScheduleBuilder(Target target) - : target_(target), - mod_eq_structural_(meta_schedule::ModuleEquality::Create("ignore-ndarray")) { - // Whether to use auto_scheduler schedule. - use_auto_scheduler_ = backend::IsAutoSchedulerEnabled(); - database_ = meta_schedule::Database::Current(); - if (backend::IsMetaScheduleEnabled()) { - CHECK(database_.defined()) << "ValueError: `use_meta_schedule` is enabled in Relay " - "build, but no `meta_schedule.Database` context is provided. "; - } - } - - CachedFunc Create(const Function& relay_func, GlobalVarSupply global_var_supply, - NameSupply constant_name_supply) { - LowerToTECompute lower_te_compute(target_, constant_name_supply); - Array tensor_outs = lower_te_compute.Lower(relay_func); - Array fn_inputs = lower_te_compute.fn_inputs_; - VisitExpr(relay_func->body); - - // TODO(mbs): This should be the definitive global by which the PrimFunc is known and - // no other GlobalVar ctors should appear inside the lowering machinery. - auto prim_fn_var = global_var_supply->FreshGlobal(lower_te_compute.candidate_name_); - prim_fn_var->checked_type_ = relay_func->checked_type(); - - // Fusion over tupled results may leave identity relationships - // between inputs and outputs, copy identity output tensors, - // since tir lowering do not support aliasing output to input buffer. - for (size_t i = 0; i < tensor_outs.size(); ++i) { - if (tensor_outs[i]->op.as()) { - tensor_outs.Set(i, topi::identity(tensor_outs[i])); - } - } - - te::Schedule schedule{nullptr}; - tir::PrimFunc prim_func{nullptr}; - // No need to register schedule for device copy op. - if (anchor_attrs_.as() == nullptr) { - if (use_auto_scheduler_) { - const auto* fauto_schedule = - runtime::Registry::Get("auto_scheduler.relay_integration.auto_schedule_topi_compute"); - ICHECK(fauto_schedule != nullptr) - << "auto_scheduler.relay_integration.auto_schedule_topi_compute is not registered"; - ObjectRef obj = (*fauto_schedule)(prim_fn_var->name_hint, tensor_outs); - if (obj.defined()) { - schedule = Downcast(obj); - } - } - if (database_) { - using tvm::meta_schedule::TuningRecord; - using tvm::tir::IndexMap; - using tvm::tir::Instruction; - using tvm::tir::InstructionKind; - using tvm::tir::PrimFunc; - using tvm::tir::Schedule; - backend::FTECompilerTIRConverter tir_converter = backend::GetTIRConverter(); - Array te_args = Concat(fn_inputs, tensor_outs); - Array constants; - for (auto [const_node, te_tensor] : lower_te_compute.constant_tensors_) { - te_args.push_back(te_tensor); - constants.push_back(const_node->data); - } - if (Optional f = tir_converter(te_args, constants)) { - IRModule query_mod = backend::PrimFuncToIRModule(f.value()); - if (Optional opt_record = database_.value()->QueryTuningRecord( - /*mod=*/query_mod, - /*target=*/target_, - /*workload_name=*/prim_fn_var->name_hint)) { - LayoutFreeConstantCollector const_collector; - const_collector(f.value()->body); - - static InstructionKind kind_transform_layout = InstructionKind::Get("TransformLayout"); - TuningRecord record = opt_record.value(); - for (const Instruction& inst : record->trace->insts) { - if (inst->kind.same_as(kind_transform_layout)) { - ICHECK_EQ(inst->inputs.size(), 2); - auto index_map = Downcast(inst->inputs[1]); - - if (!const_collector.constants.empty()) { - // In this case, RewriteLayout is acting on an AllocateConst node. - // After tuning, we reach this code path twice: First by - // the Relay MetaScheduleLayoutRewrite pass, and next by the final - // compilation (Relay to TE schedule lowering). - // - // Due to Relay MetaScheduleLayoutRewrite and FoldConstant passes, - // the Relay subgraph for which we query the database during the - // final compilation has its weight tensor transformed according to - // the index map, determined during tuning. For example, - // - // fn (%p0: Tensor[(1, 56, 56, 64), float32]) { - // %0 = nn.conv2d(%p0, meta[relay.Constant][0], - // /*ty=Tensor[(4, 2, 2, 3, 3, 32, 8), float32]*/, ...); - // add(%0, meta[relay.Constant][1]) - // } - // - // Note that the database does not have an entry corresponding to such subgraphs, - // since an input subgraph to the tuning system always has its weight tensor in - // the original layout, e.g. - // - // fn (%p0: Tensor[(1, 56, 56, 64), float32]) { - // %0 = nn.conv2d(%p0, meta[relay.Constant][0], - // /*ty=Tensor[(3, 3, 64, 64), float32]*/, ...); - // add(%0, meta[relay.Constant][1]) - // } - // - // Thus, in both of the two cases where we reach this code path, we need careful - // logic to make sure that (1) the database lookup during the final compilation - // succeeds and (2) the application of a schedule trace is well defined. - - ICHECK(const_collector.constants.size() == 1) - << "Only one layout-free constant is supported by RewriteLayout for now"; - auto constant = const_collector.constants[0]; - - auto is_constant_transformed = [index_map](runtime::NDArray c) { - if (c.Shape().size() != index_map->initial_indices.size()) { - return true; - } - size_t src_size_1d = 1; - Array orig_shape; - for (size_t i = 0; i < c.Shape().size(); ++i) { - src_size_1d *= c->shape[i]; - orig_shape.push_back(PrimExpr(static_cast((c->shape[i])))); - } - arith::Analyzer analyzer; - auto dst_shape = index_map->MapShape(orig_shape, &analyzer); - std::vector dst_shape_int; - size_t dst_size_1d = 1; - for (size_t i = 0; i < dst_shape.size(); ++i) { - dst_size_1d *= dst_shape[i].as()->value; - } - return src_size_1d != dst_size_1d; - }; - - if (!is_constant_transformed(constant)) { - // This is the first case, reached during the MetaScheduleLayoutRewrite pass. - // - // A layout-free constant having the same rank as an input to the index map - // is assumed to be transformed by this index map. - // TODO(masahi): If there are multiple layout-free constants in one - // TIR mod (e.g. conv2d -> conv2d fusion), this assumption does not hold. - // We need to determine which constant the given index map acts on. - // - // We know that, during the final compilation, we will query the database - // for a subgraph that the tuner has never seen. We workaround this problem - // by adding a dummy entry to the database. The dummy entry is carefully - // constructed so that the lookup during the final compilation would succeed. - runtime::NDArray rewritten_constant = index_map->MapNDArray(constant); - auto f_dummy = AllocateConstReplaceConstant::Rewrite( - f.value(), {{constant, rewritten_constant}}); - auto workload_dummy = - database_.value()->CommitWorkload(backend::PrimFuncToIRModule(f_dummy)); - TuningRecord rec_dummy(record->trace, workload_dummy, record->run_secs, - record->target, record->args_info); - database_.value()->CommitTuningRecord(rec_dummy); - } else { - // The constant is already transformed, so this is the second case, reached - // during the final compilation. - // - // The schedule trace is supposed to be applied to the weight in its original - // layout. But as explained above, the Relay subgraph we get in this case - // has its weight tensor transformed according to the corresponding index map. - // So effectively, we undo the layout transformation on the weight to restore - // the original PrimFunc that the schedule trace is supposed to act on. - ICHECK(index_map->inverse_index_map); - auto inverse_map = Downcast(index_map->inverse_index_map.value()); - ICHECK(constant.Shape().size() == inverse_map->initial_indices.size()); - runtime::NDArray orig_constant = inverse_map->MapNDArray(constant); - auto f_ = AllocateConstReplaceConstant::Rewrite(f.value(), - {{constant, orig_constant}}); - query_mod = backend::PrimFuncToIRModule(f_); - } - } - MetaScheduleLayoutRewriter::LayoutQueuePush(index_map); - } - } - - Schedule sch = Schedule::Traced(query_mod, /*seed=*/-1, /*debug_mask=*/0, - tir::ScheduleErrorRenderLevel::kDetail); - - if (!mod_eq_structural_->Equal(query_mod, opt_record.value()->workload->mod)) { - // When the database lookup succeeds while structural equality check fails, - // it implies that the anchor block based equality has been used during tuning. - // The trace in the record cannot directly be applied to this query module. - meta_schedule::ScheduleUsingAnchorTrace(sch, record->trace, target_); - } else { - record->trace->ApplyToSchedule(sch, /*remove_postproc=*/false); - } - - IRModule mod = sch->mod(); - ICHECK_EQ(mod->functions.size(), 1); - mod = tir::transform::RemoveWeightLayoutRewriteBlock(/*skip_ndarray_rewrite*/ false)( - std::move(mod)); - prim_func = Downcast(mod->Lookup("main")); - // Need to copy attrs from relay function over to prim func. Most notably the structural - // hash. - prim_func = WithAttrs(prim_func, relay_func->attrs->dict); - } else { - int dispatch = backend::UseMetaScheduleDispatch(); - // (dispatch & 2): controls whether to print TVMScript for missing TIR - // (dispatch & 4): controls whether to raise fatal errors for missing TIR - if (dispatch & 2) { - LOG(WARNING) << "Cannot find workload: " << prim_fn_var->name_hint << "\n" - << f.value(); - } else { - LOG(WARNING) << "Cannot find workload: " << prim_fn_var->name_hint; - } - if (dispatch & 4) { - LOG(FATAL); - } - } - } - } - // Use TOPI schedule if user specified, or the function has no auto_scheduler schedule. - if (!schedule.defined() && !prim_func.defined()) { - if (anchor_op_.defined()) { - auto anchor_impl = lower_te_compute.op_implementations_.find(anchor_op_.operator->()); - ICHECK(anchor_impl != lower_te_compute.op_implementations_.end()); - schedule = anchor_impl->second.Schedule(anchor_attrs_, tensor_outs, target_); - } else { - auto default_sched = GenericFunc::Get("schedule_injective"); - ICHECK(default_sched.defined()) << "schedule_injective not registered for " << target_; - With tctx(target_); - schedule = default_sched(tensor_outs); - } - } - if (schedule.defined()) { - for (const auto& scalar : lower_te_compute.scalars_) { - if (schedule->Contain(scalar)) { - schedule[scalar].compute_inline(); - } - } - } - } - - IRModule funcs = IRModule(Map({})); - return CachedFunc(target_, prim_fn_var, fn_inputs, tensor_outs, schedule, prim_func, {}, funcs, - lower_te_compute.constant_tensors_); - } - - void VisitExpr_(const CallNode* call_node) final { - static auto fpattern = Op::GetAttrMap("TOpPattern"); - - ICHECK(call_node->op.as()) << "Primitive function only allows call into primitive ops"; - Op op = Downcast(call_node->op); - - for (Expr arg : call_node->args) { - VisitExpr(arg); - } - - int op_pattern = fpattern[op]; - if (!use_auto_scheduler_ && !database_.defined() && op_pattern >= kCommReduce) { - ICHECK(!anchor_op_.defined() || anchor_op_pattern_ < kCommReduce) - << "Cannot apply TOPI schedule to a primitive function with two complicated ops" - << " anchor=" << anchor_op_ << " current=" << op; - } - if (op_pattern >= anchor_op_pattern_) { - anchor_op_ = op; - anchor_attrs_ = call_node->attrs; - anchor_op_pattern_ = op_pattern; - } - } - - private: - tvm::Target target_; - Op anchor_op_; - Attrs anchor_attrs_; - int anchor_op_pattern_{0}; - bool use_auto_scheduler_; - Optional database_; - std::unique_ptr mod_eq_structural_; -}; - -/*! - * \brief Create schedule for target. - * \param source_func The primitive function to be lowered. - * \param target The target we want to create schedule for. - * \return Pair of schedule and cache. - * The funcs field in cache is not yet populated. - */ -CachedFunc PrimFuncFor(const Function& source_func, const Target& target, - GlobalVarSupply global_var_supply, NameSupply constant_name_supply) { - return ScheduleBuilder(target).Create(source_func, global_var_supply, constant_name_supply); -} - -// Creates shape function from functor. -class MakeShapeFunc : public backend::MemoizedExprTranslator> { - public: - MakeShapeFunc() {} - - CachedFunc Create(const Function& prim_func, const Target& target, - GlobalVarSupply global_var_supply) { - VLOG_CONTEXT << "MakeShapeFunc"; - TShapeDataDependent shape_func_param_states; - - for (auto param : prim_func->params) { - param_states_[param] = kNoNeed; - Array data_inputs; - Array shape_inputs; - - for (const auto& ttype : FlattenTupleType(param->checked_type())) { - // Add data placeholder (in case we discover we need it below) - Shape shape = GetShape(ttype->shape); - tvm::te::Tensor data_tensor = - tvm::te::placeholder(shape, ttype->dtype, "data_" + param->vid->name_hint); - data_inputs.push_back(data_tensor); - // Add shape placeholder (in case we discover we need it below) - int64_t ndim = shape.size(); - Shape sshape; - if (ndim > 0) { - sshape.push_back(tvm::Integer(ndim)); - } - tvm::te::Tensor shape_tensor = - tvm::te::placeholder(sshape, DataType::Int(64), "shape_" + param->vid->name_hint); - shape_inputs.push_back(shape_tensor); - } - param_data_[param] = data_inputs; - param_shapes_[param] = shape_inputs; - } - - // Setup the name; - readable_name_stream_ << "shape_func"; - - // Create the tensor expressions representing the output shapes. - Array outputs = VisitExpr(prim_func->body); - - // Generate a name. - auto candidate_name = readable_name_stream_.str(); - - constexpr static size_t kMaxFuncNameLength = 80; - // WARNING: Please make sure to also update TVM_CRT_MAX_STRLEN_FUNCTION_NAME - // whenever the value of kMaxFuncNameLength changes - if (candidate_name.size() > kMaxFuncNameLength) { - std::stringstream truncated_name; - truncated_name << candidate_name.substr(0, kMaxFuncNameLength); - truncated_name << "_" << std::hex << std::hash{}(candidate_name) << "_"; - candidate_name = truncated_name.str(); - } - - // Set all the inputs correctly, and accumulate their types from the p.o.v. of the - // shape function rather than the primitive it is derived for. - Array inputs; - Array shape_function_arg_types; - for (auto param : prim_func->params) { - int state = param_states_[param]; - shape_func_param_states.push_back(IntImm(DataType::Int(32), state)); - if (state & kNeedInputData) { - // Pass the primitive arguments directly (though in flattened form and on the host) - for (auto t : param_data_[param]) { - inputs.push_back(t); - shape_function_arg_types.push_back(TensorType(t->GetShape(), t->GetDataType())); - } - } - if (state & kNeedInputShape) { - // Pass the shapes of the primitive arguments (also on the host) - for (auto t : param_shapes_[param]) { - inputs.push_back(t); - shape_function_arg_types.push_back(TensorType(t->GetShape(), t->GetDataType())); - } - } - } - - // TODO(mbs): This should be the definitive global by which the PrimFunc is known and - // no other GlobalVar ctors should appear inside the lowering machinery. - auto prim_fn_gvar = global_var_supply->FreshGlobal(candidate_name); - - // Gather the result types, again from the p.o.v. of the shape function rather than - // the primitive it is derived for. - Array shape_function_res_types; - for (const auto& t : outputs) { - shape_function_res_types.push_back(TensorType(t->GetShape(), t->GetDataType())); - } - - // Assign the shape function its true type. - FuncType type(shape_function_arg_types, TupleType(shape_function_res_types), - /*type_params=*/{}, /*type_constraints=*/{}); - VLOG(1) << "shape function '" << prim_fn_gvar->name_hint << "' has type:" << std::endl - << PrettyPrint(type) << std::endl - << "corresponding to primitive of type:" << std::endl - << PrettyPrint(prim_func->checked_type()); - prim_fn_gvar->checked_type_ = std::move(type); - - // generate schedule for shape func - Array out_ops; - for (auto t : outputs) { - out_ops.push_back(t->op); - } - te::Schedule schedule = te::create_schedule(out_ops); - tvm::te::AutoInlineInjective(schedule); - for (const auto& scalar : scalars_) { - auto scalar_op = scalar->op; - if (schedule->Contain(scalar_op)) { - schedule[scalar_op].compute_inline(); - } - } - - Array all_args = Array(inputs); - for (te::Tensor arg : outputs) { - all_args.push_back(arg); - } - - using tvm::transform::PassContext; - With fresh_pass_ctx_scope(PassContext::Create()); - - std::unordered_map binds; - IRModule lowered_module = - tvm::LowerSchedule(schedule, all_args, prim_fn_gvar->name_hint, binds, global_var_supply); - return CachedFunc(target, prim_fn_gvar, inputs, outputs, schedule, tir::PrimFunc{nullptr}, - shape_func_param_states, lowered_module); - } - - Array VisitExpr(const Expr& expr) final { - if (expr.as()) { - // Do not memoize vars because shape functions could use either the data - // or the shape of a var each time. - return ExprFunctor::VisitExpr(expr); - } - // For other case, do memoized visit - return backend::MemoizedExprTranslator>::VisitExpr(expr); - } - - Array VisitExpr_(const VarNode* var_node) final { - auto var = GetRef(var_node); - auto it = param_arg_map_.find(var); - if (it != param_arg_map_.end()) { - // This var is a parameter of a nested function. Visit the corresponding argument in the - // function call site. - return VisitExpr(it->second); - } - if (param_states_.find(var) == param_states_.end()) { - LOG(FATAL) << "Unexpected free variable " << PrettyPrint(var); - } else { - ICHECK(data_dependents_per_input_.size()); - auto data_dependent = data_dependents_per_input_.back(); - if (data_dependent) { - param_states_[var] |= kNeedInputData; - return param_data_[var]; - } else { - param_states_[var] |= kNeedInputShape; - return param_shapes_[var]; - } - } - } - - Array VisitExpr_(const ConstantNode* op) final { - using tir::make_const; - ICHECK(data_dependents_per_input_.size()); - bool data_dependent = data_dependents_per_input_.back(); - if (!op->is_scalar()) { - // This is a constant weight, extract the shape of the weight tensor. - // This can not be data dependent. - CHECK(!data_dependent); - auto ttype = op->checked_type().as(); - int ndim = static_cast(ttype->shape.size()); - Array out_shape{ndim}; - te::Tensor value = tvm::te::compute( - out_shape, - [&](const Array& indices) { - auto idx = indices[0]; - PrimExpr ret = make_const(DataType::Int(64), 0); - for (int i = 0; i < ndim; i++) { - ret = tvm::if_then_else(idx == i, ttype->shape[i], ret); - } - return ret; - }, - "shape_const", topi::kBroadcast); - scalars_.push_back(value); - return {value}; - } - if (data_dependent) { - void* data = op->data->data; - DataType dtype = DataType(op->data->dtype); - auto value = tvm::te::compute( - {}, - [&](const Array&) { - if (dtype == DataType::Int(32)) { - return make_const(dtype, static_cast(data)[0]); - } else if (dtype == DataType::Int(64)) { - return make_const(dtype, static_cast(data)[0]); - } else if (dtype == DataType::Float(32)) { - return make_const(dtype, static_cast(data)[0]); - } else if (dtype == DataType::Float(64)) { - return make_const(dtype, static_cast(data)[0]); - } else if (dtype == DataType::Bool()) { - return make_const(dtype, static_cast(data)[0]); - } else { - LOG(FATAL) << "not handled"; - } - }, - "data_const", topi::kBroadcast); - scalars_.push_back(value); - return {value}; - } else { - auto value = tvm::te::compute( - {}, [&](const Array&) { return tir::make_const(DataType::Int(64), 0); }, - "shape_const", topi::kBroadcast); - scalars_.push_back(value); - return {value}; - } - } - - Array VisitExpr_(const CallNode* call_node) final { - VLOG(1) << "considering call:" << std::endl << PrettyPrint(GetRef(call_node)); - if (auto* func = call_node->op.as()) { - VLOG(1) << "user function"; - for (size_t i = 0; i < func->params.size(); ++i) { - param_arg_map_[func->params[i]] = call_node->args[i]; - } - return VisitExpr(func->body); - } - - static auto fshape_func = Op::GetAttrMap("FShapeFunc"); - static auto tshape_data_dependent = Op::GetAttrMap("TShapeDataDependent"); - ICHECK(call_node->op.as()) << "Primitive function only allows call into primitive ops"; - Op op = Downcast(call_node->op); - ICHECK(data_dependents_per_input_.empty() || !data_dependents_per_input_.back()) - << "Error in op fusion: output of the shape func is fed to a " - << "data-dependent shape func"; - ICHECK_GT(fshape_func.count(op), 0) << "Internal error, cannot find ShapeFunc for " << op->name; - ICHECK_GT(tshape_data_dependent.count(op), 0) - << "Internal error, cannot find TShapeDataDependent for " << op->name; - - Array dep_spec = tshape_data_dependent[op]; - if (dep_spec.size() == 1) { - // This is for cases when data dependence is specified per op - // Replicate 0 or 1 flag to all arguments - for (size_t i = 1; i < call_node->args.size(); ++i) { - dep_spec.push_back(dep_spec[0]); - } - } - - // Visit all inputs - Array inputs; - int count_tuple = 0; - for (size_t i = 0; i < call_node->args.size(); ++i) { - Expr arg = call_node->args[i]; - if (arg->checked_type().as()) { - ++count_tuple; - } - data_dependents_per_input_.push_back(dep_spec[i]->value != 0); - for (te::Tensor tensor : VisitExpr(arg)) { - inputs.push_back(tensor); - } - data_dependents_per_input_.pop_back(); - } - if (count_tuple) { - ICHECK_EQ(call_node->args.size(), 1U) << "Only allow function with a single tuple input"; - } - // Get output ndims - auto ret_type = call_node->checked_type(); - Array out_ndims; - for (const auto& ttype : FlattenTupleType(ret_type)) { - out_ndims.push_back(IntImm(DataType::Int(32), ttype->shape.size())); - } - - // Call shape function - Array outputs = fshape_func[op](call_node->attrs, inputs, out_ndims); - VLOG(1) << "shape function for '" << op->name << "' with inputs:" << std::endl - << inputs << std::endl - << "yielded outputs:" << std::endl - << outputs; - readable_name_stream_ << "_" << op->name; - return outputs; - } - - Array VisitExpr_(const FunctionNode* op) final { - LOG(FATAL) << "Nested functions are not allowed to be visited."; - } - - Array VisitExpr_(const LetNode* op) final { - Array val = VisitExpr(op->value); - ICHECK(!memo_.count(op->var)); - memo_[op->var] = val; - return VisitExpr(op->body); - } - - Array VisitExpr_(const TupleNode* op) final { - Array fields; - for (Expr field : op->fields) { - ICHECK(field->checked_type().as()) - << "Expected a Tuple of Tensor, but got " << PrettyPrint(field->checked_type()); - Array res = VisitExpr(field); - ICHECK_EQ(res.size(), 1); - fields.push_back(res[0]); - } - return fields; - } - - Array VisitExpr_(const TupleGetItemNode* op) final { - Array input_shapes = VisitExpr(op->tuple); - Array out; - out.push_back(input_shapes[op->index]); - return out; - } - - private: - /*! \brief String stream for function name */ - std::ostringstream readable_name_stream_; - /*! \brief Map from parameter to its shape function usage state */ - std::unordered_map param_states_; - /*! \brief Map from parameter to list of data placeholder */ - std::unordered_map, ObjectPtrHash, ObjectPtrEqual> param_data_; - /*! \brief Map from parameter to list of shape placeholder */ - std::unordered_map, ObjectPtrHash, ObjectPtrEqual> param_shapes_; - /*! \brief Stack of data dependencies for shape function, specified per each op input */ - std::vector data_dependents_per_input_; - /*! \brief Scalars used in the shape function */ - Array scalars_; - /*! \brief Map from parameters of a nested function to corresponding arguments in a function - * call site. - */ - std::unordered_map param_arg_map_; -}; - -CachedFunc ShapeFuncFor(const Function& prim_func, const Target& target, - GlobalVarSupply global_var_supply) { - return MakeShapeFunc().Create(prim_func, target, global_var_supply); -} - -std::tuple, Array, std::string> LowerTECompute( - const Function& source_func, Target target, NameSupply constant_name_supply, - bool return_inputs) { - LowerToTECompute lower_te_compute(target, constant_name_supply); - Array outputs = lower_te_compute.Lower(source_func); - // Following ScheduleBuilder, remove placeholder ops from outputs. - tvm::Array tensor_outs; - for (const auto& tensor : outputs) { - if (!tensor->op.as()) { - tensor_outs.push_back(tensor); - } - } - - tvm::Array constants; - for (auto [const_node, te_tensor] : lower_te_compute.constant_tensors_) { - tensor_outs.push_back(te_tensor); - constants.push_back(const_node->data); - } - - if (return_inputs) { - return std::make_tuple(Concat(lower_te_compute.fn_inputs_, tensor_outs), constants, - lower_te_compute.candidate_name_); - } - return std::make_tuple(tensor_outs, constants, lower_te_compute.candidate_name_); -} - -std::pair, std::string> LowerToPrimFunc(const Function& relay_func, - Target target, - NameSupply constant_name_supply) { - ICHECK(relay_func->HasNonzeroAttr(attr::kPrimitive)) - << "The input must be a Relay primitive function."; - - auto [inputs_outputs, constants, fused_name] = - tec::LowerTECompute(relay_func, target, constant_name_supply, /*return_inputs=*/true); - auto tir_converter = backend::GetTIRConverter(); - return std::make_pair(tir_converter(inputs_outputs, constants), fused_name); -} - -tir::PrimFunc LowerToPrimFunc(const Function& relay_func, Target target) { - auto [f_opt, _] = LowerToPrimFunc(relay_func, target, NameSupply()); - (void)_; // to suppress -Werror=unused-variable warning - if (f_opt) { - return f_opt.value(); - } - LOG(FATAL) << "Failed to convert the Relay function: " << AsText(relay_func, false); - return PrimFunc(); -} - -TVM_REGISTER_GLOBAL("relay.backend.LowerToPrimFunc") - .set_body_typed([](Function relay_func, Target target) { - return LowerToPrimFunc(relay_func, target); - }); - -TVM_REGISTER_GLOBAL("relay.backend.LowerToTE").set_body_typed([](Function prim_func) { - auto tgt = tvm::Target("ext_dev"); - LowerToTECompute lower_te_compute(tgt, NameSupply()); - auto outputs = lower_te_compute.Lower(prim_func); - return CachedFunc(tgt, GlobalVar(lower_te_compute.candidate_name_), lower_te_compute.fn_inputs_, - outputs, te::Schedule(), tir::PrimFunc(), {}, - IRModule(Map({})), lower_te_compute.constant_tensors_); -}); - -} // namespace tec -} // namespace relay -} // namespace tvm diff --git a/src/relay/backend/te_compiler_cache.h b/src/relay/backend/te_compiler_cache.h deleted file mode 100644 index 502e0063220f..000000000000 --- a/src/relay/backend/te_compiler_cache.h +++ /dev/null @@ -1,292 +0,0 @@ -/* - * 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 relay/backend/tec_compiler_cache.h - * \brief Utilities for compiling tensor expressions inside of the Relay compiler. - */ -#ifndef TVM_RELAY_BACKEND_TE_COMPILER_CACHE_H_ -#define TVM_RELAY_BACKEND_TE_COMPILER_CACHE_H_ - -#include -#include -#include -#include -#include -#include -#include -#include -#include -#include - -#include -#include -#include -#include -#include - -#include "../transforms/infer_layout_utils.h" - -namespace tvm { -namespace relay { -namespace tec { - -/*! \brief Indicate whether the data or shape or both of a parameter is used in the shape func. */ -enum ShapeFuncParamState { - kNoNeed = 0, - kNeedInputData = 1, - kNeedInputShape = 2, - kNeedBoth = 3, -}; - -struct LoweredOutputNode : public Object { - /*! \brief The outputs to the function */ - tvm::Array outputs; - /*! \brief The implementation used to compute the output */ - OpImplementation implementation; - - void VisitAttrs(tvm::AttrVisitor* v) { - v->Visit("outputs", &outputs); - v->Visit("implementation", &implementation); - } - static constexpr const char* _type_key = "relay.LoweredOutput"; - TVM_DECLARE_FINAL_OBJECT_INFO(LoweredOutputNode, Object); -}; - -class LoweredOutput : public ObjectRef { - public: - TVM_DLL LoweredOutput(tvm::Array outputs, OpImplementation impl); - - TVM_DEFINE_OBJECT_REF_METHODS(LoweredOutput, ObjectRef, LoweredOutputNode); -}; - -class CCacheKey; -/*! \brief Compile cache key */ -class CCacheKeyNode : public Object { - public: - /*! \brief The source function to be lowered. */ - Function source_func; - /*! \brief The hardware target.*/ - Target target; - /*! \brief The virtual device constrains.*/ - VirtualDevice virtual_device; - - void VisitAttrs(tvm::AttrVisitor* v) { - v->Visit("source_func", &source_func); - v->Visit("target", &target); - v->Visit("virtual_device", &virtual_device); - } - /*! \return The hash value of CCacheKey. */ - inline size_t Hash() const; - /*! - * \brief check content equality - * \param other The other value. - * \return The result of equality check. - */ - inline bool Equal(const CCacheKeyNode* other) const; - - static constexpr const char* _type_key = "relay.CCacheKey"; - TVM_DECLARE_FINAL_OBJECT_INFO(CCacheKeyNode, tvm::Object); - - private: - /*! - * \brief internal cached hash value. - */ - mutable size_t hash_{0}; -}; - -/*! \brief cache entry used in compile engine */ -class CCacheKey : public ObjectRef { - public: - CCacheKey() {} - explicit CCacheKey(ObjectPtr n) : ObjectRef(n) {} - - /*! - * \brief The constructor - * \param source_func The source function. - * \param target The target device. - */ - TVM_DLL CCacheKey(Function source_func, Target target, - VirtualDevice virtual_device = VirtualDevice::FullyUnconstrained()); - - const CCacheKeyNode* operator->() const { return static_cast(get()); } - // comparator - inline bool operator==(const CCacheKey& other) const { - ICHECK(defined() && other.defined()); - return (*this)->Equal(other.operator->()); - } - using ContainerType = CCacheKeyNode; -}; - -/*! \brief Node container to represent a cached function. */ -struct CachedFuncNode : public Object { - /*! \brief compiled target */ - tvm::Target target; - /*! \brief Primitive Function Name */ - GlobalVar prim_fn_var; - /*! \brief The inputs to the function */ - tvm::Array inputs; - /*! \brief The outputs to the function */ - tvm::Array outputs; - /*! \brief The schedule to the function */ - te::Schedule schedule; - /*! \brief The TIR function if lowering in the meta schedule path */ - Optional prim_func; - /*! \brief Parameter usage states in the shape function. */ - tvm::Array shape_func_param_states; - /*! \brief The lowered functions to support the function. */ - IRModule funcs = IRModule(Map({})); - std::unordered_map constant_tensors; - - void VisitAttrs(tvm::AttrVisitor* v) { - v->Visit("target", &target); - v->Visit("prim_fn_var", &prim_fn_var); - v->Visit("inputs", &inputs); - v->Visit("outputs", &outputs); - v->Visit("schedule", &schedule); - v->Visit("prim_func", &prim_func); - v->Visit("funcs", &funcs); - v->Visit("shape_func_param_states", &shape_func_param_states); - } - - static constexpr const char* _type_key = "relay.CachedFunc"; - TVM_DECLARE_FINAL_OBJECT_INFO(CachedFuncNode, Object); -}; - -class CachedFunc : public ObjectRef { - public: - CachedFunc(tvm::Target target, GlobalVar prim_fn_name, tvm::Array inputs, - tvm::Array outputs, te::Schedule schedule, tir::PrimFunc prim_func, - tvm::Array shape_func_param_states, - IRModule funcs = IRModule(Map({})), - std::unordered_map constant_tensors = {}); - - public: - TVM_DEFINE_OBJECT_REF_METHODS(CachedFunc, ObjectRef, CachedFuncNode); -}; - -/*! \brief Node container for compile cache. */ -class CCacheValueNode : public Object { - public: - /*! \brief The corresponding function */ - CachedFunc cached_func; - /*! \brief Result of Packed function generated by JIT */ - PackedFunc packed_func; - /*! \brief usage statistics */ - int use_count{0}; - - void VisitAttrs(tvm::AttrVisitor* v) { - v->Visit("cached_func", &cached_func); - v->Visit("use_count", &use_count); - } - static constexpr const char* _type_key = "relay.CCacheValue"; - TVM_DECLARE_FINAL_OBJECT_INFO(CCacheValueNode, tvm::Object); -}; - -/*! \brief cache entry used in compile engine */ -class CCacheValue : public ObjectRef { - public: - CCacheValue() {} - explicit CCacheValue(ObjectPtr n) : ObjectRef(n) {} - CCacheValueNode* operator->() { return static_cast(get_mutable()); } - const CCacheValueNode* operator->() const { return static_cast(get()); } - using ContainerType = CCacheValueNode; -}; - -Array GetShape(const Array& shape); - -/*! - * \brief Lower Relay primitive Function to TE Compute - * \param source_func The primitive function to be lowered. - * \param target The compilation target. - * \param constant_name_supply A name supplier for constants - * across different invocations of this function. - * \param return_inputs If true, prepend input tensors to the output array of tensors. - * \return Tuple of the lowered TE compute, constant raw data, and fused function name. - */ -std::tuple, Array, std::string> LowerTECompute( - const Function& source_func, Target target, NameSupply constant_name_supply, - bool return_inputs = true); - -/*! - * \brief Lower Relay Function to TIR PrimFunc, by composing LowerTECompute and CreatePrimFunc. - * \param relay_func The primitive function to be lowered. - * \param target The compilation target. - * \param constant_name_supply A name supplier for constants - * across different invocations of this function. - * \return A pair of the created prim func and the name of the fused function. - */ -std::pair, std::string> LowerToPrimFunc(const Function& relay_func, - Target target, - NameSupply constant_name_supply); - -/*! - * \brief Create schedule for target. - * \param source_func The primitive function to be lowered. - * \param target The compilation target. - * \param global_var_supply A name supplier for global variables. - * \param constant_name_supply A name supplier for constants. - * \return Pair of schedule and cache. - * The funcs field in cache is not yet populated. - */ -CachedFunc PrimFuncFor(const Function& source_func, const Target& target, - GlobalVarSupply global_var_supply, NameSupply constant_name_supply); - -/*! \brief A specialization of PrimFuncFor, meant to be used when the names of constants do not - * matter. */ -inline CachedFunc PrimFuncFor(const Function& source_func, const Target& target) { - return PrimFuncFor(source_func, target, GlobalVarSupply(), NameSupply()); -} - -CachedFunc ShapeFuncFor(const Function& prim_func, const Target& target, - GlobalVarSupply global_var_supply); - -// implementations -inline size_t CCacheKeyNode::Hash() const { - if (hash_ != 0) return hash_; - // do structral hash, avoid 0. - hash_ = tvm::StructuralHash()(this->source_func); - hash_ = dmlc::HashCombine(hash_, std::hash()(target->str())); - if (hash_ == 0) hash_ = 1; - return hash_; -} - -inline bool CCacheKeyNode::Equal(const CCacheKeyNode* other) const { - if (Hash() != other->Hash()) return false; - return this->target->str() == other->target->str() && - this->virtual_device == other->virtual_device && - tvm::StructuralEqual()(this->source_func, other->source_func); -} - -} // namespace tec -} // namespace relay -} // namespace tvm - -namespace std { -// overload hash -template <> -struct hash<::tvm::relay::tec::CCacheKey> { - size_t operator()(const ::tvm::relay::tec::CCacheKey& key) const { - ICHECK(key.defined()); - return key->Hash(); - } -}; -} // namespace std - -#endif // TVM_RELAY_BACKEND_TE_COMPILER_CACHE_H_ diff --git a/src/relay/backend/token_allocator.cc b/src/relay/backend/token_allocator.cc deleted file mode 100644 index e974944b33b0..000000000000 --- a/src/relay/backend/token_allocator.cc +++ /dev/null @@ -1,155 +0,0 @@ -/* - * 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 relay/backend/token_allocator.cc - * \brief Token allocation classes for backend - */ - -#include "token_allocator.h" - -#include - -#include -#include - -namespace tvm { -namespace relay { -constexpr auto Is2DStorage = runtime::IsTextureStorage; - -/* - * Mixed mode memory allocator - */ -size_t TokenAllocatorMixed::GetMemorySize(StorageToken* prototype) { - TensorType ttype = prototype->ttype; - ICHECK(ttype.defined()); - size_t size = 1; - if (relay::Is2DStorage(prototype->virtual_device->memory_scope)) { - size = GetSize2D(prototype); - } else { - for (IndexExpr dim : ttype->shape) { - const int64_t* pval = tir::as_const_int(dim); - ICHECK(pval != nullptr) << "Cannot allocate memory symbolic tensor shape " << ttype->shape; - ICHECK_GE(*pval, 0) << "Cannot allocate memory for tensor with negative shape" << *pval; - size *= static_cast(pval[0]); - } - size *= DivRoundUp(ttype->dtype.bits() * ttype->dtype.lanes(), 8); - } - return size; -} - -String GetDeviceCompatibleToken(StorageToken* tok) { - Target null_tgt{nullptr}; - if (null_tgt == tok->virtual_device->target) { - return tok->virtual_device->memory_scope; - } - std::string dev_kind = tok->virtual_device->target->kind->name; - auto* device_scope_handler = tvm::runtime::Registry::Get("DeviceScopeCompatibility." + dev_kind); - if (device_scope_handler) { - String dev_scope = - (*device_scope_handler)(tok->virtual_device->target, tok->virtual_device->memory_scope); - return dev_scope; - } - return tok->virtual_device->memory_scope; -} - -StorageToken* TokenAllocatorMixed::Request(StorageToken* prototype) { - // calculate the size; - size_t size = GetMemorySize(prototype); - // search memory block in [size / match_range_, size * match_range_) - if (match_range_ == 0) { - return nullptr; - } - auto begin = free_.lower_bound(size / match_range_); - auto mid = free_.lower_bound(size); - auto end = free_.upper_bound(size * match_range_); - // search for memory blocks larger than requested - for (auto it = mid; it != end; ++it) { - StorageToken* tok = it->second; - bool dev_compatible = (GetDeviceCompatibleToken(tok) == GetDeviceCompatibleToken(prototype)); - if (tok->is_compatible(*prototype) || (dev_compatible)) { - ICHECK_EQ(tok->ref_counter, 0); - // Use exect matching strategy - if (size > tok->max_bytes) { - tok->max_bytes = size; - tok->ttype = prototype->ttype; - } - tok->ref_counter = prototype->ref_counter; - // find a exact match, erase from map and return - free_.erase(it); - return tok; - } - } - // then search for memory blocks smaller than requested space - for (auto it = mid; it != begin;) { - --it; - StorageToken* tok = it->second; - bool dev_compatible = (GetDeviceCompatibleToken(tok) == GetDeviceCompatibleToken(prototype)); - if (tok->is_compatible(*prototype) || (dev_compatible)) { - ICHECK_EQ(tok->ref_counter, 0); - // Use exect matching strategy - if (size > tok->max_bytes) { - tok->max_bytes = size; - tok->ttype = prototype->ttype; - } - tok->ref_counter = prototype->ref_counter; - // erase from map and return - free_.erase(it); - return tok; - } - } - return nullptr; -} - -StorageToken* TokenAllocatorMixed::Alloc(StorageToken* prototype, int64_t storage_id) { - size_t size = GetMemorySize(prototype); - prototype->max_bytes = size; - prototype->storage_id = storage_id; - data_.push_back(prototype); - return prototype; -} - -void TokenAllocatorMixed::CheckForRelease(StorageToken* tok) { - ICHECK_GE(tok->storage_id, 0); - ICHECK_GE(tok->ref_counter, 0); - if (tok->ref_counter == 0) { - free_.insert({tok->max_bytes, tok}); - } -} - -size_t TokenAllocatorMixed::GetSize2D(StorageToken* prototype) { - TensorType ttype = prototype->ttype; - ICHECK(ttype.defined()); - struct Shape { - const Array& shape; - int64_t operator[](size_t i) const { return *tir::as_const_int(shape[i]); } - int size() { return this->shape.size(); } - }; - auto shape = Shape{ttype->shape}; - int image_row_align = - prototype->virtual_device->target->GetAttr("image_base_address_alignment") - .value_or(Integer(64)) - ->value; - return runtime::GetTextureMemorySize(shape, ttype->dtype.bits(), ttype->dtype.lanes(), - prototype->virtual_device->memory_scope, - image_row_align); -} - -} // namespace relay -} // namespace tvm diff --git a/src/relay/backend/token_allocator.h b/src/relay/backend/token_allocator.h deleted file mode 100644 index 5524e6b2c634..000000000000 --- a/src/relay/backend/token_allocator.h +++ /dev/null @@ -1,129 +0,0 @@ -/* - * 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 relay/backend/token_allocator.h - * \brief Token allocation classes for backend - */ -#ifndef TVM_RELAY_BACKEND_TOKEN_ALLOCATOR_H_ -#define TVM_RELAY_BACKEND_TOKEN_ALLOCATOR_H_ - -#include -#include - -#include -#include -#include -#include -#include - -#include "../../runtime/texture.h" - -namespace tvm { -namespace relay { - -/*! A representation of a block of memory required at runtime on some device. */ -struct StorageToken { - /*! \brief Reference counter */ - int ref_counter{0}; - /*! \brief number of bytes */ - size_t max_bytes{0}; - /*! \brief The corresponding tensor type. */ - TensorType ttype{nullptr}; - /*! \brief VirtualDevice on which the memory will reside. */ - VirtualDevice virtual_device = VirtualDevice::FullyUnconstrained(); - /*! \brief The storage id */ - int64_t storage_id{-1}; - - bool is_valid() const { return !virtual_device->IsFullyUnconstrained(); } - - bool is_compatible(const StorageToken& that) const { - return virtual_device == that.virtual_device; - } - - std::string ToString() const { - std::ostringstream os; - os << "{storage_id: " << storage_id << ", max_bytes: " << max_bytes - << ", ttype: " << PrettyPrint(ttype) << ", virtual_device: " << virtual_device << "}"; - return os.str(); - } -}; - -/** - * @brief Memory manager for mixed mode memory types - */ -class TokenAllocatorMixed { - public: - /*! - * \brief ceil(size/word_size) to get number of words. - * \param size The original size. - * \param word_size The element size. - */ - static size_t DivRoundUp(size_t size, size_t word_size) { - return (size + word_size - 1) / word_size; - } - - /*! - * \brief Get the memory requirement. - * \param prototype The prototype token. - * \return The required memory size. - * - * TODO(mbs): Gf GetMemorySizeBytes in aot_executor_codegen.cc, - * CalculateRelayExprSizeBytes in utils.cc - */ - size_t GetMemorySize(StorageToken* prototype); - /*! - * \brief Request a storage token for a given prototype. - * \param prototype. The prototype storage token. - * \return The result token. - */ - StorageToken* Request(StorageToken* prototype); - /*! - * \brief Alloacte a storage token by consuming prototype - * \param prototype The prototype token. - * \param size The size of memory being requested. - */ - StorageToken* Alloc(StorageToken* prototype, int64_t storage_id); - /*! - * \brief Check if we can release token. - * \param tok The token to be released. - */ - void CheckForRelease(StorageToken* tok); - /*! - * \brief Get the texture 2d size requirement - * \param prototype The prototype token. - * \return The physical memory size. - */ - size_t GetSize2D(StorageToken* prototype); - - protected: - // free list of storage entry - std::multimap free_; - // all the storage resources available - std::vector data_; - - private: - // scale used for rough match - const size_t match_range_{16}; -}; - -} // namespace relay -} // namespace tvm - -#endif // TVM_RELAY_BACKEND_TOKEN_ALLOCATOR_H_ diff --git a/src/relay/backend/utils.cc b/src/relay/backend/utils.cc deleted file mode 100644 index b7453590742d..000000000000 --- a/src/relay/backend/utils.cc +++ /dev/null @@ -1,458 +0,0 @@ - -/* - * 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 relay/backend/util.cc - * \brief Relay backend utilities. - */ - -#include "utils.h" - -#include -#include -#include -#include - -#include "../../arith/scalable_expression.h" -#include "../../te/operation/create_primfunc.h" - -namespace tvm { -namespace relay { -namespace backend { - -TVM_REGISTER_NODE_TYPE(StorageInfoNode); - -TVM_STATIC_IR_FUNCTOR(ReprPrinter, vtable) - .set_dispatch([](const ObjectRef& ref, ReprPrinter* p) { - const auto* node = ref.as(); - p->stream << "StorageInfoNode(" - << "storage_ids=["; - for (auto id : node->storage_ids) { - p->stream << id << ","; - } - p->stream << "], virtual_devices=["; - for (const auto& virtual_device : node->virtual_devices) { - p->stream << virtual_device << ","; - } - p->stream << "], storage_size_in_bytes=["; - for (auto bytes : node->storage_sizes_in_bytes) { - p->stream << bytes << ","; - } - p->stream << "])"; - }); - -StorageInfo::StorageInfo(std::vector storage_ids, - std::vector virtual_devices, - std::vector storage_sizes_in_bytes) { - ICHECK_EQ(storage_ids.size(), virtual_devices.size()); - ICHECK_EQ(storage_ids.size(), storage_sizes_in_bytes.size()); - auto node = make_object(); - node->storage_ids = std::move(storage_ids); - node->virtual_devices = std::move(virtual_devices); - node->storage_sizes_in_bytes = std::move(storage_sizes_in_bytes); - data_ = std::move(node); -} - -// This is the legacy interface for devices as DLDeviceTypes (represented by integers) -TVM_REGISTER_GLOBAL("relay.ir.StorageInfo") - .set_body_typed([](const Array& sids, const Array& device_types, - const Array& sizes_in_bytes) { - std::vector sids_v; - sids_v.reserve(sids.size()); - for (auto s : sids) { - sids_v.push_back(s.IntValue()); - } - std::vector virtual_devices_v; - virtual_devices_v.reserve(device_types.size()); - for (const auto& device_type : device_types) { - virtual_devices_v.emplace_back(VirtualDevice::ForDeviceType(device_type)); - } - std::vector size_in_bytes_v; - size_in_bytes_v.reserve(sizes_in_bytes.size()); - for (auto s : sizes_in_bytes) { - size_in_bytes_v.push_back(s.IntValue()); - } - return StorageInfo(std::move(sids_v), std::move(virtual_devices_v), - std::move(size_in_bytes_v)); - }); - -TVM_REGISTER_GLOBAL("relay.ir.StorageInfoStorageIds").set_body_typed([](StorageInfo si) { - Array ids; - for (auto id : si->storage_ids) { - ids.push_back(id); - } - return ids; -}); - -// This is the legacy interface for devices as DLDeviceTypes (represented by integers) -TVM_REGISTER_GLOBAL("relay.ir.StorageInfoDeviceTypes").set_body_typed([](StorageInfo si) { - Array device_types; - for (const auto& virtual_device : si->virtual_devices) { - device_types.push_back(virtual_device->device_type()); - } - return device_types; -}); - -TVM_REGISTER_GLOBAL("relay.ir.StorageInfoStorageSizes").set_body_typed([](StorageInfo si) { - Array storage_sizes_in_bytes; - for (auto id : si->storage_sizes_in_bytes) { - storage_sizes_in_bytes.push_back(id); - } - return storage_sizes_in_bytes; -}); - -TVM_REGISTER_GLOBAL("relay.ir.StorageInfoVirtualDevices").set_body_typed([](StorageInfo si) { - Array virtual_devices; - for (auto id : si->virtual_devices) { - virtual_devices.push_back(id); - } - return virtual_devices; -}); - -TVM_REGISTER_NODE_TYPE(StaticMemoryPlanNode); - -StaticMemoryPlan::StaticMemoryPlan(Map expr_to_storage_info) { - auto n = make_object(); - n->expr_to_storage_info = std::move(expr_to_storage_info); - data_ = std::move(n); -} - -TVM_REGISTER_GLOBAL("relay.ir.StaticMemoryPlan") - .set_body_typed([](const Map& expr_to_storage_info) { - return StaticMemoryPlan(expr_to_storage_info); - }); - -size_t DivRoundUp(size_t size, size_t word_size) { return (size + word_size - 1) / word_size; } - -size_t GetMemorySizeBytes(const Array& shape, const DataType& dtype) { - size_t size = 1; - for (IndexExpr dim : shape) { - const int64_t* pval = tir::as_const_int(dim); - ICHECK(pval != nullptr) << "Cannot allocate memory symbolic tensor shape " << shape; - ICHECK_GE(*pval, 0) << "Cannot allocate memory for tensor with negative shape" << *pval; - size *= static_cast(pval[0]); - } - size *= DivRoundUp(dtype.bits() * dtype.lanes(), 8); - return size; -} - -int64_t CalculateRelayExprSizeBytes(const Type& expr_type) { - if (expr_type->IsInstance()) { - auto tuple_type = Downcast(expr_type); - int64_t size = 0; - for (const auto& field : tuple_type->fields) { - size += CalculateRelayExprSizeBytes(field); - } - return size; - } - auto tensor_type = expr_type.as(); - ICHECK(tensor_type); - auto shape = tensor_type->shape; - return GetMemorySizeBytes(tensor_type->shape, tensor_type->dtype); -} - -TVM_REGISTER_NODE_TYPE(FunctionInfoNode); - -FunctionInfo::FunctionInfo(Map workspace_sizes, Map io_sizes, - Map constant_sizes, - Map tir_primfuncs, - Map relay_primfuncs) { - ObjectPtr n = make_object(); - n->workspace_sizes = std::move(workspace_sizes); - n->io_sizes = std::move(io_sizes); - n->constant_sizes = std::move(constant_sizes); - n->tir_primfuncs = std::move(tir_primfuncs); - n->relay_primfuncs = std::move(relay_primfuncs); - data_ = std::move(n); -} - -TVM_STATIC_IR_FUNCTOR(ReprPrinter, vtable) - .set_dispatch([](const ObjectRef& ref, ReprPrinter* p) { - auto* node = static_cast(ref.get()); - p->stream << "FunctionInfoNode(\n" - << "workspace_sizes=" << node->workspace_sizes << ",\n io_sizes=" << node->io_sizes - << ",\n constant_sizes=" << node->constant_sizes - << ",\n tir_primfuncs=" << node->tir_primfuncs - << ",\n relay_primfuncs=" << node->relay_primfuncs << ")"; - }); - -ExecutorCodegenMetadata::ExecutorCodegenMetadata( - Array inputs, Array input_tensor_types, Array outputs, - Array output_tensor_types, Array pools, Array devices, - String executor, String mod_name, String interface_api, bool unpacked_api, - Integer workspace_alignment, Integer constant_alignment, - Map pool_inputs, - Map io_pool_allocations) { - auto n = make_object(); - n->inputs = inputs; - n->input_tensor_types = input_tensor_types; - n->outputs = outputs; - n->output_tensor_types = output_tensor_types; - n->pools = pools; - n->devices = devices; - n->executor = executor; - n->interface_api = interface_api; - n->unpacked_api = unpacked_api; - n->mod_name = mod_name; - n->workspace_alignment = workspace_alignment; - n->constant_alignment = constant_alignment; - n->pool_inputs = pool_inputs; - n->io_pool_allocations = io_pool_allocations; - data_ = std::move(n); -} - -TVM_REGISTER_NODE_TYPE(ExecutorCodegenMetadataNode); - -Array GetPassPrefix(bool is_homogeneous, bool is_vm) { - Array pass_seqs; - // TODO(mbs): Would be nice to get spans on all diagnostics, but since they arg forgotton - // by most passes there's little utility in including this now. Plus we'd need to only do - // this if there's no existing spans to work from. - // pass_seqs.push_back(parser::AnnotateSpans()); - Array entry_functions{"main"}; - pass_seqs.push_back(transform::RemoveUnusedFunctions(entry_functions)); - pass_seqs.push_back(transform::ToBasicBlockNormalForm()); - // Run all dialect legalization passes. - pass_seqs.push_back(relay::qnn::transform::Legalize()); - - // Legalize pass is restricted to homogeneous execution for now. - if (is_homogeneous) { - pass_seqs.push_back(transform::Legalize()); - } - - pass_seqs.push_back(transform::SimplifyInference()); - - if (is_vm) { - // eta expand to support constructors in argument position - pass_seqs.push_back(transform::EtaExpand( - /* expand_constructor */ true, /* expand_global_var */ false)); - } - - PackedFunc fskip = PackedFunc([](TVMArgs args, TVMRetValue* rv) { - Expr expr = args[0]; - if (auto* call_node = expr.as()) { - auto op_node = call_node->op.as(); - if (op_node->name == "cast") { - auto attrs = call_node->attrs.as(); - if (attrs->dtype == DataType::Int(32)) { - *rv = true; - } - } - } - *rv = false; - }); - pass_seqs.push_back(transform::EliminateCommonSubexpr(fskip)); - pass_seqs.push_back(transform::CombineParallelConv2D(3)); - pass_seqs.push_back(transform::CombineParallelDense(3)); - pass_seqs.push_back(transform::CombineParallelBatchMatmul(3)); - pass_seqs.push_back(transform::FoldConstant()); - pass_seqs.push_back(transform::FoldScaleAxis()); - pass_seqs.push_back(transform::SimplifyExpr()); - pass_seqs.push_back(transform::CanonicalizeCast()); - pass_seqs.push_back(transform::CanonicalizeOps()); - pass_seqs.push_back(transform::FlattenAtrousConv()); - - // Alter layout transformation is currently only applied to homogeneous execution. - if (is_homogeneous) { - if (!is_vm) { - pass_seqs.push_back(transform::InferType()); - } - pass_seqs.push_back(transform::AlterOpLayout()); - pass_seqs.push_back(transform::SimplifyExprPostAlterOp()); - } - - // Fast math optimizations. - pass_seqs.push_back(transform::FastMath()); - pass_seqs.push_back(transform::FoldConstant()); - - return pass_seqs; -} - -std::unordered_map -TargetModuleMapToTargetStrModuleMap(Map input_map) { - std::unordered_map std_map; - for (auto kv : input_map) { - std_map[kv.first] = kv.second; - } - return std_map; -} - -Map TargetStrModuleMapToTargetModuleMap( - std::unordered_map input_map) { - Map tvm_map; - for (auto kv : input_map) { - tvm_map.Set(kv.first, kv.second); - } - return tvm_map; -} - -void UpdateAutoSchedulerOpWeights(const IRModule& module) { - const auto* te_compiler_update_weights = - runtime::Registry::Get("auto_scheduler.relay_integration.te_compiler_update_weights"); - - ICHECK(te_compiler_update_weights != nullptr) - << "auto_scheduler.relay_integration.te_compiler_update_weights"; - - Map weight_map = - module->GetAttr>("op_weights", Map()).value(); - - (*te_compiler_update_weights)(weight_map); -} - -std::vector ShapeToJSON(tvm::Array shape) { - std::vector ret; - for (IndexExpr dim : shape) { - const int64_t* pval = tir::as_const_int(dim); - ret.push_back(*pval); - } - return ret; -} - -relay::Function BindParamsByName(relay::Function func, - const std::unordered_map& params) { - std::unordered_map name_dict; - std::unordered_set repeat_var; - for (auto arg : func->params) { - const auto& name = arg->name_hint(); - if (name_dict.count(name)) { - repeat_var.insert(name_dict[name]); - } else { - name_dict[name] = arg; - } - } - - std::unordered_map bind_dict; - for (auto& kv : params) { - if (name_dict.count(kv.first) == 0) { - continue; - } - auto arg = name_dict.at(kv.first); - if (repeat_var.count(arg)) { - LOG(FATAL) << "Multiple args in the function have name " << kv.first; - } - bind_dict[arg] = Constant(kv.second); - } - Expr bound_expr = relay::Bind(func, bind_dict); - Function ret = Downcast(bound_expr); - ICHECK(ret.defined()) << "The returning type is expected to be a Relay Function." - << "\n"; - return ret; -} - -void BindParamsInModule(IRModule mod, - const std::unordered_map& params) { - if (!params.empty()) { - BaseFunc base_func = mod->Lookup("main"); - ICHECK(base_func->IsInstance()); - auto f = relay::backend::BindParamsByName(Downcast(base_func), params); - auto gvar = mod->GetGlobalVar("main"); - mod->Add(gvar, f); - } -} - -void BindParamsInModule(IRModule mod, Map params) { - std::unordered_map params_tmp; - for (const auto& kv : params) { - params_tmp[kv.first] = kv.second; - } - BindParamsInModule(mod, params_tmp); -} - -/*! - * \brief A default TE compute to TIR compute. - * \param args The inputs/outputs of the TE compute graph. - * \param constants The constants bound to TIR - * \param allow_extern_op Whether to allow extern operation in TE. - * \return The TIR converted; NullOpt if not supported (dynamic shape) - */ -Optional DefaultTIRConverterImpl(const Array& args, - const Array& constants, - bool allow_extern_op) { - using namespace ::tvm::te; - std::vector stack; - std::unordered_set visited; - for (const Tensor& v : args) { - for (const PrimExpr& e : v->shape) { - // Dynamic shape is not supported for now - if (!e->IsInstance()) { - return NullOpt; - } - } - if (!visited.count(v.get())) { - visited.insert(v.get()); - stack.push_back(v); - } - } - while (!stack.empty()) { - Tensor tensor = stack.back(); - stack.pop_back(); - if (tensor->op->IsInstance()) { - // do nothing - } else if (tensor->op->IsInstance() || - (allow_extern_op && tensor->op->IsInstance())) { - Array inputs = tensor->op->InputTensors(); - for (const Tensor& v : inputs) { - if (!visited.count(v.get())) { - visited.insert(v.get()); - stack.push_back(v); - } - } - } else { - return NullOpt; - } - } - PrimFunc func = te::CreatePrimFuncWithConstants(args, constants, DataType::Int(64)); - bool dynamic_loop_extent = false; - tir::PostOrderVisit(func->body, [&dynamic_loop_extent](const ObjectRef& obj) -> void { - if (const auto* loop = obj.as()) { - if (!loop->extent->IsInstance() && - !tvm::arith::ContainsVscaleCall(loop->extent)) { - dynamic_loop_extent = true; - } - } - }); - if (dynamic_loop_extent) { - return NullOpt; - } - return func; -} - -TVM_REGISTER_GLOBAL("relay.backend.tir_converter.default") - .set_body_typed([](const Array& args, - const Array& constants) -> Optional { - return DefaultTIRConverterImpl(args, constants, false); - }); - -TVM_REGISTER_GLOBAL("relay.backend.tir_converter.allow_extern") - .set_body_typed([](const Array& args, - const Array& constants) -> Optional { - return DefaultTIRConverterImpl(args, constants, true); - }); - -TVM_REGISTER_GLOBAL("relay.backend.GetPassPrefixSeq") - .set_body_typed([](bool is_homogeneous, bool is_vm) { - auto pass_seqs = GetPassPrefix(is_homogeneous, is_vm); - transform::Sequential seq(pass_seqs); - return seq; - }); - -} // namespace backend -} // namespace relay -} // namespace tvm diff --git a/src/relay/backend/utils.h b/src/relay/backend/utils.h deleted file mode 100644 index acaea425d178..000000000000 --- a/src/relay/backend/utils.h +++ /dev/null @@ -1,763 +0,0 @@ -/* - * 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 relay/backend/utils.h - * \brief Utils function for backend - */ -#ifndef TVM_RELAY_BACKEND_UTILS_H_ -#define TVM_RELAY_BACKEND_UTILS_H_ - -#include -#include -#include -#include -#include -#include -#include -#include -#include -#include -#include - -#include -#include -#include -#include -#include -#include -#include -#include - -#include "../../runtime/meta_data.h" -#include "../../target/metadata.h" -#include "tvm/runtime/ndarray.h" - -namespace tvm { -namespace relay { - -namespace tec { -class TECompiler; -} - -namespace backend { -using Pass = tvm::transform::Pass; - -/*! \brief Describes the type of kernel call emitted. */ -enum CallType { - /*! - * \brief Emit PackedFunc calls bound just-in-time using TVMBackend* functions. - * - * When this type is selected, assumes all operators must be called via TVMFuncCall. Given the - * implementation of TVMFuncCall in the C++ runtime, this in practice implies that those - * functions are of type TVMBackendPackedCFunc. - * - * The following code is emitted at call sites to call a function named `func`: - * void* func_ptr = TVMBackendGetFuncFromEnv("func"); - * TVMFuncCall(func_ptr, values, tcodes, num_args, ret_values, ret_tcodes) - * - * The arguments given to the tir::Call node are encoded into `values`, `tcodes`, and `num_args` - * by LowerTVMBuiltin TIR transform. - * - * If `resource_handle` is passed to `func`, it is determined by TVMFuncCall (often, - * `resource_handle` is registered with the C++ runtime to provide a `this` equivalent when - * `func` is implemented in C). - * - * Compatible with both C++ and C runtimes, implemented with the C runtime only. - */ - kPacked, // Emit tir.call_packed and wrap all arguments in DLTensor. - - /*! - * \brief Directly call a TVMBackendPackedCFunc named according to the tir::Call. - * - * When this type is selected, assumes all operators are implemented in functions of type - * `TVMBackendPackedCFunc` and should be called directly. That is, presumes at the time of - * downstream compilation that there is a symbol named after the 0th arg to tir::Call of - * type `TVMBackendPackedCFunc`. This situation should occur when target_host == target. - * - * The following code is emitted at call sites to call a function named `func`: - * func(values, tcodes, num_args, ret_values, ret_tcodes, resource_handle) - * - * The arguments given to the tir::Call node are encoded into `values`, `tcodes`, and `num_args` - * by LowerTVMBuiltin TIR transform. - * - * `resource_handle` is encoded as the final argument to the tir::Call node. In practice, it is - * always the device context parameter when not null. At present, the implementation does not - * support forwarding device context parameters to CPacked. - * - * Compatible with the C runtime and C++ runtime (so long as target_host == target). Implemented - * in the same scenarios. - */ - kCPacked, // Emit tir.call_cpacked and wrap all arguments in DLTensor. - - /*! \brief Directly call a function accepting the `data` arrays as args. - * - * When this type is selected, assumes all operaotrs are implemented in C functions whose - * arguments are 1-to-1 with those in the tir::Call. DLTensor arguments are encoded as just the - * `data` parameters (i.e. no DLTensor object is passed along). - * - * The following code is emitted at call sites to a function named `func`: - * func(void* arg0, void* arg1, ..., void* argN) // no resource_handle - * -or- - * func(void* arg0, void* arg1, ..., void* argN, void* resource_handle) // with resource_handle - * - * `resource_handle` is encoded as the final argument to the tir::Call node. In practice, it is - * always the device context parameter when not null. - * - * Compatible with the C runtime and C++ runtime (so long as target_host == target). Implemented - * with the C runtime only. - */ - kUnpacked, // Emit tir.call_extern passing only the `data` part of DLTensors. -}; - -/*! - * \brief Structure that can be optionally used by the executor codegen - */ -class ExecutorCodegenMetadataNode : public Object { - public: - /*! \brief input information for the main function */ - Array inputs; - /*! \brief input tensor type information */ - Array input_tensor_types; - /*! \brief output information for the main function */ - Array outputs; - /*! \brief output tensor type information */ - Array output_tensor_types; - /*! \brief pool information for the main function */ - Array pools; - /*! \brief device contexts information for the main function */ - Array devices; - /*! \brief the executor to be used to run the model */ - String executor = runtime::kTvmExecutorGraph; - /*! \brief The external API (packed or c) in use */ - String interface_api; - /*! \brief The internal API (packed or unpacked) in use */ - bool unpacked_api; - /*! \brief Alginment of the workspace in bytes */ - Integer workspace_alignment; - /*! \brief Alginment of the constants in bytes */ - Integer constant_alignment; - /*! \brief the input var names that correspond to pool_inputs */ - Optional> pool_inputs; - /*! \brief the I/O tensor to PoolAllocations if any*/ - Map io_pool_allocations; - - String mod_name = ""; - - void VisitAttrs(tvm::AttrVisitor* v) { - v->Visit("inputs", &inputs); - v->Visit("input_tensor_types", &input_tensor_types); - v->Visit("outputs", &outputs); - v->Visit("output_tensor_types", &output_tensor_types); - v->Visit("pools", &pools); - v->Visit("devices", &devices); - v->Visit("executor", &executor); - v->Visit("interface_api", &interface_api); - v->Visit("unpacked_api", &unpacked_api); - v->Visit("workspace_alignment", &workspace_alignment); - v->Visit("constant_alignment", &constant_alignment); - v->Visit("pool_inputs", &pool_inputs); - v->Visit("io_pool_allocations", &io_pool_allocations); - v->Visit("mod_name", &mod_name); - } - - static constexpr const char* _type_key = "MetadataObj"; - TVM_DECLARE_FINAL_OBJECT_INFO(ExecutorCodegenMetadataNode, Object); -}; - -/*! - * \brief Managed reference to ExecutorCodegenMetadataNode. - */ -class ExecutorCodegenMetadata : public ObjectRef { - public: - TVM_DLL ExecutorCodegenMetadata(Array inputs, Array input_tensor_types, - Array outputs, Array output_tensor_types, - Array pools, Array devices, String executor, - String mod_name, String interface_api = "packed", - bool unpacked_api = false, Integer workspace_alignment = 16, - Integer constant_alignment = 16, - Map pool_inputs = - Map(), - Map io_pool_allocations = {}); - TVM_DEFINE_MUTABLE_OBJECT_REF_METHODS(ExecutorCodegenMetadata, ObjectRef, - ExecutorCodegenMetadataNode); -}; - -/*! - * \brief The static storage information for each Tensor in the result of a Relay expression - * (as per relay::FlattenTupleType). - */ -class StorageInfoNode : public Object { - public: - // TODO(mbs): Switch from struct-of-array to array-of-struct repr throughout. - /*! \brief The set of storage ids where the expression is stored. */ - std::vector storage_ids; - /* \brief The virtual devices these expressions are stored within. */ - std::vector virtual_devices; - /* \brief The sizes of each storage element, in bytes. */ - std::vector storage_sizes_in_bytes; - - // TODO(@jroesch): expose the fields - void VisitAttrs(AttrVisitor* v) {} - - static constexpr const char* _type_key = "relay.StorageInfo"; - TVM_DECLARE_FINAL_OBJECT_INFO(StorageInfoNode, Object); -}; - -/*! \brief The storage information for a single expression. */ -class StorageInfo : public ObjectRef { - public: - StorageInfo(std::vector storage_ids, std::vector virtual_devices, - std::vector storage_sizes_in_bytes); - TVM_DEFINE_OBJECT_REF_METHODS(StorageInfo, ObjectRef, StorageInfoNode); -}; - -/*! - * \brief The result of static memory planning. - */ -class StaticMemoryPlanNode : public Object { - public: - Map expr_to_storage_info; - - void VisitAttrs(AttrVisitor* v) { v->Visit("expr_to_storage_info", &expr_to_storage_info); } - - static constexpr const char* _type_key = "relay.StaticMemoryPlan"; - TVM_DECLARE_FINAL_OBJECT_INFO(StaticMemoryPlanNode, Object); -}; - -/*! \brief The result of running static memory planning. */ -class StaticMemoryPlan : public ObjectRef { - public: - explicit StaticMemoryPlan(Map expr_to_storage_info); - TVM_DEFINE_OBJECT_REF_METHODS(StaticMemoryPlan, ObjectRef, StaticMemoryPlanNode); -}; - -struct FunctionInfoNode : public Object { - Map workspace_sizes; - Map io_sizes; - Map constant_sizes; - Map tir_primfuncs; - Map relay_primfuncs; - - void VisitAttrs(tvm::AttrVisitor* v) { - v->Visit("workspace_sizes", &workspace_sizes); - v->Visit("io_sizes", &io_sizes); - v->Visit("constant_sizes", &constant_sizes); - v->Visit("tir_primfuncs", &tir_primfuncs); - v->Visit("relay_primfuncs", &relay_primfuncs); - } - - static constexpr const char* _type_key = "relay.backend.FunctionInfo"; - TVM_DECLARE_FINAL_OBJECT_INFO(FunctionInfoNode, Object); -}; - -class FunctionInfo : public ObjectRef { - public: - FunctionInfo(Map workspace_sizes, Map io_sizes, - Map constant_sizes, Map tir_primfuncs, - Map relay_primfuncs); - - TVM_DEFINE_MUTABLE_OBJECT_REF_METHODS(FunctionInfo, ObjectRef, FunctionInfoNode); -}; - -/*! - * \brief Calculate the bytes of memory needed to hold a tensor of a given shape and data type. - * \param shape The shape of the tensor - * \param dtype The data type of the tensor - */ -size_t GetMemorySizeBytes(const Array& shape, const DataType& dtype); - -/*! - * \brief Calculate the storage required to store the type of relay.Expr - * - * \param func The relay expr for which the storage is calculated - */ -int64_t CalculateRelayExprSizeBytes(const Type& expr_type); - -/*! - * \brief Executor generator artifacts. Those artifacts are subsequently - * used by the relay build process. - */ -struct LoweredOutput { - std::string graph_json; - Map lowered_funcs; - Array external_mods; - Map function_metadata; - /*! - * \brief Map from constant names (allocated by the codegen as constants are encountered) - * to the constant's value. - */ - std::unordered_map params; - ExecutorCodegenMetadata metadata; -}; - -/*! - * \brief This class is needed to avoid a GCC 5 bug that prevents maps containing enums from being - compiled. If i386 GCC version is increased, we can remove it. - */ -struct EnumClassHash { - template - std::size_t operator()(T t) const { - return static_cast(t); - } -}; - -/*! - * \brief A helper to expand the params by adding the ones used in a given expression. - */ -struct ConstantUpdater : public ExprVisitor { - public: - ConstantUpdater(const std::string& symbol, - std::unordered_map* params) - : symbol_(symbol), params_(params) {} - - void VisitExpr_(const ConstantNode* cn) final { - std::string name = symbol_ + "_const_" + std::to_string(const_idx_++); - VLOG(1) << "binding '" << name << "' to constant of type " << PrettyPrint(cn->checked_type()); - (*params_)[name] = cn->data; - } - - private: - int const_idx_{0}; - std::string symbol_; - std::unordered_map* params_; -}; - -/*! - * \brief A function to update the params with constants found in an external function. - * \param func The function from which to get the constant params. - * \param params The params to update with the constants. - */ -inline void UpdateConstants(BaseFunc func, - std::unordered_map* params) { - VLOG_CONTEXT << "UpdateConstants"; - VLOG(1) << "updating constants for:" << std::endl << PrettyPrint(func); - auto codegen = func->GetAttr(attr::kCompiler); - ICHECK(codegen.defined()) << "No external codegen is set"; - std::string codegen_name = codegen.value(); - const auto name_node = func->GetAttr(tvm::attr::kGlobalSymbol); - std::string symbol = std::string(name_node.value()); - std::string const_update_name = "relay.ext." + codegen_name + ".constant_updater"; - // Get the constant updater for the external codegen - auto pf = tvm::runtime::Registry::Get(const_update_name); - // If the backend hasn't registered a constant updater, use a default one - if (pf == nullptr) { - ConstantUpdater const_visit(symbol, params); - const_visit(func); - } else { - Map constants = (*pf)(func, symbol); - for (const auto& it : constants) { - std::string const_name(it.first); - // Constant names should begin this the compiler name (to avoid conflicts) - ICHECK(const_name.find(codegen_name) == 0) - << "External constant names must start with compiler name"; - (*params)[const_name] = it.second; - } - } - for (const auto& pair : *params) { - VLOG(1) << "Constants: " << pair.first << " = " << PrettyPrint(pair.second); - } -} - -/*! - * \brief A simple wrapper around ExprFunctor for a single argument case. - * The result of visit is memoized. - */ -template -class MemoizedExprTranslator : public ::tvm::relay::ExprFunctor { - using BaseFunctor = ::tvm::relay::ExprFunctor; - - public: - /*! \brief virtual destructor */ - virtual ~MemoizedExprTranslator() {} - - /*! - * \brief The memoized call. - * \param n The expression node. - * \return The result of the call - */ - virtual OutputType VisitExpr(const Expr& n) { - ICHECK(n.defined()); - auto it = memo_.find(n); - if (it != memo_.end()) { - return it->second; - } - auto res = BaseFunctor::VisitExpr(n); - memo_[n] = res; - return res; - } - - protected: - /*! \brief Internal map used for memoization. */ - std::unordered_map memo_; -}; - -/*! - * \brief Get the Packed Func - * - * \param func_name - * \return const PackedFunc* - */ -inline const PackedFunc* GetPackedFunc(const std::string& func_name) { - return tvm::runtime::Registry::Get(func_name); -} - -/*! - * \brief Get a typed packed function. - * - * \param func_name - * \return const PackedFunc* - */ -template -inline const runtime::TypedPackedFunc GetTypedPackedFunc(const std::string& func_name) { - auto* pf = GetPackedFunc(func_name); - ICHECK(pf != nullptr) << "can not find packed function"; - return runtime::TypedPackedFunc(*pf); -} - -/*! - * \brief Extract shape from an IndexExpr array to std::vector - * - * \param shape The shape in Array - * \return The converted shape in std::vector - */ -inline std::vector GetIntShape(const Array& shape) { - std::vector ret; - for (const auto& dim : shape) { - const int64_t* pval = tir::as_const_int(dim); - ret.push_back(pval ? *pval : -1); - } - return ret; -} - -/*! - * \brief Convert type to string - * - * \param typ - * \return std::string string format of type - */ -inline std::string DType2String(const tvm::DataType dtype) { - std::ostringstream os; - if (dtype.is_float()) { - os << "float"; - } else if (dtype.is_int()) { - os << "int"; - } else if (dtype.is_uint()) { - os << "uint"; - } else if (dtype.is_bfloat16()) { - os << "bfloat"; - } else if ((*GetPackedFunc("runtime._datatype_get_type_registered"))(dtype.code())) { - os << "custom[" - << (*GetPackedFunc("runtime._datatype_get_type_name"))(dtype.code()).operator std::string() - << "]"; - } else { - LOG(FATAL) << "Unknown type with code " << static_cast(dtype.code()); - } - os << dtype.bits(); - return os.str(); -} - -/*! - * \brief Bind params to function by using name - * \param func Relay function - * \param params params dict - * \return relay::Function - */ -relay::Function BindParamsByName(relay::Function func, - const std::unordered_map& params); - -/*! - * \brief Bind params to the main function in Relay module, using BindParamsByName - * \param mod Relay module - * \param params params dict - */ -void BindParamsInModule(IRModule mod, - const std::unordered_map& params); - -void BindParamsInModule(IRModule mod, Map params); - -/*! - * \brief Extract the shape from a Relay tensor type. - * \param type The provided type. - * \return The extracted shape in a list. - */ -inline std::vector GetShape(const Type& type) { - const auto* ttype = type.as(); - ICHECK(ttype) << "Expect TensorTypeNode"; - std::vector shape; - for (size_t i = 0; i < ttype->shape.size(); ++i) { - auto* val = ttype->shape[i].as(); - ICHECK(val); - shape.push_back(val->value); - } - return shape; -} - -/*! - * \brief Check if a call has the provided name. - * \param call A Relay call node. - * \param op_name The name of the expected call. - * \return true if the call's name is equivalent to the given name. Otherwise, - * false. - */ -inline bool IsOp(const CallNode* call, const std::string& op_name) { - const auto* op_node = call->op.as(); - ICHECK(op_node) << "Expects a single op."; - Op op = GetRef(op_node); - return op == Op::Get(op_name); -} - -/*! - * \brief Retrieve the "root" op nested inside a fused call, such as conv2d in relu(add(conv2d)) - * \param call A Relay call node. Typically nn.relu when called the first time. - * \param depth The number of calls before the root op, counting from current_call. - * \param expected_op_names The names of ops in this fused call. Example: {"nn.conv2d", "add", - * "nn.relu"} - * \return A CallNode corresponding to the root op, whose name is expected_op_names[0] - */ -inline const CallNode* GetRootCall(const CallNode* current_call, int depth, - const std::vector& expected_op_names) { - ICHECK(current_call && depth >= 0 && static_cast(depth) < expected_op_names.size() && - IsOp(current_call, expected_op_names[depth])); - - if (depth == 0) { - return current_call; - } - - ICHECK_GT(current_call->args.size(), 0); - size_t valid_node_idx = 0; - while (valid_node_idx < current_call->args.size() && - current_call->args[valid_node_idx].as()) { - valid_node_idx++; - } - while (valid_node_idx < current_call->args.size() && - !(IsOp(current_call->args[valid_node_idx].as(), expected_op_names[depth - 1]))) { - valid_node_idx++; - } - const auto* next_call = current_call->args[valid_node_idx].as(); - return GetRootCall(next_call, depth - 1, expected_op_names); -} - -/*! - * \brief Retrieve the "root" op nested inside a fused call, such as conv2d in relu(add(conv2d)) - * Unlike the previous definition, it does not verify operator names of intermediate nodes. Instead, - * it recursively visit child nodes until it finds a call node with the given op_name. - * \param call A Relay call node. - * \param op_name The name of an op to look for, such as ""nn.conv2d". - * \return A CallNode corresponding to the root op with the given op_name - */ -inline const CallNode* GetRootCall(const CallNode* current_call, const std::string& op_name) { - if (current_call == nullptr) return nullptr; - if (IsOp(current_call, op_name)) return current_call; - - ICHECK_GT(current_call->args.size(), 0); - - const auto* next_call = current_call->args[0].as(); - return GetRootCall(next_call, op_name); -} - -/*! - * \brief Retrieve the expected "root" op nested inside a fused call, such as conv2d in - * relu(add(conv2d)) - * \param call A Relay call node. Typically nn.relu when called the first time. - * \param max_depth The maximum number of calls before the root op, counting from current_call. - * \param op_name The name of expected "root" op in this fused call. - * \return A CallNode corresponding to the root op - */ -inline const CallNode* GetRootCall(const CallNode* current_call, int max_depth, - const std::string& op_name) { - ICHECK(current_call && max_depth >= 0); - - if (max_depth == 0) { - ICHECK(current_call && IsOp(current_call, op_name)); - return current_call; - } - if (IsOp(current_call, op_name)) { - return current_call; - } - - ICHECK_GT(current_call->args.size(), 0); - - size_t valid_node_idx = 0; - while (valid_node_idx < current_call->args.size() && - current_call->args[valid_node_idx].as()) { - valid_node_idx++; - } - - const auto* next_call = current_call->args[valid_node_idx].as(); - return GetRootCall(next_call, max_depth - 1, op_name); -} - -/*! - * \brief Get the external symbol of the Relay function name. - * - * \param func The provided function. - * \return An external symbol. - */ -inline std::string GetExtSymbol(const Function& func) { - const auto name_node = func->GetAttr(tvm::attr::kGlobalSymbol); - ICHECK(name_node.defined()) << "Fail to retrieve external symbol."; - return std::string(name_node.value()); -} - -/*! - * \brief Return whether the auto scheduler is enabled in the pass context. - */ -inline bool IsAutoSchedulerEnabled() { - return transform::PassContext::Current() - ->GetConfig("relay.backend.use_auto_scheduler", Bool(false)) - .value(); -} - -/*! - * \brief Return whether the meta schedule is enabled in the pass context. - */ -inline bool IsMetaScheduleEnabled() { - return transform::PassContext::Current() - ->GetConfig("relay.backend.use_meta_schedule", Bool(false)) - .value(); -} - -/*! \brief Consider MetaSchedule's dispatch option. */ -inline int UseMetaScheduleDispatch() { - return transform::PassContext::Current() - ->GetConfig("relay.backend.use_meta_schedule_dispatch", Integer(0)) - .value() - ->value; -} -/*! - * \brief Method in TECompiler to convert TE compute to scheduleable TIR - * \param args The arguments of the TE compute - * \param constants The constants used in AllocateConst - * \return NullOpt if conversion fails; Otherwise the converted TIR - * \note This method could be further used as a task filtering mechanism in task extraction - */ -using FTECompilerTIRConverter = runtime::TypedPackedFunc< // - Optional( // - const Array& args, // - const Array& constants)>; - -/*! \brief Return a task filter for AutoTIR according to `relay.backend.tir_converter` */ -inline FTECompilerTIRConverter GetTIRConverter() { - String name = transform::PassContext::Current() - ->GetConfig("relay.backend.tir_converter", "default") - .value(); - const PackedFunc* f = runtime::Registry::Get("relay.backend.tir_converter." + name); - ICHECK(f != nullptr) << "IndexError: Cannot find TIR converter: " << name; - return FTECompilerTIRConverter(*f); -} - -/*! \brief Converts a PrimFunc to IRModule. */ -inline IRModule PrimFuncToIRModule(tir::PrimFunc f) { - f = WithAttrs(f, Map{ - {tvm::attr::kGlobalSymbol, String("main")}, - {tvm::tir::attr::kNoAlias, Bool(1)}, - }); - return IRModule({{GlobalVar("main"), f}}); -} - -/*! - * \brief Get the sequence of Relay optimization passes based on backend type. - * The prefix of the Relay passes almost overlaps between the vm and graph backend, with some slight - * difference. This function unifies the shared optimization pass prefix between vm and graph - * runtime, and returns the pass prefix given the backend type. - * - * \param is_homogeneous True if all primitives are to be executed on the same device and target. - * \param is_vm True if passes are to be used for the vm executor. - * \return An array of passes. - */ -Array GetPassPrefix(bool is_homogeneous, bool is_vm); - -/*! \brief Target hash function */ -struct TargetStrHash { - /*! - * \brief Calculate the hash code of a Target based on the string value of the Target KIND. - Note that this hash should NOT be used in new usecases, equality of targets based on their - value is not well-defined. - This will be removed when maps from Targets to IRModules are removed from the codebase. - * \param target The Target to hash - * \return String hash of the target - */ - size_t operator()(const Target& target) const { - std::string s(target->kind->name); - return String::StableHashBytes(s.c_str(), s.size()); - } -}; - -/*! \brief Target equality function based on the string value of Target -Note that this equality function should NOT be used in new usecases, equality of targets based on -their value is not well-defined. This will be removed when maps from Targets to IRModules are -removed from the codebase.*/ -struct TargetStrEqual { - /*! - * \brief Check if the two Targets are equal - * \param target One Target - * \param other_target The other Target - * \return String equality of the targets - */ - const bool operator()(const Target& target, const Target& other_target) const { - TargetStrHash target_hash = TargetStrHash(); - return target_hash(target) == target_hash(other_target); - } -}; - -/*! - * \brief Convert a Map to std::unordered_map Target equality is currently based on pointer equality, which is a problem since - * we have a lot of Map in the codebase. This function converts the map to a - * version that is keyed based on string value of the Target instead. Note that once we remove - * Map, this function will be removed. - * \param input_map The map to convert - * \return The converted map - */ -std::unordered_map -TargetModuleMapToTargetStrModuleMap(Map input_map); - -/*! - * \brief Convert a std::unordered_map to - * Map This function is a helper that undoes TargetModuleMapToTargetStr. Note that - * once we remove Map, this function will be removed. - * \param input_map The map to convert - * \return The converted map - */ -Map TargetStrModuleMapToTargetModuleMap( - std::unordered_map input_map); - -/*! - * \brief Call "weight update callback" to communicate op weights seen during Relay module - * lowering back to the auto scheduler. - * Op weights refer to the number of times each distinct op/workload appears in a given module. - * It is called "use_count" in TECompiler. - * \param IRModule after lowering by LowerTEPass. - */ -void UpdateAutoSchedulerOpWeights(const IRModule& module); - -/*! - * \brief Extract shape from expr to vector - * - * \param shape - * \return std::vector - */ -std::vector ShapeToJSON(tvm::Array shape); - -} // namespace backend -} // namespace relay -} // namespace tvm - -#endif // TVM_RELAY_BACKEND_UTILS_H_ diff --git a/src/relay/backend/vm/compiler.cc b/src/relay/backend/vm/compiler.cc deleted file mode 100644 index 848c23eba63b..000000000000 --- a/src/relay/backend/vm/compiler.cc +++ /dev/null @@ -1,1221 +0,0 @@ -/* - * 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 src/relay/backend/vm/compiler.cc - * \brief A compiler from relay::Module to the VM byte code. - */ - -#include "compiler.h" - -#include -#include -#include -#include -#include -#include -#include -#include -#include -#include -#include -#include -#include -#include - -#include -#include -#include -#include -#include -#include - -#include "../../../driver/internal_driver_api.h" -#include "../../../target/metadata_module.h" -#include "../../../target/source/codegen_source_base.h" -#include "../../op/annotation/annotation.h" -#include "../../op/memory/device_copy.h" -#include "../../op/op_common.h" -#include "../../transforms/device_aware_visitors.h" -#include "../../transforms/pass_utils.h" -#include "../utils.h" -#include "./compiler.h" - -namespace tvm { -namespace relay { - -namespace transform { - -Pass LambdaLift(); -Pass LabelOps(); - -Pass MemoryPlan() { - auto f = tvm::runtime::Registry::Get("relay.transform.MemoryPlan"); - ICHECK(f != nullptr) << "unable to load the memory planning pass"; - return (*f)(); -} - -Pass LiftConstants() { - auto f = tvm::runtime::Registry::Get("relay.transform.LiftConstants"); - ICHECK(f != nullptr) << "unable to load the constant lifting pass"; - return (*f)(); -} - -} // namespace transform - -namespace vm { - -using namespace tvm::runtime; -using namespace tvm::runtime::vm; -using namespace relay::transform; - -/*! \brief The host device is always stored at device index 0. */ -constexpr Index kHostDeviceIndex = 0; - -// (@jroesch): VM passes, eventually declare as passes. -bool IsClosure(const Function& func); - -// Represent a runtime object that's going to be matched by pattern match expressions -struct MatchValue { - virtual ~MatchValue() {} -}; -using MatchValuePtr = std::shared_ptr; - -// A runtime object that resides in a register -struct RegisterValue : MatchValue { - // The register num - RegName register_num; - - explicit RegisterValue(RegName reg) : register_num(reg) {} - - ~RegisterValue() {} -}; - -// The value is a field of another runtime object -struct AccessField : MatchValue { - MatchValuePtr parent; - // Field index - size_t index; - // Runtime register num after compiling the access field path - RegName reg{-1}; - - AccessField(MatchValuePtr parent, size_t index) : parent(parent), index(index) {} - - ~AccessField() {} -}; - -/*! - * \brief Condition in a decision tree - */ -struct ConditionNode { - virtual ~ConditionNode() {} -}; - -using ConditionObjectPtr = std::shared_ptr; - -/*! - * \brief A var binding condition - */ -struct VarBinding : ConditionNode { - Var var; - MatchValuePtr val; - - VarBinding(Var var, MatchValuePtr val) : var(var), val(val) {} - - ~VarBinding() {} -}; - -/*! - * \brief Compare the tag of the object - */ -struct TagCompare : ConditionNode { - /*! \brief The object to be examined */ - MatchValuePtr obj; - - /*! \brief The expected tag */ - int target_tag; - - TagCompare(MatchValuePtr obj, size_t target) : obj(obj), target_tag(target) {} - - ~TagCompare() {} -}; - -using TreeObjectPtr = typename relay::TreeNode::pointer; -using TreeLeafNode = relay::TreeLeafNode; -using TreeLeafFatalNode = relay::TreeLeafFatalNode; -using TreeBranchNode = relay::TreeBranchNode; - -TreeObjectPtr BuildDecisionTreeFromPattern(MatchValuePtr data, Pattern pattern, - TreeObjectPtr then_branch, TreeObjectPtr else_branch) { - if (pattern.as()) { - // We ignore wildcard binding since it's not producing new vars - return then_branch; - } else if (const auto* pvn = pattern.as()) { - auto cond = std::make_shared(pvn->var, data); - return TreeBranchNode::Make(cond, then_branch, else_branch); - } else if (const auto* pcn = pattern.as()) { - auto tag = pcn->constructor->tag; - - size_t field_index = 0; - for (auto& p : pcn->patterns) { - auto d = std::make_shared(data, field_index); - then_branch = BuildDecisionTreeFromPattern(d, p, then_branch, else_branch); - field_index++; - } - auto cond = std::make_shared(data, tag); - return TreeBranchNode::Make(cond, then_branch, else_branch); - } else { - const auto* pt = pattern.as(); - ICHECK(pt) << "unhandled case: " << AsText(pattern, false); - size_t field_index = 0; - for (auto& p : pt->patterns) { - auto d = std::make_shared(data, field_index++); - then_branch = BuildDecisionTreeFromPattern(d, p, then_branch, else_branch); - } - return then_branch; - } -} - -TreeObjectPtr BuildDecisionTreeFromClause(MatchValuePtr data, Clause clause, - TreeObjectPtr else_branch) { - return BuildDecisionTreeFromPattern(data, clause->lhs, TreeLeafNode::Make(clause->rhs), - else_branch); -} - -TreeObjectPtr BuildDecisionTreeFromClauses(MatchValuePtr data, tvm::Array clauses) { - // When nothing matches, the VM throws fatal error - TreeObjectPtr else_branch = TreeLeafFatalNode::Make(); - // Start from the last clause - for (auto it = clauses.rbegin(); it != clauses.rend(); ++it) { - else_branch = BuildDecisionTreeFromClause(data, *it, else_branch); - } - return else_branch; -} - -std::vector ToAllocTensorShape(NDArray shape) { - std::vector raw_shape; - if (shape->ndim == 0) { - return raw_shape; - } - ICHECK_EQ(shape->ndim, 1u); - ICHECK_EQ(shape->dtype.code, 0U) << "The dtype of constant shape must be int32 or int64, but got " - << DLDataType2String(shape->dtype); - ICHECK(shape->dtype.bits == 64 || shape->dtype.bits == 32) - << "The dtype of constant shape must be int32 or int64, but got" - << DLDataType2String(shape->dtype); - - if (shape->dtype.bits == 64) { - int64_t* int_ptr = reinterpret_cast(shape->data); - for (auto i = 0; i < shape->shape[0]; i++) { - raw_shape.push_back(int_ptr[i]); - } - } else { // int32 - int32_t* int_ptr = reinterpret_cast(shape->data); - for (auto i = 0; i < shape->shape[0]; i++) { - raw_shape.push_back(static_cast(int_ptr[i])); - } - } - return raw_shape; -} - -class VMFunctionCompiler : DeviceAwareExprFunctor { - public: - VMFunctionCompiler(VMCompilerContext* context, VirtualDevice host_virtual_device) - : DeviceAwareExprFunctor(context->module), - last_register_(0), - registers_num_(0), - context_(context), - host_virtual_device_(std::move(host_virtual_device)) {} - - VMFunction Compile(const GlobalVar& var, const Function& func) { - VLOG(1) << "Compiling:" << std::endl << PrettyPrint(func); - std::vector param_device_indexes; - if (IsClosure(func)) { - // After lifting we'll have functions of the form: - // fn(closure args) { fn(lifted function args) { body } } - // But we want the closure's function to be: - // fn(closure args, lifter function args) { body } - // Do that flattening on-the-fly here. - Function inner_func = Downcast(func->body); - std::vector params; - params.reserve(func->params.size() + inner_func->params.size()); - param_device_indexes.reserve(func->params.size() + inner_func->params.size()); - for (size_t i = 0; i < func->params.size(); ++i) { - params.emplace_back(func->params[i]); - param_device_indexes.push_back(GetDeviceIndex(func->params[i]->virtual_device())); - } - for (size_t i = 0; i < inner_func->params.size(); ++i) { - params.emplace_back(inner_func->params[i]); - - param_device_indexes.push_back(GetDeviceIndex(inner_func->params[i]->virtual_device())); - } - std::vector type_params; - type_params.reserve(func->type_params.size() + inner_func->type_params.size()); - for (const auto& tyvar : func->type_params) { - type_params.push_back(tyvar); - } - for (const auto& tyvar : inner_func->type_params) { - type_params.push_back(tyvar); - } - Function flattened_func = Function(params, inner_func->body, inner_func->ret_type, - type_params, func->attrs, func->span); - flattened_func->virtual_device_ = inner_func->virtual_device(); - VisitExpr(flattened_func); - } else { - param_device_indexes.reserve(func->params.size()); - for (size_t i = 0; i < func->params.size(); ++i) { - param_device_indexes.push_back(GetDeviceIndex(func->params[i]->virtual_device())); - } - VisitExpr(func); - } - return VMFunction(var->name_hint, params_, instructions_, registers_num_, - std::move(param_device_indexes)); - } - - /*! \brief Attrs objects for each op. */ - std::map> op_attrs; - - /*! \brief Attrs objects for each callsite. */ - std::map> callsite_attrs; - - protected: - size_t NewRegister() { return registers_num_++; } - - inline void Emit(const Instruction& instr) { - size_t instruction_index = instructions_.size(); - VLOG(2) << "instruction[" << instruction_index << "] = " << instr; - ICHECK((int)instr.op < 100) << "Invalid opcode " << (int)instr.op; - switch (instr.op) { - case Opcode::AllocADT: - case Opcode::AllocTensor: - case Opcode::AllocTensorReg: - case Opcode::GetField: - case Opcode::GetTag: - case Opcode::LoadConst: - case Opcode::LoadConsti: - case Opcode::Invoke: - case Opcode::AllocClosure: - case Opcode::AllocStorage: - case Opcode::ShapeOf: - case Opcode::ReshapeTensor: - case Opcode::Move: - case Opcode::InvokeClosure: - case Opcode::DeviceCopy: - last_register_ = instr.dst; - break; - case Opcode::InvokePacked: - case Opcode::If: - case Opcode::Ret: - case Opcode::Goto: - case Opcode::Fatal: - case Opcode::KillRegister: - break; - } - instructions_.push_back(instr); - } - - /*! - * \brief Returns the "device index" to represent \p virtual_device for primitives - * in emitted code. Note that the host device is always at index 0. - */ - Index GetDeviceIndex(const VirtualDevice& virtual_device) { - ICHECK(!virtual_device->IsFullyUnconstrained()); - auto itr = std::find(context_->virtual_devices_.begin(), context_->virtual_devices_.end(), - virtual_device); - if (itr != context_->virtual_devices_.end()) { - return std::distance(context_->virtual_devices_.begin(), itr); - } - - ICHECK_GT(context_->virtual_devices_.size(), 0); - ICHECK_NE(virtual_device, host_virtual_device_); // the host scope is always at index 0 - - if (virtual_device->device_type() == context_->virtual_devices_.front()->device_type()) { - // It's ok if we see distinct scopes which share the host device type. This is because - // we allow the VirtualDevice for the host to be different from the VirtualDevice for - // primitive operations which both happen to be on the same device (typically CPU). - return 0; - } - - ICHECK(virtual_device != host_virtual_device_); - Index index = context_->virtual_devices_.size(); - VLOG(2) << "virtual_device[" << index << "] = " << virtual_device; - context_->virtual_devices_.push_back(virtual_device); - - return index; - } - - using DeviceAwareExprFunctor::VisitExpr_; - - void VisitExpr_(const ConstantNode* const_node) final { - // Check the shape is valid - NDArray data = const_node->data; - size_t const_index = context_->constants.size(); - auto con = GetRef(const_node); - Index device_index = GetDeviceIndex(GetVirtualDevice(con)); - VLOG(2) << "constant[" << const_index << "] on device[" << device_index << "]"; - context_->const_device_indexes.push_back(device_index); - context_->constants.push_back(const_node->data); - Emit(Instruction::LoadConst(const_index, device_index, NewRegister())); - } - - void VisitExpr_(const VarNode* var_node) final { - auto var = GetRef(var_node); - auto reg_it = this->var_register_map_.find(var); - ICHECK(reg_it != this->var_register_map_.end()); - last_register_ = reg_it->second; - } - - void VisitExpr_(const TupleNode* tuple_node) final { - auto tuple = GetRef(tuple_node); - std::vector fields_registers; - - for (auto& field : tuple->fields) { - this->VisitExpr(field); - fields_registers.push_back(last_register_); - } - - // TODO(@jroesch): use correct tag - Emit(Instruction::AllocADT(0, tuple->fields.size(), fields_registers, NewRegister())); - } - - void VisitExpr_(const MatchNode* match_node) final { - auto match = GetRef(match_node); - - this->VisitExpr(match->data); - CompileMatch(match); - } - - void PreVisitLetBinding_(const Var& var, const Expr& value) final { - ICHECK(!value.as()) - << "unexpected function:" << std::endl - << PrettyPrint(value) << std::endl - << "bound to var '" << var->name_hint() << "'. Did you set opt_level = 2?"; - VisitExpr(value); - var_register_map_.emplace(var, this->last_register_); - } - - void VisitExpr_(const TupleGetItemNode* get_node) final { - auto get = GetRef(get_node); - this->VisitExpr(get->tuple); - auto tuple_register = last_register_; - Emit(Instruction::GetField(tuple_register, get->index, NewRegister())); - } - - void VisitExpr_(const GlobalVarNode* gvar) final { - auto var = GetRef(gvar); - auto func = context_->module->Lookup(var); - auto it = context_->global_map.find(var); - ICHECK(it != context_->global_map.end()) << PrettyPrint(var); - // Allocate closure with zero free vars - Emit(Instruction::AllocClosure(it->second, 0, {}, NewRegister())); - } - - void VisitExpr_(const IfNode* if_node) final { - this->VisitExpr(if_node->cond); - - size_t test_register = last_register_; - - this->Emit(Instruction::LoadConsti(1, NewRegister())); - auto after_cond = instructions_.size(); - auto target_register = last_register_; - this->Emit(Instruction::If(test_register, target_register, 0, 0)); - this->VisitExpr(if_node->true_branch); - - // It saves the result of If-Else expression. - auto merge_register = NewRegister(); - Emit(Instruction::Move(last_register_, merge_register)); - Emit(Instruction::Goto(0)); - - // Finally store how many instructions there are in the - // true branch. - auto after_true = this->instructions_.size(); - - this->VisitExpr(if_node->false_branch); - - size_t false_register = last_register_; - - // In else-branch, override the then-branch register - Emit(Instruction::Move(false_register, merge_register)); - // Compute the total number of instructions - // after generating false. - auto after_false = this->instructions_.size(); - - // Now we will compute the jump targets in order - // to properly patch the instruction with the - // the requiste targets. - - // After we emit the true body, and false body, - // we patch up the if instruction, and goto. - auto true_offset = 1; - auto false_offset = after_true - after_cond; - instructions_[after_cond].if_op.true_offset = true_offset; - instructions_[after_cond].if_op.false_offset = false_offset; - - // Patch the Goto. - this->instructions_[after_true - 1].pc_offset = (after_false - after_true) + 1; - - this->last_register_ = merge_register; - } - - void EmitInvokeTVMOp(const Expr& func, const Expr& inputs, const Expr& outputs, - const DictAttrs& attrs) { - std::vector argument_registers; - - const auto* global_var_node = func.as(); - ICHECK(global_var_node) << "Expecting function in invoke_tvm_op to be a global"; - - auto input_tuple = inputs.as(); - ICHECK(input_tuple) << "internal error: invoke_tvm_op inputs must be a tuple," - << "please file a bug in the memory manifestation pass"; - - auto output_tuple = outputs.as(); - ICHECK(output_tuple) << "internal error: invoke_tvm_op outputs must be a tuple," - << "please file a bug in the memory manifestation pass"; - - for (auto input : input_tuple->fields) { - VisitExpr(input); - argument_registers.push_back(last_register_); - } - - for (auto output : output_tuple->fields) { - ICHECK(output->IsInstance()) << "output should be var, found:" << std::endl - << PrettyPrint(output); - auto reg = var_register_map_.find(Downcast(output)); - ICHECK(reg != var_register_map_.end()) - << "internal error: all variables should be in the register mapping"; - argument_registers.push_back(reg->second); - } - - Index op_index; - auto itr = context_->primitive_map.find(global_var_node->name_hint); - if (itr == context_->primitive_map.end()) { - op_index = context_->primitive_map.size(); - context_->primitive_map.emplace(global_var_node->name_hint, op_index); - } else { - op_index = itr->second; - } - - if (attrs.defined() && attrs->dict.defined()) { - // Capture the dictionary of attributes from the original primitive function so that they - // can contribute to the hash of the compiled primitive. This way we can distinguish - // primitives with the same body expression but different attributes which may arbitrarily - // influence code generation. - op_attrs[op_index] = attrs->dict; - } - - Emit(Instruction::InvokePacked(op_index, argument_registers.size(), output_tuple->fields.size(), - argument_registers)); - } - - void DeviceAwareVisitExpr_(const CallNode* call_node) final { - DeviceCopyProps device_copy_props = GetDeviceCopyProps(call_node); - CallLoweredProps call_lowered_props = GetCallLoweredProps(call_node); - ICHECK(!call_lowered_props.lowered_func.defined()); - if (device_copy_props.body.defined()) { - // TODO(mbs): device_copy cleanup. - VisitExpr(device_copy_props.body); - RegName src_reg = last_register_; - Index src_index = GetDeviceIndex(device_copy_props.src_virtual_device); - Index dst_index = GetDeviceIndex(device_copy_props.dst_virtual_device); - // Since scopes distinguish by targets (including any target hosts) but at runtime we - // deal only with devices, the copy may be unnecessary. - if (src_index != dst_index) { - Emit(Instruction::DeviceCopy(src_reg, src_index, dst_index, NewRegister())); - } - return; - } - - // Now we handle the case in which we are using an opaque operator used to define a - // sub-dialect, such as memory allocation operations. - if (call_node->op.as()) { - OpMatch matcher; - matcher - .Match("vm.invoke_tvm_op", - [this](const Array& args, const Attrs& attrs, const Array& type_arg) { - ICHECK_EQ(args.size(), 3); - EmitInvokeTVMOp(args[0], args[1], args[2], Downcast(attrs)); - }) - .Match("memory.alloc_tensor", - [this](const Array& args, const Attrs& attrs, const Array& type_arg) { - ICHECK_EQ(args.size(), 3); - - // Get the attributes. - auto alloc_attrs = attrs.as(); - ICHECK(alloc_attrs != nullptr) << "must be the alloc tensor attrs"; - auto dtype = alloc_attrs->dtype; - - // The storage will be passed dynamically. - this->VisitExpr(args[0]); - auto storage_register = last_register_; - - // The storage will be passed dynamically. - this->VisitExpr(args[1]); - auto offset_register = last_register_; - - // If the shape is constant then we will emit a static tensor allocation - // instruction. It may be wrapped by an on_device, but it will be on the host - // which is assumed by the alloc_tensor instruction anyway. - auto const_shape = AsIgnoringOnDevice(args[2]); - - if (const_shape) { - NDArray shape = const_shape->data; - // TODO(@jroesch): we need to get an RFC done to standarize shape dtype - std::vector raw_shape = ToAllocTensorShape(shape); - // Add context field. - Emit(Instruction::AllocTensor(storage_register, offset_register, raw_shape, - dtype, NewRegister())); - } else { - this->VisitExpr(args[2]); - auto shape_register = last_register_; - Emit(Instruction::AllocTensorReg(storage_register, offset_register, - shape_register, dtype, NewRegister())); - } - }) - .Match("memory.alloc_storage", - [this](const Array& args, const Attrs& attrs, const Array& type_arg) { - ICHECK_EQ(args.size(), 3); - // Compute the size of the allocation. - this->VisitExpr(args[0]); - auto size_register = last_register_; - - auto const_shape = AsIgnoringOnDevice(args[1]); - std::vector raw_shape; - if (const_shape) { - NDArray shape = const_shape->data; - // TODO(@jroesch): we need to get an RFC done to standarize shape dtype - raw_shape = ToAllocTensorShape(shape); - } - - ICHECK(args[2].as()); // Always a literal. - NDArray alignment_arr = args[2].as()->data; - ICHECK_EQ(alignment_arr->dtype.code, 0U) - << "The dtype of constant shape must be int32 or int64, but got " - << DLDataType2String(alignment_arr->dtype); - ICHECK_EQ(alignment_arr->dtype.bits, 64U); - Index alignment = reinterpret_cast(alignment_arr->data)[0]; - - // Get the dtype hint from the attributes. - auto alloc_attrs = attrs.as(); - ICHECK(alloc_attrs != nullptr) << "must be the AllocStorage attrs"; - auto dtype = alloc_attrs->dtype; - - Emit(Instruction::AllocStorage(size_register, alignment, dtype, - GetDeviceIndex(alloc_attrs->virtual_device), - raw_shape, NewRegister())); - }) - .Match("vm.shape_of", - [this](const Array& args, const Attrs& attrs, const Array& type_arg) { - ICHECK_EQ(args.size(), 1U); - // Get the attributes. - const auto* shape_of_attrs = attrs.as(); - ICHECK(shape_of_attrs) << "Must be the shape_of attrs"; - ICHECK_EQ(shape_of_attrs->dtype.bits(), 64) - << "The dtype of shape of must be int64, but got" - << DLDataType2String(shape_of_attrs->dtype); - this->VisitExpr(args[0]); - Emit(Instruction::ShapeOf(last_register_, NewRegister())); - }) - .Match("vm.reshape_tensor", - [this](const Array& args, const Attrs& attrs, const Array& type_arg) { - ICHECK_EQ(args.size(), 2u); - this->VisitExpr(args[0]); - auto tensor_reg = last_register_; - this->VisitExpr(args[1]); - auto shape_reg = last_register_; - Emit(Instruction::ReshapeTensor(tensor_reg, shape_reg, NewRegister())); - }) - .Match("memory.kill", - [this](const Array& args, const Attrs& attrs, const Array& type_arg) { - ICHECK_EQ(args.size(), 1u); - this->VisitExpr(args[0]); - Emit(Instruction::KillRegister(this->last_register_)); - }); - matcher(GetRef(call_node)); - return; - } - - // In the case it's not one of these specialized operators we will generate code - // for one of the "standard" cases. - std::vector args_registers; - - // Evaluate the call arguments. - for (auto arg : call_node->args) { - VisitExpr(arg); - args_registers.push_back(last_register_); - } - - if (const auto* global_var_node = call_node->op.as()) { - // In the case we are invoking a global we need to find its - // global ID, and then check whether it is closure invocation - // or whether it is a standard global, and emit the correct - // calling convention. - auto global = GetRef(global_var_node); - auto it = context_->global_map.find(global); - ICHECK(it != context_->global_map.end()) << PrettyPrint(global); - VLOG(2) << "VisitExpr_: generating invoke for " << global->name_hint - << " with func_index=" << it->second; - - // TODO(tvm-team): - // Think about mixed call into global that is not a relay::Function - // perhaps establish as an invariance(all functions in mod must be relay::Function) - auto func = Downcast(context_->module->Lookup(global)); - - if (IsClosure(func)) { - auto arity = func->params.size(); - Emit(Instruction::AllocClosure(it->second, arity, args_registers, NewRegister())); - } else { - Emit(Instruction::Invoke(it->second, args_registers, NewRegister())); - } - } else if (const auto* constructor_node = call_node->op.as()) { - // In the constructor case, we simply need to find its tag - // and emit a call to allocate the data structure. - auto constructor = GetRef(constructor_node); - Emit(Instruction::AllocADT(constructor->tag, call_node->args.size(), args_registers, - NewRegister())); - } else if (auto var = call_node->op.as()) { - // If we are calling a variable, it must be the case that it is a closure so we - // emit invoke closure here. - VisitExpr(var.value()); - Emit(Instruction::InvokeClosure(last_register_, args_registers, NewRegister())); - } else if (auto inner_call = call_node->op.as()) { - VisitExpr(inner_call.value()); - Emit(Instruction::InvokeClosure(last_register_, args_registers, NewRegister())); - } else { - // Finally if there are any other cases this is a bug. - LOG(FATAL) << "internal error: unreachable code," - << "should be transformed away by previous passes:" << std::endl - << PrettyPrint(GetRef(call_node)); - } - } - - void DeviceAwareVisitExpr_(const FunctionNode* func_node) final { - if (function_nesting() > 1) { - ICHECK(func_node->HasNonzeroAttr(attr::kPrimitive)) - << "local functions should have been removed by lambda lifting:" << std::endl - << "Program: " << AsText(GetRef(func_node), false) << std::endl - << "AST: " << GetRef(func_node); - return; - } - - // We're processing a top-level function which has possibly been rejigged to capture - // both closure and function arguments. Those functions retain their 'Closure' attribute, - // but we can just process them like any other function here. - - // Assign a register num to each parameter. - size_t i = 0; - for (auto param : func_node->params) { - auto arg_register = NewRegister(); - ICHECK_EQ(i, arg_register); - var_register_map_.insert({param, arg_register}); - params_.push_back(param->name_hint()); - ++i; - } - - VisitExpr(func_node->body); - - instructions_.push_back(Instruction::Ret(last_register_)); - } - - /*! - * \brief Compile a match value - * Generate byte code that compute the value specified in val - * - * \return The register number assigned for the final value - */ - RegName CompileMatchValue(MatchValuePtr val) { - if (std::dynamic_pointer_cast(val)) { - auto r = std::dynamic_pointer_cast(val); - return r->register_num; - } else { - auto path = std::dynamic_pointer_cast(val); - auto p = CompileMatchValue(path->parent); - Emit(Instruction::GetField(p, path->index, NewRegister())); - path->reg = last_register_; - return path->reg; - } - } - - void CompileTreeNode(TreeObjectPtr tree) { - if (auto node = std::dynamic_pointer_cast(tree)) { - VisitExpr(node->body); - } else if (std::dynamic_pointer_cast(tree)) { - Emit(Instruction::Fatal()); - } else if (auto node = std::dynamic_pointer_cast(tree)) { - if (auto cond = std::dynamic_pointer_cast(node->cond)) { - // For Tag compariton, generate branches - auto r = CompileMatchValue(cond->obj); - Emit(Instruction::GetTag(r, NewRegister())); - auto operand1 = last_register_; - Emit(Instruction::LoadConsti(cond->target_tag, NewRegister())); - auto operand2 = last_register_; - - Emit(Instruction::If(operand1, operand2, 1, 0)); - auto cond_offset = instructions_.size() - 1; - CompileTreeNode(node->then_branch); - auto if_reg = last_register_; - Emit(Instruction::Goto(1)); - auto goto_offset = instructions_.size() - 1; - CompileTreeNode(node->else_branch); - auto else_reg = last_register_; - Emit(Instruction::Move(else_reg, if_reg)); - last_register_ = if_reg; - auto else_offset = instructions_.size() - 1; - // Fixing offsets - instructions_[cond_offset].if_op.false_offset = goto_offset - cond_offset + 1; - instructions_[goto_offset].pc_offset = else_offset - goto_offset + 1; - } else { - // For other non-branch conditions, move to then_branch directly - auto var_bind = std::dynamic_pointer_cast(node->cond); - var_register_map_[var_bind->var] = CompileMatchValue(var_bind->val); - CompileTreeNode(node->then_branch); - } - } - } - - /*! - * \brief Compile a pattern match expression - * It first converts the pattern match expression into a decision tree, the condition - * could be object comparison or variable binding. If any of the condition fails in a clause, - * the decision tree switches to check the conditions of next clause and so on. If no clause - * matches the value, a fatal node is inserted. - * - * After the decision tree is built, we convert it into bytecodes using If/Goto. - */ - void CompileMatch(Match match) { - auto data = std::make_shared(last_register_); - auto decision_tree = BuildDecisionTreeFromClauses(data, match->clauses); - CompileTreeNode(decision_tree); - } - - protected: - /*! \brief Store the expression a variable points to. */ - std::unordered_map expr_map_; - /*! \brief Instructions in the VMFunction. */ - std::vector instructions_; - /*! \brief Parameter names of the function. */ - std::vector params_; - /*! \brief Map from var to register number. */ - std::unordered_map var_register_map_; - /*! \brief Last used register number. */ - size_t last_register_; - /*! \brief Total number of virtual registers allocated. */ - size_t registers_num_; - /*! \brief Global shared meta data */ - VMCompilerContext* context_; - /*! \brief VirtualDevice for data and computation which must reside on a CPU. */ - VirtualDevice host_virtual_device_; -}; - -PackedFunc VMCompiler::GetFunction(const String& name, const ObjectPtr& sptr_to_self) { - if (name == "lower") { - return PackedFunc([sptr_to_self, this](TVMArgs args, TVMRetValue* rv) { - ICHECK_EQ(args.num_args, 2); - this->Lower(args[0], args[1]); - }); - } else if (name == "codegen") { - return PackedFunc([sptr_to_self, this](TVMArgs args, TVMRetValue* rv) { - ICHECK_EQ(args.num_args, 0); - this->Codegen(); - }); - } else if (name == "get_executable") { - return PackedFunc([sptr_to_self, this](TVMArgs args, TVMRetValue* rv) { - ICHECK_EQ(args.num_args, 0); - *rv = this->GetExecutable(); - }); - } else if (name == "set_params") { - return PackedFunc([sptr_to_self, this](TVMArgs args, TVMRetValue* rv) { - Map params = args[0]; - for (const auto& kv : params) { - this->SetParam(kv.first, kv.second->data); - } - }); - } else if (name == "get_params") { - return PackedFunc([sptr_to_self, this](TVMArgs args, TVMRetValue* rv) { - Map ret; - for (const auto& kv : params_) { - ret.Set(kv.first, Constant(kv.second)); - } - *rv = ret; - }); - } else if (name == "optimize") { - return PackedFunc([sptr_to_self, this](TVMArgs args, TVMRetValue* rv) { - ICHECK_EQ(args.num_args, 2); - *rv = this->OptimizeModule(args[0], args[1]); - }); - } else { - LOG(FATAL) << "Unknown packed function: " << name; - } -} - -void VMCompiler::SetParam(const std::string& name, runtime::NDArray data_in) { - params_[name] = data_in; -} - -void VMCompiler::Lower(IRModule mod, const Array& raw_targets) { - VLOG_CONTEXT << "VM Lower"; - Setup(raw_targets); - LowerImpl(std::move(mod)); -} - -IRModule VMCompiler::OptimizeModule(IRModule mod, const Array& raw_targets) { - VLOG_CONTEXT << "VM Optimize"; - Setup(raw_targets); - return OptimizeModuleImpl(std::move(mod)); -} - -runtime::Module VMCompiler::GetExecutable() const { - if (exec_ == nullptr) { - LOG(WARNING) << "No executable to return. Did you forget to call VMCompiler::Lower?"; - } - if (exec_->imports().empty()) { - LOG(WARNING) << "Executable is empty. Did you forget to call VMCompiler::Codegen?"; - } - return runtime::Module(exec_); -} - -void VMCompiler::Setup(const Array& raw_targets) { - ICHECK(exec_ == nullptr) << "Can't reuse VMComplier object for multiple modules"; - exec_ = make_object(); - ICHECK(!config_.defined()); - config_ = CompilationConfig(PassContext::Current(), raw_targets); - VLOG(1) << "Using compilation config:" << std::endl << config_; - - // The first device is always for the host. - CHECK(context_.virtual_devices_.empty()); - VLOG(1) << "virtual_device[0] = " << config_->host_virtual_device << " (host)"; - context_.virtual_devices_.push_back(config_->host_virtual_device); -} - -void VMCompiler::LowerImpl(IRModule mod) { - // Run the optimizations necessary to target the VM. - context_.module = OptimizeModuleImpl(std::move(mod)); - - // Build the map from global variables bound to Functions to a global index in the - // VMFunction table. - size_t num_functions = PopulateGlobalMap(); - - // Next we get ready by allocating space for - // the global state. - exec_->functions.resize(num_functions); - - for (const auto& pair : context_.module->functions) { - auto gvar = pair.first; - if (auto opt = pair.second.as()) { - auto func = opt.value(); - if (func->HasNonzeroAttr(attr::kExtern)) { - // Already compiled during lowering. - continue; - } - - VMFunctionCompiler func_compiler(&context_, config_->host_virtual_device); - auto vm_func = func_compiler.Compile(gvar, func); - - size_t func_index = context_.global_map.at(gvar); - ICHECK(func_index < exec_->functions.size()); - exec_->functions[func_index] = vm_func; - - // update structural hashes for tvm ops - for (auto p : func_compiler.op_attrs) { - exec_->op_attrs.insert(p); - } - } - } - - // Populate virtual devices and the host device index. - for (const auto& virtual_device : context_.virtual_devices_) { - ICHECK(!virtual_device->IsFullyUnconstrained()); - ICHECK_GT(virtual_device->device_type(), 0); - exec_->virtual_devices.push_back( - std::make_pair(Device{/*device_type=*/virtual_device->device_type(), - /*device_id=*/virtual_device->virtual_device_id}, - virtual_device->memory_scope)); - } - exec_->host_device_index = kHostDeviceIndex; - - // populate constants - for (const auto& data : context_.constants) { - exec_->constants.push_back(data); - } - - for (auto index : context_.const_device_indexes) { - exec_->const_device_indexes.push_back(index); - } - - // update global function map - for (const auto& gv : context_.global_map) { - exec_->global_map.insert({gv.first->name_hint, gv.second}); - } - - // update primitive function map - for (const auto& pair : context_.primitive_map) { - exec_->primitive_map.insert(pair); - } - - VLOG(1) << "Compiled to:" << std::endl - << "-------------------------------------------------" << std::endl - << exec_->GetVirtualDevices() // - << exec_->GetConstants() // - << exec_->GetPrimitives() // - << exec_->GetBytecode() // - << "-------------------------------------------------"; - - if (backend::IsAutoSchedulerEnabled()) { - backend::UpdateAutoSchedulerOpWeights(context_.module); - } -} - -transform::Sequential VMCompiler::MemoryOpt(const CompilationConfig& config) { - Array pass_seqs; - // Remove unused functions - Array entry_functions{"main"}; - pass_seqs.push_back(transform::RemoveUnusedFunctions(entry_functions)); - // Manifest the allocations. - pass_seqs.push_back(transform::ManifestAlloc(config->host_virtual_device)); - - // Compute away possibly introduced constant computation. - pass_seqs.push_back(transform::FoldConstant()); - - // Fuse & lower any new shape functions and device_copies. - pass_seqs.push_back(FuseAndLowerOperators(config)); - - // Manifest the allocations needed for the shape functions. - pass_seqs.push_back(transform::ManifestAlloc(config->host_virtual_device)); - - // Fuse & lower any new allocations. - pass_seqs.push_back(FuseAndLowerOperators(config)); - - // TODO(mbrookhart, jroesch, masahi): this pass is very slow, and is - // incomplete to provide memory resuse optimizations. Disable it until we can - // rewrite it in C++ and complete it. - // // Perform memory planning in order to coalesce/reduce allocations. - // pass_seqs.push_back(transform::MemoryPlan()); - - // Compute away constant computation introduced by coalescing allocations. - pass_seqs.push_back(transform::FoldConstant()); - - // Fuse & lower yet again - pass_seqs.push_back(FuseAndLowerOperators(config)); - - // Create allocations for math introduced by dynamic region math. - pass_seqs.push_back(transform::ManifestAlloc(config->host_virtual_device)); - - // Compute away possibly introduced constant computation. - pass_seqs.push_back(transform::FoldConstant()); - - // Insert kills to free memory. - pass_seqs.push_back(transform::ManifestLifetimes()); - - // Lift constants to the top-level of the block to simplify VM code generation. - // TODO(@icemelon9, @jroesch): Remove this pass for now because some - // instructions need to access to constant - // pass_seqs.push_back(transform::LiftConstants()); - - return transform::Sequential(std::move(pass_seqs)); -} - -transform::Sequential VMCompiler::FuseAndLowerOperators(const CompilationConfig& config) { - Array pass_seqs; - // Hoist operators to "primitive" Functions. - pass_seqs.push_back(FuseOps()); - // Give each "primitive" Function a hash. - pass_seqs.push_back(LabelOps()); - // Lower "primitive" Functions to PrimFuncs and rewrite calls. - pass_seqs.push_back(tec::LowerTE(/*module_name=*/"vm_mod", config, [this](const BaseFunc& func) { - if (func->GetAttr(attr::kCompiler).defined()) { - backend::UpdateConstants(func, ¶ms_); - } - })); - // Since lowered functions are bound in the IRModule, we can now eliminate any unused - // let-bound functions. - pass_seqs.push_back(DeadCodeElimination(/*inline_once=*/false)); - return transform::Sequential(std::move(pass_seqs)); -} - -IRModule VMCompiler::OptimizeModuleImpl(IRModule mod) { - backend::BindParamsInModule(mod, params_); - Array pass_seqs = relay::backend::GetPassPrefix( - /*is_homogeneous=*/config_->optional_homogeneous_target.defined(), /*is_vm=*/true); - - // Always plan devices so the remaining passes don't need to distinguish homogeneous vs - // heterogeneous execution. - pass_seqs.push_back(transform::PlanDevices(config_)); - if (config_->optional_homogeneous_target.defined()) { - // This pass currently only supports the homogeneous case. - pass_seqs.push_back(transform::SplitArgs( - config_->optional_homogeneous_target->GetAttr("max_function_args", 0) - .value() - .IntValue())); - } - - pass_seqs.push_back(transform::FuseOps()); - pass_seqs.push_back(transform::AnnotateMemoryScope()); - - // Do layout rewrite for auto-scheduler. - transform::PassContext pass_ctx = PassContext::Current(); - if (backend::IsAutoSchedulerEnabled() && config_->optional_homogeneous_target.defined()) { - Pass major_pass = transform::AutoSchedulerLayoutRewrite(); - bool enable_layout_rewrite_targets = - config_->optional_homogeneous_target->GetTargetDeviceType() == kDLCPU || - config_->optional_homogeneous_target->GetAttr("device", "") == "mali"; - if (enable_layout_rewrite_targets && pass_ctx.PassEnabled(major_pass->Info())) { - With tctx(config_->optional_homogeneous_target); - pass_seqs.push_back(major_pass); - // Defuse ops to fold constants, then fuse them again - pass_seqs.push_back(transform::DefuseOps()); - pass_seqs.push_back(transform::FoldConstant()); - pass_seqs.push_back(transform::FuseOps()); - } - } - if (backend::IsMetaScheduleEnabled() && config_->optional_homogeneous_target.defined()) { - Pass major_pass = transform::MetaScheduleLayoutRewrite(); - bool enable_layout_rewrite_targets = - config_->optional_homogeneous_target->GetTargetDeviceType() == kDLCPU || - config_->optional_homogeneous_target->GetAttr("device", "") == "mali"; - if (enable_layout_rewrite_targets && pass_ctx.PassEnabled(major_pass->Info())) { - With tctx(config_->optional_homogeneous_target); - pass_seqs.push_back(major_pass); - // Defuse ops to fold constants, then fuse them again - pass_seqs.push_back(transform::DefuseOps()); - pass_seqs.push_back(transform::FoldConstant()); - pass_seqs.push_back(transform::FuseOps()); - } - } - - pass_seqs.push_back(transform::ToANormalForm()); - pass_seqs.push_back(transform::InferType()); - pass_seqs.push_back(transform::LambdaLift()); - - // Eliminate dead-code before we lower. We don't track the purity of PrimFuncs, thus after - // lowering all calls to lowered functions will be kept. - pass_seqs.push_back(DeadCodeElimination(/*inline_once=*/false)); - pass_seqs.push_back(transform::LabelOps()); - - // Lower all functions annotated as "primitive" by FuseOps. - pass_seqs.push_back(tec::LowerTE(/*module_name=*/"vm_mod", config_, [this](const BaseFunc& func) { - if (func->GetAttr(attr::kCompiler).defined()) { - backend::UpdateConstants(func, ¶ms_); - } - })); - - // Since lowered functions are bound in the IRModule, we can now eliminate any unused - // let-bound functions. - pass_seqs.push_back(DeadCodeElimination(/*inline_once=*/false)); - - // At this point it's possible to run PlanDevices again to pick up any additional constraints - // introduced during lowering. However we'll not do this until more testing has been done. - - // Inline the functions that are lifted to the module scope. We perform this - // pass after all other optimization passes but before the memory allocation - // pass. This is because memory allocation pass will insert `invoke_tvm_op` - // and we use these ops to invoke the symbols in the module generated by - // external codegen. - pass_seqs.push_back(transform::Inline()); - - pass_seqs.push_back(MemoryOpt(config_)); - pass_seqs.push_back(transform::InferType()); - - transform::Sequential seq(pass_seqs); - tvm::With ctx(pass_ctx); - if (config_->optional_homogeneous_target.defined()) { - With tctx(config_->optional_homogeneous_target); - return seq(std::move(mod)); - } else { - return seq(std::move(mod)); - } -} - -size_t VMCompiler::PopulateGlobalMap() { - // Allocate a VMFunction index for every Relay Function we could call. - // Excludes PrimFuncs and externs, which are managed by the primitive_map_. - for (const auto& kv : context_.module->functions) { - if (const auto* function_node = kv.second.as()) { - if (!function_node->HasNonzeroAttr(attr::kExtern)) { - context_.global_map.emplace(kv.first, context_.global_map.size()); - } - } - } - return context_.global_map.size(); -} - -void VMCompiler::Codegen() { - VLOG_CONTEXT << "VM Codegen"; - if (!context_.module.defined()) { - LOG(WARNING) << "No compiled module to codegen from. Did you forget to call VMCompiler::Lower?"; - return; - } - - // At this point context_.module will contain only: - // - non-external Relay functions, which we've compiled into VMFunctions. - // - external Relay functions, which will have definitions within some external runtime module - // in the "external_mods" attribute - // - PrimFuncs annotated with their targets. - // Only the PrimFuncs will appear in per_target_modules, and there may legitimately be none. - Map per_tvm_target_modules = tec::GetPerTargetModules(context_.module); - for (const auto& kv : per_tvm_target_modules) { - ICHECK(kv.first->GetTargetDeviceType() != kDLExtDev); - } - - // Retrieve all external runtime modules accumulated by external codegen (both function-at-a-time - // and IRModule-at-a-time). - Array external_mods = - context_.module->GetAttr>(tvm::attr::kExternalMods).value_or({}); - - // Retrieve any constant bindings accumulated by external codegen (by IRModule-at-a-time passes). - Map const_name_to_constant = - context_.module->GetAttr>(tvm::attr::kConstNameToConstant) - .value_or({}); - - VLOG(0) << "have " << per_tvm_target_modules.size() << " targets to build, " - << external_mods.size() << " external runtime modules, " << const_name_to_constant.size() - << " external constants, and " << params_.size() << " local constants"; - - // Any constant bindings must be merged into the overall 'params' map we've directly accumulated - // via the TECompiler callback. - for (const auto& kv : const_name_to_constant) { - ICHECK_EQ(params_.count(kv.first), 0); - params_.emplace(kv.first, kv.second); - } - - runtime::Module lib; - if (per_tvm_target_modules.empty()) { - // There is no function handled by TVM. We create a virtual main module - // to make sure a DSO module will be also available. - LOG(INFO) << "All lowered functions have been build by BYOC -- generating an empty TVM module"; - lib = codegen::CSourceModuleCreate(";", "", Array{}); - } else { - lib = tvm::TIRToRuntime(per_tvm_target_modules, config_->host_target); - } - - lib = - codegen::CreateMetadataModule(params_, lib, external_mods, config_->host_target, - Runtime::Create("cpp"), Executor::Create("graph"), // DNS HACK - relay::backend::ExecutorCodegenMetadata()); - exec_->SetLib(lib); -} - -runtime::Module CreateVMCompiler() { - auto exec = make_object(); - return runtime::Module(std::move(exec)); -} - -TVM_REGISTER_GLOBAL("relay._vm._VMCompiler").set_body_typed(CreateVMCompiler); - -} // namespace vm -} // namespace relay -} // namespace tvm diff --git a/src/relay/backend/vm/compiler.h b/src/relay/backend/vm/compiler.h deleted file mode 100644 index d22fb3d4d5ca..000000000000 --- a/src/relay/backend/vm/compiler.h +++ /dev/null @@ -1,179 +0,0 @@ -/* - * 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 src/relay/backend/vm/compiler.h - * \brief A compiler from relay::Module to the VM byte code. - */ - -#ifndef TVM_RELAY_BACKEND_VM_COMPILER_H_ -#define TVM_RELAY_BACKEND_VM_COMPILER_H_ - -#include -#include -#include -#include -#include -#include -#include - -#include -#include -#include -#include -#include -#include -#include - -#include "../../../runtime/memory/naive_allocator.h" -#include "../../../runtime/vm/profiler/vm.h" -#include "../../transforms/pass_utils.h" -#include "../te_compiler.h" -#include "../te_compiler_cache.h" - -namespace tvm { -namespace relay { -namespace vm { - -using tvm::runtime::ModulePropertyMask; -using tvm::runtime::NDArray; -using namespace tvm::runtime::vm; -using namespace relay::transform; - -template -using NodeMap = std::unordered_map; -using TagMap = NodeMap; -using TagNameMap = std::unordered_map; -using GlobalMap = NodeMap; -using ConstMap = NodeMap; -using ConstTensorShapeMap = NodeMap>; - -struct VMCompilerContext { - // The module context for the compilation - IRModule module; - // Error reporter - ErrorReporter err_reporter; - // Map from a unique integer to ADT constructor tag - TagNameMap tag_index_map; - // Map from ADT constructor tag to a unique integer - TagMap tag_map; - // Map from global var to a unique integer - GlobalMap global_map; - // List of constants - std::vector constants; - // Device indexes for constants - std::vector const_device_indexes; - // Map from names of primitive functions already allocated to their primitive function index. - std::unordered_map primitive_map; - // The virtual devices corresponding to each device index. - std::vector virtual_devices_; -}; - -class VMCompiler : public runtime::ModuleNode { - public: - VMCompiler() = default; - virtual ~VMCompiler() = default; - - virtual PackedFunc GetFunction(const String& name, const ObjectPtr& sptr_to_self); - - const char* type_key() const final { return "VMCompiler"; } - - /*! \brief Get the property of the runtime module .*/ - int GetPropertyMask() const final { return ModulePropertyMask::kRunnable; } - - /*! - * \brief Set the parameters - * - * \param name name of parameter - * \param data_in input DLTensor - */ - void SetParam(const std::string& name, runtime::NDArray data_in); - - /*! - * \brief Lower the functions in a Module. - * - * ---------------------------------------------------------------------------------- - * | This is the main entry point for the VM compilation flow. | - * | - Preceded by \p SetParam for the global params. | - * | - Followed by \p Codegen() to finalize the executable. | - * | - Then the result runtime::Module can be constructed by GetExecutable. | - * ---------------------------------------------------------------------------------- - * - * \param mod Relay Module - * \param raw_targets List of available targets for running kernels. Any host target should - * be conveyed by the 'host' target field. - */ - void Lower(IRModule mod, const Array& raw_targets); - - /* - * \brief Perform a series of optimizations on the input IR module. Can be used instead - * of Lower if wish to stop and observe optimized IRModule. Otherwise not needed on - * regular compilation flow. - * - * \param mod The input IRModule. - * \param raw_targets List of available target for running kernels. - * - * \return The optimized IRModule. - */ - IRModule OptimizeModule(IRModule mod, const Array& raw_targets); - - /*! \brief Generate the machine code for lowered functions. */ - void Codegen(); - - /*! \brief Returns the runtime::Module containing the compiled VM code. */ - runtime::Module GetExecutable() const; - - protected: - /*! \brief Builds the executor and compilation config to match \p raw_targets. */ - void Setup(const Array& raw_targets); - - /*! \brief Internal implementation of \p Lower. */ - void LowerImpl(IRModule mod); - - /*! \brief Internal implementation of \p OptimizeModule. */ - IRModule OptimizeModuleImpl(IRModule mod); - - /*! \brief Returns the passes which layout memory. */ - transform::Sequential MemoryOpt(const CompilationConfig& config); - - /*! \brief Returns the passes which fuse then lower Relay primitive operators. */ - transform::Sequential FuseAndLowerOperators(const CompilationConfig& config); - - /*! - * \brief Populate the global function names in a map where the value is used - * as the index by the VMFunctions. Returns the number of functions. - */ - size_t PopulateGlobalMap(); - - protected: - /*! \brief Targets and scopes needed for compilation. */ - CompilationConfig config_; - /*! \brief Global shared meta data */ - VMCompilerContext context_; - /*! \brief Compiled executable. */ - ObjectPtr exec_; - /*! \brief parameters */ - std::unordered_map params_; -}; - -} // namespace vm -} // namespace relay -} // namespace tvm - -#endif // TVM_RELAY_BACKEND_VM_COMPILER_H_ diff --git a/src/relay/backend/vm/lambda_lift.cc b/src/relay/backend/vm/lambda_lift.cc deleted file mode 100644 index 48449eb02149..000000000000 --- a/src/relay/backend/vm/lambda_lift.cc +++ /dev/null @@ -1,262 +0,0 @@ -/* - * 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 tvm/relay/backend/vm/lambda_lift.cc - * \brief Lift all local functions into global functions. - */ - -#include -#include -#include -#include -#include -#include - -#include -#include - -#include "../../op/annotation/annotation.h" -#include "../../transforms/device_aware_visitors.h" - -using namespace tvm::runtime; - -namespace tvm { -namespace relay { -namespace vm { - -inline std::string GenerateName(const Function& func) { - size_t hash = tvm::StructuralHash()(func); - return std::string("lifted_name") + std::to_string(hash); -} - -bool IsClosure(const Function& func) { return func->HasNonzeroAttr(attr::kClosure); } - -Function MarkClosure(Function func) { - return WithAttr(std::move(func), attr::kClosure, tvm::Integer(1)); -} - -/* The goal of this class is to lift out any nested functions into top-level - * functions. - * - * We will lift a function out into a global which takes the set of the free - * vars and then return the new created function. - */ -class LambdaLifter : public transform::DeviceAwareExprMutator { - public: - explicit LambdaLifter(const IRModule& module) - : transform::DeviceAwareExprMutator(module), module_(module) {} - - std::pair PreVisitLetBinding_(const Var& var, const Expr& value) final { - bool is_lambda = false; - if (const auto* func_node = value.as()) { - if (!func_node->HasNonzeroAttr(attr::kPrimitive)) { - is_lambda = true; - this->letrec_.push_back(var); - } - } - Expr new_value = this->VisitExpr(value); - - if (is_lambda) { - this->letrec_.pop_back(); - } - return {var, new_value}; - } - - Expr DeviceAwareVisitExpr_(const CallNode* call_node) final { - auto call = Downcast(DeviceAwareExprMutator::DeviceAwareVisitExpr_(call_node)); - if (auto opt = call_node->op.as()) { - auto var = opt.value(); - if (!letrec_.empty() && var == letrec_.back()) { - auto it = lambda_map_.find(var); - ICHECK(it != lambda_map_.end()); - return Call(it->second, call->args, call_node->attrs, call_node->type_args); - } - } - return std::move(call); - } - - Expr DeviceAwareVisitExpr_(const FunctionNode* func_node) final { - auto func = GetRef(func_node); - - if (func->HasNonzeroAttr(attr::kPrimitive)) { - // We should not transform primitive functions. - return std::move(func); - } - - if (function_nesting() == 1) { - // We don't need to lift global functions. - return WithFields(GetRef(func_node), func_node->params, VisitExpr(func_node->body)); - } - - auto name = GenerateName(func); - auto global = GlobalVar(name); - auto free_vars = FreeVars(func); - auto free_type_vars = FreeTypeVars(func, module_); - - Array captured_vars; - bool recursive = false; - for (const auto& var : free_vars) { - if (!letrec_.empty() && var == letrec_.back()) { - recursive = true; - continue; - } - captured_vars.push_back(var); - } - - // Freshen all the captured vars. - Array typed_captured_vars; - Map rebinding_map; - for (auto free_var : captured_vars) { - auto var = Var(free_var->name_hint(), free_var->checked_type()); - var->virtual_device_ = GetVirtualDevice(free_var); - typed_captured_vars.push_back(var); - rebinding_map.Set(free_var, var); - } - - VirtualDevice result_virtual_device = GetVirtualDevice(func_node->body); - - if (recursive) { - if (!captured_vars.empty()) { - Array fvs; - for (auto fv : captured_vars) { - fvs.push_back(fv); - } - lambda_map_.emplace(letrec_.back(), Call(global, fvs)); - } else { - lambda_map_.emplace(letrec_.back(), global); - } - } - - auto body = Downcast(DeviceAwareExprMutator::DeviceAwareVisitExpr_(func_node)); - - // When performing this optimization there are two cases. - // - // The first case in which we have no free variables - // we can just lift the function into the global - // environment without needing to allocate a closure. - // - // - // The second case requires that we generate a special - // function which makes a distinction between allocating - // a closure, and then the code for the closure. - // - // We represent a closure allocation by lifting the - // closure to a global function which takes its - // captured arguments and then directly returns - // the function representing the closure's code. - // - // When we generate code later on a call to the "outer" - // function marked as a closure is used to emit allocation - // code for the closure's environment. - // - // The "inner" function should be used to generate the - // code for the closure. - Function lifted_func; - if (captured_vars.empty() && free_type_vars.empty()) { - lifted_func = Function(body->params, body->body, body->ret_type, body->type_params, - body->attrs, body->span); - // We also need to copy the virtual device - lifted_func->virtual_device_ = body->virtual_device(); - } else { - // When a closure is locally bound in a program, we have its full type information - // avalible to us. - // - // If we lift the closure out of its bound context it may have free variables which - // do not have type annotations. - // - // In this case we first type check the program assigning a type to all sub-expressions. - // - // We then change the un-annotated free variables into annotated free variables, use - // bind to go from unannotated free variables -> annotated free variables and then - // construct the "closure" function with fully annotated arguments, no longer relying - // on type inference. - size_t before_arity = body->params.size(); - VLOG(9) << "Binding " << rebinding_map << " into\n" << PrettyPrint(body->body); - auto rebound_body = WithFields(func, func->params, Bind(body->body, rebinding_map)); - size_t after_arity = rebound_body->params.size(); - CHECK_EQ(before_arity, after_arity); - lifted_func = - Function(typed_captured_vars, rebound_body, /*ret_type=*/func->func_type_annotation(), - free_type_vars, DictAttrs(), func->span); - lifted_func->virtual_device_ = result_virtual_device; - lifted_func = MarkClosure(lifted_func); - } - - ICHECK(lifted_func.defined()); - - if (module_->ContainGlobalVar(name)) { - const auto existing_func = module_->Lookup(name); - ICHECK(tvm::StructuralEqual()(lifted_func, existing_func)) - << "lifted function hash collision"; - // If an identical function already exists, use its global var. - global = module_->GetGlobalVar(name); - } else { - // Add the lifted function to the module. - module_->Add(global, lifted_func); - } - - if (captured_vars.empty()) { - return std::move(global); - } else { - // If we need to allocate a closure, - // we pass the variables in its environment here. - Array fvs; - for (auto fv : captured_vars) { - fvs.push_back(fv); - } - return Call(global, fvs); - } - } - - IRModule Lift() { - // There is an ordering bug here. - auto glob_funcs = module_->functions; - for (auto pair : glob_funcs) { - if (auto* n = pair.second.as()) { - if (n->GetAttr(attr::kCompiler).defined()) continue; - auto func = GetRef(n); - module_->Add(pair.first, Downcast(Mutate(func)), /*update=*/true); - } - } - return module_; - } - - private: - std::unordered_map lambda_map_; - std::vector letrec_; - IRModule module_; -}; - -} // namespace vm - -namespace transform { - -Pass LambdaLift() { - runtime::TypedPackedFunc pass_func = - [=](IRModule m, PassContext pc) { return relay::vm::LambdaLifter(m).Lift(); }; - return CreateModulePass(pass_func, 1, "LambdaLift", {}); -} - -TVM_REGISTER_GLOBAL("relay._transform.LambdaLift").set_body_typed(LambdaLift); - -} // namespace transform - -} // namespace relay -} // namespace tvm diff --git a/src/relay/backend/vm/manifest_lifetimes.cc b/src/relay/backend/vm/manifest_lifetimes.cc deleted file mode 100644 index 892648d67844..000000000000 --- a/src/relay/backend/vm/manifest_lifetimes.cc +++ /dev/null @@ -1,262 +0,0 @@ -/* - * 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 src/relay/backend/vm/manifest_lifetimes.cc - * \brief Analysis and explicit manifestation of variable lifetimes. NOTE: the input IR should be in - * ANF and post-memory-lowering (explicit manifestation of allocations). - */ - -#include - -#include "../../../support/arena.h" -#include "../../op/memory/device_copy.h" -#include "../../transforms/device_aware_visitors.h" -#include "../../transforms/let_list.h" -#include "../liveness_analysis.h" - -namespace tvm { -namespace relay { -namespace transform { - -/*! - * \brief Helper class to insert kills using liveness information. - */ -class KillInserter : public ExprMutator { - public: - KillInserter(const ControlFlowGraph* cfg, const LivenessAnalysis* lva) : cfg_(cfg), lva_(lva) {} - - // Limitations - // ----------- - // (1) For simplicity, we only insert kills when visiting Let bindings, and always emit the kill - // as a single subsequent binding. This is slightly inaccurate; for example, if the condition of - // an If is dead after the test, we can immediately kill the condition in each branch: - // let %x = if (%dead_cond) { - // let %_0 = memory.kill(%dead_cond); - // ... - // } else { - // let %_1 = memory.kill(%dead_cond); - // ... - // } - // as opposed to: - // let %x = if (%dead_cond) ... - // let %_0 = memory.kill(%dead_cond); - // - // (2) Killed variables are calculated as live in - live out, which misses variables that are - // actually dead but not in a live-in set. Example: - // @f(%x: int, %y: int, %c: bool) { - // let %w = if (%c) { - // let %z = %y + %y; - // %z - // } else { - // %y - // }; - // %w - // } - // After inserting kills: - // @f(%x: int, %y: int, %c: bool) { - // /* %x is always dead, so never in any live in or live out set */ - // let %w = if (%c) { - // let %z = %y + %y; - // let %_0 = memory.kill(%y); - // %z - // } else { - // %y - // /* %y is dead at this point */ - // }; - // let %_1 = memory.kill(%c); - // /* no kill for %y since it's not in the live-in of %w AND %w isn't a let binding */ - // %w - // } - // - // (3) When the result expr of an If branch is a variable, and this expr is the last use of the - // var, we cannot "kill" the var since it is being returned. The VM compiler also emits a Move - // instruction to merge the branch results, which creates another ObjectRef to the Object held - // by the var. The var is also not in the subsequent live-in (since it is indeed dead by this - // point), so it won't be killed. An example can be seen in the previous code block for (2), where - // %y is not killed if the else-branch is taken (and indeed it can be killed, as %w is mapped to - // a new register and holds a fresh reference to the object referenced by %y). - // - // However, these limitations are unlikely to cause large leaks in practice. - - Expr VisitExpr_(const LetNode* let_node) override { - Expr expr = GetRef(let_node); - LetList ll; - - while (const LetNode* inner_let_node = expr.as()) { - ll.Push(inner_let_node->var, VisitExpr(inner_let_node->value)); - - ICHECK(!inner_let_node->value.as()) << "aliasing should have been eliminated."; - ICHECK(cfg_->let_map.count(expr)) << "all Let exprs should be mapped in the CFG"; - - const ControlFlowGraph::NodePtr n = cfg_->let_map.at(expr); - - const VarSet& li = lva_->live_in.at(n); - const VarSet& lo = lva_->live_out.at(n); - - // Killed vars = live in - live out. - VarSet kills; - for (const Var& v : li) { - if (!lo.count(v)) { - kills.insert(v); - } - } - - for (const Var& v : kills) { - ll.Push(Call(Op::Get("memory.kill"), {v})); - } - - expr = inner_let_node->body; - } - - return ll.Get(VisitExpr(expr)); - } - - private: - const ControlFlowGraph* cfg_; - const LivenessAnalysis* lva_; -}; - -/*! - * \brief Helper class to eliminate variable aliasing. This pass anticipates the VM compiler's - * register aliasing behavior so as to avoid killing vars that point to the same register. An - * alternative approach would be to track aliasing within the VM compiler itself, so that kill - * instructions are only emitted when all aliases are killed. - */ -class AliasEliminator : public MixedModeMutator { - public: - using MixedModeMutator::VisitExpr_; - - Expr VisitExpr_(const LetNode* let_node) override { - Expr expr = GetRef(let_node); - LetList ll; - std::vector aliased_vars; - - while (const LetNode* inner_let_node = expr.as()) { - const Var& var = inner_let_node->var; - const Expr& val = inner_let_node->value; - bool aliased = false; - ICHECK(!alias_.count(var)); - - if (const VarNode* alias_of_n = AsIgnoringOnDevice(val)) { - alias_[var] = Downcast(VisitExpr_(alias_of_n)); - aliased = true; - } else if (AsIgnoringOnDevice(val)) { - // Copying to the same device is aliasing. - // WARNING: this must be kept in sync with the VM compiler logic in - // src/relay/backend/vm/compiler.cc, line 541, in DeviceAwareVisitExpr_(const CallNode*). - Expr unwrapped = IgnoreOnDevice(val); - DeviceCopyProps copy_props = GetDeviceCopyProps(unwrapped); - if (copy_props.body.defined()) { - if (copy_props.src_virtual_device->device_type() == - copy_props.dst_virtual_device->device_type() && - copy_props.src_virtual_device->virtual_device_id == - copy_props.dst_virtual_device->virtual_device_id && - copy_props.src_virtual_device->memory_scope == - copy_props.dst_virtual_device->memory_scope) { - Expr to_copy = Downcast(unwrapped)->args[0]; - if (const VarNode* alias_of_n = to_copy.as()) { - alias_[var] = Downcast(VisitExpr_(alias_of_n)); - aliased = true; - } - } - } - } - - if (!aliased) { - ll.Push(var, VisitExpr(val)); - } else { - aliased_vars.push_back(var); - } - - expr = inner_let_node->body; - } - - Expr body = ll.Get(VisitExpr(expr)); - - // remove the aliased vars so that alias_ only tracks things in scope - for (const Var& v : aliased_vars) { - alias_.erase(v); - } - - return body; - } - - Expr VisitExpr_(const VarNode* var_node) override { - Var var = GetRef(var_node); - if (alias_.count(var)) { - return alias_[var]; - } - return std::move(var); - } - - Expr VisitExpr_(const FunctionNode* func_node) override { - Expr new_body = VisitExpr(func_node->body); - return WithFields(GetRef(func_node), /*opt_params=*/NullOpt, /*opt_body=*/new_body); - } - - // The only register-level aliasing that occurs in Match expressions is when - // the deconstructed expression is a Var, and the matched pattern is also a Var. - Expr VisitExpr_(const MatchNode* match_node) override { - if (const VarNode* data_var_node = AsIgnoringOnDevice(match_node->data)) { - Var data_var = Downcast(VisitExpr_(data_var_node)); - std::vector new_clauses; - for (const Clause& clause : match_node->clauses) { - const PatternVarNode* pv_node = nullptr; - if ((pv_node = clause->lhs.as())) { - alias_[pv_node->var] = data_var; - } - new_clauses.push_back(Clause(clause->lhs, VisitExpr(clause->rhs))); - if (pv_node) { - alias_.erase(pv_node->var); - } - } - return Match(data_var, new_clauses, match_node->complete, match_node->span); - } else { - return ExprMutator::VisitExpr_(match_node); - } - } - - private: - /*! - * \brief Mapping of var -> var it's an alias of. Note that transitive aliases - * (e.g. x = 0; y = x; z = y) are mapped to the non-aliased variable (in this example "x"). - */ - std::unordered_map alias_; -}; - -Pass ManifestLifetimes() { - auto pass_func = [](Function f, IRModule m, PassContext pc) -> Function { - f = Downcast(AliasEliminator().Mutate(f)); - Arena arena; - ControlFlowGraph cfg = ControlFlowGraph::Create(&arena, f); - UseDefAnalysis use_def = UseDefAnalysis::Analyze(cfg); - LivenessAnalysis lva = LivenessAnalysis::Analyze(cfg, use_def); - KillInserter ki(&cfg, &lva); - Function nf = Downcast(ki.Mutate(f)); - return nf; - }; - return CreateFunctionPass(pass_func, 0, "ManifestLifetimes", {}); -} - -TVM_REGISTER_GLOBAL("relay._transform.ManifestLifetimes").set_body_typed(ManifestLifetimes); - -} // namespace transform -} // namespace relay -} // namespace tvm diff --git a/src/relay/backend/vm/removed_unused_funcs.cc b/src/relay/backend/vm/removed_unused_funcs.cc deleted file mode 100644 index 67ba1bd8594f..000000000000 --- a/src/relay/backend/vm/removed_unused_funcs.cc +++ /dev/null @@ -1,137 +0,0 @@ -/* - * 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 tvm/relay/backend/vm/remove_unused_funcs.cc - * \brief Remove unused global relay functions in a relay module. - */ - -#include -#include -#include -#include -#include -#include - -#include -#include -#include - -#include "../../op/call/call.h" - -namespace tvm { -namespace relay { -namespace vm { - -/** - * \brief Detects all the functions that can be possibly called by entry function. - */ -struct CallTracer : ExprVisitor { - IRModule module_; - - // Record the names of all encountered functions - std::unordered_set called_funcs_; - - // Record the expressions that are being visited - std::unordered_set visiting_; - - explicit CallTracer(const IRModule& module) : module_{module}, called_funcs_{}, visiting_{} {} - - void VisitExpr_(const GlobalVarNode* op) final { - called_funcs_.insert(op->name_hint); - auto func = module_->Lookup(op->name_hint); - if (auto function_node = func.as()) { - VisitExpr(function_node.value()); - } - // else: Don't visit PrimFuncs -- we don't need to collect any tir.Calls therein. - } - - void VisitExpr_(const CallNode* call_node) final { - // TODO(mbs): Cleanup shape functions. - CallLoweredProps props = GetCallLoweredProps(call_node); - if (props.lowered_func.defined() && props.attrs.metadata.count("prim_shape_fn_var")) { - auto callee = Downcast(props.attrs.metadata["prim_shape_fn_var"]); - // We are implicitly calling the shape function *in addition to* the callee. - called_funcs_.insert(callee->name_hint); - } - ExprVisitor::VisitExpr_(call_node); - } - - void VisitExpr_(const FunctionNode* func_node) final { - auto func = GetRef(func_node); - if (visiting_.find(func) == visiting_.end()) { - visiting_.insert(func); - for (auto param : func_node->params) { - ExprVisitor::VisitExpr(param); - } - ExprVisitor::VisitExpr(func_node->body); - } - } - - std::unordered_set Trace(const std::string& entry) { - called_funcs_.insert(entry); - auto main_func = module_->Lookup(entry); - VisitExpr(main_func); - return called_funcs_; - } -}; - -/*! - * \brief Remove functions that are not used. - * - * \param module The Relay module. - * \param entry_funcs The set of functions that can be entry function. - * - * \return The module with dead functions removed. - */ -IRModule RemoveUnusedFunctions(const IRModule& module, Array entry_funcs) { - std::unordered_set called_funcs{}; - for (auto entry : entry_funcs) { - auto funcs = CallTracer(module).Trace(entry); - called_funcs.insert(funcs.cbegin(), funcs.cend()); - } - auto existing_functions = module->functions; - for (auto f : existing_functions) { - auto it = called_funcs.find(f.first->name_hint); - if (it == called_funcs.end()) { - module->Remove(f.first); - } - } - return module; -} - -} // namespace vm - -namespace transform { - -Pass RemoveUnusedFunctions(Array entry_functions) { - runtime::TypedPackedFunc pass_func = [=](IRModule m, - PassContext pc) { - return relay::vm::RemoveUnusedFunctions(m, entry_functions); - }; - return CreateModulePass(pass_func, 1, "RemoveUnusedFunctions", {}); -} - -TVM_REGISTER_GLOBAL("relay._transform.RemoveUnusedFunctions").set_body_typed(RemoveUnusedFunctions); - -} // namespace transform - -} // namespace relay -} // namespace tvm diff --git a/src/relay/collage/candidate_function_cache.cc b/src/relay/collage/candidate_function_cache.cc deleted file mode 100644 index 32982dc08f3d..000000000000 --- a/src/relay/collage/candidate_function_cache.cc +++ /dev/null @@ -1,49 +0,0 @@ -/* - * 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 src/relay/collage/candidate_function_cache.cc - * \brief A cache of the unique global name and costs for partitioned functions. - */ - -#include "./candidate_function_cache.h" - -namespace tvm { -namespace relay { -namespace collage { - -CandidateFunctionCache::Entry& CandidateFunctionCache::GetEntry(const std::string& label, - const Function& function) { - auto itr = cache_.find(function); - if (itr == cache_.end()) { - String compiler = function->GetAttr(attr::kCompiler, String("tvm")).value(); - std::string global_symbol_name = name_supply_->Fresh({compiler, label}); - GlobalVar global_symbol(std::move(global_symbol_name), function->checked_type()); - itr = cache_.emplace(function, Entry(std::move(global_symbol))).first; - } - return itr->second; -} - -GlobalVar CandidateFunctionCache::GetGlobalSymbol(const Function& function) { - return GetEntry(/*label=*/"", function).global_symbol; -} - -} // namespace collage -} // namespace relay -} // namespace tvm diff --git a/src/relay/collage/candidate_function_cache.h b/src/relay/collage/candidate_function_cache.h deleted file mode 100644 index 8734f5a8e1af..000000000000 --- a/src/relay/collage/candidate_function_cache.h +++ /dev/null @@ -1,79 +0,0 @@ -/* - * 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 src/relay/collage/candidate_function_cache.h - * \brief A cache of the unique global symbol name and cost for partitioned functions. - */ - -#ifndef TVM_RELAY_COLLAGE_CANDIDATE_FUNCTION_CACHE_H_ -#define TVM_RELAY_COLLAGE_CANDIDATE_FUNCTION_CACHE_H_ - -#include - -#include -#include -#include -#include - -#include "../transforms/compiler_function_utils.h" -#include "./cost.h" -#include "./name_supply.h" - -namespace tvm { -namespace relay { -namespace collage { - -/*! - * \brief A cache of the unique global symbol and cost for functions extracted to represent - * partitions. If two functions are structurally equal (which includes equality of their "Compiler" - * attributes) then they will share the same global symbol and estimated cost. We rely on the - * function's attributes to distinguish partitions which are structurally the same graph but - * intended for different targets. - */ -class CandidateFunctionCache : public transform::GlobalSymbolCache { - public: - explicit CandidateFunctionCache(std::shared_ptr name_supply) - : name_supply_(std::move(name_supply)) {} - - struct Entry { - GlobalVar global_symbol; - Cost cost = Cost::Unknown(); // Filled in when have estimated cost. - - explicit Entry(GlobalVar global_symbol) : global_symbol(std::move(global_symbol)) {} - }; - - /*! - * \brief Returns the unique entry for \p function. If no such entry already exists, create it - * and assign it a unique global symbol name. - */ - Entry& GetEntry(const std::string& label, const Function& function); - - GlobalVar GetGlobalSymbol(const Function& function) final; - - private: - std::shared_ptr name_supply_; - std::unordered_map cache_; -}; - -} // namespace collage -} // namespace relay -} // namespace tvm - -#endif // TVM_RELAY_COLLAGE_CANDIDATE_FUNCTION_CACHE_H_ diff --git a/src/relay/collage/candidate_partition.cc b/src/relay/collage/candidate_partition.cc deleted file mode 100644 index 2050fbddb16b..000000000000 --- a/src/relay/collage/candidate_partition.cc +++ /dev/null @@ -1,357 +0,0 @@ -/* - * 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 src/relay/collage/candidate_partition.cc - * \brief A potential partition in the Collage search. - */ - -#include "./candidate_partition.h" - -#include -#include -#include - -#include "../transforms/compiler_function_utils.h" -#include "./candidate_function_cache.h" -#include "./candidate_set.h" -#include "./partition_rule.h" -#include "./partition_spec.h" -#include "./utils.h" - -namespace tvm { -namespace relay { -namespace collage { - -TVM_REGISTER_NODE_TYPE(CandidatePartitionNode); - -void CandidatePartitionNode::VisitAttrs(AttrVisitor* v) { - v->Visit("rule_name", &rule_name_); - v->Visit("sub_graph", &sub_graph_); - v->Visit("spec", &spec_); - // TODO(mbs): cost_ -} - -PartitionSpec CandidatePartitionNode::partition_spec() const { - return Downcast(spec_); -} - -std::string CandidatePartitionNode::partition_spec_name() const { - return Downcast(spec_)->spec_name_; -} - -Target CandidatePartitionNode::target() const { return Downcast(spec_)->target_; } - -std::string CandidatePartitionNode::ToSummary(const DataflowGraph& dataflow_graph) const { - std::ostringstream os; - os << sub_graph_->label_; - os << " | ("; - bool first = true; - for (PostDfsIndex index : sub_graph_->input_) { - Expr sub_expr = dataflow_graph.index_to_node(index)->ref(); - if (CanInline(sub_expr)) { - continue; - } - if (first) { - first = false; - } else { - os << ", "; - } - os << PrettyPrint(sub_expr->checked_type()); - } - os << ") -> ("; - first = true; - for (PostDfsIndex index : sub_graph_->exit_) { - Expr sub_expr = dataflow_graph.index_to_node(index)->ref(); - if (CanInline(sub_expr)) { - continue; - } - if (first) { - first = false; - } else { - os << ", "; - } - os << PrettyPrint(sub_expr->checked_type()); - } - os << ") | "; - os << sub_graph_->inside_.ToString(); - os << " | "; - os << partition_spec_name(); - os << " | "; - os << cost_.ToString(); - return os.str(); -} - -std::string CandidatePartitionNode::ToString() const { - std::ostringstream os; - os << "{rule_name=" << rule_name_; - os << ",sub_graph=" << sub_graph_->ToString(); - os << ",spec_name=" << partition_spec_name(); - if (!cost_.is_unknown()) { - os << ",cost=" << cost_.ToString(); - } - os << "}"; - return os.str(); -} - -namespace { -/*! - * \brief If function's body is a call to an inlined "Primitive" function, return it. - * Otherwise return function directly. - */ -Function GetPrimitiveFunction(const Function& function) { - if (const auto* call_node = function->body.as()) { - if (const auto* function_node = call_node->op.as()) { - if (function_node->HasNonzeroAttr(attr::kPrimitive)) { - return GetRef(function_node); - } - } - } - return function; -} - -/*! - * \brief Eta-expand any tuple arguments of \p function. Ie rewrite: - * \code - * f(x: (t1, t2)) { ... x ... } - * \endcode - * to - * \code - * f(x_1: t1, x_2: t2) { ... (x_1, x_2) ... } - * \endcode - */ -Function EtaExpandTuples(const Function& function) { - Map subst; - Array new_params; - for (const auto& param : function->params) { - std::vector tensor_types = FlattenTupleType(param->type_annotation); - if (tensor_types.size() == 1) { - new_params.push_back(param); - } else { - Array fields; - for (size_t i = 0; i < tensor_types.size(); ++i) { - Var new_param(param->name_hint() + "_" + std::to_string(i), tensor_types[i], param->span); - new_param->checked_type_ = tensor_types[i]; - new_params.push_back(new_param); - fields.push_back(new_param); - } - Tuple new_tuple(fields); - subst.Set(param, new_tuple); - } - } - if (subst.empty()) { - return function; - } - return WithFields(function, new_params, Bind(function->body, subst)); -} - -} // namespace - -Cost CandidatePartitionNode::EstimatedCost( - const DataflowGraph& dataflow_graph, const CostEstimator& cost_estimator, - const std::shared_ptr& cache) const { - if (cost_.is_unknown()) { - VLOG_CONTEXT << "spec " << partition_spec_name(); - Function extracted_function = sub_graph_->ExtractAsFunction(dataflow_graph); - VLOG(2) << "Extracted function:" << std::endl << PrettyPrint(extracted_function); - extracted_function = EtaExpandTuples(extracted_function); - VLOG(2) << "Validating function:" << std::endl << PrettyPrint(extracted_function); - String error = partition_spec()->validate_sub_graph_func_(extracted_function); - if (!error.empty()) { - cost_ = Cost::Invalid(); - VLOG(1) << "Unable to rewrite function: " << error; - } else { - // The extracted function may be the eta-expansion of a "Primitive" function. - // If so we want the cached external name and cost to be w.r.t. that function - // rather than the outer so that we'll get a cache hit when we outline functions - // in the final program. - Function primitive_function = GetPrimitiveFunction(extracted_function); - CandidateFunctionCache::Entry& entry = - cache->GetEntry(sub_graph_->label_, primitive_function); - if (entry.cost.is_unknown()) { - IRModule mod = IRModule::FromExpr(extracted_function); - VLOG(1) << "Outlining:" << std::endl << PrettyPrint(mod); - mod = OutlineCompilerFunctions(cache)(mod); - VLOG(1) << "Estimating cost of:" << std::endl - << PrettyPrint(mod) << std::endl - << "using target " << target()->ToDebugString(); - entry.cost = cost_estimator->Estimate(mod, target()); - VLOG(1) << "Measured cost as " << entry.cost.ToString(); - } else { - VLOG(1) << "Reusing cost " << entry.cost.ToString() - << " cached in candidate function cache"; - } - cost_ = entry.cost; - } - } else { - VLOG(1) << "Reusing cost " << cost_.ToString() << " cached in candidate"; - } - return cost_; -} - -CandidatePartition::CandidatePartition(String rule_name, SubGraph sub_graph, - ObjectRef /* actually PartitionSpec */ spec, Cost cost) { - auto node = runtime::make_object(); - node->rule_name_ = std::move(rule_name); - node->sub_graph_ = std::move(sub_graph); - node->spec_ = std::move(spec); - node->cost_ = cost; - data_ = std::move(node); -} - -CandidatePartition WithRuleName(CandidatePartition candidate, String rule_name) { - if (rule_name == candidate->rule_name_) { - return candidate; - } - auto* node = candidate.CopyOnWrite(); - node->rule_name_ = std::move(rule_name); - return GetRef(node); -} - -CandidatePartition WithSubGraph(CandidatePartition candidate, SubGraph sub_graph) { - if (sub_graph == candidate->sub_graph_) { - return candidate; - } - auto* node = candidate.CopyOnWrite(); - node->sub_graph_ = std::move(sub_graph); - return GetRef(node); -} - -bool CandidatePartition::operator<(const CandidatePartition& that) const { - // Order lexicographically on sub-graphs. - if (*get()->sub_graph_.get() < *that->sub_graph_.get()) { - return true; - } - if (*that->sub_graph_.get() < *get()->sub_graph_.get()) { - return false; - } - // Break ties by rule name. - return get()->rule_name_ < that->rule_name_; -} - -bool CandidatePartition::AreTouching(const DataflowGraph& dataflow_graph, - const CandidatePartition& that) const { - return get()->spec_ == that->spec_ && - get()->sub_graph_.AreTouching(dataflow_graph, that->sub_graph_); -} - -CandidatePartition CandidatePartition::DisjointUnion(const DataflowGraph& dataflow_graph, - const CandidatePartition& that) const { - ICHECK_EQ(get()->spec_, that->spec_); - return CandidatePartition(UnionLabels(get()->rule_name_, that->rule_name_), - get()->sub_graph_.DisjointUnion(dataflow_graph, that->sub_graph_), - get()->spec_, get()->cost_ + that->cost_); -} - -/*static*/ -CandidatePartition CandidatePartition::DisjointUnion(const DataflowGraph& dataflow_graph, - std::vector candidates) { - ICHECK_GT(candidates.size(), 1); - CandidatePartition result = candidates.front(); - for (size_t i = 1; i < candidates.size(); ++i) { - result = result.DisjointUnion(dataflow_graph, candidates[i]); - } - return result; -} - -/*static*/ -Expr CandidatePartition::ParallelRewrite(const DataflowGraph& dataflow_graph, - const std::vector& candidates) { - std::vector sub_graphs; - sub_graphs.reserve(candidates.size()); - for (const auto& candidate : candidates) { - sub_graphs.emplace_back(candidate->sub_graph_); - } - return SubGraph::ParallelRewrite(dataflow_graph, sub_graphs); -} - -/*static*/ -std::vector CandidatePartition::MaxCoalesce( - const DataflowGraph& dataflow_graph, std::vector candidates) { - VLOG(1) << "Running MaxCoalesce over " << candidates.size() << " candidates"; - // This is an eager version of using the simple (kOpaque, kOpaque) combiner. - - // Switch to set representation. - CandidateSet result_set(std::move(candidates)); - - // Until fixed point... - size_t num_rounds = 0; - while (result_set.PrepareForNextRound()) { - VLOG_CONTEXT << "round " << ++num_rounds; - VLOG(1) << "checking " << result_set.size() << " candidates (" << result_set.first_new_index() - << " existing)"; - IndexSet removed_this_round(result_set.size()); // over candidate indexes! - - // Build map from post-dfs indices to the indices of candidates with corresponding entry node. - // NOTE: the index set is over candidate indices not post-dfs indices! - std::vector entry_map(dataflow_graph.size(), IndexSet(result_set.size())); - for (size_t i = 0; i < result_set.size(); ++i) { - CandidatePartition candidate = result_set.at(i); - for (PostDfsIndex entry_index : candidate->sub_graph_->entry_) { - entry_map[entry_index].Add(i); - } - } - - for (size_t i = 0; i < result_set.size(); ++i) { - if (removed_this_round[i]) { - // Already merged. - continue; - } - CandidatePartition upstream = result_set.at(i); - // Narrow our search to just those candidates which could touch. - IndexSet possible_downstream(result_set.size()); // over candidate indexes! - for (PostDfsIndex output_index : upstream->sub_graph_->output_) { - possible_downstream = possible_downstream | entry_map[output_index]; - } - for (size_t j : possible_downstream) { - if (removed_this_round[j]) { - // Already merged. - continue; - } - if (i == j) { - // Ignore self. - continue; - } - CandidatePartition downstream = result_set.at(j); - if (!upstream.AreTouching(dataflow_graph, downstream)) { - continue; - } - CandidatePartition new_candidate = upstream.DisjointUnion(dataflow_graph, downstream); - VLOG(2) << "Merging upstream candidate " << upstream->ToString() - << " and downstream candidate " << downstream->ToString() << " to yield " - << new_candidate->ToString(); - result_set.Add(dataflow_graph, new_candidate); - result_set.Remove(upstream); - removed_this_round.Add(i); - result_set.Remove(downstream); - removed_this_round.Add(j); - } - } - } - - // Restore canonical order. - result_set.sort(); - - VLOG(1) << "MaxCoalesce produced " << result_set.size() << " candidates"; - return result_set.MovedCurrentCandidates(); -} - -} // namespace collage -} // namespace relay -} // namespace tvm diff --git a/src/relay/collage/candidate_partition.h b/src/relay/collage/candidate_partition.h deleted file mode 100644 index 36a23f14bc53..000000000000 --- a/src/relay/collage/candidate_partition.h +++ /dev/null @@ -1,190 +0,0 @@ -/* - * 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 src/relay/collage/candidate_partition.cc - * \brief A potential partition in the Collage search. - */ - -#ifndef TVM_RELAY_COLLAGE_CANDIDATE_PARTITION_H_ -#define TVM_RELAY_COLLAGE_CANDIDATE_PARTITION_H_ - -#include -#include - -#include -#include -#include - -#include "./candidate_function_cache.h" -#include "./cost.h" -#include "./cost_estimator.h" -#include "./name_supply.h" -#include "./sub_graph.h" - -namespace tvm { -namespace relay { -namespace collage { - -class PartitionSpec; - -/*! - * \brief A candidate partition w.r.t. the overall Relay model. - * - * We represent the partition as a sub-graph. This means not only can we represent the scope - * of Relay sub-expressions intended for a particular partition (or kernel), but we can also - * represent various conventions for encoding how the operators within the partition should be - * tagged for downstream processing. - */ -class CandidatePartitionNode : public Object { - public: - CandidatePartitionNode() = default; - - /*! - * \brief Combination of all the partition rule names which produced this candidate. - * For debugging and explainability. - */ - String rule_name_; - - /*! - * \brief The sub-graph of the overall expression matched by the partition rule. - */ - SubGraph sub_graph_; - - /*! - * \brief The partition specification which produced this candidate. - */ - ObjectRef /* actually PartitionSpec */ spec_; - - /*! - * \brief The (cached) cost of the partition. - * - * Initially Cost::Unknown, calculated and cached by EstimateCost. - */ - mutable Cost cost_ = Cost::Unknown(); - - void VisitAttrs(AttrVisitor* v); - - /*! - * \brief Returns the partition specification which produced this candidate. - */ - PartitionSpec partition_spec() const; - - /*! - * \brief Returns the name of the partition specification which produced this candidate. - */ - std::string partition_spec_name() const; - - /*! - * \brief Returns the target of the partition specification which produced this candidate. - */ - Target target() const; - - /*! - * \brief Return the estimated cost of the candidate partition, using \p cost_estimator and - * \p cache. - */ - Cost EstimatedCost(const DataflowGraph& dataflow_graph, const CostEstimator& cost_estimator, - const std::shared_ptr& cache) const; - - /*! - * \brief Returns a brief description of candidate suitable for debugging output. - */ - std::string ToSummary(const DataflowGraph& dataflow_graph) const; - - std::string ToString() const; - - static constexpr const char* _type_key = "relay.collage.CandidatePartition"; - TVM_DECLARE_FINAL_OBJECT_INFO(CandidatePartitionNode, Object); -}; - -class CandidatePartition : public ObjectRef { - public: - CandidatePartition(String rule_name, SubGraph sub_graph, - ObjectRef /* actually PartitionSpec */ spec, Cost cost = Cost::Unknown()); - - bool operator<(const CandidatePartition& that) const; - - /*! - * \brief Returns true if this and \p that candidate are disjoint, have the same (or no) target, - * and touch. This does not imply the \p DisjointUnion of this and that will be valid. For - * example, the result may be too deep or have too many outputs. - */ - bool AreTouching(const DataflowGraph& dataflow_graph, const CandidatePartition& that) const; - - /*! - * \brief Returns the disjoint union of this and \p that. - */ - CandidatePartition DisjointUnion(const DataflowGraph& dataflow_graph, - const CandidatePartition& that) const; - - /*! - * \brief Returns the disjoint union of all \p candidates. - */ - static CandidatePartition DisjointUnion(const DataflowGraph& dataflow_graph, - std::vector candidates); - - /*! - * \brief Returns the root expression of \p dataflow_graph rewritten to apply all the partitions - * implied by \p candidates. The candidates can be in any order but must be disjoint. - */ - static Expr ParallelRewrite(const DataflowGraph& dataflow_graph, - const std::vector& candidates); - - /*! - * Eagerly merge all touching candidates for the same target. The candidates must be disjoint - * and have their Targets filled in. This is typically called on the optimal list of candidate - * partitions found by the Collage search in order to remove unnecessary partition boundaries. - * Ideally the search would never produce such candidates however to keep the search space - * manageable Collage may only consider candidate partitions up to a particular depth. - */ - static std::vector MaxCoalesce(const DataflowGraph& dataflow_graph, - std::vector candidates); - - TVM_DEFINE_OBJECT_REF_METHODS(CandidatePartition, ObjectRef, CandidatePartitionNode); - TVM_DEFINE_OBJECT_REF_COW_METHOD(CandidatePartitionNode); -}; - -CandidatePartition WithRuleName(CandidatePartition candidate, String rule_name); -CandidatePartition WithTarget(CandidatePartition candidate, Target target); -CandidatePartition WithSubGraph(CandidatePartition candidate, SubGraph sub_graph); - -struct CandidatePartitionHash { - size_t operator()(const CandidatePartition& candidate) const { - return candidate->sub_graph_->hash(); - } -}; - -struct CandidatePartitionEquals { - bool operator()(const CandidatePartition& left, const CandidatePartition& right) const { - return *left->sub_graph_.get() == *right->sub_graph_.get(); - } -}; - -struct CandidatePartitionCompare { - bool operator()(const CandidatePartition& left, const CandidatePartition& right) const { - return *left->sub_graph_.get() < *right->sub_graph_.get(); - } -}; - -} // namespace collage -} // namespace relay -} // namespace tvm - -#endif // TVM_RELAY_COLLAGE_CANDIDATE_PARTITION_H_ diff --git a/src/relay/collage/candidate_partition_index.cc b/src/relay/collage/candidate_partition_index.cc deleted file mode 100644 index 4a9cd65bff2f..000000000000 --- a/src/relay/collage/candidate_partition_index.cc +++ /dev/null @@ -1,148 +0,0 @@ -/* - * 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 relay/collage/candidate_partition_index.h - * \brief Index for finding relevant candidate partitions for a particular search state. - */ - -#include "./candidate_partition_index.h" - -#include "./gather_partition_specs.h" -#include "./prune_candidates.h" -#include "./utils.h" - -namespace tvm { -namespace relay { -namespace collage { - -CandidatePartitionIndex::CandidatePartitionIndex( - const std::unordered_map* virtual_devices, - DataflowGraph* dataflow_graph) - : virtual_devices_(virtual_devices), - dataflow_graph_(dataflow_graph), - first_inside_index_to_candidates_(dataflow_graph->size()) {} - -void CandidatePartitionIndex::Index(const Array& partition_specs) { - std::vector candidates = Collect(partition_specs); - candidates = PruneCandidates(*dataflow_graph_, candidates); - // Index the candidates by their first inside index. - for (auto& candidate : candidates) { - first_inside_index_to_candidates_[candidate->sub_graph_->first_inside_index_].emplace_back( - candidate); - } - size_ = candidates.size(); -} - -void CandidatePartitionIndex::EstimateAllCosts( - const CostEstimator cost_estimator, const std::shared_ptr& cache) { - size_t n = 0; - for (PostDfsIndex index = 0; index < dataflow_graph_->size(); ++index) { - for (const auto& candidate : first_inside_index_to_candidates_[index]) { - LOG(INFO) << "Estimating cost of candidate " << candidate->ToSummary(*dataflow_graph_) << " [" - << n++ << "/" << size_ << "]"; - // Cost will be cached in candidate as a side effect. - Cost cost = candidate->EstimatedCost(*dataflow_graph_, cost_estimator, cache); - LOG(INFO) << "Candidate has cost " << cost.ToString(); - } - } -} - -std::string CandidatePartitionIndex::ToSummary() const { - std::vector lines; - for (const auto& candidates : first_inside_index_to_candidates_) { - for (const auto& candidate : candidates) { - if (candidate->partition_spec_name() == kHostSpecName) { - continue; - } - lines.emplace_back(candidate->ToSummary(*dataflow_graph_)); - } - } - std::sort(lines.begin(), lines.end()); - std::ostringstream os; - bool first = true; - for (const auto& line : lines) { - if (first) { - first = false; - } else { - os << std::endl; - } - os << line; - } - return os.str(); -} - -bool CandidatePartitionIndex::IsCompatibleWithVirtualDevice(const CandidatePartition& candidate) { - for (PostDfsIndex index : candidate->sub_graph_->inside_) { - const ExprNode* sub_expr_node = dataflow_graph_->index_to_node(index)->node_ref_; - if (sub_expr_node->IsInstance() || sub_expr_node->IsInstance()) { - // These nodes are target/device polymorphic. - continue; - } - auto itr = virtual_devices_->find(sub_expr_node); - ICHECK(itr != virtual_devices_->end()) << PrettyPrint(GetRef(sub_expr_node)); - const Target& existing_target = itr->second->target; - if (!existing_target.defined()) { - // No constraint. - continue; - } - if (StructuralEqual()(existing_target, candidate->target())) { - // No disagreement. - continue; - } - if (!candidate->target().IsExternalCodegenFor(itr->second->target)) { - // The candidate's target is not an external codegen target compatible with the existing - // target. - // TODO(mbs): There's a conflict here between Collage's desire to leave some expression nodes - // 'behind' on the VM and PlanDevice's desire to assign a primitive Target to every node. - // I think PlanDevices is the one that needs to give here by leaving such nodes - // unconstrained. - VLOG(1) << "Ignoring candidate " << candidate->ToString() - << " since incompatible with existing virtual device assignment of:" << std::endl - << itr->second << std::endl - << "to sub-graph:" << std::endl - << PrettyPrint(GetRef(sub_expr_node)); - return false; - } - } - return true; -} - -std::vector CandidatePartitionIndex::Collect( - const Array& partition_specs) { - VLOG_CONTEXT << "collecting"; - std::vector result; - for (const auto& spec : partition_specs) { - VLOG_CONTEXT << "spec " << spec->spec_name_; - VLOG(1) << "collecting candidates"; - std::vector candidates = spec->AllCandidates(*dataflow_graph_); - for (auto& candidate : candidates) { - if (!IsCompatibleWithVirtualDevice(candidate)) { - continue; - } - result.push_back(candidate); - } - } - VLOG(1) << "Found " << result.size() << " candidates"; - return result; -} - -} // namespace collage -} // namespace relay -} // namespace tvm diff --git a/src/relay/collage/candidate_partition_index.h b/src/relay/collage/candidate_partition_index.h deleted file mode 100644 index aa3f7d4fcd81..000000000000 --- a/src/relay/collage/candidate_partition_index.h +++ /dev/null @@ -1,102 +0,0 @@ -/* - * 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 relay/collage/candidate_partition_index.h - * \brief Index for finding relevant candidate partitions for a particular search state. - */ -#ifndef TVM_RELAY_COLLAGE_CANDIDATE_PARTITION_INDEX_H_ -#define TVM_RELAY_COLLAGE_CANDIDATE_PARTITION_INDEX_H_ - -#include - -#include -#include -#include -#include - -#include "./partition_spec.h" - -namespace tvm { -namespace relay { -namespace collage { - -/*! - * \brief Collects and indexes all the candidate partitions for the overall expression. This index - * is used during partitioning search to find the next valid candidate partition to explore from the - * current search state. We do not yet attempt to estimate the cost of each candidate partition, and - * when we do so during the search we may discover it to be infeasible. - */ -class CandidatePartitionIndex { - public: - CandidatePartitionIndex(const std::unordered_map* virtual_devices, - DataflowGraph* dataflow_graph); - - /*! \brief Constructs the index. */ - void Index(const Array& partition_specs); - - /*! \brief Returns all the candidates which may begin at \p index. */ - const std::vector& candidates_at(PostDfsIndex index) const { - ICHECK_LT(index, dataflow_graph_->size()); - return first_inside_index_to_candidates_[index]; - } - - /*! \brief Estimates the casts of all candidates in the index. Each candidate caches its cost. */ - void EstimateAllCosts(const CostEstimator cost_estimator, - const std::shared_ptr& cache); - - size_t size() const { return size_; } - - std::string ToSummary() const; - - private: - /*! - * \brief Returns true if \p candidate's desired target is compatible with any existing target - * constraints on the candidate's sub-expressions. - */ - bool IsCompatibleWithVirtualDevice(const CandidatePartition& candidate); - - /*! \brief Returns all valid candidates found from \p partition_specs. */ - std::vector Collect(const Array& partition_specs); - - /*! - * \brief The \p VirtualDevice for every sub-expression in the overall expression. Needed to - * ensure candidates do not contradict the target/device placement already determined by - * device planning. - */ - const std::unordered_map* virtual_devices_; - - /*! \brief Dataflow graph for overall expression. */ - DataflowGraph* dataflow_graph_; - - /*! - * \brief Maps post-dfs indexes to the all the candidates which have that as their first inside - * index, and which should be considered in the Collage search. - */ - std::vector> first_inside_index_to_candidates_; - - /*! \brief Number of entries in above. */ - size_t size_ = 0; -}; - -} // namespace collage -} // namespace relay -} // namespace tvm - -#endif // TVM_RELAY_COLLAGE_CANDIDATE_PARTITION_INDEX_H_ diff --git a/src/relay/collage/candidate_set.cc b/src/relay/collage/candidate_set.cc deleted file mode 100644 index 2c2a7eaf8d54..000000000000 --- a/src/relay/collage/candidate_set.cc +++ /dev/null @@ -1,76 +0,0 @@ -/* - * 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 src/relay/collage/candidate_set.cc - * \brief Collects a set of candidate partitions. - */ - -#include "./candidate_set.h" - -namespace tvm { -namespace relay { -namespace collage { - -CandidateSet::CandidateSet(std::vector candidates_to_add) - : candidates_to_add_(std::move(candidates_to_add)) { - for (const auto& candidate : candidates_to_add_) { - seen_.emplace(candidate); - } -} - -void CandidateSet::Add(const DataflowGraph& dataflow_graph, - const CandidatePartition& new_candidate) { - VLOG(2) << "adding " << new_candidate->ToString(); - if (seen_.count(new_candidate)) { - VLOG(2) << "already seen candidate, ignoring"; - return; - } - seen_.emplace(new_candidate); - candidates_to_add_.emplace_back(new_candidate); -} - -void CandidateSet::Remove(const CandidatePartition& old_candidate) { - ICHECK(seen_.count(old_candidate)); - VLOG(2) << "removing " << old_candidate->ToString(); - candidates_to_remove_.emplace_back(old_candidate); -} - -bool CandidateSet::PrepareForNextRound() { - size_t init_size = current_candidates_.size(); - for (const auto& candidate_to_remove : candidates_to_remove_) { - current_candidates_.erase( - std::remove(current_candidates_.begin(), current_candidates_.end(), candidate_to_remove), - current_candidates_.end()); - } - size_t num_removed = init_size - current_candidates_.size(); - candidates_to_remove_.clear(); - first_new_index_ = current_candidates_.size(); - for (const auto& new_candidate : candidates_to_add_) { - current_candidates_.push_back(new_candidate); - } - size_t num_added = candidates_to_add_.size(); - candidates_to_add_.clear(); - VLOG(1) << "removed " << num_removed << " and added " << num_added << " candidates"; - return num_removed + num_added > 0; -} - -} // namespace collage -} // namespace relay -} // namespace tvm diff --git a/src/relay/collage/candidate_set.h b/src/relay/collage/candidate_set.h deleted file mode 100644 index 4cb2c40e9500..000000000000 --- a/src/relay/collage/candidate_set.h +++ /dev/null @@ -1,99 +0,0 @@ -/* - * 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 src/relay/collage/candidate_set.h - * \brief Collects a set of candidate partitions. - */ - -#ifndef TVM_RELAY_COLLAGE_CANDIDATE_SET_H_ -#define TVM_RELAY_COLLAGE_CANDIDATE_SET_H_ - -#include -#include -#include -#include - -#include "./candidate_partition.h" -#include "./dataflow_graph.h" - -namespace tvm { -namespace relay { -namespace collage { - -/*! - * \brief Holds a vector of current candidates and the additions/removals to apply to them. - */ -struct CandidateSet { - CandidateSet() = default; - - explicit CandidateSet(std::vector candidates_to_add); - - /*! - * \brief Schedule \p new_candidate for addition before the next round (unless it is not valid). - */ - void Add(const DataflowGraph& dataflow_graph, const CandidatePartition& new_candidate); - - /*! \brief Schedule \p old_candidate for removal before the next round. */ - void Remove(const CandidatePartition& old_candidate); - - /*! - * \brief Update \p current_candidates and \p first_new_index. Return false if no - * new candidates were added, in which case we have reached a fixed point. - */ - bool PrepareForNextRound(); - - size_t size() const { return current_candidates_.size(); } - - CandidatePartition operator[](size_t i) const { - ICHECK_LT(i, current_candidates_.size()); - return current_candidates_[i]; - } - CandidatePartition at(size_t i) const { return (*this)[i]; } - - size_t first_new_index() const { return first_new_index_; } - - void sort() { std::sort(current_candidates_.begin(), current_candidates_.end()); } - - std::vector MovedCurrentCandidates() { - return std::move(current_candidates_); - } - - private: - /*! - * \brief Index of first candidate in current_candidates added in last round. This can be used to - * avoid considering candidates or candidate combinations which have already been considered in an - * earlier round. - */ - size_t first_new_index_ = 0; - /*! \brief Candidates gathered in previous rounds. */ - std::vector current_candidates_; - /*! \brief New candidates gathered in the current round. */ - std::vector candidates_to_add_; - /*! \brief Existing candidates to remove before starting the next round. */ - std::vector candidates_to_remove_; - /*! \brief Which candidates have been seen so far and should not be added again. */ - std::unordered_set seen_; -}; - -} // namespace collage -} // namespace relay -} // namespace tvm - -#endif // TVM_RELAY_COLLAGE_CANDIDATE_SET_H_ diff --git a/src/relay/collage/collage_partitioner.cc b/src/relay/collage/collage_partitioner.cc deleted file mode 100644 index 54fc6c45ca70..000000000000 --- a/src/relay/collage/collage_partitioner.cc +++ /dev/null @@ -1,352 +0,0 @@ -/* - * 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 src/relay/collage/collage_partitioner.cc - * \brief Search for an optimal partitioning of a Relay model. - */ - -#include "./collage_partitioner.h" - -#include -#include -#include -#include -#include -#include -#include -#include - -#include "../ir/dataflow_matcher_impl.h" -#include "../transforms/compiler_function_utils.h" -#include "../transforms/device_aware_visitors.h" -#include "./candidate_partition.h" -#include "./candidate_partition_index.h" -#include "./cost.h" -#include "./cost_estimator.h" -#include "./gather_partition_specs.h" -#include "./name_supply.h" -#include "./partition_rule.h" -#include "./partition_spec.h" -#include "./priority_queue.h" -#include "./sub_graph.h" -#include "./utils.h" - -namespace tvm { -namespace relay { -namespace collage { -namespace { - -TVM_REGISTER_PASS_CONFIG_OPTION("relay.collage.tvm_max_depth", Integer); -TVM_REGISTER_PASS_CONFIG_OPTION("relay.collage.byoc_max_depth", Integer); -TVM_REGISTER_PASS_CONFIG_OPTION("relay.collage.byoc_fusion_style", Array); -/*! - * \brief Represents the overall expression after some number of non-overlapping candidate - * partitions have been applied. - */ -class SearchState { - public: - explicit SearchState(IndexSet covered) : covered_(std::move(covered)) {} - - /*! - * \brief Order states by increasing best cost, breaking ties by lexicographic order on - * the covering sub graph. - */ - bool operator<(const SearchState& that) const { - return std::tie(best_cost_, covered_) < std::tie(that.best_cost_, that.covered_); - } - - const IndexSet& covered() const { return covered_; } - - std::string ToString() const { - std::ostringstream os; - os << "State("; - os << "covered=" << covered_.ToString(); - os << ",best_cost=" << best_cost_.ToString(); - if (best_candidate_.defined()) { - os << ",best_candidate=" << best_candidate_->ToString(); - } - os << ")"; - return os.str(); - } - - private: - /*! \brief Which nodes of overall expression have been placed on all paths to this state. */ - IndexSet covered_; - /*! \brief Predecessor state for sequence of candidates reaching this state with least - * cost. Null if initial search state. */ - SearchState* pred_state_ = nullptr; - /*! - * \brief Cost of reaching this state using placement implied by path given by pred_state fields. - * Includes estimated/measured cost of all candidates plus any candidate launch penalty. - * Initially invalid cost. - */ - Cost best_cost_ = Cost::Invalid(); - /*! \brief Candidate partition selected in transition from pred_state to this state. */ - CandidatePartition best_candidate_; - - friend class Partitioner; -}; - -struct CompareSearchStatePtrs { - bool operator()(const SearchState* left, const SearchState* right) const { - return *left < *right; - } -}; - -struct EqualSearchStatePtrs { - bool operator()(const SearchState* left, const SearchState* right) const { - return left->covered() == right->covered(); - } -}; - -/*! - * \brief Finds the optimal partitioning of an expression to candidate partitions. - * Though no candidate partitions overlap, it is possible some sub-expressions end up in - * no candidate. Those sub-expressions must be evaluated by the host executor (eg VM). - */ -class Partitioner { - public: - explicit Partitioner(Array partition_specs, - const std::unordered_map* virtual_devices, - CostEstimator cost_estimator, std::shared_ptr cache, - Expr expr) - : partition_specs_(std::move(partition_specs)), - virtual_devices_(virtual_devices), - cost_estimator_(std::move(cost_estimator)), - cache_(std::move(cache)), - expr_(std::move(expr)) {} - - Expr Partition() { - // Establish core data structures. - dataflow_graph_ = std::make_unique(expr_); - VLOG(1) << "Created dataflow graph with " << dataflow_graph_->size() << " nodes"; - - // Build the candidate index. This is where all the partition rules are invoked . - index_ = std::make_unique(virtual_devices_, dataflow_graph_.get()); - index_->Index(partition_specs_); - VLOG(1) << "All candidates before search:" << std::endl << index_->ToSummary(); - - // 'Eagerly' estimate the cost of all candidates. - // - // Note if this is not done costs will simply be estimated 'lazily' as the search proceeds. - // Typically, some candidates are never explored during the search because: - // - There are no paths in which the candidate does not intersect candidates already - // applied on the path. - // - The Dijkstra search terminates early with a least cost path. - // So eager may result in more estimation overhead. However, eager could be made - // embarrassingly parallel. - VLOG(1) << "Beginning eager cost estimation"; - index_->EstimateAllCosts(cost_estimator_, cache_); - VLOG(1) << "Finished eager cost estimation"; - - // Setup initial state. - SearchState* init_state = GetState(IndexSet(dataflow_graph_->size())); - init_state->best_cost_ = Cost::Zero(); - pq_.Push(init_state); - - size_t num_transitions = 0; - - VLOG(1) << "#### Commencing Collage search over " << index_->size() << " candidates ####"; - while (!pq_.empty()) { - SearchState* curr_state = pq_.Pop(); - VLOG(1) << "Looking at state " << curr_state->covered_.ToString(); - PostDfsIndex next_index = curr_state->covered_.FirstOutsideIndex(); - - if (next_index >= dataflow_graph_->size()) { - // The entire expression has been explored. Collect the candidates on the optimal path. - VLOG(1) << "#### Finished Collage search after exploring " << num_transitions - << " transitions ####"; - std::vector best_candidates; - while (curr_state != init_state) { - ICHECK(curr_state->best_candidate_.defined()); - best_candidates.emplace_back(curr_state->best_candidate_); - curr_state = curr_state->pred_state_; - ICHECK(curr_state != nullptr); - } - return Finalize(best_candidates); - } - - size_t num_fires = 0; - Expr sub_expr = dataflow_graph_->index_to_node(next_index)->ref(); - VLOG(1) << "Looking at index " << next_index << " for sub-expression " - << SubExprKindAndLabel(sub_expr).second << " out of " << dataflow_graph_->size() - << " total dataflow nodes"; - - // Explore all the outgoing candidates from the current state. - for (const auto& candidate : index_->candidates_at(next_index)) { - VLOG(1) << "Considering candidate " << candidate->ToSummary(*dataflow_graph_) - << " for transition " << ++num_transitions << " over " << index_->size() - << " total candidates"; - if (!candidate->sub_graph_->inside_.AreDisjoint(curr_state->covered_)) { - LOG(INFO) << "Candidate overlaps with already partitioned nodes"; - continue; - } - IndexSet next_covered = curr_state->covered_ | candidate->sub_graph_->inside_; - SearchState* next_state = GetState(next_covered); - Relax(curr_state, next_state, candidate); - ++num_fires; - } - ICHECK_GT(num_fires, 0) - << "No candidate was found covering sub-expression at index " << next_index - << ", suggesting the partition rules are incomplete for the given targets."; - } - - ICHECK(false) << "should have reached end state in which all sub-expressions are covered"; - return {}; - } - - /*! \brief Returns the unique state corresponding to the \p covered sub-graph. */ - SearchState* GetState(const IndexSet& covered) { - auto itr = covered_to_state_.find(covered); - if (itr != covered_to_state_.end()) { - return itr->second.get(); - } - auto state = std::make_unique(covered); - SearchState* raw_ptr = state.get(); - covered_to_state_.emplace(covered, std::move(state)); - return raw_ptr; - } - - /*! - * \brief Record that it is possible to reach \p next_state by choosing \p candidate - * in \p curr_state. If the resulting cost is better than the best known so far, update - * \p next_state's best cost, predecessor and candidate to match. - */ - void Relax(SearchState* curr_state, SearchState* next_state, - const CandidatePartition& candidate) { - // Note this may already be cached if the candidate partition costs were 'eagerly' estimated. - Cost candidate_cost = candidate->EstimatedCost(*dataflow_graph_, cost_estimator_, cache_); - VLOG(1) << "Candidate has cost " << candidate_cost.ToString(); - Cost new_state_cost = candidate_cost + curr_state->best_cost_; - const bool is_new = next_state->best_cost_.is_invalid(); - CandidatePartition previously_best_candidate = next_state->best_candidate_; - if (is_new || new_state_cost < next_state->best_cost_) { - next_state->pred_state_ = curr_state; - Cost previously_best_cost = next_state->best_cost_; - next_state->best_cost_ = new_state_cost; - next_state->best_candidate_ = candidate; - if (is_new) { - VLOG(1) << "transition " << curr_state->ToString() << " --> " << next_state->ToString() - << " (New state for spec " << candidate->partition_spec_name() << ")"; - pq_.Push(next_state); - } else { - VLOG(1) << "transition " << curr_state->ToString() << " --> " << next_state->ToString() - << " (Spec " << candidate->partition_spec_name() << " beats previous spec " - << previously_best_candidate->partition_spec_name() << " by " - << (previously_best_cost - curr_state->best_cost_).ToString() << ")"; - pq_.Update(next_state); - } - } else { - VLOG(1) << "transition " << curr_state->ToString() << " --> " << next_state->ToString() - << " (Spec " << candidate->partition_spec_name() << " does not beat existing spec " - << previously_best_candidate->partition_spec_name() << ")"; - } - } - - /*! - * \brief Returns the result of partitioning \p expr according to 'optimal' candidates found - * by the search. - */ - Expr Finalize(std::vector best_candidates) { - best_candidates = CandidatePartition::MaxCoalesce(*dataflow_graph_, best_candidates); - - Cost total_cost = Cost::Zero(); - std::ostringstream os; - os << "Optimal partitioning:" << std::endl; - for (const auto& best_candidate : best_candidates) { - if (best_candidate->partition_spec_name() == kHostSpecName) { - continue; - } - os << best_candidate->ToSummary(*dataflow_graph_); - os << std::endl; - total_cost = total_cost + best_candidate->cost_; - } - os << "Estimated overall cost is " << total_cost.ToString(); - LOG(INFO) << os.str(); - - LOG(INFO) << "All candidates after search:" << std::endl << index_->ToSummary(); - - return CandidatePartition::ParallelRewrite(*dataflow_graph_, best_candidates); - } - - private: - /*! \brief Available partition specs to use during search. */ - Array partition_specs_; - /*! - * \brief The virtual devices for every sub-expression so we can respect any existing target - * constraints. - */ - const std::unordered_map* virtual_devices_; - /*! \brief Cost estimator to use for candidates. */ - CostEstimator cost_estimator_; - /*! \brief Cached names and costs for all partition functions. */ - std::shared_ptr cache_; - /*! \brief The expression we will be partitioning. */ - Expr expr_; - /*! \brief Dataflow graph for overall expression. */ - std::unique_ptr dataflow_graph_; - /*! \brief Index of all avoilable candidates we are searching over. */ - std::unique_ptr index_; - /*! \brief Map from covered sub-graphs to the corresponding state. */ - std::unordered_map, IndexSetHash, IndexSetEqual> - covered_to_state_; - /*! \brief Priority queue of states, ordered by increasing cost. */ - PriorityQueue pq_; -}; - -} // namespace - -transform::Pass CollagePartition(CompilationConfig config, CostEstimator cost_estimator) { - runtime::TypedPackedFunc pass_func = - [config = std::move(config), cost_estimator = std::move(cost_estimator)]( - IRModule mod, transform::PassContext ctxt) { - VLOG(1) << "CollagePartition input:" << std::endl << PrettyPrint(mod); - - Array partition_specs = GatherPartitionSpecs(config); - VLOG(1) << "Gathered " << partition_specs.size() << " partition specs"; - - auto cache = - std::make_shared(std::make_shared("collage")); - - IRModule out_mod = mod->ShallowCopy(); - for (const auto& kv : mod->functions) { - if (const auto* function_node = AsOptimizableFunctionNode(kv.second)) { - auto function = GetRef(function_node); - std::unordered_map virtual_devices = - transform::RecoverVirtualDeviceMap(mod, function); - Partitioner partitioner(partition_specs, &virtual_devices, cost_estimator, cache, - function); - Function result = Downcast(partitioner.Partition()); - out_mod->Add(kv.first, result); - } - } - - out_mod = OutlineCompilerFunctions(cache)(std::move(out_mod)); - VLOG(1) << "CollagePartition result:" << std::endl << PrettyPrint(out_mod); - return out_mod; - }; - return tvm::transform::CreateModulePass(pass_func, /*opt_level=*/0, "CollagePartition", {}); -} - -TVM_REGISTER_GLOBAL("relay._transform.CollagePartition").set_body_typed(CollagePartition); - -} // namespace collage -} // namespace relay -} // namespace tvm diff --git a/src/relay/collage/collage_partitioner.h b/src/relay/collage/collage_partitioner.h deleted file mode 100644 index 7c8de87ffe0a..000000000000 --- a/src/relay/collage/collage_partitioner.h +++ /dev/null @@ -1,50 +0,0 @@ -/* - * 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 relay/collage/collage_partitioner.h - * \brief Search for an optimal partitioning of a Relay model. - * - * See: - * Collage: Automated Integration of Deep Learning Backends - * Byungsoo Jeon, Sunghyun Park, Peiyuan Liao, Sheng Xu, Tianqi Chen, Zhihao Jia - * https://arxiv.org/pdf/2111.00655.pdf - */ -#ifndef TVM_RELAY_COLLAGE_COLLAGE_PARTITIONER_H_ -#define TVM_RELAY_COLLAGE_COLLAGE_PARTITIONER_H_ - -#include - -#include "./cost_estimator.h" - -namespace tvm { -namespace relay { -namespace collage { - -/*! - * \brief Explores the space of all possible (sub-graph, target) pairs which cover the - * model, and applies the globally optimal choice (assuming partition costs are additive). - */ -transform::Pass CollagePartition(CompilationConfig config, CostEstimator cost_estimator); - -} // namespace collage -} // namespace relay -} // namespace tvm - -#endif // TVM_RELAY_COLLAGE_COLLAGE_PARTITIONER_H_ diff --git a/src/relay/collage/combiner_rule.cc b/src/relay/collage/combiner_rule.cc deleted file mode 100644 index bcfef0477292..000000000000 --- a/src/relay/collage/combiner_rule.cc +++ /dev/null @@ -1,395 +0,0 @@ -/* - * 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 src/relay/collage/combiner_rule.cc - * \brief Helpers for the \p CombinePartitionRule - */ - -#include "./combiner_rule.h" - -#include "./partition_spec.h" - -namespace tvm { -namespace relay { -namespace collage { - -TVM_REGISTER_NODE_TYPE(SimpleCombinerRuleNode); - -void SimpleCombinerRuleNode::VisitAttrs(AttrVisitor* v) { - // TODO(mbs) -} - -bool SimpleCombinerRuleNode::Fires(const DataflowGraph& dataflow_graph, - const CandidatePartition& upstream, - const CandidatePartition& downstream) const { - return false; -} - -std::string SimpleCombinerRuleNode::ToString() const { - return "SimpleCombinerRule(" + rule_name_ + ")"; -} - -SimpleCombinerRule::SimpleCombinerRule(String rule_name) { - auto node = runtime::make_object(); - node->rule_name_ = std::move(rule_name); - data_ = std::move(node); -} - -TVM_REGISTER_NODE_TYPE(ByKindSimpleCombinerRuleNode); - -void ByKindSimpleCombinerRuleNode::VisitAttrs(AttrVisitor* v) { - // TODO(mbs) -} - -bool ByKindSimpleCombinerRuleNode::Fires(const DataflowGraph& dataflow_graph, - const CandidatePartition& upstream, - const CandidatePartition& downstream) const { - return upstream->sub_graph_->kind_ <= upstream_kind_ && - downstream->sub_graph_->kind_ <= downstream_kind_; -} - -std::string ByKindSimpleCombinerRuleNode::ToString() const { - std::ostringstream os; - os << "ByKindSimpleCombinerRule(" << rule_name_ << ")"; - return os.str(); -} - -ByKindSimpleCombinerRule::ByKindSimpleCombinerRule(OpPatternKind upstream_kind, - OpPatternKind downstream_kind) { - auto node = runtime::make_object(); - String rule_name = KindToString(upstream_kind) + "->" + KindToString(downstream_kind); - node->rule_name_ = std::move(rule_name); - node->upstream_kind_ = upstream_kind; - node->downstream_kind_ = downstream_kind; - data_ = std::move(node); -} - -TVM_REGISTER_NODE_TYPE(CombinerRuleNode); - -void CombinerRuleNode::VisitAttrs(AttrVisitor* v) { - // TODO(mbs) -} - -void CombinerRuleNode::AppendAllResults(AppendAllResultsContext* ctxt) const {} - -std::string CombinerRuleNode::ToString() const { return "CombinerRuleNode(" + rule_name_ + ")"; } - -CombinerRule::CombinerRule(String rule_name) { - auto node = runtime::make_object(); - node->rule_name_ = std::move(rule_name); - data_ = std::move(node); -} - -TVM_REGISTER_NODE_TYPE(AllSimpleCombinerRuleNode); - -void AllSimpleCombinerRuleNode::VisitAttrs(AttrVisitor* v) { - // TODO(mbs) -} - -void AllSimpleCombinerRuleNode::AppendAllResults(AppendAllResultsContext* ctxt) const { - VLOG(1) << "running AllSimpleCombinerRule(" << rule_name_ << ")"; - // Build map from post-dfs indices to the indices of candidates with corresponding entry node. - // NOTE: the index set is over candidate indices not post-dfs indices! - std::vector entry_map(ctxt->dataflow_graph->size(), - IndexSet(ctxt->candidate_set->size())); - for (size_t i = 0; i < ctxt->candidate_set->size(); ++i) { - CandidatePartition candidate = ctxt->candidate_set->at(i); - for (PostDfsIndex entry_index : candidate->sub_graph_->entry_) { - entry_map[entry_index].Add(i); - } - } - - for (size_t i = 0; i < ctxt->candidate_set->size(); ++i) { - CandidatePartition upstream = ctxt->candidate_set->at(i); - // Narrow our search to just those candidates which could touch. - IndexSet possible_downstream(ctxt->candidate_set->size()); - for (PostDfsIndex output_index : upstream->sub_graph_->output_) { - possible_downstream = possible_downstream | entry_map[output_index]; - } - size_t start_j = - i < ctxt->candidate_set->first_new_index() ? ctxt->candidate_set->first_new_index() : 0; - for (size_t j : possible_downstream) { - if (i == j) { - continue; - } - if (i < start_j) { - // We already explored the cross-product of candidates [0, first_new_index), so don't - // do it again. - continue; - } - // Note that the rules are not commutative so we can't just ignore if j < i. - CandidatePartition downstream = ctxt->candidate_set->at(j); - if (ctxt->max_depth > 0 && - upstream->sub_graph_->depth_ + downstream->sub_graph_->depth_ > ctxt->max_depth) { - continue; - } - if (!upstream.AreTouching(*ctxt->dataflow_graph, downstream)) { - continue; - } - for (const auto& simple_rule : simple_rules_) { - if (simple_rule->Fires(*ctxt->dataflow_graph, upstream, downstream)) { - CandidatePartition new_candidate = - upstream.DisjointUnion(*ctxt->dataflow_graph, downstream); - VLOG(2) << "Fired " << simple_rule->rule_name_ << " on upstream candidate " - << upstream->ToString() << " and downstream candidate " << downstream->ToString() - << " to yield " << new_candidate->ToString(); - ctxt->candidate_set->Add(*ctxt->dataflow_graph, new_candidate); - } - } - } - } -} - -std::string AllSimpleCombinerRuleNode::ToString() const { - std::ostringstream os; - os << "AllSimpleCombinerRule(" << rule_name_; - for (const auto& simple : simple_rules_) { - os << ", " << simple->ToString(); - } - os << ")"; - return os.str(); -} - -AllSimpleCombinerRule::AllSimpleCombinerRule(String rule_name, - Array simple_rules) { - auto node = runtime::make_object(); - node->rule_name_ = std::move(rule_name); - node->simple_rules_ = std::move(simple_rules); - data_ = std::move(node); -} - -TVM_REGISTER_NODE_TYPE(TupleArgCombinerRuleNode); - -void TupleArgCombinerRuleNode::VisitAttrs(AttrVisitor* v) { - // TODO(mbs) -} - -void TupleArgCombinerRuleNode::AppendAllResults(AppendAllResultsContext* ctxt) const { - VLOG(1) << "running TupleArgCombinerRule(" << rule_name_ << ")"; - // Build map from post-dfs index to the indices of injective candidates with corresponding entry - // node. NOTE: the index set is over candidate indices not post-dfs indices! - std::vector exit_map(ctxt->dataflow_graph->size(), - IndexSet(ctxt->candidate_set->size())); - for (size_t i = 0; i < ctxt->candidate_set->size(); ++i) { - CandidatePartition candidate = ctxt->candidate_set->at(i); - if (candidate->sub_graph_->kind_ > kInjective) { - continue; - } - for (PostDfsIndex exit_index : candidate->sub_graph_->exit_) { - exit_map[exit_index].Add(i); - } - } - - // The two-step I -> tuple -> I rule. - // Look all possible tuple consumers... - for (size_t i = 0; i < ctxt->candidate_set->size(); ++i) { - CandidatePartition tuple_consumer_candidate = ctxt->candidate_set->at(i); - if (tuple_consumer_candidate->sub_graph_->kind_ > kInjective) { - continue; - } - // For all possible tuples feeding into candidate... - for (PostDfsIndex input_index : tuple_consumer_candidate->sub_graph_->input_) { - auto node = ctxt->dataflow_graph->index_to_node(input_index); - Expr sub_expr = node->ref(); - const auto* tuple_node = sub_expr.as(); - if (tuple_node == nullptr) { - continue; - } - // The tuple_consumer_candidate candidate consumes (at least one) tuple, eg as an argument - // to an operator. - // eg: concatenate((field1, ..., fieldn)) - auto tuple_dataflow_node = ctxt->dataflow_graph->item_to_node(tuple_node); - - // Collect all the possible unions. There may be more than one if different candidates - // could supply the same tuple field. - std::vector> all_possible_unions; - - // Obviously we must include the consumer. - all_possible_unions.emplace_back(); - all_possible_unions.back().emplace_back(tuple_consumer_candidate); - - // We must include the tuple itself. - SubGraph tuple_sub_graph(*ctxt->dataflow_graph, - IndexSet(ctxt->dataflow_graph->size(), {node->index_}), kInjective, - "tuple"); - CandidatePartition tuple_candidate("", std::move(tuple_sub_graph), - tuple_consumer_candidate->partition_spec()); - all_possible_unions.back().emplace_back(std::move(tuple_candidate)); - - // For all tuple fields... - bool all_tuple_fields_have_producer = true; - for (auto* tuple_field_dataflow_node : tuple_dataflow_node->inputs_) { - // Collect all the candidates which could produce this tuple field. - std::vector to_appends; - size_t start_j = - i < ctxt->candidate_set->first_new_index() ? ctxt->candidate_set->first_new_index() : 0; - for (size_t j : exit_map[tuple_field_dataflow_node->index_]) { - if (i == j) { - continue; - } - if (i < start_j) { - // We already explored the cross-product of candidates [0, first_new_index), so don't - // do it again. - continue; - } - CandidatePartition tuple_field_producer = ctxt->candidate_set->at(j); - // The tuple_field_producer candidate can provide this tuple field. - // eg concatenate((..., producer, ...)) - to_appends.emplace_back(tuple_field_producer); - } - if (to_appends.empty()) { - // At least one of the tuple's fields does not have a producer candidate we can - // union in, so we need to give up. - all_tuple_fields_have_producer = false; - break; - } else { - // If to_appends = [A, B] and we already have possible unions [C, D] and [E, F] then - // the new possible unions are [C, D, A], [C, D, B], [E, F, A] and [E, F, B]. - std::vector> new_all_possible_unions; - for (const auto& to_append : to_appends) { - for (const auto& possible_union : all_possible_unions) { - new_all_possible_unions.emplace_back(possible_union); - new_all_possible_unions.back().emplace_back(to_append); - } - } - all_possible_unions = std::move(new_all_possible_unions); - } - } - - if (!all_tuple_fields_have_producer) { - continue; - } - - // Actually build the candidates which union according to all_possible_unions. - for (const auto& possible_union : all_possible_unions) { - if (possible_union.size() > 2) { - CandidatePartition new_candidate = - CandidatePartition::DisjointUnion(*ctxt->dataflow_graph, possible_union); -#if TVM_LOG_DEBUG - std::ostringstream os; - bool first = true; - for (const auto& candidate : possible_union) { - if (first) { - first = false; - } else { - os << ", "; - } - os << candidate->ToString(); - } - VLOG(2) << "Fired rule " << rule_name_ << " on {" << os.str() << "} to yield " - << new_candidate->ToString(); -#endif - ctxt->candidate_set->Add(*ctxt->dataflow_graph, new_candidate); - } - } - } - } -} - -std::string TupleArgCombinerRuleNode::ToString() const { - return "TupleArgCombinerRule(" + rule_name_ + ")"; -} - -TupleArgCombinerRule::TupleArgCombinerRule(String rule_name) { - auto node = runtime::make_object(); - node->rule_name_ = std::move(rule_name); - data_ = std::move(node); -} - -TVM_REGISTER_NODE_TYPE(TupleProjCombinerRuleNode); - -void TupleProjCombinerRuleNode::VisitAttrs(AttrVisitor* v) { - // TODO(mbs) -} - -void TupleProjCombinerRuleNode::AppendAllResults(AppendAllResultsContext* ctxt) const { - VLOG(1) << "running TupleProjCombinerRule(" << rule_name_ << ")"; - // We already explored [0, first_new_index), so don't do it again. - for (size_t i = ctxt->candidate_set->first_new_index(); i < ctxt->candidate_set->size(); ++i) { - CandidatePartition base = ctxt->candidate_set->at(i); - for (PostDfsIndex index : base->sub_graph_->output_) { - auto node = ctxt->dataflow_graph->index_to_node(index); - if (node->ref().as()) { - IndexSet index_set(ctxt->dataflow_graph->size(), {node->index_}); - SubGraph sub_graph(*ctxt->dataflow_graph, std::move(index_set), kInjective, "proj"); - CandidatePartition proj_candidate("", std::move(sub_graph), base->spec_); - CandidatePartition new_candidate = - base.DisjointUnion(*ctxt->dataflow_graph, proj_candidate); - VLOG(2) << "Fired rule " << rule_name_ << " on " << proj_candidate->ToString() << " and " - << base->ToString() << " to yield " << new_candidate->ToString(); - ctxt->candidate_set->Add(*ctxt->dataflow_graph, new_candidate); - } - } - } -} - -std::string TupleProjCombinerRuleNode::ToString() const { - return "TupleProjCombinerRule(" + rule_name_ + ")"; -} - -TupleProjCombinerRule::TupleProjCombinerRule(String rule_name) { - auto node = runtime::make_object(); - node->rule_name_ = std::move(rule_name); - data_ = std::move(node); -} - -TVM_REGISTER_NODE_TYPE(ConstantCombinerRuleNode); - -void ConstantCombinerRuleNode::VisitAttrs(AttrVisitor* v) { - // TODO(mbs) -} - -void ConstantCombinerRuleNode::AppendAllResults(AppendAllResultsContext* ctxt) const { - VLOG(1) << "running ConstantCombinerRule(" << rule_name_ << ")"; - // We already explored [0, first_new_index), so don't do it again. - for (size_t i = ctxt->candidate_set->first_new_index(); i < ctxt->candidate_set->size(); ++i) { - CandidatePartition base = ctxt->candidate_set->at(i); - IndexSet new_constants(ctxt->dataflow_graph->size()); - for (PostDfsIndex index : base->sub_graph_->input_) { - auto node = ctxt->dataflow_graph->index_to_node(index); - if (node->ref().as()) { - new_constants.Add(index); - } - } - if (!new_constants.IsZero()) { - SubGraph sub_graph(*ctxt->dataflow_graph, new_constants, kElemWise, "const"); - CandidatePartition new_const_candidate("", std::move(sub_graph), base->spec_); - CandidatePartition new_candidate = - base.DisjointUnion(*ctxt->dataflow_graph, new_const_candidate); - VLOG(2) << "Fired rule " << rule_name_ << " on " << new_const_candidate->ToString() << " and " - << base->ToString() << " to yield " << new_candidate->ToString(); - ctxt->candidate_set->Add(*ctxt->dataflow_graph, new_candidate); - } - } -} - -std::string ConstantCombinerRuleNode::ToString() const { - return "ConstantCombinerRule(" + rule_name_ + ")"; -} - -ConstantCombinerRule::ConstantCombinerRule(String rule_name) { - auto node = runtime::make_object(); - node->rule_name_ = std::move(rule_name); - data_ = std::move(node); -} - -} // namespace collage -} // namespace relay -} // namespace tvm diff --git a/src/relay/collage/combiner_rule.h b/src/relay/collage/combiner_rule.h deleted file mode 100644 index 04ea2a9cc127..000000000000 --- a/src/relay/collage/combiner_rule.h +++ /dev/null @@ -1,229 +0,0 @@ -/* - * 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 src/relay/collage/combiner_rule.h - * \brief Helpers for the \p CombinePartitionRule - */ - -#ifndef TVM_RELAY_COLLAGE_COMBINER_RULE_H_ -#define TVM_RELAY_COLLAGE_COMBINER_RULE_H_ - -#include -#include - -#include - -#include "./candidate_partition.h" -#include "./candidate_set.h" -#include "./sub_graph.h" - -namespace tvm { -namespace relay { -namespace collage { - -/*! - * \brief Base class for all 'simple' combiner rules. - * - * Given \p upstream and \p downstream candidates which touch, a simple combiner rule returns - * true if their union should also be considered a candidate. - */ -class SimpleCombinerRuleNode : public Object { - public: - String rule_name_; - - void VisitAttrs(AttrVisitor* v); - - virtual bool Fires(const DataflowGraph& dataflow_graph, const CandidatePartition& upstream, - const CandidatePartition& downstream) const; - - virtual std::string ToString() const; - - static constexpr const char* _type_key = "relay.collage.SimpleCombinerRule"; - static constexpr const uint32_t _type_child_slots = 1; - TVM_DECLARE_BASE_OBJECT_INFO(SimpleCombinerRuleNode, Object); -}; - -class SimpleCombinerRule : public ObjectRef { - public: - explicit SimpleCombinerRule(String rule_name); - - TVM_DEFINE_OBJECT_REF_METHODS(SimpleCombinerRule, ObjectRef, SimpleCombinerRuleNode); -}; - -/*! - * \brief A simple combiner rule which fires if the \p upstream and \p downstream candidates have - * the given \p upstream_kind and \p downstream_kind (or less) respectively. - */ -class ByKindSimpleCombinerRuleNode : public SimpleCombinerRuleNode { - public: - OpPatternKind upstream_kind_; - OpPatternKind downstream_kind_; - - void VisitAttrs(AttrVisitor* v); - - bool Fires(const DataflowGraph& dataflow_graph, const CandidatePartition& upstream, - const CandidatePartition& downstream) const override; - std::string ToString() const override; - - static constexpr const char* _type_key = "relay.collage.ByKindSimpleCombinerRule"; - TVM_DECLARE_FINAL_OBJECT_INFO(ByKindSimpleCombinerRuleNode, SimpleCombinerRuleNode); -}; - -class ByKindSimpleCombinerRule : public SimpleCombinerRule { - public: - ByKindSimpleCombinerRule(OpPatternKind upstream_kind, OpPatternKind downstream_kind); - - TVM_DEFINE_OBJECT_REF_METHODS(ByKindSimpleCombinerRule, SimpleCombinerRule, - ByKindSimpleCombinerRuleNode); -}; - -/*! \brief Context required by CombineRuleNode::AppendAllResultsContext. */ -struct AppendAllResultsContext { - AppendAllResultsContext(const DataflowGraph* dataflow_graph, size_t max_depth, - CandidateSet* candidate_set) - : dataflow_graph(dataflow_graph), max_depth(max_depth), candidate_set(candidate_set) {} - - const DataflowGraph* dataflow_graph; - size_t max_depth; - CandidateSet* candidate_set; -}; - -/*! - * \brief Base class for all 'combiner' rules. - * - * Given the current candidate set, a combiner rule looks for opportunities to form larger - * candidates, optionally removing existing candidates in the process. - */ -class CombinerRuleNode : public Object { - public: - String rule_name_; - - void VisitAttrs(AttrVisitor* v); - - virtual void AppendAllResults(AppendAllResultsContext* ctxt) const; - virtual std::string ToString() const; - - static constexpr const char* _type_key = "relay.collage.CombinerRule"; - static constexpr const uint32_t _type_child_slots = 4; - TVM_DECLARE_BASE_OBJECT_INFO(CombinerRuleNode, Object); -}; - -class CombinerRule : public ObjectRef { - public: - explicit CombinerRule(String rule_name); - - TVM_DEFINE_OBJECT_REF_METHODS(CombinerRule, ObjectRef, CombinerRuleNode); -}; - -/*! - * \brief A combiner rule which runs one or more simple combiner rules over the current - * touching candidates. - */ -class AllSimpleCombinerRuleNode : public CombinerRuleNode { - public: - Array simple_rules_; - - void VisitAttrs(AttrVisitor* v); - - void AppendAllResults(AppendAllResultsContext* ctxt) const override; - std::string ToString() const override; - - static constexpr const char* _type_key = "relay.collage.AllSimpleCombinerRule"; - TVM_DECLARE_FINAL_OBJECT_INFO(AllSimpleCombinerRuleNode, CombinerRuleNode); -}; - -class AllSimpleCombinerRule : public CombinerRule { - public: - AllSimpleCombinerRule(String rule_name, Array simple_rules); - - TVM_DEFINE_OBJECT_REF_METHODS(AllSimpleCombinerRule, CombinerRule, AllSimpleCombinerRuleNode); -}; - -/*! - * \brief A combiner rule which combines injective sub-groups which appear inside tuples which are - * themselves inputs to injective sub-groups. - */ -class TupleArgCombinerRuleNode : public CombinerRuleNode { - public: - void VisitAttrs(AttrVisitor* v); - - void AppendAllResults(AppendAllResultsContext* ctxt) const override; - std::string ToString() const override; - - static constexpr const char* _type_key = "relay.collage.TupleArgCombinerRule"; - TVM_DECLARE_FINAL_OBJECT_INFO(TupleArgCombinerRuleNode, CombinerRuleNode); -}; - -class TupleArgCombinerRule : public CombinerRule { - public: - explicit TupleArgCombinerRule(String rule_name); - - TVM_DEFINE_OBJECT_REF_METHODS(TupleArgCombinerRule, CombinerRule, TupleArgCombinerRuleNode); -}; - -/*! - * \brief A combiner rule which combines tuple projection if it's an output of an injective - * group. - */ -class TupleProjCombinerRuleNode : public CombinerRuleNode { - public: - void VisitAttrs(AttrVisitor* v); - - void AppendAllResults(AppendAllResultsContext* ctxt) const override; - std::string ToString() const override; - - static constexpr const char* _type_key = "relay.collage.TupleProjCombinerRule"; - TVM_DECLARE_FINAL_OBJECT_INFO(TupleProjCombinerRuleNode, CombinerRuleNode); -}; - -class TupleProjCombinerRule : public CombinerRule { - public: - explicit TupleProjCombinerRule(String rule_name); - - TVM_DEFINE_OBJECT_REF_METHODS(TupleProjCombinerRule, CombinerRule, TupleProjCombinerRuleNode); -}; - -/*! - * \brief A combiner rule which combines constants in argument positions to existing candidates. - * Note that scalars are always inlined, so this rule only combines tensor constant arguments. - */ -class ConstantCombinerRuleNode : public CombinerRuleNode { - public: - void VisitAttrs(AttrVisitor* v); - - void AppendAllResults(AppendAllResultsContext* ctxt) const override; - std::string ToString() const override; - - static constexpr const char* _type_key = "relay.collage.ConstantCombinerRule"; - TVM_DECLARE_FINAL_OBJECT_INFO(ConstantCombinerRuleNode, CombinerRuleNode); -}; - -class ConstantCombinerRule : public CombinerRule { - public: - explicit ConstantCombinerRule(String rule_name); - - TVM_DEFINE_OBJECT_REF_METHODS(ConstantCombinerRule, CombinerRule, ConstantCombinerRuleNode); -}; - -} // namespace collage -} // namespace relay -} // namespace tvm - -#endif // TVM_RELAY_COLLAGE_COMBINER_RULE_H_ diff --git a/src/relay/collage/cost.cc b/src/relay/collage/cost.cc deleted file mode 100644 index ae2eb8600ebd..000000000000 --- a/src/relay/collage/cost.cc +++ /dev/null @@ -1,45 +0,0 @@ -/* - * 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 src/relay/collage/cost.cc - * \brief Represents the estimated cost of a candidate partition. - */ - -#include "./cost.h" - -namespace tvm { -namespace relay { -namespace collage { - -std::string Cost::ToString() const { - if (is_invalid()) { - return "invalid"; - } else if (is_unknown()) { - return "unknown"; - } else if (value_ == 0.0) { - return "0"; - } else { - return std::to_string(value_ * 1e6) + "us"; - } -} - -} // namespace collage -} // namespace relay -} // namespace tvm diff --git a/src/relay/collage/cost.h b/src/relay/collage/cost.h deleted file mode 100644 index 723c5b58ac94..000000000000 --- a/src/relay/collage/cost.h +++ /dev/null @@ -1,108 +0,0 @@ -/* - * 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 src/relay/collage/cost.h - * \brief Represents the estimated cost of a candidate partition. - */ -#ifndef TVM_RELAY_COLLAGE_COST_H_ -#define TVM_RELAY_COLLAGE_COST_H_ - -#include - -#include -#include -#include - -namespace tvm { -namespace relay { -namespace collage { - -/*! - * \brief The assumed cost for a candidate partition. Generally average execution time in seconds. - * However other cost functions are possible, for example to introduce a penalty for high memory - * use, etc. - */ -class Cost { - public: - Cost() = delete; - - static Cost Zero() { return Cost(0.0); } - - /*! - * \brief Returns the distinguished 'invalid' cost signaling a candidate partition is not - * supported by the intended target, for example because the sub-graph has an unsupported operator - * or the intermediate memory required exceeds some system limit. - */ - static Cost Invalid() { return Cost(std::numeric_limits::infinity()); } - - bool is_invalid() const { return std::isinf(value_) && value_ > 0.0; } - - /*! - * \brief Returns the distinguished 'unknown' cost, signaling fixed priorities should be used to - * choose the best partitions. This can be used to disable tuning and fallback to fixed rules, - * much as TVM will use an un-tuned kernel if no tuning records are available. - */ - static Cost Unknown() { return Cost(std::numeric_limits::quiet_NaN()); } - - bool is_unknown() const { return std::isnan(value_); } - - /*! \brief Returns cost with given finite, non-negative value. */ - static Cost Value(double value) { - ICHECK(!std::isnan(value) && !std::isinf(value) && value >= 0.0); - return Cost(value); - } - - bool is_value() const { return !std::isnan(value_) && !std::isinf(value_); } - - double value() const { - ICHECK(is_value()); - return value_; - } - - /*! \brief Return true if the less-than relation is defined for this and that. */ - bool are_comparable(Cost that) const { return !std::isnan(value_) && !std::isnan(that.value_); } - - /*! \brief Returns sum of this and that. */ - Cost operator+(Cost that) const { return Cost(value_ + that.value_); } - - /*! \brief Returns difference of this and that. */ - Cost operator-(Cost that) const { return Cost(value_ - that.value_); } - - /*! \brief Returns true if this is cheaper than that, assuming they are comparable. */ - bool operator<(Cost that) const { return value_ < that.value_; } - - std::string ToString() const; - - private: - explicit Cost(double value) : value_(value) {} - - /*! - * \brief Non-negative value or: - * - +inf if candidate partition is not feasible. - * - NaN if candidate partition has an unknown cost (priority may be used to break ties). - */ - double value_ = 0.0; -}; - -} // namespace collage -} // namespace relay -} // namespace tvm - -#endif // TVM_RELAY_COLLAGE_COST_H_ diff --git a/src/relay/collage/cost_estimator.cc b/src/relay/collage/cost_estimator.cc deleted file mode 100644 index 8197e58f67a4..000000000000 --- a/src/relay/collage/cost_estimator.cc +++ /dev/null @@ -1,59 +0,0 @@ -/* - * 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 src/relay/collage/cost_estimator.cc - * \brief Interface for measuring candidate partition cost. - */ - -#include "./cost_estimator.h" - -#include - -namespace tvm { -namespace relay { -namespace collage { - -TVM_REGISTER_OBJECT_TYPE(CostEstimatorNode); - -CostEstimator::CostEstimator() { - auto node = make_object(); - data_ = std::move(node); -} - -Cost CostEstimatorNode::Estimate(const IRModule& mod, const Target& target) const { - // TODO(mbs): Eventually should be abstract. For now bounce to the Python local impl. - static const runtime::PackedFunc* estimate_seconds = - runtime::Registry::Get("tvm.relay.collage.estimate_seconds"); - ICHECK(estimate_seconds); - const double value = (*estimate_seconds)(mod, target); - if (std::isinf(value)) { - return Cost::Invalid(); - } else if (std::isnan(value)) { - return Cost::Unknown(); - } else { - return Cost::Value(value); - } -} - -TVM_REGISTER_GLOBAL("relay.collage.CostEstimator").set_body_typed([]() { return CostEstimator(); }); - -} // namespace collage -} // namespace relay -} // namespace tvm diff --git a/src/relay/collage/cost_estimator.h b/src/relay/collage/cost_estimator.h deleted file mode 100644 index 55f389be685a..000000000000 --- a/src/relay/collage/cost_estimator.h +++ /dev/null @@ -1,71 +0,0 @@ -/* - * 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 src/relay/collage/cost_estimator.cc - * \brief Interface for measuring candidate partition cost. - */ - -#ifndef TVM_RELAY_COLLAGE_COST_ESTIMATOR_H_ -#define TVM_RELAY_COLLAGE_COST_ESTIMATOR_H_ - -#include - -#include "./cost.h" - -namespace tvm { -namespace relay { -namespace collage { - -/*! - * \brief An (abstract) estimator for the cost of executing "main" in an \p IRModule representing - * a candidate partition, using the given target for lowering and codegen. - * - * Generally the implementation will compile to a \p runtime::Module (possibly on a target-specific - * worker if cross-compilation is not available), repeatedly invoke "main" with random data until - * measure variance is acceptable (on a target-specific worker), and return the summarized costs. - * - * If using a TVM native \p Target, it is possible compilation will itself invoke TVM tuning. - * - * TODO(mbs): Actually, currently not abstract so can get some local measurements. - */ -class CostEstimatorNode : public Object { - public: - /*! - * \brief Returns the estimated cost (possibly after many many minutes of training time) of - * running "main" in \p mod using \p target, which represents a possible partitioning of - * some overall Relay expression. - */ - virtual Cost Estimate(const IRModule& mod, const Target& target) const; - - static constexpr const char* _type_key = "relay.collage.CostEstimator"; - TVM_DECLARE_BASE_OBJECT_INFO(CostEstimatorNode, Object); -}; - -class CostEstimator : public ObjectRef { - public: - CostEstimator(); - TVM_DEFINE_NOTNULLABLE_OBJECT_REF_METHODS(CostEstimator, ObjectRef, CostEstimatorNode); -}; - -} // namespace collage -} // namespace relay -} // namespace tvm - -#endif // TVM_RELAY_COLLAGE_COST_ESTIMATOR_H_ diff --git a/src/relay/collage/custom_cost_estimator.cc b/src/relay/collage/custom_cost_estimator.cc deleted file mode 100644 index dea4df072cac..000000000000 --- a/src/relay/collage/custom_cost_estimator.cc +++ /dev/null @@ -1,60 +0,0 @@ -/* - * 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 src/relay/collage/custom_cost_estimator.cc - * \brief A custom CostEstimator to support alternative cost functions. - */ - -#include "./custom_cost_estimator.h" - -#include - -namespace tvm { -namespace relay { -namespace collage { - -TVM_REGISTER_OBJECT_TYPE(CustomCostEstimatorNode); - -Cost CustomCostEstimatorNode::Estimate(const IRModule& mod, const Target& target) const { - static const runtime::PackedFunc* estimate_seconds = runtime::Registry::Get(py_fn_estimator_); - ICHECK(estimate_seconds); - const double value = (*estimate_seconds)(mod, target); - if (std::isinf(value)) { - return Cost::Invalid(); - } else if (std::isnan(value)) { - return Cost::Unknown(); - } else { - return Cost::Value(value); - } -} - -CustomCostEstimator::CustomCostEstimator(String py_fn_estimator) { - auto node = make_object(); - node->py_fn_estimator_ = std::move(py_fn_estimator); - data_ = std::move(node); -} - -TVM_REGISTER_GLOBAL("relay.collage.CustomCostEstimator").set_body_typed([](String py_fn_estimator) { - return CustomCostEstimator(std::move(py_fn_estimator)); -}); - -} // namespace collage -} // namespace relay -} // namespace tvm diff --git a/src/relay/collage/custom_cost_estimator.h b/src/relay/collage/custom_cost_estimator.h deleted file mode 100644 index 4e6b45832eb2..000000000000 --- a/src/relay/collage/custom_cost_estimator.h +++ /dev/null @@ -1,67 +0,0 @@ -/* - * 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 src/relay/collage/custom_cost_estimator.cc - * \brief A custom CostEstimator to support target-specific cost functions. - */ - -#ifndef TVM_RELAY_COLLAGE_CUSTOM_COST_ESTIMATOR_H_ -#define TVM_RELAY_COLLAGE_CUSTOM_COST_ESTIMATOR_H_ - -#include - -#include "./cost.h" -#include "./cost_estimator.h" - -namespace tvm { -namespace relay { -namespace collage { - -/*! - * \brief A cost estimator that uses a target-specific cost function. - */ -class CustomCostEstimatorNode : public CostEstimatorNode { - public: - Cost Estimate(const IRModule& mod, const Target& target) const override; - - static constexpr const char* _type_key = "relay.collage.CustomCostEstimator"; - TVM_DECLARE_FINAL_OBJECT_INFO(CustomCostEstimatorNode, CostEstimatorNode); - - protected: - /*! - * \brief Python implemented cost function name. - */ - String py_fn_estimator_; - - friend class CustomCostEstimator; -}; - -class CustomCostEstimator : public CostEstimator { - public: - explicit CustomCostEstimator(String py_fn_estimator); - - TVM_DEFINE_OBJECT_REF_METHODS(CustomCostEstimator, CostEstimator, CustomCostEstimatorNode); -}; - -} // namespace collage -} // namespace relay -} // namespace tvm - -#endif // TVM_RELAY_COLLAGE_CUSTOM_COST_ESTIMATOR_H_ diff --git a/src/relay/collage/dataflow_graph.cc b/src/relay/collage/dataflow_graph.cc deleted file mode 100644 index b4e19a73f04d..000000000000 --- a/src/relay/collage/dataflow_graph.cc +++ /dev/null @@ -1,48 +0,0 @@ -/* - * 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 src/relay/collage/dataflow_graph.cc - * \brief A representation of the dataflow for an overall Relay expression. - */ - -#include "./dataflow_graph.h" - -namespace tvm { -namespace relay { -namespace collage { - -DataflowGraph::DataflowGraph(Expr expr) : expr_(std::move(expr)) { - indexed_graph_ = CreateIndexedGraph(expr_); - downstream_map_.reserve(indexed_graph_->size()); - for (PostDfsIndex index = 0; index < indexed_graph_->size(); ++index) { - const Node* node = indexed_graph_->index_to_node(index); - std::unordered_set downstream_nodes; - node->AccumulateDownstreamNodes(&downstream_nodes); - IndexSet index_set(indexed_graph_->size()); - for (const Node* downstream_node : downstream_nodes) { - index_set.Add(downstream_node->index_); - } - downstream_map_.emplace_back(std::move(index_set)); - } -} - -} // namespace collage -} // namespace relay -} // namespace tvm diff --git a/src/relay/collage/dataflow_graph.h b/src/relay/collage/dataflow_graph.h deleted file mode 100644 index c3c22381a889..000000000000 --- a/src/relay/collage/dataflow_graph.h +++ /dev/null @@ -1,77 +0,0 @@ -/* - * 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 src/relay/collage/dataflow_graph.h - * \brief A representation of the dataflow for an overall Relay expression. - */ -#ifndef TVM_RELAY_COLLAGE_DATAFLOW_GRAPH_H_ -#define TVM_RELAY_COLLAGE_DATAFLOW_GRAPH_H_ - -#include - -#include -#include - -#include "../ir/indexed_graph.h" -#include "./index_set.h" - -namespace tvm { -namespace relay { -namespace collage { - -/*! - * \brief Represents the dataflow of an overall Relay expression. - */ -class DataflowGraph { - public: - using Node = IndexedGraph::Node; - - explicit DataflowGraph(Expr expr); - - size_t size() const { return indexed_graph_->size(); } - const Node* index_to_node(PostDfsIndex index) const { - return indexed_graph_->index_to_node(index); - } - const Node* item_to_node(const Expr& expr) const { return indexed_graph_->item_to_node(expr); } - const Node* item_to_node(const ExprNode* expr_node) const { - return indexed_graph_->item_to_node(expr_node); - } - const Expr& expr() const { return expr_; } - const IndexedGraph& indexed_graph() const { return *indexed_graph_; } - - const IndexSet& downstream_of(PostDfsIndex index) const { - ICHECK_LT(index, indexed_graph_->size()); - return downstream_map_[index]; - } - - private: - /*! \brief The overall expression. */ - Expr expr_; - /*! \brief The indexed graph which captures the main dataflow. */ - std::unique_ptr> indexed_graph_; - /*! \brief Map from a node's PostDfsIndex to the set of its downstream dataflow node indexes. */ - std::vector downstream_map_; -}; - -} // namespace collage -} // namespace relay -} // namespace tvm - -#endif // TVM_RELAY_COLLAGE_DATAFLOW_GRAPH_H_ diff --git a/src/relay/collage/gather_partition_specs.cc b/src/relay/collage/gather_partition_specs.cc deleted file mode 100644 index ad451673341d..000000000000 --- a/src/relay/collage/gather_partition_specs.cc +++ /dev/null @@ -1,241 +0,0 @@ -/* - * 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 src/relay/collage/gather_partition_specs.cc - * \brief Gather the relevant \p PartitionSpecs from the available \p Targets. - */ - -#include "./gather_partition_specs.h" - -#include "./utils.h" - -namespace tvm { -namespace relay { -namespace collage { - -namespace { - -PartitionRule MakeCombinePartitionRule(PartitionRule sub_rule, Array combiner_rules, - size_t max_depth) { - if (combiner_rules.empty()) { - return sub_rule; - } else { - return CombinePartitionRule("", std::move(sub_rule), std::move(combiner_rules), max_depth); - } -} - -/*! \brief Returns the primitive combiner rules which mimic TVM's \p FuseOps. */ -Array TVMCombinerRules() { - Array simple_rules; - // Mimic the FuseOps rules. - simple_rules.push_back(ByKindSimpleCombinerRule(kOutEWiseFusable, kBroadcast)); - simple_rules.push_back(ByKindSimpleCombinerRule(kBroadcast, kCommReduce)); - simple_rules.push_back(ByKindSimpleCombinerRule(kInjective, kInjective)); - - Array combiner_rules; - // Fire the simple fusion rules - combiner_rules.push_back(AllSimpleCombinerRule("combiner", std::move(simple_rules))); - // Fuse tuple arguments - combiner_rules.push_back(TupleArgCombinerRule("tuple")); - // Fuse tuple projection - combiner_rules.push_back(TupleProjCombinerRule("proj")); - - return combiner_rules; -} - -size_t GetMaxDepth(std::string key) { - tvm::transform::PassContext ctxt = tvm::transform::PassContext::Current(); - std::string config_key = "relay.collage." + key; - Optional opt_max_depth = ctxt->GetConfig(config_key, Optional()); - ICHECK(opt_max_depth.defined()) << "missing binding for '" << config_key << " in pass context"; - ICHECK(opt_max_depth.value()->value > 0) - << "invalid value for '" << config_key << " in pass context"; - return static_cast(opt_max_depth.value()->value); -} - -/*! \brief Returns partition rule mimicking TVM FuseOps. */ -PartitionRule MakeTVMPartitionRule() { - size_t max_depth = GetMaxDepth("tvm_max_depth"); - // Build singleton candidates for all calls to ops <= kOutEWiseFusable. - OpCallByKindPartitionRule op_call_by_kind(""); - // Combine candidates according to the TVM fusion rules. - PartitionRule combine = - MakeCombinePartitionRule(std::move(op_call_by_kind), TVMCombinerRules(), max_depth); - // Discard invalid candidates. - SubGraphConfig sub_graph_config; - sub_graph_config.allow_taps = false; - sub_graph_config.max_depth = max_depth; - sub_graph_config.max_exits = 1; - return OnlyValidPartitionRule("", std::move(combine), sub_graph_config); - // NOTE: We don't wrap by a "Primitive" since we want to defer making TVM fusion decisions until - // after running more Relay passes. -} - -/*! - * \brief Returns the fusion style for default compiler. - */ -BYOCStyle DefaultBYOCFusionStyleForCompiler(const String& compiler) { - if (compiler == "cutlass" || compiler == "cublas" || compiler == "cudnn") { - return kNoFusionBYOCStyle; - } else if (compiler == "tensorrt") { - return kTVMFusionBYOCStyle; - } else { - return kArbitraryFusionBYOCStyle; - } -} - -/*! - * \brief Returns the fusion style for given compiler. - */ -BYOCStyle BYOCFusionStyleForCompiler(const String& compiler) { - tvm::transform::PassContext ctxt = tvm::transform::PassContext::Current(); - std::string config_key = "relay.collage.byoc_fusion_style"; - Optional> byoc_configs = ctxt->GetConfig(config_key, Optional>()); - BYOCStyle byoc_fusion_style = DefaultBYOCFusionStyleForCompiler(compiler); - if (!byoc_configs) { - return byoc_fusion_style; - } - for (auto config_ : byoc_configs.value()) { - std::vector byoc_cfg = SplitString(config_, "."); - if (byoc_cfg[0] == compiler) { - if (byoc_cfg[1] == "NoFusion") { - byoc_fusion_style = kNoFusionBYOCStyle; - } else if (byoc_cfg[1] == "TVMFusion") { - byoc_fusion_style = kTVMFusionBYOCStyle; - } else if (byoc_cfg[1] == "ArbitraryFusion") { - byoc_fusion_style = kArbitraryFusionBYOCStyle; - } else { - ICHECK(false) << "Invalid fusion name for compiler " << byoc_cfg[0] << " in pass context"; - } - break; - } - } - return byoc_fusion_style; -} - -/*! - * \brief Returns the primitive combiner rules which allow for any touching candidates - * to be fused provided they don't have kind \p kOpaque. - */ -Array BYOCCombinerRules(const String& compiler) { - Array simple_rules; - Array combiner_rules; - switch (BYOCFusionStyleForCompiler(compiler)) { - case kNoFusionBYOCStyle: - break; - case kTVMFusionBYOCStyle: - // Conservatively assume the BYOC toolchain follows the same rules as for TVM's FuseOps. - simple_rules.push_back(ByKindSimpleCombinerRule(kOutEWiseFusable, kBroadcast)); - simple_rules.push_back(ByKindSimpleCombinerRule(kBroadcast, kCommReduce)); - simple_rules.push_back(ByKindSimpleCombinerRule(kInjective, kInjective)); - combiner_rules.push_back(AllSimpleCombinerRule("combiner", std::move(simple_rules))); - break; - case kArbitraryFusionBYOCStyle: - // Just try all combinations up to the max_depth limit. - simple_rules.push_back(ByKindSimpleCombinerRule(kOutEWiseFusable, kOutEWiseFusable)); - combiner_rules.push_back(AllSimpleCombinerRule("combiner", std::move(simple_rules))); - break; - } - return combiner_rules; -} - -/*! - * \brief Returns partition rule mimicking one entry in the patterns list passed to the - * MergeComposite pass. - */ -PartitionRule MakeLabelledDFPatternPartitionRule( - const std::string& compiler, String rule_name, DFPattern dataflow_pattern, - TPatternPredicate predicate = DefaultPatternPredicate) { - DFPatternPartitionRule patterns("", std::move(dataflow_pattern), std::move(predicate)); - return CompositePartitionRule(std::move(rule_name), std::move(patterns)); -} - -/*! - * \brief Returns partition rule mimicking - * MergeComposite/AnnotateTarget/MergeCompilerRegions/PartitionGraph passes for "compiler" - * attribute of \p target. - */ -PartitionRule MakePatternBYOCPartitionRule(const std::string& compiler, - Array sub_rules) { - size_t max_depth = GetMaxDepth("byoc_max_depth"); - // Union all the individual pattern rules. - UnionPartitionRule unioned("", std::move(sub_rules)); - PartitionRule combine = - MakeCombinePartitionRule(std::move(unioned), BYOCCombinerRules(compiler), max_depth); - // Ignore invalid candidates. - SubGraphConfig sub_graph_config; - sub_graph_config.allow_taps = false; - sub_graph_config.max_depth = max_depth; - sub_graph_config.max_exits = 1; - OnlyValidPartitionRule valid("", std::move(combine), sub_graph_config); - // Wrap the candidates in a "Primitive" function with a "Compiler" attribute. - return PrimitivePartitionRule("", std::move(valid)); -} - -TVM_REGISTER_GLOBAL("relay.collage.MakeLabelledDFPatternPartitionRule") - .set_body_typed(MakeLabelledDFPatternPartitionRule); - -TVM_REGISTER_GLOBAL("relay.collage.MakeLabelledDFPatternPartitionRuleWithPredicate") - .set_body_typed(MakeLabelledDFPatternPartitionRule); - -TVM_REGISTER_GLOBAL("relay.collage.MakePatternBYOCPartitionRule") - .set_body_typed(MakePatternBYOCPartitionRule); - -/*! - * \brief Returns the rule to pick out expression nodes which can be 'left behind' for execution - * on the host. - */ -PartitionRule MakeHostPartitionRule() { return HostPartitionRule(""); } - -} // namespace - -Array GatherPartitionSpecs(const CompilationConfig& config) { - Array result; - for (const auto& primitive_target : config->primitive_targets) { - String spec_name = GetSpecName(primitive_target); - PartitionRule rule; - if (primitive_target.IsExternalCodegen()) { - // Transition to the Python side so we can get access to the BYOC pattern registry. - // That will bounce right back into the above construction helpers. - static const runtime::PackedFunc* make_byoc_partition_rule = - runtime::Registry::Get("tvm.relay.collage.make_byoc_partition_rule"); - ICHECK(make_byoc_partition_rule); - rule = (*make_byoc_partition_rule)(spec_name); // spec_name == primitive_target->kind->name - VLOG(1) << "Target " << primitive_target->ToDebugString() << " is for BYOC spec_name " - << spec_name << " and has default partition rule:\n" - << rule->ToString(); - } else { - rule = MakeTVMPartitionRule(); - VLOG(1) << "Target " << primitive_target->ToDebugString() << " is for TVM spec_name " - << spec_name << " and has default partition rule:\n" - << rule->ToString(); - } - result.push_back(PartitionSpec(spec_name, primitive_target, rule)); - } - - // Add one more spec to cover the host target. - result.push_back(PartitionSpec(kHostSpecName, config->host_target, MakeHostPartitionRule())); - - return result; -} - -} // namespace collage -} // namespace relay -} // namespace tvm diff --git a/src/relay/collage/gather_partition_specs.h b/src/relay/collage/gather_partition_specs.h deleted file mode 100644 index 62ffca27d635..000000000000 --- a/src/relay/collage/gather_partition_specs.h +++ /dev/null @@ -1,71 +0,0 @@ -/* - * 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 src/relay/collage/gather_partition_specs.h - * \brief Gather the relevant \p PartitionSpecs from the available \p Targets. - */ -#ifndef TVM_RELAY_COLLAGE_GATHER_PARTITION_SPECS_H_ -#define TVM_RELAY_COLLAGE_GATHER_PARTITION_SPECS_H_ - -#include - -#include "./partition_spec.h" - -namespace tvm { -namespace relay { -namespace collage { - -/*! - * \brief The 'styles' of BYOC integrations. Used to influence how their corresponding - * partition rule is constructed. - */ -enum BYOCStyle { - /*! - * \brief The BYOC patterns pick out 'ideal' candidates directly, either because: - * - the BYOC toolchain does not perform any fusion so each matched sub-expression maps 1:1 to a - * BYOC-provided operator, or - * - the BYOC toolchain does perform fusion, however the patterns have been written to pick out - * fusable sub-graphs. - */ - kNoFusionBYOCStyle, - - /*! - * \brief The BYOC patterns pick out supported operators, but the BYOC backend may perform - * fusion over those operators in much the same way TVM does. - */ - kTVMFusionBYOCStyle, - - /*! - * \brief The BYOC patterns pick out supported operators, but the BYOC backend may perform - * arbitrary fusion over those operators. - */ - kArbitraryFusionBYOCStyle, -}; - -/*! - * \brief Returns all the partition specifications gathered from the \p Targets in \p config. - */ -Array GatherPartitionSpecs(const CompilationConfig& config); - -} // namespace collage -} // namespace relay -} // namespace tvm - -#endif // TVM_RELAY_COLLAGE_GATHER_PARTITION_SPECS_H_ diff --git a/src/relay/collage/index_set.cc b/src/relay/collage/index_set.cc deleted file mode 100644 index 55bec80820a4..000000000000 --- a/src/relay/collage/index_set.cc +++ /dev/null @@ -1,231 +0,0 @@ -/* - * 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 src/relay/collage/index_set.cc - * \brief Efficient representation of a set of post-dfs indexes. - */ - -#include "./index_set.h" - -namespace tvm { -namespace relay { -namespace collage { - -// TODO(mbs): These should operate one-word-at-a-time - -IndexSet::IndexSet(size_t size, const std::vector& indexes) : bitvec_(size, false) { - for (size_t index : indexes) { - ICHECK_LT(index, bitvec_.size()); - ICHECK(!bitvec_[index]); - bitvec_[index] = true; - } -} - -IndexSet IndexSet::operator&(const IndexSet& that) const { - ICHECK_EQ(bitvec_.size(), that.bitvec_.size()); - std::vector result(bitvec_.size(), false); - for (size_t index = 0; index < bitvec_.size(); ++index) { - result[index] = bitvec_[index] && that.bitvec_[index]; - } - return IndexSet(result); -} - -IndexSet IndexSet::operator|(const IndexSet& that) const { - ICHECK_EQ(bitvec_.size(), that.bitvec_.size()); - std::vector result(bitvec_.size(), false); - for (size_t index = 0; index < bitvec_.size(); ++index) { - result[index] = bitvec_[index] || that.bitvec_[index]; - } - return IndexSet(result); -} - -IndexSet IndexSet::operator-(const IndexSet& that) const { - ICHECK_EQ(bitvec_.size(), that.bitvec_.size()); - std::vector result(bitvec_.size()); - for (size_t index = 0; index < bitvec_.size(); ++index) { - result[index] = bitvec_[index] && !that.bitvec_[index]; - } - return IndexSet(result); -} - -bool IndexSet::AreDisjoint(const IndexSet& that) const { - ICHECK_EQ(bitvec_.size(), that.bitvec_.size()); - for (size_t index = 0; index < bitvec_.size(); index++) { - if (bitvec_[index] && that.bitvec_[index]) { - return false; - } - } - return true; -} - -bool IndexSet::IsSubset(const IndexSet& that) const { - ICHECK_EQ(bitvec_.size(), that.bitvec_.size()); - for (size_t index = 0; index < bitvec_.size(); index++) { - if (bitvec_[index] && !that.bitvec_[index]) { - return false; - } - } - return true; -} - -bool IndexSet::Intersects(const IndexSet& that) const { - ICHECK_EQ(bitvec_.size(), that.bitvec_.size()); - for (size_t index = 0; index < bitvec_.size(); index++) { - if (bitvec_[index] && that.bitvec_[index]) { - return true; - } - } - return false; -} - -IndexSet IndexSet::Subst(size_t new_size, const IndexSubst& subst) const { - std::vector result(new_size, false); - for (PostDfsIndex index = 0; index < bitvec_.size(); ++index) { - if (!bitvec_[index]) { - continue; - } - auto itr = subst.find(index); - ICHECK(itr != subst.end()); - PostDfsIndex new_index = itr->second; - ICHECK(new_index < new_size); - ICHECK(!result[new_index]); - result[new_index] = true; - } - return IndexSet(result); -} - -size_t IndexSet::PopCount() const { - size_t n = 0; - for (size_t index = 0; index < bitvec_.size(); index++) { - if (bitvec_[index]) { - ++n; - } - } - return n; -} - -bool IndexSet::IsZero() const { - for (size_t index = 0; index < bitvec_.size(); index++) { - if (bitvec_[index]) { - return false; - } - } - return true; -} - -size_t IndexSet::FirstInsideIndex() const { - for (size_t index = 0; index < bitvec_.size(); index++) { - if (bitvec_[index]) { - return index; - } - } - return bitvec_.size(); -} - -size_t IndexSet::LastInsideIndex() const { - for (size_t i = bitvec_.size(); i > 0; i--) { - const size_t index = i - 1; - if (bitvec_[index]) { - return index; - } - } - return bitvec_.size(); -} - -size_t IndexSet::NextIndex(size_t index) const { - ICHECK_LT(index, bitvec_.size()); - for (index++; index < bitvec_.size(); index++) { - if (bitvec_[index]) { - return index; - } - } - return bitvec_.size(); -} - -size_t IndexSet::FirstOutsideIndex() const { - for (size_t index = 0; index < bitvec_.size(); index++) { - if (!bitvec_[index]) { - return index; - } - } - return bitvec_.size(); -} - -bool IndexSet::operator==(const IndexSet& that) const { - ICHECK_EQ(bitvec_.size(), that.bitvec_.size()); - return bitvec_ == that.bitvec_; -} - -bool IndexSet::operator!=(const IndexSet& that) const { - ICHECK_EQ(bitvec_.size(), that.bitvec_.size()); - return bitvec_ != that.bitvec_; -} - -bool IndexSet::operator<(const IndexSet& that) const { - ICHECK_EQ(bitvec_.size(), that.bitvec_.size()); - for (size_t index = 0; index < bitvec_.size(); index++) { - if (bitvec_[index] && !that.bitvec_[index]) { - return true; - } - if (!bitvec_[index] && that.bitvec_[index]) { - return false; - } - } - return false; -} - -size_t IndexSet::hash() const { - std::hash> h; - return h(bitvec_); -} - -std::string IndexSet::ToString() const { - std::ostringstream os; - os << "{"; - bool first = true; - for (size_t start = 0; start < bitvec_.size(); /*no-op*/) { - if (!bitvec_[start]) { - ++start; - continue; - } - size_t end; - for (end = start + 1; end < bitvec_.size() && bitvec_[end]; ++end) { - /*no-op*/ - } - if (first) { - first = false; - } else { - os << ","; - } - os << start; - if (end > start + 2) { - os << ".." << (end - 1); - start = end; - } else { - ++start; - } - } - os << "}"; - return os.str(); -} - -} // namespace collage -} // namespace relay -} // namespace tvm diff --git a/src/relay/collage/index_set.h b/src/relay/collage/index_set.h deleted file mode 100644 index f24b695cc76c..000000000000 --- a/src/relay/collage/index_set.h +++ /dev/null @@ -1,128 +0,0 @@ -/* - * 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 src/relay/collage/index_set.h - * \brief Efficient representation of a set of post-dfs indexes. - */ - -#ifndef TVM_RELAY_COLLAGE_INDEX_SET_H_ -#define TVM_RELAY_COLLAGE_INDEX_SET_H_ - -#include -#include -#include -#include - -#include "../ir/dataflow_matcher_impl.h" -#include "../ir/indexed_graph.h" - -namespace tvm { -namespace relay { -namespace collage { - -using IndexSubst = std::unordered_map; - -class IndexSet { - public: - IndexSet() = default; - explicit IndexSet(size_t size) : bitvec_(size, false) {} - IndexSet(size_t size, const std::vector& indexes); - - IndexSet operator&(const IndexSet& that) const; - IndexSet operator|(const IndexSet& that) const; - IndexSet operator-(const IndexSet& that) const; - bool AreDisjoint(const IndexSet& that) const; - bool IsSubset(const IndexSet& that) const; - bool Intersects(const IndexSet& that) const; - - bool operator[](size_t index) const { - ICHECK_LT(index, bitvec_.size()); - return bitvec_[index]; - } - - IndexSet& Add(size_t index) { - ICHECK_LT(index, bitvec_.size()); - bitvec_[index] = true; - return *this; - } - - IndexSet Subst(size_t new_size, const IndexSubst& subst) const; - - size_t end_index() const { return bitvec_.size(); } - size_t PopCount() const; - bool IsZero() const; - size_t FirstInsideIndex() const; - size_t LastInsideIndex() const; - size_t NextIndex(size_t index) const; - size_t FirstOutsideIndex() const; - bool operator==(const IndexSet& that) const; - bool operator!=(const IndexSet& that) const; - bool operator<(const IndexSet& that) const; - size_t hash() const; - std::string ToString() const; - - struct IndexSetIterator { - const IndexSet* set; - size_t i; - - size_t operator*() const { - ICHECK_LT(i, set->end_index()); - return i; - } - - const IndexSetIterator& operator++() { - ICHECK_LT(i, set->end_index()); - i = set->NextIndex(i); - return *this; - } - - bool operator==(const IndexSetIterator& that) const { - ICHECK(set == that.set); - return i == that.i; - } - - bool operator!=(const IndexSetIterator& that) const { - ICHECK(set == that.set); - return i != that.i; - } - }; - - IndexSetIterator begin() const { return IndexSetIterator{this, FirstInsideIndex()}; } - IndexSetIterator end() const { return IndexSetIterator{this, end_index()}; } - - private: - explicit IndexSet(std::vector bitvec) : bitvec_(std::move(bitvec)) {} - - std::vector bitvec_; -}; - -struct IndexSetEqual { - bool operator()(const IndexSet& left, const IndexSet& right) const { return left == right; } -}; - -struct IndexSetHash { - size_t operator()(const IndexSet& set) const { return set.hash(); } -}; - -} // namespace collage -} // namespace relay -} // namespace tvm - -#endif // TVM_RELAY_COLLAGE_INDEX_SET_H_ diff --git a/src/relay/collage/mock_cost_estimator.cc b/src/relay/collage/mock_cost_estimator.cc deleted file mode 100644 index 78fd24840517..000000000000 --- a/src/relay/collage/mock_cost_estimator.cc +++ /dev/null @@ -1,120 +0,0 @@ -/* - * 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 src/relay/collage/mock_cost_estimator.cc - * \brief A mock CostEstimator to support unit tests. - */ - -#include "./mock_cost_estimator.h" - -#include - -namespace tvm { -namespace relay { -namespace collage { - -TVM_REGISTER_OBJECT_TYPE(MockCostEstimatorNode); - -namespace { - -/*! - * \brief Visitor to accumulate the costs of all calls to operators in an expression. - */ -class MockEstimationVisitor : private ExprVisitor { - public: - MockEstimationVisitor(double op_cost, double fusion_benefit) - : op_cost_(op_cost), fusion_benefit_(fusion_benefit) {} - - double EstimateCost(const Expr& body) { - VisitExpr(body); - return cost_; - } - - private: - /*! \brief The assumed baseline cost of each operator call. */ - double op_cost_; - /*! - * \brief The factor by which each operator call cost is to be changed for every other - * operator call in the same group. - */ - double fusion_benefit_; - /*! \brief The number of operator calls seen so far. */ - size_t num_ops_ = 0; - /*! \brief Accumulate overall cost. */ - double cost_ = 0.0; - - void VisitExpr_(const CallNode* call_node) final { - if (call_node->op->IsInstance()) { - // Account for number of ops seens os far. - cost_ += op_cost_ * pow(fusion_benefit_, static_cast(num_ops_)); - num_ops_++; - } - ExprVisitor::VisitExpr_(call_node); - } - - void VisitExpr_(const FunctionNode* function_node) final { - // No "Compiler" functions can be inlined. - ICHECK(!function_node->GetAttr(attr::kCompiler).defined()) - << "All Compiler functions should have been outlined when preparing to estimate costs"; - ExprVisitor::VisitExpr_(function_node); - } -}; - -} // namespace - -Cost MockCostEstimatorNode::Estimate(const IRModule& mod, const Target& target) const { - // Limit the number of estimations. - ICHECK(max_estimates_->value == 0 || num_estimates_ < static_cast(max_estimates_->value)) - << "At most " << max_estimates_->value - << " non-trivial distinct candidates should have been generated."; - ++num_estimates_; - double op_cost = static_cast(target_costs_.at(target->kind->name)->value); - double cost = 0.0; - for (const auto& kv : mod->functions) { - if (const auto* function = kv.second.as()) { - if (kv.first->name_hint == "main") { - // Only tensor args are allowed to main. - for (const auto& param : function->params) { - ICHECK(param->type_annotation->IsInstance()) - << "Any tuple-of-tensor arguments should have been eta-exanded when preparing to " - "estimate costs"; - } - } - cost += MockEstimationVisitor(op_cost, /*fusion_benefit=*/0.9).EstimateCost(function->body); - } - } - return Cost::Value(cost); -} - -MockCostEstimator::MockCostEstimator(Map target_costs, Integer max_estimates) { - auto node = make_object(); - node->target_costs_ = std::move(target_costs); - node->max_estimates_ = std::move(max_estimates); - data_ = std::move(node); -} - -TVM_REGISTER_GLOBAL("relay.collage.MockCostEstimator") - .set_body_typed([](Map target_costs, Integer max_estimates) { - return MockCostEstimator(std::move(target_costs), std::move(max_estimates)); - }); - -} // namespace collage -} // namespace relay -} // namespace tvm diff --git a/src/relay/collage/mock_cost_estimator.h b/src/relay/collage/mock_cost_estimator.h deleted file mode 100644 index 3aa97923a201..000000000000 --- a/src/relay/collage/mock_cost_estimator.h +++ /dev/null @@ -1,94 +0,0 @@ -/* - * 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 src/relay/collage/mock_cost_estimator.cc - * \brief A mock CostEstimator to support unit tests. - */ - -#ifndef TVM_RELAY_COLLAGE_MOCK_COST_ESTIMATOR_H_ -#define TVM_RELAY_COLLAGE_MOCK_COST_ESTIMATOR_H_ - -#include - -#include "./cost.h" -#include "./cost_estimator.h" - -namespace tvm { -namespace relay { -namespace collage { - -// Clang (15.0.3, at least) validly complains about `@main`, but it invalidly -// complains even about `\c @main`. -#if __clang__ -#pragma clang diagnostic push -#pragma clang diagnostic ignored "-Wdocumentation-unknown-command" -#endif - -/*! - * \brief A mock cost estimator which can determine the cost of a candidate based on both - * the candidate's target and the number of operator calls inside it. - * - * The help unit tests the estimator also ICHECK fails if: - * - the module has inlined "Compiler" functions - * - @main has non-tensor arguments (eg a tuple) - * - more than the given number of candidate modules are measured - * - * To support unit testing only. - */ -class MockCostEstimatorNode : public CostEstimatorNode { - public: - Cost Estimate(const IRModule& mod, const Target& target) const override; - - static constexpr const char* _type_key = "relay.collage.MockCostEstimator"; - TVM_DECLARE_FINAL_OBJECT_INFO(MockCostEstimatorNode, CostEstimatorNode); - - protected: - /*! - * \brief Map from target kind name to assumed baseline cost (in integer seconds) for all - * operator calls. - */ - Map target_costs_; - - /*! - * \brief If non-zero, the maximum number of distinct modules which may be estimated. - */ - Integer max_estimates_; - - /*! \brief Number of calls to Estimate. */ - mutable size_t num_estimates_ = 0; - - friend class MockCostEstimator; -}; -#if __clang__ -#pragma clang diagnostic pop -#endif - -class MockCostEstimator : public CostEstimator { - public: - explicit MockCostEstimator(Map target_costs, Integer max_estimates = 0); - - TVM_DEFINE_OBJECT_REF_METHODS(MockCostEstimator, CostEstimator, MockCostEstimatorNode); -}; - -} // namespace collage -} // namespace relay -} // namespace tvm - -#endif // TVM_RELAY_COLLAGE_MOCK_COST_ESTIMATOR_H_ diff --git a/src/relay/collage/name_supply.cc b/src/relay/collage/name_supply.cc deleted file mode 100644 index 4b7d497b0d57..000000000000 --- a/src/relay/collage/name_supply.cc +++ /dev/null @@ -1,90 +0,0 @@ -/* - * 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 src/relay/collage/name_supply.cc - * \brief A source of fresh variable names. - */ - -#include "./name_supply.h" - -#include -#include - -namespace tvm { -namespace relay { -namespace collage { - -namespace { -void AppendCSafe(bool* first, std::ostringstream& os, const std::string& str) { - for (size_t i = 0; i < str.size(); ++i) { - const char c = str[i]; - if (i == 0 && first && (!std::isalpha(c) && c != '_')) { - os << "_"; - } - if (c == '_' || std::isalnum(c)) { - os << c; - } else { - os << "_"; - } - *first = false; - } -} -} // namespace - -NameSupply NameSupply::MakeSubNameSupply() { - NameSupply result(prefix_); - for (const auto& kv : next_free_index_) { - result.next_free_index_.emplace(kv.first, kv.second); - } - return result; -} - -std::string NameSupply::Fresh(const std::initializer_list& hints) { - std::ostringstream os; - bool first = true; - bool need_sep = false; - if (!prefix_.empty()) { - AppendCSafe(&first, os, prefix_); - need_sep = true; - } - for (const auto& hint : hints) { - if (hint.empty()) { - continue; - } - if (need_sep) { - os << "_"; - } - AppendCSafe(&first, os, hint); - need_sep = true; - } - std::string name = os.str(); - auto itr = next_free_index_.find(name); - if (itr == next_free_index_.end()) { - next_free_index_.emplace(name, 1); - } else { - os << "_" << itr->second++; - name = os.str(); - } - return name; -} - -} // namespace collage -} // namespace relay -} // namespace tvm diff --git a/src/relay/collage/name_supply.h b/src/relay/collage/name_supply.h deleted file mode 100644 index d37023ab6f81..000000000000 --- a/src/relay/collage/name_supply.h +++ /dev/null @@ -1,58 +0,0 @@ -/* - * 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 src/relay/collage/name_supply.h - * \brief A source of fresh variable names. - */ - -#ifndef TVM_RELAY_COLLAGE_NAME_SUPPLY_H_ -#define TVM_RELAY_COLLAGE_NAME_SUPPLY_H_ - -#include -#include -#include - -namespace tvm { -namespace relay { -namespace collage { - -/*! \brief A supply of fresh names. */ -class NameSupply { - public: - explicit NameSupply(std::string prefix) : prefix_(std::move(prefix)) {} - - NameSupply MakeSubNameSupply(); - - void Reserve(const std::string& existing) { next_free_index_.emplace(existing, 1); } - - std::string Fresh(const std::initializer_list& hints); - - private: - /*! \brief Prefix for all names. May be empty. */ - std::string prefix_; - /*! \brief Next unused index for variables with given basename. */ - std::unordered_map next_free_index_; -}; - -} // namespace collage -} // namespace relay -} // namespace tvm - -#endif // TVM_RELAY_COLLAGE_NAME_SUPPLY_H_ diff --git a/src/relay/collage/partition_rule.cc b/src/relay/collage/partition_rule.cc deleted file mode 100644 index 1d8c5e9723ee..000000000000 --- a/src/relay/collage/partition_rule.cc +++ /dev/null @@ -1,426 +0,0 @@ -/* - * 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 src/relay/collage/partition_rule.cc - * \brief Compositional partitioning rules. - */ - -#include "./partition_rule.h" - -#include - -#include "./partition_rule.h" -#include "./partition_spec.h" -#include "./utils.h" - -namespace tvm { -namespace relay { -namespace collage { - -TVM_REGISTER_NODE_TYPE(PartitionRuleNode); - -void PartitionRuleNode::VisitAttrs(AttrVisitor* v) { - // TODO(mbs) -} - -std::vector PartitionRuleNode::AllCandidates( - const DataflowGraph& dataflow_graph, const PartitionSpec& spec) const { - ICHECK(false) << "PartitionRuleNode::AllCandidates should be overridden in sub-class"; - return {}; -} - -std::string PartitionRuleNode::ToString() const { return ToDoc().str(); } - -Doc PartitionRuleNode::ToDoc() const { - Doc doc; - doc << GetTypeKey() << "(" << Doc::NewLine(2); - std::vector body_items; - AppendBodyItems(&body_items); - doc << Doc::Indent(2, Doc::Concat(body_items, Doc::NewLine())) << Doc::NewLine(); - doc << ")"; - return doc; -} - -void PartitionRuleNode::AppendBodyItems(std::vector* body_items) const { - body_items->emplace_back(); - body_items->back() << "rule_name=" << Doc::StrLiteral(rule_name_); -} - -PartitionRule::PartitionRule(String rule_name) { - auto node = runtime::make_object(); - node->rule_name_ = std::move(rule_name); - data_ = std::move(node); -} - -bool DefaultPatternPredicate(const Expr& matched_sub_expr) { return true; } - -TVM_REGISTER_NODE_TYPE(DFPatternPartitionRuleNode); - -void DFPatternPartitionRuleNode::VisitAttrs(AttrVisitor* v) { - // TODO(mbs) -} - -std::vector DFPatternPartitionRuleNode::AllCandidates( - const DataflowGraph& dataflow_graph, const PartitionSpec& spec) const { - VLOG(1) << "running DFPatternPartitionRule(" << rule_name_ << ")"; - std::vector result; - DFPatternMatcher matcher(&dataflow_graph.indexed_graph()); - for (PostDfsIndex index = 0; index < dataflow_graph.size(); ++index) { - Expr sub_expr = dataflow_graph.index_to_node(index)->ref(); - if (!matcher.Match(pattern_, sub_expr)) { - continue; - } - if (!predicate_(sub_expr)) { - VLOG(1) << "DFPatternPartitionRule(" << rule_name_ << ") has failing predicate"; - continue; - } - IndexSet inside = MatcherToIndexSet(matcher); - auto [kind, label] = SubGraphKindAndLabel(dataflow_graph, inside); - SubGraph sub_graph(dataflow_graph, std::move(inside), kind, std::move(label)); - String rule_name = rule_name_.empty() ? sub_graph->label_ : rule_name_; - CandidatePartition candidate(std::move(rule_name), std::move(sub_graph), spec); - VLOG(2) << "DFPatternPartitionRule(" << rule_name_ << ") yields " << candidate->ToString(); - result.emplace_back(std::move(candidate)); - } - VLOG(1) << "DFPatternPartitionRule(" << rule_name_ << ") produced " << result.size() - << " candidates"; - return result; -} - -void DFPatternPartitionRuleNode::AppendBodyItems(std::vector* body_items) const { - PartitionRuleNode::AppendBodyItems(body_items); - body_items->emplace_back(); - body_items->back() << "pattern=" << PrettyPrint(pattern_); -} - -DFPatternPartitionRule::DFPatternPartitionRule(String rule_name, DFPattern pattern, - TPatternPredicate predicate) { - auto node = runtime::make_object(); - node->rule_name_ = std::move(rule_name); - node->pattern_ = std::move(pattern); - node->predicate_ = std::move(predicate); - data_ = std::move(node); -} - -TVM_REGISTER_NODE_TYPE(CompositePartitionRuleNode); - -void CompositePartitionRuleNode::VisitAttrs(AttrVisitor* v) { - // TODO(mbs) -} - -std::vector CompositePartitionRuleNode::AllCandidates( - const DataflowGraph& dataflow_graph, const PartitionSpec& spec) const { - std::vector candidates = sub_rule_->AllCandidates(dataflow_graph, spec); - VLOG(1) << "running CompositePartitionRule(" << rule_name_ << ") over " << candidates.size() - << " sub-candidates"; - std::vector result; - FunctionAttrsMap attrs; - attrs.Set(attr::kComposite, rule_name_); - for (auto& candidate : candidates) { - String rule_name = NestLabels(rule_name_, candidate->rule_name_); - SubGraph sub_graph = candidate->sub_graph_.WithAttrs(dataflow_graph, attrs); - CandidatePartition new_candidate = WithSubGraph( - WithRuleName(std::move(candidate), std::move(rule_name)), std::move(sub_graph)); - VLOG(2) << "CompositePartitionRule(" << rule_name_ << ") yields " << new_candidate->ToString(); - result.emplace_back(std::move(new_candidate)); - } - VLOG(1) << "CompositePartitionRule(" << rule_name_ << ") produced " << result.size() - << " candidates"; - return result; -} - -void CompositePartitionRuleNode::AppendBodyItems(std::vector* body_items) const { - PartitionRuleNode::AppendBodyItems(body_items); - body_items->emplace_back(); - body_items->back() << "sub_rule=" << sub_rule_->ToDoc(); -} - -CompositePartitionRule::CompositePartitionRule(String rule_name, PartitionRule sub_rule) { - auto node = runtime::make_object(); - node->rule_name_ = std::move(rule_name); - node->sub_rule_ = std::move(sub_rule); - data_ = std::move(node); -} - -TVM_REGISTER_NODE_TYPE(PrimitivePartitionRuleNode); - -void PrimitivePartitionRuleNode::VisitAttrs(AttrVisitor* v) { - // TODO(mbs) -} - -std::vector PrimitivePartitionRuleNode::AllCandidates( - const DataflowGraph& dataflow_graph, const PartitionSpec& spec) const { - std::vector candidates = sub_rule_->AllCandidates(dataflow_graph, spec); - VLOG(1) << "running PrimitivePartitionRule(" << rule_name_ << ") over " << candidates.size() - << " sub-candidates"; - std::vector result; - FunctionAttrsMap attrs; - attrs.Set(attr::kPrimitive, Integer(1)); - if (spec->target_.IsExternalCodegen()) { - // The spec name will be the target kind name which is 1:1 with the "Compiler" attribute name. - attrs.Set(attr::kCompiler, spec->spec_name_); - } - for (auto& candidate : candidates) { - String rule_name = NestLabels(rule_name_, candidate->rule_name_); - SubGraph sub_graph = candidate->sub_graph_.WithAttrs(dataflow_graph, attrs); - CandidatePartition new_candidate = WithSubGraph( - WithRuleName(std::move(candidate), std::move(rule_name)), std::move(sub_graph)); - VLOG(2) << "PrimitivePartitionRule(" << rule_name_ << ") yields " << new_candidate->ToString(); - result.emplace_back(std::move(new_candidate)); - } - VLOG(1) << "PrimitivePartitionRule(" << rule_name_ << ") produced " << result.size() - << " candidates"; - return result; -} - -void PrimitivePartitionRuleNode::AppendBodyItems(std::vector* body_items) const { - PartitionRuleNode::AppendBodyItems(body_items); - body_items->emplace_back(); - body_items->back() << "sub_rule=" << sub_rule_->ToDoc(); -} - -PrimitivePartitionRule::PrimitivePartitionRule(String rule_name, PartitionRule sub_rule) { - auto node = runtime::make_object(); - node->rule_name_ = std::move(rule_name); - node->sub_rule_ = std::move(sub_rule); - data_ = std::move(node); -} - -TVM_REGISTER_NODE_TYPE(UnionPartitionRuleNode); - -void UnionPartitionRuleNode::VisitAttrs(AttrVisitor* v) { - // TODO(mbs) -} - -std::vector UnionPartitionRuleNode::AllCandidates( - const DataflowGraph& dataflow_graph, const PartitionSpec& spec) const { - std::vector result; - for (const auto& sub_rule : sub_rules_) { - std::vector candidates = sub_rule->AllCandidates(dataflow_graph, spec); - for (auto& candidate : candidates) { - String rule_name = NestLabels(rule_name_, candidate->rule_name_); - CandidatePartition new_candidate = WithRuleName(std::move(candidate), std::move(rule_name)); - VLOG(2) << "UnionPartitionRule(" << rule_name_ << ") yields " << new_candidate->ToString(); - result.emplace_back(std::move(new_candidate)); - } - } - VLOG(1) << "UnionPartitionRule(" << rule_name_ << ") produced " << result.size() << " candidates"; - return result; -} - -void UnionPartitionRuleNode::AppendBodyItems(std::vector* body_items) const { - PartitionRuleNode::AppendBodyItems(body_items); - for (const auto& sub_rule : sub_rules_) { - body_items->emplace_back(); - body_items->back() << "sub_rule=" << sub_rule->ToDoc(); - } -} - -UnionPartitionRule::UnionPartitionRule(String rule_name, Array sub_rules) { - auto node = runtime::make_object(); - node->rule_name_ = std::move(rule_name); - node->sub_rules_ = std::move(sub_rules); - data_ = std::move(node); -} - -TVM_REGISTER_NODE_TYPE(OpCallByKindPartitionRuleNode); - -void OpCallByKindPartitionRuleNode::VisitAttrs(AttrVisitor* v) { - // TODO(mbs) -} - -std::vector OpCallByKindPartitionRuleNode::AllCandidates( - const DataflowGraph& dataflow_graph, const PartitionSpec& spec) const { - VLOG(1) << "running OpCallByKindPartitionRule(" << rule_name_ << ")"; - std::vector result; - for (PostDfsIndex index = 0; index < dataflow_graph.size(); ++index) { - auto node = dataflow_graph.index_to_node(index); - Expr sub_expr = node->ref(); - if (sub_expr->IsInstance()) { - auto [kind, label] = SubExprKindAndLabel(sub_expr); - if (kind <= kOutEWiseFusable) { - IndexSet inside(dataflow_graph.size(), {index}); - SubGraph sub_graph(dataflow_graph, std::move(inside), kind, std::move(label)); - String rule_name = NestLabels(rule_name_, sub_graph->label_); - CandidatePartition candidate(std::move(rule_name), std::move(sub_graph), spec); - VLOG(2) << "OpCallByKindPartitionRule(" << rule_name_ << ") yields " - << candidate->ToString(); - result.emplace_back(std::move(candidate)); - } - } - } - VLOG(1) << "OpCallByKindPartitionRule(" << rule_name_ << ") produced " << result.size() - << " candidates"; - return result; -} - -void OpCallByKindPartitionRuleNode::AppendBodyItems(std::vector* body_items) const { - PartitionRuleNode::AppendBodyItems(body_items); -} - -OpCallByKindPartitionRule::OpCallByKindPartitionRule(String rule_name) { - auto node = runtime::make_object(); - node->rule_name_ = std::move(rule_name); - data_ = std::move(node); -} - -TVM_REGISTER_NODE_TYPE(CombinePartitionRuleNode); - -void CombinePartitionRuleNode::VisitAttrs(AttrVisitor* v) { - // TODO(mbs) -} - -std::vector CombinePartitionRuleNode::AllCandidates( - const DataflowGraph& dataflow_graph, const PartitionSpec& spec) const { - // We'll accumulate all the candidates here, starting with those from the sub-rule. - // Once a candidate is added to this vector it is immutable. - std::vector candidates = sub_rule_->AllCandidates(dataflow_graph, spec); - VLOG(1) << "running CombinePartitionRule(" << rule_name_ << ") over " << candidates.size() - << " sub-candidates"; - CandidateSet result_set(std::move(candidates)); - - size_t num_rounds = 0; - AppendAllResultsContext ctxt(&dataflow_graph, max_depth_, &result_set); - while (result_set.PrepareForNextRound()) { - VLOG_CONTEXT << "round " << ++num_rounds; - VLOG(1) << "checking " << result_set.size() << " candidates (" << result_set.first_new_index() - << " existing)"; - for (const auto& combiner_rule : combiner_rules_) { - combiner_rule->AppendAllResults(&ctxt); - } - } - - std::vector result; - for (auto& candidate : result_set.MovedCurrentCandidates()) { - String rule_name = NestLabels(rule_name_, candidate->rule_name_); - CandidatePartition new_candidate = WithRuleName(std::move(candidate), std::move(rule_name)); - VLOG(2) << "CombinePartitionRule(" << rule_name_ << ") yields " << new_candidate->ToString(); - result.emplace_back(std::move(new_candidate)); - } - VLOG(1) << "CombinePartitionRule(" << rule_name_ << ") produced " << result.size() - << " candidates"; - return result; -} - -void CombinePartitionRuleNode::AppendBodyItems(std::vector* body_items) const { - PartitionRuleNode::AppendBodyItems(body_items); - body_items->emplace_back(); - body_items->back() << "sub_rule=" << sub_rule_->ToDoc(); - for (const auto& combiner_rule : combiner_rules_) { - body_items->emplace_back(); - body_items->back() << "combiner_rule=" << combiner_rule->ToString(); - } - body_items->emplace_back(); - body_items->back() << "max_depth=" << max_depth_; -} - -CombinePartitionRule::CombinePartitionRule(String rule_name, PartitionRule sub_rule, - Array combiner_rules, size_t max_depth_) { - auto node = runtime::make_object(); - node->rule_name_ = std::move(rule_name); - node->sub_rule_ = std::move(sub_rule); - node->combiner_rules_ = std::move(combiner_rules); - node->max_depth_ = max_depth_; - data_ = std::move(node); -} - -TVM_REGISTER_NODE_TYPE(OnlyValidPartitionRuleNode); - -void OnlyValidPartitionRuleNode::VisitAttrs(AttrVisitor* v) { - // TODO(mbs) -} - -std::vector OnlyValidPartitionRuleNode::AllCandidates( - const DataflowGraph& dataflow_graph, const PartitionSpec& spec) const { - std::vector candidates = sub_rule_->AllCandidates(dataflow_graph, spec); - VLOG(1) << "running OnlyValidPartitionRule(" << rule_name_ << ") over " << candidates.size() - << " sub-candidates"; - std::vector result; - for (auto& candidate : candidates) { - if (!candidate->sub_graph_->IsValid(dataflow_graph, config_)) { - VLOG(2) << "Ignoring invalid candidate " << candidate->ToString(); - continue; - } - String rule_name = NestLabels(rule_name_, candidate->rule_name_); - CandidatePartition new_candidate = WithRuleName(std::move(candidate), std::move(rule_name)); - VLOG(2) << "OnlyValidPartitionRule(" << rule_name_ << ") yields " << new_candidate->ToString(); - result.emplace_back(std::move(new_candidate)); - } - VLOG(1) << "OnlyValidPartitionRule(" << rule_name_ << ") produced " << result.size() - << " candidates"; - return result; -} - -void OnlyValidPartitionRuleNode::AppendBodyItems(std::vector* body_items) const { - PartitionRuleNode::AppendBodyItems(body_items); - body_items->emplace_back(); - body_items->back() << "sub_rule=" << sub_rule_->ToDoc(); - body_items->emplace_back(); - body_items->back() << "config=" << config_.ToString(); -} - -OnlyValidPartitionRule::OnlyValidPartitionRule(String rule_name, PartitionRule sub_rule, - const SubGraphConfig& config) { - auto node = runtime::make_object(); - node->rule_name_ = std::move(rule_name); - node->sub_rule_ = std::move(sub_rule); - node->config_ = config; - data_ = std::move(node); -} - -TVM_REGISTER_NODE_TYPE(HostPartitionRuleNode); - -void HostPartitionRuleNode::VisitAttrs(AttrVisitor* v) { - // TODO(mbs) -} - -std::vector HostPartitionRuleNode::AllCandidates( - const DataflowGraph& dataflow_graph, const PartitionSpec& spec) const { - VLOG(1) << "running HostPartitionRule(" << rule_name_ << ")"; - std::vector result; - for (PostDfsIndex index = 0; index < dataflow_graph.size(); ++index) { - if (MustBeLowered(dataflow_graph.index_to_node(index)->ref())) { - continue; - } - IndexSet inside(dataflow_graph.size(), {index}); - auto [kind, label] = SubGraphKindAndLabel(dataflow_graph, inside); - SubGraph sub_graph(dataflow_graph, std::move(inside), kind, label); - String rule_name = NestLabels(rule_name_, sub_graph->label_); - // We'll a zero cost for the candidate since we'll never want to actually estimate the cost - // of this 'partition'. - CandidatePartition candidate(std::move(rule_name), std::move(sub_graph), spec, Cost::Zero()); - VLOG(2) << "HostPartitionRule(" << rule_name_ << ") yields " << candidate->ToString(); - result.push_back(candidate); - } - VLOG(1) << "HostPartitionRule(" << rule_name_ << ") produced " << result.size() << " candidates"; - return result; -} - -void HostPartitionRuleNode::AppendBodyItems(std::vector* body_items) const {} - -HostPartitionRule::HostPartitionRule(String rule_name) { - auto node = runtime::make_object(); - node->rule_name_ = std::move(rule_name); - data_ = std::move(node); -} - -} // namespace collage -} // namespace relay -} // namespace tvm diff --git a/src/relay/collage/partition_rule.h b/src/relay/collage/partition_rule.h deleted file mode 100644 index c9b7e93d7138..000000000000 --- a/src/relay/collage/partition_rule.h +++ /dev/null @@ -1,487 +0,0 @@ -/* - * 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 src/relay/collage/partition_rule.h - * \brief Compositional partitioning rules. - */ - -#ifndef TVM_RELAY_COLLAGE_PARTITION_RULE_H_ -#define TVM_RELAY_COLLAGE_PARTITION_RULE_H_ - -#include -#include - -#include -#include - -#include "../printer/doc.h" -#include "./candidate_partition.h" -#include "./combiner_rule.h" -#include "./sub_graph.h" - -namespace tvm { -namespace relay { -namespace collage { - -/*! - * \brief Type of function to check if a matched sub-expression should be accepted by a rule. This - * can be used to, eg, reject operators of unsupported shape or dtype, or otherwise implement rules - * which are difficult to express in the dataflow pattern language directly. - */ -using TPatternPredicate = TypedPackedFunc; - -/*! - * \brief The default pattern predicate. Always returns true. - */ -bool DefaultPatternPredicate(const Expr& matched_sub_expr); - -/*! - * \brief Base class of all partition rules. - * - * A \p PartitionRule describes how to find a set of \p CandidatePartitions for a \p DataflowGraph. - * The candidates are allowed to overlap, and ultimately it is the job of the Collage searcher to - * find a selection of candidates which covers the whole Relay expression without overlap. Partition - * rules are paired with their \p Target and other 'top level' configuration in a \p PartitionSpec. - * - * We provide a set of 'base' partition rules which produce candidates from the dataflow graph - * directly. We also provide a set of 'combinator' partition rules which can produce new candidates - * from the results of an arbitrary sub-rule or sub-rules. By mixing these base and combinator - * rules we can express a wide variety of partition strategies and encoding conventions. - * - * There may be many thousands of candidates in flight during the Collage search. We take care to - * defer constructing or rewriting Relay expressions until absolutely necessary. We only pay for - * extracting a function to represent a candidate when we need to measure it's cost. And we only - * pay for rewriting the overall Relay expression to commit to a partitioning when the Collage - * search has completed. - * - * The base rules implemented so far: - * - \p DFPatternPartitionRule: Given a \p DFPattern and expression predicate, produces a candidate - * for every sub-graph matched by the pattern and predicate. Unlike the \p PatternRewriter, - * candidates are free to overlap. Used to bring BYOC patterns into the Collage framework. - * - \p OpCallByKindPartitionRule: Uses the "TOpPattern" attribute provided for every Relay - * operator to produce a candidate for every call to a 'fusable Relay operator'. Used to - * look ahead to how TVM will fuse sub-graphs. - * - * The combinator rules implemented so far: - * - \p CompositePartitionRule: Indicates all candidates matched by the sub-rule should be wrapped - * by a "Composite" function. The "Composite" name is taken from the rule name. Used to indicate - * Relay operators (or groups of Relay operators) should be mapped to target-specific operators, - * both for BYOC and TVM external library integrations. - * - \p PrimitivePartitionRule: Indicates all candidates matched by the sub-rule should be wrapped - * by a "Primitive" function, possibly with an additional "Compiler" attribute. Used to - * delineate a partition (or kernel). - * - \p UnionPartitionRule: Simply unions all the candidates from all sub-rules together. Used to - * combine individual \p DFPatternPartitionRules. - * - \p CombinePartitionRule: Given a sub-rule and a list of 'combiner' rules, finds - * all possible ways of combining the sub-rule's candidates to yield even larger candidates. - * Note that the sub-rule's candidates may also be directly included in the results. The - * 'combiner' rules allow combining by \p OpPatternKinds, combining the arguments to tuples - * which themselves are arguments to Relay operator calls, and so on. This rule is intended to - * mimic the existing TVM \p FuseOps pass, though: - * i) all candidates are found rather than just the largest, ii) the starting set of candidates - * can be provided by any other rule, and iii) we rely on \p SubGraph validity checking to weed - * out infeasible candidates. - * - \p OnlyValidPartitionRule: Given a \p SubGraphConfig, ignores candidates with 'invalid' - * sub-graphs. Used to limit the maximum candidate depth, the number of independent outputs, - * and whether intermediate 'taps' are allowed. - * - \p HostPartitionRule: Produces candidates for all Relay expressions which could be - * 'left behind' for execution by the host (eg on the VM). This rule lets us simplify the - * overall Collage search algorithm. - * - * (Though not yet implemented, we'd like to allow a combinator rule which will union candidate - * based on their 'anchor' operators. This can be used to implement 'vertical' and 'horizontal' - * partition on more primitive candidates. Note that the \p SubGraph machinery supports - * multiple-input and -output sub-graphs and their validation, so horizontal partition is easy - * implement.) - * - * Here are some typical ways to combine \p PartitionRules for different partition/fusion - * strategies: - * - * - Classic pattern-based BYOC with \p MergeComposite/AnnotateTarget/PartitionGraph passes: - * \code - * PrimitivePartitionRule - * OnlyValidPartitionRule - * CombinePartitionRule (with join-anything combiner rule) - * UnionPartitionRule - * CompositePartitionRule(label1) - * DFPatternPartitionRule(pattern1) - * : - * CompositePartitionRule(labeln) - * DFPatternPartitionRule(patternn) - * \endcode - * - * - "Consider this library implementation for these sub-expressions", using \p DFPatterns to - * pick out which Relay operators are supported: - * \code - * OnlyValidPartitionRule - * CombinePartitionRule (with default TVM combiner rules) - * UnionPartitionRule - * OpCallByKindPartitionRule - * CompositePartitionRule(lable1) - * DFPatternPartitionRule(pattern1) - * : - * CompositePartitionRule(lablen) - * DFPatternPartitionRule(patternn) - * \endcode - * - * - Classic TVM \p FuseOps - * \code - * PrimitivePartitionRule - * OnlyValidPartitionRule - * CombinePartitionRule (with default TVM combiner rules) - * OpCallByKindPartitionRule - * \endcode - * - * - "Just fuse what I tell you to fuse", using \p DFPatterns to directly select candidates: - * \code - * PrimitivePartitionRule - * OnlyValidPartitionRule - * UnionPartitionRule - * DFPatternPartitionRule(pattern1) - * : - * DFPatternPartitionRule(patternn) - * \endcode - */ -class PartitionRuleNode : public Object { - public: - /*! - * \brief A unique (over all rules for the same target) name for the rule. Rule names are - * combined and captured with \p PartitionCandidate rule names for debuggability and - * explainability. Some rules will copy the rule name into function attributes. - * - */ - String rule_name_; - - void VisitAttrs(AttrVisitor* v); - - /*! - * \brief Returns all the possible candidate partitions according to this rule for the overall - * expression corresponding to \p dataflow_graph. The candidates will generally have unknown - * target and cost: the target will be filled in by the \p PartitionSpec, while the cost will - * be filled in lazily. - */ - virtual std::vector AllCandidates(const DataflowGraph& dataflow_graph, - const PartitionSpec& spec) const; - - std::string ToString() const; - Doc ToDoc() const; - - protected: - virtual void AppendBodyItems(std::vector* body_items) const; - - public: - static constexpr const char* _type_key = "relay.collage.PartitionRule"; - static constexpr const uint32_t _type_child_slots = 10; - TVM_DECLARE_BASE_OBJECT_INFO(PartitionRuleNode, Object); -}; - -class PartitionRule : public ObjectRef { - public: - explicit PartitionRule(String rule_name); - - TVM_DEFINE_OBJECT_REF_METHODS(PartitionRule, ObjectRef, PartitionRuleNode); -}; - -/*! - * \brief Partition rule which fires on all sub-expressions matching a dataflow-pattern and pattern - * predicate. It is valid for matching candidates to overlap. - */ -class DFPatternPartitionRuleNode : public PartitionRuleNode { - public: - /*! - * \brief Relay pattern. - */ - DFPattern pattern_; - - /*! - * \brief Predicate on matched sub-expression to decide if partition rule should fire. - */ - TPatternPredicate predicate_; - - void VisitAttrs(AttrVisitor* v); - - std::vector AllCandidates(const DataflowGraph& dataflow_graph, - const PartitionSpec& spec) const override; - - void AppendBodyItems(std::vector* body_items) const override; - - static constexpr const char* _type_key = "relay.collage.DFPatternPartitionRule"; - TVM_DECLARE_FINAL_OBJECT_INFO(DFPatternPartitionRuleNode, PartitionRuleNode); -}; - -class DFPatternPartitionRule : public PartitionRule { - public: - DFPatternPartitionRule(String rule_name, DFPattern pattern, - TPatternPredicate predicate = DefaultPatternPredicate); - - TVM_DEFINE_OBJECT_REF_METHODS(DFPatternPartitionRule, PartitionRule, DFPatternPartitionRuleNode); -}; - -/*! - * \brief Partition rule which wraps candidates within a function with the "Composite" attribute - * bound to the given rule name. - * - * This is the standard way by which operators or operator groups are tagged as being supported - * by a particular externally provided function. It is up to the BYOC lowering function to - * recognize the "Composite" name and emit the appropriate code or call. - */ -class CompositePartitionRuleNode : public PartitionRuleNode { - public: - /*! \brief The sub-partition rule. */ - PartitionRule sub_rule_; - - void VisitAttrs(AttrVisitor* v); - - std::vector AllCandidates(const DataflowGraph& dataflow_graph, - const PartitionSpec& spec) const override; - - void AppendBodyItems(std::vector* body_items) const override; - - static constexpr const char* _type_key = "relay.collage.CompositePartitionRule"; - TVM_DECLARE_FINAL_OBJECT_INFO(CompositePartitionRuleNode, PartitionRuleNode); -}; - -class CompositePartitionRule : public PartitionRule { - public: - CompositePartitionRule(String rule_name, PartitionRule sub_rule); - - TVM_DEFINE_OBJECT_REF_METHODS(CompositePartitionRule, PartitionRule, CompositePartitionRuleNode); -}; - -/*! - * \brief Partition rule which wraps candidates within a function with the "Primitive" attribute - * bound to 1. If the partition spec target(s) have the "compiler" attribute then that name is - * also added to the function as a "Compiler" attribute. - * - * This is the standard way by which sub-graphs are marked as being in a 'partition' who's - * compilation will be managed by an external BYOC toolchain. It can also be used to mark - * sub-graphs for lowering to a single kernel by the built-in TVM lowering machinery. - */ -class PrimitivePartitionRuleNode : public PartitionRuleNode { - public: - /*! \brief The sub-partition rule. */ - PartitionRule sub_rule_; - - void VisitAttrs(AttrVisitor* v); - - std::vector AllCandidates(const DataflowGraph& dataflow_graph, - const PartitionSpec& spec) const override; - - void AppendBodyItems(std::vector* body_items) const override; - - static constexpr const char* _type_key = "relay.collage.PrimitivePartitionRule"; - TVM_DECLARE_FINAL_OBJECT_INFO(PrimitivePartitionRuleNode, PartitionRuleNode); -}; - -class PrimitivePartitionRule : public PartitionRule { - public: - PrimitivePartitionRule(String rule_name, PartitionRule sub_rule); - - TVM_DEFINE_OBJECT_REF_METHODS(PrimitivePartitionRule, PartitionRule, PrimitivePartitionRuleNode); -}; - -/*! - * \brief Partition rule which simply unions all matches from all sub-partition rules. - * - * This can be used to combine the results of a set of, eg, DFPatternPartitionRules. - */ -class UnionPartitionRuleNode : public PartitionRuleNode { - public: - Array sub_rules_; - - void VisitAttrs(AttrVisitor* v); - - std::vector AllCandidates(const DataflowGraph& dataflow_graph, - const PartitionSpec& spec) const override; - - void AppendBodyItems(std::vector* body_items) const override; - - static constexpr const char* _type_key = "relay.collage.UnionPartitionRule"; - TVM_DECLARE_FINAL_OBJECT_INFO(UnionPartitionRuleNode, PartitionRuleNode); -}; - -class UnionPartitionRule : public PartitionRule { - public: - UnionPartitionRule(String rule_name, Array sub_rules); - - TVM_DEFINE_OBJECT_REF_METHODS(UnionPartitionRule, PartitionRule, UnionPartitionRuleNode) -}; - -/* - *! \brief Partition rule which places calls to Relay operators with a "TOpPattern" attribute of - * \p kOutEWiseFusable or less in their own singleton sub-graph. No other Relay sub-expressions - * (such as tuples or tuple projection) are selected, and it is up to outer partition rules to - * account for them. - */ -class OpCallByKindPartitionRuleNode : public PartitionRuleNode { - public: - void VisitAttrs(AttrVisitor* v); - - std::vector AllCandidates(const DataflowGraph& dataflow_graph, - const PartitionSpec& spec) const override; - - void AppendBodyItems(std::vector* body_items) const override; - - static constexpr const char* _type_key = "relay.collage.OpCallByKindPartitionRule"; - TVM_DECLARE_FINAL_OBJECT_INFO(OpCallByKindPartitionRuleNode, PartitionRuleNode); -}; - -class OpCallByKindPartitionRule : public PartitionRule { - public: - explicit OpCallByKindPartitionRule(String rule_name); - - TVM_DEFINE_OBJECT_REF_METHODS(OpCallByKindPartitionRule, PartitionRule, - OpCallByKindPartitionRuleNode); -}; - -/*! - * \brief Partition rule which combines sub-graphs to exploit optimizations commonly available in - * backends (including the TVM lowering backend). Those optimization rules are in turn described by - * one or more primitive \p CombinerRules. - * - * For TVM these primitive combiner rules are guided by the \p OpPatternKind associated with every - * sub-graph. That in turn is the maximum of the kind of each expression node in the sub-graph, - * using the rules: - * - Constants are \p kElemwise. - * - A call to a Relay operator has the kind of its callee. - * - Tuple construction and projection are injective provided all tuple fields are of tensor type. - * - All other sub-expressions are opaque. - * - * The available \p OpPatternKinds (and our abbreviations for them) are: - * - E: kElemWise, eg nn.relu - * - B: kBroadcast, eg add - * - I: kInjective, eg concatenate - * - R: kCommReduce, eg sum - * - A: kOutEWiseFusable, eg nn.conv2d (often called 'anchor nodes', hence the A abbreviation) - * - O: kOpaque, everything else - * (The kTuple kind is not used by this machinery.) - * - * Kinds are ordered as above from least- to most-constraining w.r.t. possible partition - * opportunities. When we write a kind abbreviation below we intend it to mean that kind *or less*. - * And when write 'kl -> kr' we mean it to match a sub-expression of kind kr or less who's - * dataflow inputs are all of kind kl or less. - * - * We can then mimic the classic \p FuseOps TVM Pass with the following more primitive combiner - * rules: - * - Sub-groups cannot have taps. In the classic \p FuseOps pass taps are avoided by construction - * by always considering all node->dominator paths. Here we naively allow taps on all candidates, - * but reject them using SubGraph::IsValid with a SubGraphConfig with allow_taps = false. - * - Combine A -> B - * - Combine B -> R - * - Combine I -> I - * - Combine I -> tuple -> I. That is, if an I sub-graph has a tuple as input, and at least one - * tuple field can be provided by an I sub-graph exit, then both the tuple and all such fields - * may be joined. - gt* - * Note that \p FuseOps only considers the largest possible sub-graphs. However this partition rule - * considers all possibilities so as to 'make room' for other targets supplying other - * overlapping candidates. - * - * See combiner_rule.h for the more primitive combiner rules which implement the above. - */ -class CombinePartitionRuleNode : public PartitionRuleNode { - public: - /*! \brief The sub-rule supplying the initial set of candidates. */ - PartitionRule sub_rule_; - /*! \brief The more primitive rules to use to combine the candidates found by the above rule. */ - Array combiner_rules_; - /*! \brief Maximum max_depth for candidates. */ - size_t max_depth_; - - void VisitAttrs(AttrVisitor* v); - - std::vector AllCandidates(const DataflowGraph& dataflow_graph, - const PartitionSpec& spec) const override; - - void AppendBodyItems(std::vector* body_items) const override; - - public: - static constexpr const char* _type_key = "relay.collage.CombinePartitionRule"; - TVM_DECLARE_FINAL_OBJECT_INFO(CombinePartitionRuleNode, PartitionRuleNode); -}; - -class CombinePartitionRule : public PartitionRule { - public: - CombinePartitionRule(String rule_name, PartitionRule sub_rule, Array combiner_rules, - size_t max_depth_); - - TVM_DEFINE_OBJECT_REF_METHODS(CombinePartitionRule, PartitionRule, CombinePartitionRuleNode); -}; - -/*! - * \brief Partition rules which keeps only candidates from the sub-rule whose sub-groups are valid - * w.r.t. the given \p SubGraphConfig. - */ -class OnlyValidPartitionRuleNode : public PartitionRuleNode { - public: - PartitionRule sub_rule_; - SubGraphConfig config_; - - void VisitAttrs(AttrVisitor* v); - - std::vector AllCandidates(const DataflowGraph& dataflow_graph, - const PartitionSpec& spec) const override; - - void AppendBodyItems(std::vector* body_items) const override; - - public: - static constexpr const char* _type_key = "relay.collage.OnlyValidPartitionRule"; - TVM_DECLARE_FINAL_OBJECT_INFO(OnlyValidPartitionRuleNode, PartitionRuleNode); -}; - -class OnlyValidPartitionRule : public PartitionRule { - public: - OnlyValidPartitionRule(String rule_name, PartitionRule sub_rule, const SubGraphConfig& config); - - TVM_DEFINE_OBJECT_REF_METHODS(OnlyValidPartitionRule, PartitionRule, OnlyValidPartitionRuleNode); -}; - -/*! - * \brief Partition rule which selects nodes which can be 'left behind' to be executed by the host - * (eg on the VM). This includes most of the 'interstitial' Relay constructs, such a let bindings, - * operators on references, calls to non-operator functions, and so on. It can also include the - * construction of and projection from tuples which may not be supported within a partition. - */ -class HostPartitionRuleNode : public PartitionRuleNode { - public: - void VisitAttrs(AttrVisitor* v); - - std::vector AllCandidates(const DataflowGraph& dataflow_graph, - const PartitionSpec& spec) const override; - - void AppendBodyItems(std::vector* body_items) const override; - - public: - static constexpr const char* _type_key = "relay.collage.HostPartitionRule"; - TVM_DECLARE_FINAL_OBJECT_INFO(HostPartitionRuleNode, PartitionRuleNode); -}; - -class HostPartitionRule : public PartitionRule { - public: - explicit HostPartitionRule(String rule_name); - - TVM_DEFINE_OBJECT_REF_METHODS(HostPartitionRule, PartitionRule, HostPartitionRuleNode); -}; - -} // namespace collage -} // namespace relay -} // namespace tvm - -#endif // TVM_RELAY_COLLAGE_PARTITION_RULE_H_ diff --git a/src/relay/collage/partition_spec.cc b/src/relay/collage/partition_spec.cc deleted file mode 100644 index b2095d0a594e..000000000000 --- a/src/relay/collage/partition_spec.cc +++ /dev/null @@ -1,87 +0,0 @@ -/* - * 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 src/relay/collage/partition_spec.cc - * \brief Combine a \p PartitionRule with a \p Target. - */ - -#include "./partition_spec.h" - -#include "./utils.h" - -namespace tvm { -namespace relay { -namespace collage { - -String DefaultValidateSubGraphFunc(const Function& function) { return String(); } - -TVM_REGISTER_NODE_TYPE(PartitionSpecNode); - -void PartitionSpecNode::VisitAttrs(AttrVisitor* v) { - // TODO(mbs) -} - -std::vector PartitionSpecNode::AllCandidates( - const DataflowGraph& dataflow_graph) const { - std::vector result; - // Make sure the target is in scope for inspection by any predicates in - // DFPatternPartitionRuleNode rules. - With target_scope(target_); - // Gather all the candidates. - std::vector candidates = - rule_->AllCandidates(dataflow_graph, GetRef(this)); - // Update the rules names. - for (const auto& candidate : candidates) { - ICHECK_EQ(candidate->spec_, GetRef(this)); - String rule_name = NestLabels(spec_name_, candidate->rule_name_); - CandidatePartition new_candidate = WithRuleName(candidate, std::move(rule_name)); - result.emplace_back(std::move(new_candidate)); - } - return result; -} - -std::string PartitionSpecNode::ToString() const { - Doc doc; - doc << "PartitionSpec(" << Doc::NewLine(2); - std::vector body_items; - body_items.emplace_back(); - body_items.back() << "spec_name=" << Doc::StrLiteral(spec_name_); - body_items.emplace_back(); - body_items.back() << "target=" << target_->ToDebugString(); - body_items.emplace_back(); - body_items.back() << "rule=" << rule_->ToDoc(); - doc << Doc::Indent(2, Doc::Concat(body_items, Doc::NewLine())) << Doc::NewLine(); - doc << ")"; - return doc.str(); -} - -PartitionSpec::PartitionSpec(String spec_name, Target target, PartitionRule rule, - TValidateSubGraphFunc validate_sub_graph_func) { - auto node = runtime::make_object(); - node->spec_name_ = std::move(spec_name); - node->target_ = std::move(target); - node->rule_ = std::move(rule); - node->validate_sub_graph_func_ = std::move(validate_sub_graph_func); - data_ = std::move(node); -} - -} // namespace collage -} // namespace relay -} // namespace tvm diff --git a/src/relay/collage/partition_spec.h b/src/relay/collage/partition_spec.h deleted file mode 100644 index e8ce64c68468..000000000000 --- a/src/relay/collage/partition_spec.h +++ /dev/null @@ -1,120 +0,0 @@ -/* - * 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 src/relay/collage/partition_spec.h - * \brief Combine a \p PartitionRule with a \p Target. - */ - -#ifndef TVM_RELAY_COLLAGE_PARTITION_SPEC_H_ -#define TVM_RELAY_COLLAGE_PARTITION_SPEC_H_ - -#include -#include -#include - -#include -#include - -#include "./partition_rule.h" -#include "./sub_graph.h" - -namespace tvm { -namespace relay { -namespace collage { - -/*! - * \brief Type of functions for checking the validity of partitions before they proceed to lowering - * and codegen. The argument is the function extracted from the overall expression to represent - * the partition. The result is a non-empty error message string if the candidate should be - * rejected. - */ -using TValidateSubGraphFunc = TypedPackedFunc; - -/*! - * \brief The default validation function. Always returns the empty string, ie no error. - */ -String DefaultValidateSubGraphFunc(const Function& function); - -/*! - * \brief Pairs a \p PartitionRule with one or more \p Targets it can be used for. - */ -class PartitionSpecNode : public Object { - public: - /*! - * \brief Specification name to distinguish this spec from all others. Typically the BYOC - * 'compiler' name, "tvm", or "host". - */ - String spec_name_; - - /*! - * \brief The target all candidate partitions should be compiled for. - * - * It's tempting to support multiple targets here since. Eg the partitioning rules for - * TVM are the same irrespective of whether the target is "cuda" or "llvm", so it would make - * sense to build the candidate partitions first without committing to any target, then 'stamp' - * them for each target as the final step. - * - * However, we want to make sure any predicate in \p DFPatternPartitionRuleNode instances - * can have access to the current target instance. Eg the predicate may need to consult - * build-time configuration to decide what operators, shapes etc are actually supported. - * That implies the specific target is known when the candidate partitions are being constructed. - * - * So for now we'll just force each spec to have exactly one target. - */ - Target target_; - - /*! - * \brief The partition rule to use to gather candidates. - */ - PartitionRule rule_; - - /*! - * \brief The validation function to apply to each candidate's the extracted function before - * proceeding to lowering/codegen. - */ - TValidateSubGraphFunc validate_sub_graph_func_ = DefaultValidateSubGraphFunc; - - void VisitAttrs(AttrVisitor* v); - - /*! - * \brief Returns all the candidate partitions found by this specification. The candidates - * will be for a specific target, but will not yet have an extracted function or cost. - */ - std::vector AllCandidates(const DataflowGraph& dataflow_graph) const; - - std::string ToString() const; - - static constexpr const char* _type_key = "relay.collage.PartitionSpec"; - TVM_DECLARE_FINAL_OBJECT_INFO(PartitionSpecNode, Object); -}; - -class PartitionSpec : public ObjectRef { - public: - PartitionSpec(String spec_name, Target target, PartitionRule rule, - TValidateSubGraphFunc validate_sub_graph_func = DefaultValidateSubGraphFunc); - - TVM_DEFINE_OBJECT_REF_METHODS(PartitionSpec, ObjectRef, PartitionSpecNode); -}; - -} // namespace collage -} // namespace relay -} // namespace tvm - -#endif // TVM_RELAY_COLLAGE_PARTITION_SPEC_H_ diff --git a/src/relay/collage/priority_queue.h b/src/relay/collage/priority_queue.h deleted file mode 100644 index 1d30fe5d96af..000000000000 --- a/src/relay/collage/priority_queue.h +++ /dev/null @@ -1,72 +0,0 @@ -/* - * 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 src/relay/collage/priority_queue.h - * \brief An updatable priority queue. - */ - -#ifndef TVM_RELAY_COLLAGE_PRIORITY_QUEUE_H_ -#define TVM_RELAY_COLLAGE_PRIORITY_QUEUE_H_ - -#include - -namespace tvm { -namespace relay { -namespace collage { - -/*! \brief Priority queue of search states, ordered by increasing cost. */ -template -class PriorityQueue { - public: - PriorityQueue() = default; - - /*! \brief Pushes \p item onto the queue. */ - void Push(T* item) { set_.emplace(item); } - - /*! \brief Pops the item with the least cost off the queue. */ - T* Pop() { - ICHECK(!set_.empty()); - T* item = *set_.begin(); - set_.erase(set_.begin()); - return item; - } - - /*! \brief Updates the queue to account for \p item's best cost being lowered. */ - void Update(T* item) { - auto itr = std::find_if(set_.begin(), set_.end(), - [item](const T* that) { return EqTPtr()(that, item); }); - ICHECK(itr != set_.end()); - set_.erase(itr); - set_.emplace(item); - } - - bool empty() const { return set_.empty(); } - size_t size() const { return set_.size(); } - - private: - // TODO(mbs): Actually use a pri-queue datastructure! - std::set set_; -}; - -} // namespace collage -} // namespace relay -} // namespace tvm - -#endif // TVM_RELAY_COLLAGE_PRIORITY_QUEUE_H_ diff --git a/src/relay/collage/prune_candidates.cc b/src/relay/collage/prune_candidates.cc deleted file mode 100644 index 91baa6bb4dfe..000000000000 --- a/src/relay/collage/prune_candidates.cc +++ /dev/null @@ -1,218 +0,0 @@ -/* - * 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 src/relay/collage/prune_candidates.cc - * \brief Try to remove candidates which will never contribute to an optimal partitioning. - */ - -#include "./prune_candidates.h" - -#include "./dataflow_graph.h" -#include "./gather_partition_specs.h" - -namespace tvm { -namespace relay { -namespace collage { - -namespace { - -/*! - * \brief Returns a map from post-dfs dataflow node indices to the indices within \p candidates for - * those candidates which intersect that dataflow node. - * - * NOTE: The index set in the vector results is over candidate indices not post-dfs indices! - */ -std::vector MakeInsideMap(const DataflowGraph& dataflow_graph, - const std::vector& candidates) { - std::vector result(dataflow_graph.size(), IndexSet(candidates.size())); - for (size_t i = 0; i < candidates.size(); ++i) { - CandidatePartition candidate = candidates[i]; - for (PostDfsIndex index : candidate->sub_graph_->inside_) { - result[index].Add(i); - } - } - return result; -} - -/*! - * \brief Returns the maximal candidates within \p candidates. A candidate is maximal if it is not - * contained by any super-candidate for the same target. - */ -std::vector MaximalCandidates( - const DataflowGraph& dataflow_graph, const std::vector& candidates) { - std::vector inside_map = MakeInsideMap(dataflow_graph, candidates); - std::vector result; - for (size_t i = 0; i < candidates.size(); ++i) { - CandidatePartition maximal_candidate = candidates[i]; - bool has_super_candidate = false; - IndexSet explored_candidates(candidates.size()); // over candidates! - for (PostDfsIndex index : maximal_candidate->sub_graph_->inside_) { - for (size_t j : inside_map[index]) { - if (i == j) { - // Ignore self. - continue; - } - if (explored_candidates[j]) { - // Already checked. - continue; - } - explored_candidates.Add(j); - CandidatePartition super_candidate = candidates[j]; - if (maximal_candidate->spec_ == super_candidate->spec_ && - maximal_candidate->sub_graph_->inside_.IsSubset(super_candidate->sub_graph_->inside_)) { - has_super_candidate = true; - break; - } - } - if (has_super_candidate) { - break; - } - } - if (!has_super_candidate) { - VLOG(2) << "Found maximal candidate " << maximal_candidate->ToString(); - result.emplace_back(maximal_candidate); - } - } - VLOG(1) << "Have " << result.size() << " maximal candidates"; - return result; -} - -/*! - * \brief Returns all the candidates in \p candidates which intersect without being equal. - */ -std::vector IntersectingCandidates( - const DataflowGraph& dataflow_graph, const std::vector& candidates) { - std::vector inside_map = MakeInsideMap(dataflow_graph, candidates); - IndexSet intersecting(candidates.size()); // over candidates! - for (size_t i = 0; i < candidates.size(); ++i) { - CandidatePartition intersecting_candidate = candidates[i]; - IndexSet explored_candidates(candidates.size()); // over candidates! - for (PostDfsIndex index : intersecting_candidate->sub_graph_->inside_) { - for (size_t j : inside_map[index]) { - if (j < i) { - // Intersection is commutative. - continue; - } - if (i == j) { - // Ignore self. - continue; - } - if (explored_candidates[j]) { - // Already checked. - continue; - } - explored_candidates.Add(j); - CandidatePartition other_candidate = candidates[j]; - if (intersecting_candidate->sub_graph_->inside_ == other_candidate->sub_graph_->inside_) { - // Have same inside set. - continue; - } - VLOG(2) << "Candidate " << intersecting_candidate->ToString() << " intersects with " - << other_candidate->ToString(); - intersecting.Add(i); - intersecting.Add(j); - } - } - } - std::vector result; - for (size_t i : intersecting) { - CandidatePartition candidate = candidates[i]; - VLOG(2) << "Found intersecting candidate " << candidate->ToString(); - result.emplace_back(candidate); - } - VLOG(1) << "Have " << result.size() << " intersecting candidates"; - return result; -} - -/*! - * \brief Returns the set operation left - right. - */ -std::vector SetDifference(const std::vector& left, - const std::vector& right) { - std::unordered_set - right_set(right.begin(), right.end()); - std::vector result; - for (const auto& candidate : left) { - if (right_set.count(candidate) == 0) { - result.emplace_back(candidate); - } - } - return result; -} - -/*! - * \brief Adds everything in right to left. Returns the number of elements added. - */ -size_t SetUnionInPlace( - std::unordered_set* left, - const std::vector& right) { - size_t init_size = left->size(); - for (const auto& candidate : right) { - left->emplace(candidate); - } - return left->size() - init_size; -} - -} // namespace - -std::vector PruneCandidates( - const DataflowGraph& dataflow_graph, - const std::vector& initial_candidates) { - VLOG_CONTEXT << "prune"; - // Start with all candidates available. - std::vector candidates = initial_candidates; - std::unordered_set pruned; - size_t initial_num_candidates = candidates.size(); - size_t num_rounds = 0; - while (true) { - VLOG_CONTEXT << "round " << ++num_rounds; - VLOG(1) << "checking " << candidates.size() << " candidates"; - // Add all the maximal candidates to the pruned set. - std::vector maximal_candidates = - MaximalCandidates(dataflow_graph, candidates); - size_t num_new_pruned = SetUnionInPlace(&pruned, maximal_candidates); - VLOG(1) << "Added " << num_new_pruned << " new pruned candidates"; - if (num_new_pruned == 0) { - // We've reached a fixed point. - break; - } - // If two pruned candidates intersect without being equal then we may miss valid - // paths during search. So remove those intersecting candidates from the available candidates - // and try again so as to find smaller candidates to 'bridge the gaps'. - std::vector pruned_vec(pruned.begin(), pruned.end()); - std::vector intersecting_candidates = - IntersectingCandidates(dataflow_graph, pruned_vec); - // We need more maximal candidates to fill in the gaps between the current pruned candidates. - // Force that by removing the intersecting candidates from the set of available candidates - // and going around again. - candidates = SetDifference(candidates, intersecting_candidates); - } - - std::vector result(pruned.begin(), pruned.end()); - // Re-establish a canonical order of candidates. - std::sort(result.begin(), result.end()); - VLOG(1) << "Pruned " << initial_num_candidates - result.size() << " candidates (ie from " - << initial_num_candidates << " to " << result.size() << ")"; - return result; -} - -} // namespace collage -} // namespace relay -} // namespace tvm diff --git a/src/relay/collage/prune_candidates.h b/src/relay/collage/prune_candidates.h deleted file mode 100644 index 294acbb1fefe..000000000000 --- a/src/relay/collage/prune_candidates.h +++ /dev/null @@ -1,72 +0,0 @@ -/* - * 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 src/relay/collage/prune_candidates.h - * \brief Try to remove candidates which will never contribute to an optimal partitioning. - */ - -#ifndef TVM_RELAY_COLLAGE_PRUNE_CANDIDATES_H_ -#define TVM_RELAY_COLLAGE_PRUNE_CANDIDATES_H_ - -#include - -#include "./candidate_partition.h" -#include "./dataflow_graph.h" - -namespace tvm { -namespace relay { -namespace collage { - -/*! - * \brief Returns \p initial_candidates with all unnecessary candidates pruned. - * - * We prune according to the following two heuristics: - * 1. Given partitions (A, target) and (B, target) then - * cost(A union B, target) < cost(A, target) + cost(B, target). - * That is, there's no use estimating the cost of small partitions when a larger partition - * containing them is also available. More precisely, call a partition 'maximal' if it is - * not contained by any other partition for the same target. Then we want to prefer maximal - * candidates when searching. - * 2. Given maximal partitions (A union B, target) and (A union B, target') where - * target != target', then min(cost(A union B, target), cost(A union B, target')) < - * min(cost(A, target) + cost(B, target'), cost(A, target') + cost(B, target)). - * That is, there's no use estimating cross-combinations of partitions which are not maximal. - * - * However, we can't prune a non-maximal candidate if it will make some other maximal candidate - * unreachable during the Collage search. We achieve this by iterating until fixed point: - * - Find maximal candidates of current set of candidates. - * - Add those maximal candidates to the output 'pruned' set. - * - If any two candidates in the 'pruned' set intersect without being equal, remove those from - * the current set of candidates and go around again. That will force more candidates to - * be considered 'maximal'. - * That over-approximates the true necessary candidates but is at least simple. - * - * CAUTION: This is pretty experimental. The above heuristics won't always be safe, and I don't - * have a proof the pruned candidate set won't lead to 'No candidate was found covering - * sub-expression...' errors in Partitioner::Partition(). - */ -std::vector PruneCandidates( - const DataflowGraph& dataflow_graph, const std::vector& initial_candidates); - -} // namespace collage -} // namespace relay -} // namespace tvm - -#endif // TVM_RELAY_COLLAGE_PRUNE_CANDIDATES_H_ diff --git a/src/relay/collage/sub_graph.cc b/src/relay/collage/sub_graph.cc deleted file mode 100644 index a6559ff5fdb5..000000000000 --- a/src/relay/collage/sub_graph.cc +++ /dev/null @@ -1,1030 +0,0 @@ -/* - * 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 src/relay/collage/sub_graph.cc - * \brief Represents a sub-graph of an overall Relay expression. - */ - -#include "./sub_graph.h" - -#include - -#include "../../support/scalars.h" -#include "../transforms/pass_utils.h" -#include "./utils.h" - -namespace tvm { -namespace relay { -namespace collage { - -namespace { - -class Extractor; - -/*! - * \brief Helper class for rewriting expressions to replace a sub-graph according to the - * given extractor. - */ -class Rewriter : public ExprMutator { - public: - explicit Rewriter(const Extractor* extractor) : extractor_(extractor) {} - - Expr VisitExpr(const Expr& expr) final; - - private: - /*! \brief Already prepared extractor which will guide the rewrite. */ - const Extractor* extractor_; -}; - -/*! \brief Helper class for extracting matched sub-graphs from the overall expression. */ -class Extractor : public ExprMutator { - public: - Extractor(const DataflowGraph* dataflow_graph, const SubGraphNode* sub_graph, - FunctionAttrsMap opt_attrs) - : dataflow_graph_(dataflow_graph), sub_graph_(sub_graph), opt_attrs_(std::move(opt_attrs)) { - ICHECK_EQ(dataflow_graph_->size(), sub_graph_->overall_size()); - } - - const DataflowGraph& dataflow_graph() const { return *dataflow_graph_; } - - /*! - * \brief Collect the parameters and output expressions for the function representing - * the sub-graph. - */ - void Extract() { - ICHECK(!sub_graph_->IsEmpty()); - VLOG(2) << "Extracting " << sub_graph_->ToString(); - const bool for_function = opt_attrs_.defined(); - - // In reverse dataflow order... - for (PostDfsIndex i = dataflow_graph_->size(); i > 0; --i) { - PostDfsIndex index = i - 1; - if (!sub_graph_->inside_[index]) { - // Node is outside sub-graph. - continue; - } - VLOG(2) << "index " << index; - auto node = dataflow_graph_->index_to_node(index); - if (sub_graph_->exit_[node->index_] || node->is_external_ || memo_.count(node->ref()) == 0) { - // This sub-expression is: - // - inside the sub-graph and needed outside the sub-graph. So it must contribute to an - // output (even if we've already visited it while constructing an output from a - // downstream sub-expression). - // - not yet visited, in which case it must still be considered an 'output' so it will - // be evaluated for any possible side effects. - Expr output = VisitExpr(GetRef(node->node_ref_)); - VLOG(2) << "index " << index << " added as output:\n" - << PrettyPrint(output) << "\nat " << outputs_.size(); - expr_to_output_index_.emplace(node->node_ref_, outputs_.size()); - outputs_.emplace_back(std::move(output)); - output_types_.emplace_back(node->node_ref_->checked_type()); - } - } - ICHECK(!outputs_.empty()); - - // Reverse the outputs so as to preserve the original evaluation order. - std::reverse(outputs_.begin(), outputs_.end()); - std::reverse(output_types_.begin(), output_types_.end()); - for (auto& kv : expr_to_output_index_) { - kv.second = static_cast(outputs_.size()) - 1 - kv.second; - } - - // Build a 'body' expression to represent the extracted sub-graph. If we have multiple - // outputs we'll place them in a tuple. - Type body_type; - Expr body; - if (outputs_.size() > 1) { - body_type = TupleType(output_types_); - body = Tuple(outputs_); - body->checked_type_ = body_type; - } else { - body_type = output_types_.front(); - body = outputs_.front(); - } - - // Re-express all the nested sub-graphs in terms of the body. - DataflowGraph body_dataflow_graph(body); - std::vector nested_sub_graphs; - IndexSubst subst = MakeIndexSubst(body_dataflow_graph); - for (const auto& nested_sub_graph : sub_graph_->nested_sub_graphs_) { - nested_sub_graphs.emplace_back(nested_sub_graph.Subst(body_dataflow_graph, subst)); - } - - // Sweep backwards through the body, rewriting to account for each nested sub-graph. - body = NestedSubGraph::ParallelRewrite(body_dataflow_graph, body, std::move(nested_sub_graphs)); - - if (for_function) { - // Rewrite so all input nodes are now conveyed via call arguments to a new function. - Array arg_types; - arg_types.reserve(params_.size()); - for (const auto& param : params_) { - arg_types.push_back(param->checked_type()); - } - extracted_ = Function(std::move(params_), std::move(body), body_type, - /*ty_params=*/{}, DictAttrs(opt_attrs_)); - extracted_->checked_type_ = - FuncType(std::move(arg_types), body_type, /*type_params=*/{}, /*type_constraints=*/{}); - body = Call(extracted_, std::move(args_)); - body->checked_type_ = body_type; - } else { - // Don't do anything with the inputs. - extracted_ = body; - } - - // Setup the output substitution. - for (const auto& kv : expr_to_output_index_) { - Expr expr; - if (outputs_.size() == 1) { - expr = body; - } else if (for_function) { - expr = TupleGetItem(body, kv.second); - expr->checked_type_ = output_types_[kv.second]; - } else { - const auto* tuple_node = body.as(); - ICHECK(tuple_node); - expr = tuple_node->fields[kv.second]; - } - VLOG(2) << "output " << dataflow_graph_->item_to_node(kv.first)->index_ << " is at index " - << kv.second << " (of " << outputs_.size() << " outputs)"; - output_substitution_.emplace(kv.first, std::move(expr)); - } - } - - ////// Following members are valid only after Extract() has returned. - - /*! - * \brief Returns the expression representing the extracted sub-graph. If opt_attrs_ is - * defined then will be a function. - */ - Expr extracted() const { return extracted_; } - - /*! - * \brief Returns the substitution to apply to all expression nodes in the overall expression - * so as to replace references to outputs of the sub-graph with their rewritten form. - */ - const std::unordered_map& output_substitution() const { - return output_substitution_; - } - - private: - /*! - * \brief Returns a map from original index to new index for each node inside the sub-graph. Only - * valid after \p Extract has made its backwards dataflow sweep. - */ - IndexSubst MakeIndexSubst(const DataflowGraph& new_dataflow_graph) const { - VLOG(2) << "building extractor substitution"; - IndexSubst subst; - for (PostDfsIndex index : sub_graph_->inside_) { - auto orig_node = dataflow_graph_->index_to_node(index); - ICHECK_EQ(orig_node->index_, index); - auto itr = memo_.find(orig_node->ref()); - ICHECK(itr != memo_.end()); - auto new_node = new_dataflow_graph.item_to_node(itr->second); - VLOG(2) << orig_node->index_ << " |-> " << new_node->index_; - subst.emplace(orig_node->index_, new_node->index_); - } - return subst; - } - - /*! \brief Returns true if \p expr is inside the sub-graph. */ - bool inside(const Expr& expr) { - return sub_graph_->inside_[dataflow_graph_->item_to_node(expr)->index_]; - } - - /*! - * \brief Returns the variable uniquely representing \p expr, which should be - * an input node (ie outside the sub-graph but feeding into a node inside the sub-graph). - * - * It is valid for: - * - An expression outside the sub-graph to be used multiple times inside the sub-graph. - * - An expression outside the sub-graph to be used both inside and outside the sub-graph. - */ - Var VarFor(const Expr& expr) { - ICHECK(!inside(expr)); - ICHECK(opt_attrs_.defined()); - auto itr = expr_to_param_.find(expr.get()); - if (itr != expr_to_param_.end()) { - return itr->second; - } - auto fresh_var = Var("FunctionVar_" + std::to_string(params_.size()), expr->checked_type()); - fresh_var->checked_type_ = expr->checked_type(); - params_.push_back(fresh_var); - args_.push_back(expr); - expr_to_param_.emplace(expr.get(), fresh_var); - return fresh_var; - } - - /*! - * \brief If \p expr is inside the sub-graph then return it's rewritten form. - * If \p expr is outside the sub-graph then it must correspond to an input node. - * - If opt_attrs_ is defined return the variable to represent it. - * - Otherwise just return the expression directly. - * - * Should be called only on inputs to nodes which are inside the sub-graph. - */ - Expr VisitExpr(const Expr& expr) final { - if (inside(expr)) { - return ExprMutator::VisitExpr(expr); - } else if (CanInline(expr)) { - // Implicitly include inlinable input sub-expressions. - return expr; - } else if (opt_attrs_.defined()) { - // Map to a function parameter. - return VarFor(expr); - } else { - // Stop rewriting. - return expr; - } - } - - Expr VisitExpr_(const FunctionNode* function_node) override { - if (function_node->HasNonzeroAttr(attr::kPrimitive)) { - return GetRef(function_node); - } - return ExprMutator::VisitExpr_(function_node); - } - - //// Context fields, passed in constructor. - - /*! \brief The dataflow graph corresponding to the overall expression. */ - const DataflowGraph* dataflow_graph_; - /*! \brief The sub-graph of the above we are extracting. */ - const SubGraphNode* sub_graph_; - /*! \brief Optional attributes if the sub-graph should be extracted as a function. */ - FunctionAttrsMap opt_attrs_; - - //// Result fields, available after Extract() called. - - /*! - * \brief The extracted expression. If opt_attrs_ is defined this will be a function. - */ - Expr extracted_; - /*! - * \brief Map from output nodes to corresponding expressions. If the sub-graph has more than - * one exit node then each entry will be a tuple projection. - */ - std::unordered_map output_substitution_; - - //// Accumulator fields, built as we visit expressions. - - /*! \brief (If opt_attrs_ is defined) Parameters representing input expression nodes. */ - Array params_; - /*! - * \brief (If opt_attrs_ is defined) The input expression nodes for each of the above params_. - */ - Array args_; - /*! - * \brief (If opt_attrs_ is defined) Map from existing input expression nodes to the parameters - * in params_ which now representing them. - */ - std::unordered_map expr_to_param_; - /*! - * \brief Accumulated new expressions which represent the exit nodes of the rewritten sub-graph. - * It is possible to have multiple outputs. It is possible one output also contributes to other - * outputs (ie the output is a 'tap'). - */ - std::vector outputs_; - /*! \brief (If opt_attrs_ is defined) Types of original expressions corresponding to outputs_. */ - std::vector output_types_; - /*! - * \brief Map from existing exit expression nodes to the index in outputs_ which should - * represent them in the rewritten overall expression. - */ - std::unordered_map expr_to_output_index_; -}; - -Expr Rewriter::VisitExpr(const Expr& expr) { - auto itr = extractor_->output_substitution().find(expr.get()); - if (itr == extractor_->output_substitution().end()) { - return ExprMutator::VisitExpr(expr); - } else { - return itr->second; - } -} - -} // namespace - -std::pair SubExprKindAndLabel(const Expr& sub_expr) { - class Visitor : public ExprFunctor(const Expr&)> { - private: - std::pair VisitExpr_(const CallNode* call_node) final { - if (auto optional = call_node->op.as()) { - auto op = optional.value(); - static auto fpattern = Op::GetAttrMap("TOpPattern"); - if (fpattern.count(op) == 0) { - VLOG(1) << "no TOpPattern known for " << op->name << ", considering opaque"; - return {kOpaque, op->name}; - } else if (IsDynamic(call_node->checked_type()) && IsDataDependent(call_node)) { - VLOG(1) << "call has dynamic shape which is data-dependent, considering opaque"; - return {kOpaque, op->name}; - } else { - OpPatternKind kind = static_cast(fpattern[op]); - VLOG(2) << "TOpPattern for " << op->name << " is " << KindToString(kind); - return {kind, op->name}; - } - } else if (const auto* function_node = call_node->op.as()) { - Optional opt_i = - function_node->GetAttr("TOpPattern", Optional()); - if (opt_i.defined()) { - OpPatternKind kind = static_cast(opt_i.value()->value); - VLOG(1) << "TOpPattern for function is " << KindToString(kind); - return {kind, "call_prim"}; - } else { - VLOG(1) << "calling function without TOpPattern, considering opaque"; - return {kOpaque, "call_fun"}; - } - } else { - VLOG(1) << "unsupported call, considering opaque"; - return {kOpaque, "call_any"}; - } - } - - std::pair VisitExpr_(const ConstantNode* constant_node) final { - VLOG(2) << "TOpPattern for constant is " << KindToString(kElemWise); - if (support::IsSimpleScalar(constant_node)) { - return {kElemWise, "scalar"}; - } else { - return {kElemWise, "const"}; - } - } - - std::pair VisitExpr_(const TupleNode* tuple_node) final { - const auto* tuple_type_node = tuple_node->checked_type().as(); - ICHECK(tuple_type_node != nullptr); - if (std::all_of(tuple_type_node->fields.begin(), tuple_type_node->fields.end(), - [](const Type& type) { return type.as() != nullptr; })) { - VLOG(2) << "TOpPattern for tuple is " << KindToString(kInjective); - return {kInjective, "tuple"}; - } else { - VLOG(1) << "tuple contains non-tensors, considering opaque"; - return {kOpaque, "tuple"}; - } - } - - std::pair VisitExpr_( - const TupleGetItemNode* tuple_get_item_node) final { - const auto* tuple_type_node = tuple_get_item_node->tuple->checked_type().as(); - ICHECK(tuple_type_node != nullptr); - if (std::all_of(tuple_type_node->fields.begin(), tuple_type_node->fields.end(), - [](const Type& type) { return type.as() != nullptr; })) { - VLOG(2) << "TOpPattern for tuple projection is " << KindToString(kInjective); - return {kInjective, "proj"}; - } else { - VLOG(1) << "tuple being projected contains non-tensors, considering opaque"; - return {kOpaque, "proj"}; - } - } - - // TODO(mbs): We implement the following mostly so we have a lightweight way of describing - // the current sub-expression. If partitioning is ever extended beyond the usual call/tuple/proj - // sub-language we should revise the returned operator kinds to match. - - std::pair VisitExpr_(const VarNode* var_node) final { - return {kOpaque, "%" + var_node->name_hint()}; - } - std::pair VisitExpr_(const GlobalVarNode* global_var_node) final { - return {kOpaque, "@" + global_var_node->name_hint}; - } - std::pair VisitExpr_(const OpNode* op_node) final { - return {kOpaque, "`" + op_node->name}; - } - std::pair VisitExpr_(const FunctionNode* function_node) final { - return {kOpaque, "fn"}; - } - std::pair VisitExpr_(const LetNode* let_node) final { - return {kOpaque, "let"}; - } - std::pair VisitExpr_(const IfNode* if_node) final { - return {kOpaque, "if"}; - } - std::pair VisitExpr_(const RefCreateNode* ref_create_node) final { - return {kOpaque, "ref"}; - } - std::pair VisitExpr_(const RefReadNode* op) final { - return {kOpaque, "ref_read"}; - } - std::pair VisitExpr_(const RefWriteNode* op) final { - return {kOpaque, "ref_write"}; - } - std::pair VisitExpr_(const ConstructorNode* op) final { - return {kOpaque, "`" + op->name_hint}; - } - std::pair VisitExpr_(const MatchNode* op) final { - return {kOpaque, "match"}; - } - }; - return Visitor().VisitExpr(sub_expr); -} - -std::pair SubGraphKindAndLabel(const DataflowGraph& dataflow_graph, - const IndexSet& inside) { - std::ostringstream os; - bool first = true; - OpPatternKind max_kind = kElemWise; - for (PostDfsIndex index : inside) { - auto [sub_kind, sub_label] = SubExprKindAndLabel(dataflow_graph.index_to_node(index)->ref()); - if (!sub_label.empty()) { - if (first) { - first = false; - } else { - os << "+"; - } - os << sub_label; - } - max_kind = CombineKinds(max_kind, sub_kind); - } - return {max_kind, os.str()}; -} - -IndexSet MatcherToIndexSet(const DFPatternMatcher& matcher) { - IndexSet result(matcher.size()); - for (const auto& kv : matcher.memo()) { - for (const auto& matched_sub_expr : kv.second) { - if (CanInline(matched_sub_expr)) { - // Trivial sub-expressions can just be included in the extracted function body - // when we construct it and don't need to be considered part of the sub-graph. - continue; - } - if (kv.first.as()) { - // Don't consider the expressions matched by a wildcard to be part of the sub-graph. - continue; - } - result.Add(matcher.expr_to_node(matched_sub_expr)->index_); - } - } - return result; -} - -std::string SubGraphConfig::ToString() const { - std::ostringstream os; - os << "{max_exits=" << max_exits; - os << ", allow_taps=" << allow_taps; - os << ", max_depth=" << max_depth; - os << "}"; - return os.str(); -} - -TVM_REGISTER_NODE_TYPE(NestedSubGraphNode); - -void NestedSubGraphNode::VisitAttrs(AttrVisitor* v) { - // TODO(mbs) -} - -SubGraph NestedSubGraphNode::sub_graph() const { return Downcast(sub_graph_obj_); } - -bool NestedSubGraphNode::operator==(const NestedSubGraphNode& that) const { - return *sub_graph().get() == *that.sub_graph().get(); -} - -bool NestedSubGraphNode::operator<(const NestedSubGraphNode& that) const { - return *sub_graph().get() < *that.sub_graph().get(); -} - -size_t NestedSubGraphNode::hash() const { - size_t h = StructuralHash()(attrs_); - h ^= sub_graph()->hash() + 0x9e3779b9 + (h << 6) + (h >> 2); - return h; -} - -std::string NestedSubGraphNode::ToString() const { - std::ostringstream os; - os << "{sub_graph=" << sub_graph()->ToString(); - os << ", attrs=" << PrettyPrint(attrs_); - os << "}"; - return os.str(); -} - -Function NestedSubGraphNode::Extract(const DataflowGraph& dataflow_graph) const { - Extractor extractor(&dataflow_graph, sub_graph().get(), attrs_); - extractor.Extract(); - return Downcast(extractor.extracted()); -} - -Expr NestedSubGraphNode::Rewrite(const DataflowGraph& dataflow_graph, const Expr& expr) const { - Extractor extractor(&dataflow_graph, sub_graph().get(), attrs_); - extractor.Extract(); - Rewriter rewriter(&extractor); - return rewriter.VisitExpr(expr); -} - -NestedSubGraph::NestedSubGraph(SubGraph sub_graph, FunctionAttrsMap attrs) { - auto data = runtime::make_object(); - data->sub_graph_obj_ = std::move(sub_graph); - data->attrs_ = std::move(attrs); - data_ = std::move(data); -} - -NestedSubGraph NestedSubGraph::Subst( - const DataflowGraph& new_dataflow_graph, - const std::unordered_map& subst) const { - return NestedSubGraph(get()->sub_graph().Subst(new_dataflow_graph, subst), get()->attrs_); -} - -bool NestedSubGraph::TriviallyUnionable(const NestedSubGraph& that) const { - if (get()->attrs_.size() != that->attrs_.size()) { - return false; - } - for (const auto& kv : get()->attrs_) { - if (kv.first == "Composite") { - // Even if all the attributes agree we don't consider "Composite" functions to - // ever be unionable. - // TODO(mbs): Find a cleaner way to do this. - return false; - } - auto itr = that->attrs_.find(kv.first); - if (itr == that->attrs_.end()) { - return false; - } - if (!StructuralEqual()(kv.second, (*itr).second)) { - return false; - } - } - return true; -} - -NestedSubGraph NestedSubGraph::DisjointUnion(const DataflowGraph& dataflow_graph, - const NestedSubGraph& that) const { - ICHECK(TriviallyUnionable(that)); - return NestedSubGraph(get()->sub_graph().DisjointUnion(dataflow_graph, that->sub_graph()), - get()->attrs_); -} - -/*static*/ -Expr NestedSubGraph::ParallelRewrite(const DataflowGraph& dataflow_graph, const Expr& expr, - std::vector nested_sub_graphs) { - // IMPORTANT: See the corresponding comment in SubGraph::ParallelRewrite. - std::sort(nested_sub_graphs.begin(), nested_sub_graphs.end(), - [](const NestedSubGraph& left, const NestedSubGraph& right) { - return left->sub_graph()->last_inside_index_ > right->sub_graph()->last_inside_index_; - }); - - Expr result = expr; - for (const auto& nested_sub_graph : nested_sub_graphs) { - result = nested_sub_graph->Rewrite(dataflow_graph, result); - } - return result; -} - -TVM_REGISTER_NODE_TYPE(SubGraphNode); - -void SubGraphNode::VisitAttrs(AttrVisitor* v) { - // TODO(mbs) -} - -IndexSet SubGraphNode::Downstream(const DataflowGraph& dataflow_graph) const { - IndexSet downstream(dataflow_graph.size()); - for (PostDfsIndex exit_index : exit_) { - downstream = downstream | dataflow_graph.downstream_of(exit_index); - } - return downstream; -} - -bool SubGraphNode::IsValid(const DataflowGraph& dataflow_graph, - const SubGraphConfig& config) const { - // Check we don't have too many exit nodes. - if (config.max_exits > 0 && exit_.PopCount() > config.max_exits) { - VLOG(1) << "Subgraph " << ToString() << " is invalid: " << exit_.PopCount() - << " exits exceeds maximum " << config.max_exits; - return false; - } - - // Check the maximum path depth is in limit. - if (config.max_depth > 0 && depth_ > config.max_depth) { - VLOG(1) << "Subgraph " << ToString() << " is invalid: maximum depth " << depth_ - << " exceeds limit " << config.max_depth; - return false; - } - - // All inside nodes must be in the same basic block. - const DataflowGraph::Node* basic_block = nullptr; - for (PostDfsIndex index : inside_) { - auto node = dataflow_graph.index_to_node(index); - if (basic_block == nullptr) { - basic_block = node->basic_block_; - } - if (node->basic_block_ != basic_block) { - VLOG(1) << "Subgraph " << ToString() << " is invalid: nodes are from different basic blocks"; - return false; - } - } - - // The nested sub-graphs must be subsets and non-overlapping. - IndexSet union_inside(dataflow_graph.size()); - for (const auto& nested_sub_graph : nested_sub_graphs_) { - if (!nested_sub_graph->sub_graph()->inside_.AreDisjoint(union_inside)) { - VLOG(1) << "Subgraph " << ToString() << " is invalid: nested sub-graphs overlap"; - return false; - } - if (!nested_sub_graph->sub_graph()->inside_.IsSubset(inside_)) { - VLOG(1) << "Subgraph " << ToString() - << " is invalid: nested sub-graph is not subset of overall sub-graph"; - return false; - } - } - - if (!config.allow_taps) { - // Exit nodes cannot also contribute to inside nodes. - for (PostDfsIndex index : exit_) { - auto node = dataflow_graph.index_to_node(index); - if (AnyOutputInside(node)) { - VLOG(1) << "Subgraph " << ToString() - << " is invalid: inner node is 'tapped' and also contributes to output, but taps " - "are disabled"; - return false; - } - } - } - - // Check no output would end up feeding into any entry node. - for (PostDfsIndex output_index : output_) { - if (dataflow_graph.downstream_of(output_index).Intersects(entry_)) { - VLOG(1) << "Subgraph " << ToString() << " is invalid: output node " << output_index - << " feeds back into this sub-graph"; - return false; - } - } - - // Looks legit! - return true; -} - -Function SubGraphNode::ExtractAsFunction(const DataflowGraph& dataflow_graph) const { - NestedSubGraph nested_sub_graph(GetRef(this), FunctionAttrsMap()); - return nested_sub_graph->Extract(dataflow_graph); -} - -Expr SubGraphNode::Rewrite(const DataflowGraph& dataflow_graph, const Expr& expr) const { - if (nested_sub_graphs_.empty()) { - // Nothing to rewrite. - return expr; - } - Extractor extractor(&dataflow_graph, this, NullValue()); - extractor.Extract(); - Rewriter rewriter(&extractor); - return rewriter.VisitExpr(expr); -} - -std::string SubGraphNode::ToString() const { - std::ostringstream os; - os << "{inside=" << inside_.ToString(); - os << ", entry=" << entry_.ToString(); - os << ", exit=" << exit_.ToString(); - os << ", input=" << input_.ToString(); - os << ", output=" << output_.ToString(); - os << ", depth=" << depth_; - os << ", kind=" << KindToString(kind_); - if (!label_.empty()) { - os << ", label=" << label_; - } - for (const auto& nested_sub_graph : nested_sub_graphs_) { - os << ", nested_sub_graph=" << nested_sub_graph->ToString(); - } - os << "}"; - return os.str(); -} - -bool SubGraphNode::operator==(const SubGraphNode& that) const { - ICHECK_EQ(inside_.end_index(), that.inside_.end_index()); - if (inside_ != that.inside_) { - return false; - } - if (nested_sub_graphs_.size() != that.nested_sub_graphs_.size()) { - return false; - } - for (size_t i = 0; i < nested_sub_graphs_.size(); ++i) { - if (*nested_sub_graphs_[i].get() != *that.nested_sub_graphs_[i].get()) { - return false; - } - } - return true; -} - -bool SubGraphNode::operator<(const SubGraphNode& that) const { - if (first_inside_index_ < that.first_inside_index_) { - return true; - } - if (that.first_inside_index_ < first_inside_index_) { - return false; - } - return inside_ < that.inside_; -} - -size_t SubGraphNode::hash() const { - size_t h = inside_.hash(); - for (const auto& nested_sub_graph : nested_sub_graphs_) { - h ^= nested_sub_graph->hash() + 0x9e3779b9 + (h << 6) + (h >> 2); - } - return h; -} - -void SubGraphNode::Init(const DataflowGraph& dataflow_graph) { - for (PostDfsIndex index = 0; index < inside_.end_index(); ++index) { - auto node = dataflow_graph.index_to_node(index); - if (inside_[index]) { - if (AnyInputOutside(node)) { - entry_.Add(index); - } - if (AnyOutputOutside(node) || node->is_external_) { - exit_.Add(index); - } - } else { - if (AnyInputInside(node)) { - output_.Add(index); - } - if (AnyOutputInside(node) && !CanInline(node->ref())) { - input_.Add(index); - } - } - } - depth_ = Depth(dataflow_graph); -} - -size_t SubGraphNode::Depth(const DataflowGraph& dataflow_graph) const { - std::unordered_map max_depths; - std::vector stack; - size_t max_depth = 0; - // All the entry nodes have max depth 0. - for (PostDfsIndex index : entry_) { - auto node = dataflow_graph.index_to_node(index); - max_depths.emplace(node, 0); - stack.push_back(node); - } - while (!stack.empty()) { - const DataflowGraph::Node* node = stack.back(); - stack.pop_back(); - size_t next_depth = max_depths[node] + 1; - if (exit_[node->index_]) { - // If this node is external then it will have no outputs but we still wish to consider - // the path to the implied output as requiring one more step. - // Otherwise we're accounting for reaching one of the external outputs belowe. - max_depth = std::max(max_depth, next_depth); - } - for (const DataflowGraph::Node* output_node : node->outputs_) { - if (!inside_[output_node->index_]) { - continue; - } - if (max_depths.count(output_node) == 0) { - max_depths.emplace(output_node, next_depth); - stack.push_back(output_node); - } else if (next_depth > max_depths[output_node]) { - // We found a deeper path to an already expanded node. We'll expand again. - max_depths[output_node] = next_depth; - stack.push_back(output_node); - } - } - } - return max_depth; -} - -/*! \brief Returns true if any (input/output) of node is (outside/inside) the sub-graph. */ -bool SubGraphNode::AnyInputOutside(const DataflowGraph::Node* node) const { - return std::any_of(node->inputs_.begin(), node->inputs_.end(), - [this](const DataflowGraph::Node* sub_node) { - return !inside_[sub_node->index_] && !CanInline(sub_node->ref()); - }); -} - -bool SubGraphNode::AnyInputInside(const DataflowGraph::Node* node) const { - return std::any_of( - node->inputs_.begin(), node->inputs_.end(), - [this](const DataflowGraph::Node* sub_node) { return inside_[sub_node->index_]; }); -} - -bool SubGraphNode::AnyOutputOutside(const DataflowGraph::Node* node) const { - return std::any_of( - node->outputs_.begin(), node->outputs_.end(), - [this](const DataflowGraph::Node* sub_node) { return !inside_[sub_node->index_]; }); -} - -bool SubGraphNode::AnyOutputInside(const DataflowGraph::Node* node) const { - return std::any_of( - node->outputs_.begin(), node->outputs_.end(), - [this](const DataflowGraph::Node* sub_node) { return inside_[sub_node->index_]; }); -} - -SubGraph::SubGraph(const DataflowGraph& dataflow_graph, IndexSet inside, OpPatternKind kind, - String label, std::vector nested_sub_graphs) { - std::sort(nested_sub_graphs.begin(), nested_sub_graphs.end(), - [](const NestedSubGraph& left, const NestedSubGraph& right) { - return *left.get() < *right.get(); - }); - auto node = runtime::make_object(); - node->inside_ = std::move(inside); - node->first_inside_index_ = node->inside_.FirstInsideIndex(); - node->last_inside_index_ = node->inside_.LastInsideIndex(); - node->entry_ = IndexSet(node->inside_.end_index()); - node->exit_ = IndexSet(node->inside_.end_index()); - node->input_ = IndexSet(node->inside_.end_index()); - node->output_ = IndexSet(node->inside_.end_index()); - node->kind_ = kind; - node->label_ = std::move(label); - node->nested_sub_graphs_ = nested_sub_graphs; - node->Init(dataflow_graph); - data_ = std::move(node); -} - -SubGraph::SubGraph(const DataflowGraph& dataflow_graph) - : SubGraph(dataflow_graph, IndexSet(dataflow_graph.size())) {} - -bool SubGraph::AreDisjoint(const SubGraph& that) const { - return get()->inside_.AreDisjoint(that->inside_); -} - -namespace { -/*! \brief Returns true if an output of \p left not in \p right ultimately flows into \p right. */ -bool FlowsInto(const DataflowGraph& dataflow_graph, const SubGraph& left, const SubGraph& right) { - for (PostDfsIndex output_index : left->output_) { - if (!right->inside_[output_index] && - dataflow_graph.downstream_of(output_index).Intersects(right->entry_)) { - return true; - } - } - return false; -} -} // namespace - -bool SubGraph::AreTouching(const DataflowGraph& dataflow_graph, const SubGraph& that) const { - if (!get()->inside_.AreDisjoint(that->inside_)) { - // Easy rejection. - return false; - } - if (!get()->output_.Intersects(that->entry_)) { - // Not touching. - return false; - } - if (FlowsInto(dataflow_graph, *this, that) || FlowsInto(dataflow_graph, that, *this)) { - // Unioning would create a cycle. - return false; - } - return true; -} - -bool SubGraph::AreSelfContained(const SubGraph& that) const { - return get()->output_.IsSubset(that->entry_) && that->input_.IsSubset(get()->exit_); -} - -SubGraph SubGraph::DisjointUnion(const DataflowGraph& dataflow_graph, const SubGraph& that) const { - ICHECK(AreDisjoint(that)); - IndexSet inside = get()->inside_ | that->inside_; - std::vector nested_sub_graphs; - for (const auto& nested_sub_graph : get()->nested_sub_graphs_) { - nested_sub_graphs.push_back(nested_sub_graph); - } - for (const auto& nested_sub_graph : that->nested_sub_graphs_) { - auto existing_itr = std::find_if(nested_sub_graphs.begin(), nested_sub_graphs.end(), - [&nested_sub_graph](const NestedSubGraph& existing) { - return existing.TriviallyUnionable(nested_sub_graph); - }); - if (existing_itr != nested_sub_graphs.end()) { - *existing_itr = existing_itr->DisjointUnion(dataflow_graph, nested_sub_graph); - } else { - nested_sub_graphs.push_back(nested_sub_graph); - } - } - return SubGraph(dataflow_graph, std::move(inside), CombineKinds(get()->kind_, that->kind_), - UnionLabels(get()->label_, that->label_), std::move(nested_sub_graphs)); -} - -SubGraph SubGraph::WithAttrs(const DataflowGraph& dataflow_graph, FunctionAttrsMap attrs) const { - std::vector nested_sub_graphs; - nested_sub_graphs.push_back(NestedSubGraph(*this, attrs)); - return SubGraph(dataflow_graph, get()->inside_, get()->kind_, get()->label_, - std::move(nested_sub_graphs)); -} - -SubGraph SubGraph::Subst(const DataflowGraph& new_dataflow_graph, const IndexSubst& subst) const { - IndexSet new_inside = get()->inside_.Subst(new_dataflow_graph.size(), subst); - std::vector new_nested_sub_graphs; - for (const auto& nested_sub_graph : get()->nested_sub_graphs_) { - new_nested_sub_graphs.push_back(nested_sub_graph.Subst(new_dataflow_graph, subst)); - } - return SubGraph(new_dataflow_graph, std::move(new_inside), get()->kind_, get()->label_, - std::move(new_nested_sub_graphs)); -} - -/*static*/ -Expr SubGraph::ParallelRewrite(const DataflowGraph& dataflow_graph, - std::vector sub_graphs) { - // IMPORTANT: - // - All the sub-graphs will be w.r.t. the dataflow graph for the original expression. - // Each time we call Rewrite on one of those graphs the result expression will be rewritten - // from the final output back to the inputs. The inputs will then be shared with the original - // expression. Thus it is safe to iteratively rewrite all the sub-graphs without redoing the - // dataflow_graph and substituting indexes provided we work in reverse dataflow order. - // - We rely on the dataflow_graph expression reference holding the original expression alive - // so that the dataflow_graph will never contain dangling pointers (even though as per above - // we'll never dereference them). - std::sort(sub_graphs.begin(), sub_graphs.end(), [](const SubGraph& left, const SubGraph& right) { - return left->last_inside_index_ > right->last_inside_index_; - }); - Expr result = dataflow_graph.expr(); - for (const auto& sub_graph : sub_graphs) { - result = sub_graph->Rewrite(dataflow_graph, result); - } - return result; -} - -/*! - * \brief A pass which partitions (the unique) global function in the module according to the - * post-dfs indexes in \p indexes. The partitioning must respect the configuration with \p max_exits - * and \p allow_taps. - * - * Each index is also paired with a label. A non-empty label denotes the index should also be - * included in a nested sub-graph which will be extracted as a function with the label as its - * "Composite" attribute. An empty label denotes the index should go into the overall partitioned - * "Compiler" function. In this way we can simulate the usual partitioning needed by external - * codegen integrations. - * - * This function is intended to support \p SubGraph unit tests and is not used by the regular - * compilation flow. - */ -transform::Pass PartitionForTesting(Integer max_exits, Bool allow_taps, String compiler, - Array indexes, Array labels) { - auto pass_func = [=](Function function, IRModule mod, transform::PassContext ctxt) { - ICHECK(max_exits.defined() && max_exits->value >= 0); - ICHECK(allow_taps.defined()); - ICHECK(indexes.size() == labels.size()); - VLOG(1) << "Partitioning:" << std::endl << PrettyPrint(function); - DataflowGraph dataflow_graph(function); - VLOG(1) << "Dataflow graph is:" << std::endl << dataflow_graph.indexed_graph().ToString(); - - // Collect the 'inside' indexes and any nested sub-graph indexes and labels. - std::vector node_indexes; - std::unordered_map> nested_sub_graph_indexes; - node_indexes.reserve(indexes.size()); - for (size_t i = 0; i < indexes.size(); ++i) { - const Integer& index = indexes[i]; - ICHECK_GE(index->value, 0); - ICHECK_LT(index->value, dataflow_graph.size()); - auto index_int = static_cast(index->value); - node_indexes.push_back(index_int); - const String& label = labels[i]; - if (!label.empty()) { - nested_sub_graph_indexes[label].push_back(index_int); - } - } - - // Build the nested sub-graphs representing the "Composite" functions (if any). - std::vector nested_sub_graphs; - for (const auto& kv : nested_sub_graph_indexes) { - FunctionAttrsMap composite_attrs; - composite_attrs.Set("Composite", kv.first); - nested_sub_graphs.emplace_back( - SubGraph(dataflow_graph, IndexSet(dataflow_graph.size(), kv.second)), composite_attrs); - } - - // Build the overall sub-graph, which will include any "Composite" functions as - // well as any nodes without a label. - IndexSet inside(dataflow_graph.size(), node_indexes); - auto [kind, label] = SubGraphKindAndLabel(dataflow_graph, inside); - SubGraph sub_graph(dataflow_graph, inside, kind, label, std::move(nested_sub_graphs)); - - // Push the overall sub-graph into the final "Compiler" function. - FunctionAttrsMap compiler_attrs; - compiler_attrs.Set("Compiler", compiler); - NestedSubGraph overall_nested_sub_graph(sub_graph, compiler_attrs); - SubGraph overall_sub_graph(dataflow_graph, inside, kind, label, {overall_nested_sub_graph}); - - // Check the sub-graph is valid. - SubGraphConfig config; - config.max_exits = static_cast(max_exits->value); - config.allow_taps = allow_taps; - if (overall_sub_graph->IsValid(dataflow_graph, config)) { - VLOG(1) << "Sub-graph " << overall_sub_graph->ToString() << " is considered valid"; - } else { - VLOG(1) << "Sub-graph " << overall_sub_graph->ToString() - << " is NOT considered valid, not partitioning"; - return function; - } - - // Do the partitioning. - Function result = Downcast(overall_sub_graph->Rewrite(dataflow_graph, function)); - VLOG(1) << "Extracted as:" << std::endl << PrettyPrint(result); - - return result; - }; - return transform::CreateFunctionPass(pass_func, /*opt_level=*/0, "PartitionForTesting", {}); -} - -TVM_REGISTER_GLOBAL("relay.collage.PartitionForTesting").set_body_typed(PartitionForTesting); - -} // namespace collage -} // namespace relay -} // namespace tvm diff --git a/src/relay/collage/sub_graph.h b/src/relay/collage/sub_graph.h deleted file mode 100644 index f7d4354d5483..000000000000 --- a/src/relay/collage/sub_graph.h +++ /dev/null @@ -1,452 +0,0 @@ -/* - * 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 src/relay/collage/sub_graph.h - * \brief Represents a sub-graph of an overall Relay expression. - */ - -#ifndef TVM_RELAY_COLLAGE_SUB_GRAPH_H_ -#define TVM_RELAY_COLLAGE_SUB_GRAPH_H_ - -#include -#include - -#include -#include -#include -#include - -#include "../ir/dataflow_matcher_impl.h" -#include "../ir/indexed_graph.h" -#include "./dataflow_graph.h" -#include "./index_set.h" - -namespace tvm { -namespace relay { -namespace collage { - -/*! \brief Returns operator pattern kind as single-letter string. */ -std::string KindToString(OpPatternKind kind); - -/*! - * \brief Returns a kind and label for the single \p sub_expr, ignoring its nested sub expressions. - */ -std::pair SubExprKindAndLabel(const Expr& sub_expr); - -/*! - * \brief Returns a kind and label for all the nodes in \p inside. - */ -std::pair SubGraphKindAndLabel(const DataflowGraph& dataflow_graph, - const IndexSet& inside); - -/*! - * \brief Returns the index set representing all the sub-expression matched by \p matcher. - */ -IndexSet MatcherToIndexSet(const DFPatternMatcher& matcher); - -/*! - * \brief Configuration controlling which sub-graphs are considered valid. - */ -struct SubGraphConfig { - /*! \brief Maximum number of exit nodes in the sub-graph, or zero if no limit. */ - size_t max_exits = 0; - /*! - * \brief Whether a node inside the sub-graph may flow to nodes both inside and outside - * the sub-graph (which we call a 'tap'). Note that it is still possible to have multiple outputs - * even with this flag false. - */ - bool allow_taps = false; - /*! - * \brief Maximum allowed sub-graph depth, or zero if no-limit. - */ - size_t max_depth = 0; - - std::string ToString() const; -}; - -class SubGraph; -using FunctionAttrsMap = Map; - -/*! - * \brief A nested sub-graph is a sub-graph which is to be nested inside a function as part of some - * enclosing sub-graph. - * - * Extraction yields a function with input nodes replaced by parameters and exit nodes in the - * function result. Rewriting replaces the sub-graph with a call to that function, and all - * outputs with (projections from) the call result. - * - * (Note that it's tempting to move attrs_ into \p SubGraphNode and thus avoid this class. - * However we found the implementation was easier to understand in this form since it makes - * the result of \p Extract unambiguous.) - */ -class NestedSubGraphNode : public Object { - public: - /*! \brief The nested sub-graph. */ - ObjectRef /* actually SubGraph */ sub_graph_obj_; - /*! \brief Attributes (possibly empty) to attach to the extracted function. */ - FunctionAttrsMap attrs_; - - void VisitAttrs(AttrVisitor* v); - - SubGraph sub_graph() const; - - bool operator==(const NestedSubGraphNode& that) const; - bool operator!=(const NestedSubGraphNode& that) const { return !(*this == that); } - bool operator<(const NestedSubGraphNode& that) const; - size_t hash() const; - - std::string ToString() const; - - /*! - * \brief Returns the function representing this nested sub-graph within the overall expression - * represented by \p dataflow_graph: - * - All sub-graph inputs become parameters. - * - All sub-graph outputs become function results (either directly or as a field in a tuple). - * - The function has attrs_ for attributes (which may be empty). - * - The function body accounts for any rewrites implied by the nested sub-graph. - */ - Function Extract(const DataflowGraph& dataflow_graph) const; - - /*! - * \brief Returns \p expr rewritten to encode the partitioning implied by this nested sub-graph. - * - * It is valid for \p expr to not be the same as \p dataflow_graph.expr(), however all nodes - * inside this nested sub-graph must correspond to nodes shared between \p dataflow_graph.expr() - * and \p expr. See \p SubGraph::ParallelRewrite below. - */ - Expr Rewrite(const DataflowGraph& dataflow_graph, const Expr& expr) const; - - static constexpr const char* _type_key = "relay.collage.NestedSubGraph"; - TVM_DECLARE_FINAL_OBJECT_INFO(NestedSubGraphNode, Object); -}; - -class NestedSubGraph : public ObjectRef { - public: - NestedSubGraph(SubGraph sub_graph, FunctionAttrsMap attrs); - - /*! - * \brief Returns copy of this nested sub-graph with all indexes substituted according to - * \p subst, whose range is w.r.t. \p new_dataflow_graph. - */ - NestedSubGraph Subst(const DataflowGraph& new_dataflow_graph, - const std::unordered_map& subst) const; - - /*! - * \brief Returns true if this can be safely unioned. - */ - bool TriviallyUnionable(const NestedSubGraph& that) const; - - /*! - * \brief Returns the disjoint union of this and \p that nested sub-graphs, which must agree on - * their attributes. - */ - NestedSubGraph DisjointUnion(const DataflowGraph& dataflow_graph, - const NestedSubGraph& that) const; - - /*! - * \brief Returns \p expr rewritten according to all the given nested sub-graphs. The - * nested sub-graphs can be given in any order, but must be disjoint. - * - * It is valid for \p expr to not be the same as \p dataflow_graph.expr(), however all nodes - * inside the nested sub-graphs must correspond to nodes shared between \p dataflow_graph.expr() - * and \p expr. See \p SubGraph::ParallelRewrite below. - */ - static Expr ParallelRewrite(const DataflowGraph& dataflow_graph, const Expr& expr, - std::vector nested_sub_graphs); - - TVM_DEFINE_OBJECT_REF_METHODS(NestedSubGraph, ObjectRef, NestedSubGraphNode); -}; - -using NestedSubGraphs = Array; - -/*! - * \brief A compact representation of a sub-graph within an (implied) overall Relay expression. - * - * Sub-graphs can be used to represent partitions/kernels/composite functions without having to - * pay the cost of constructing or rewriting any expressions. We also allow 'extracting' a - * function to use for measuring a partition/kernel's latency independently from 'rewriting' - * the overall Relay expression since only a tiny subset of candidate partitions will end up being - * needed after Collage has completed its search. - * - * We expect O(thousands) of sub-graphs to be in flight while processing a given model, so we are - * mindful of space overhead. - * - * A sub-graph classifies every dataflow node of the overall expression as either 'inside' or - * 'outside' the sub-graph. Obviously not all such divisions make sense, for example it is not - * valid for an inside node to feed into another inside node via outside nodes. We provide the - * \p IsValid method to check for validity, and \p SubGraphConfig to control which validity rules - * apply (such as maximum depth). - * - * We generally work with the \p DataflowGraph representation of the overall Relay expression - * rather than the expression itself. We use the post-dfs visit index to uniquely refer to - * expression nodes. - * - * As well as 'inside' and 'outside' we have four other flavors of dataflow nodes, all uniquely - * determined from the 'inside' nodes: - * - 'entry' nodes are those inside with at least one dataflow input outside. - * - 'exit' nodes are those inside with at least one dataflow output outside, or which - * are considered 'external' in the underlying dataflow graph (eg because they represent - * the result of the overall function). - * - 'input' nodes are those outside with at least one dataflow output inside. - * - 'output' nodes are those outside with at least one dataflow input inside. - * Index sets for these are cached with the sub-graph for performance. - * - * It is valid to have multiple entry nodes (we can bind a parameter for each). It may be valid to - * have multiple exit nodes (we can build a tuple of all such). It may be valid to have exit nodes - * which also contribute to other inside nodes (ie represent a 'tap' on an intermediate result). - * - * Sub-graphs are closed under: - * - Disjoint union. - * - Wrapping by a function with given attributes (see \p NestedSubGraph above). This can be used - * to encode "Composite" functions, or to represent a candidate kernel within a "Primitive" - * function. (By combining 'wrapping' with 'union' we can encode, eg, 'this sub-graph should - * be placed inside a primitive function which itself may have calls to composite functions). - * - Substitution, which allows a sub-graph w.r.t. one dataflow graph to be transformed to - * match some other (typically smaller) dataflow graph. - * - * See the subclasses of \p PartitionRule for how sub-graphs are built and combined during Collage - * search. - * - * To support some of the \p OpPatternKind-based fusion rule processing we give sub-graphs - * a kind, which is generally the maximum of the kinds of all the operator calls appearing - * inside it. We also given sub-graphs a (not necessarily unique) label to help debugging - * and guide the selection of global symbol names. - */ -class SubGraphNode : public Object { - public: - /*! - * \brief Which sub-expressions are inside the sub-graph (using their post-dfs indexes w.r.t. - * the implied DataflowGraph). - */ - IndexSet inside_; - - /*! - * \brief Index of first and last inside nodes. - * - * Cached for performance, uniquely determined by inside_. - */ - PostDfsIndex first_inside_index_ = 0; - PostDfsIndex last_inside_index_ = 0; - - /*! - * \brief Which sub-expressions are entry/exit/input/output for this sub-graph. - * - * Cached for performance, uniquely determined by inside_. - */ - IndexSet entry_; - IndexSet exit_; - IndexSet input_; - IndexSet output_; - - /*! - * \brief Maximum depth of any dataflow path from an entry to an output sub-expression. - * - * Cached for performance, uniquely determined by inside_. - */ - size_t depth_ = 0; - - /*! - * \brief The \p OpPatternKind summarizing the input/output behavior of the sub-graph. - * - * A sub-graph consisting of a single Relay expression node is given kind: - * - For Call to a Relay operator, the "TOpPattern" attribute of that operator (provided the - * call does not involve data-dependent dynamic shapes). - * - For Call to Relay Function, the "TOpPattern" attribute of the function (provided it has - * that attribute) - * - For Constants, \p kElemWise. - * - For Tuple and tuple projections, \p kInjective (provided all tuple fields are of tensor - * type) - * - All other nodes \p kOpaque. - * Sub-graphs with more than one node have the maximum of the kind of each node. - * - * Cached for performance, uniquely determined by inside_. - */ - OpPatternKind kind_ = kOpaque; - - /*! - * \brief A label for the sub-graph. Not guaranteed to be unique, but is a human-readable summary - * of the sub-graph which can help with debugging and guide the selection of global symbol names. - */ - String label_; - - /*! - * \brief Nested sub-graphs of this sub-graph which must be represented by functions. These must - * be disjoint, but it's ok for this sub-graph to have nodes not inside any nested sub-graph. - */ - NestedSubGraphs nested_sub_graphs_; - - void VisitAttrs(AttrVisitor* v); - - // TODO(mbs): 'Anchor nodes' and rules for unioning them. - // In FuseOps it's just the unique kEWiseFusable node, if any. - // I'd like to allow writing vertical fusion rules, eg if two candidates are directly - // connected and have nn.conv2d anchors allow their join. - // I'd also like to allow horizontal fusion rules, eg if two candidates are not directly - // connected but could be joined without producing invalid (eg cyclic) and have nn.conv2d anchors - // then do so. Come back to this. - - /*! \brief Number of nodes in overall dataflow graph. */ - size_t overall_size() const { return inside_.end_index(); } - - bool IsEmpty() const { return inside_.IsZero(); } - - /*! \brief Number of nodes in sub-graph. */ - size_t Size() const { return inside_.PopCount(); } - - /*! - * \brief Returns the dataflow nodes downstream of all exit nodes. - */ - IndexSet Downstream(const DataflowGraph& dataflow_graph) const; - - /*! - * \brief Returns true if this sub-graph is valid. Ie: - * - no output of the sub-graph can flow to any input of the sub-graph (otherwise we'd end up - * with a dataflow cycle when we partition). - * - all inputs and outputs of the sub-graph are in the same scope, ie not separated by - * control flow (otherwise there'd be no consistent program point at which to eval the - * partitioned function). - * - no more than config.max_outputs outputs are required. - * - if config.allow_taps is false, no inside node has outputs to nodes both inside and - * outside the sub-graph. - */ - bool IsValid(const DataflowGraph& dataflow_graph, const SubGraphConfig& config) const; - - /*! - * \brief Returns this sub-graph extracted as a stand-alone function. The function will have - * no attributes, and is suitable for building and profiling by the \p CostEstimator. - */ - Function ExtractAsFunction(const DataflowGraph& dataflow_graph) const; - - /*! - * \brief Returns \p expr rewritten to encode the partitioning implied by this sub-graph. - * - * It is valid for \p expr to not be the same as \p dataflow_graph.expr(), however all nodes - * inside this sub-graph must correspond to nodes shared between \p dataflow_graph.expr() and - * \p expr. See \p SubGraph::ParallelRewrite below. - */ - Expr Rewrite(const DataflowGraph& dataflow_graph, const Expr& expr) const; - - std::string ToString() const; - - bool operator==(const SubGraphNode& that) const; - bool operator!=(const SubGraphNode& that) const { return !(*this == that); } - bool operator<(const SubGraphNode& that) const; - size_t hash() const; - - private: - /*! \brief Initialize the entry/exit/input/output sets given the inside and \p dataflow_graph. */ - void Init(const DataflowGraph& dataflow_graph); - - /*! \brief Calculates and returns the maximum path depth. */ - size_t Depth(const DataflowGraph& dataflow_graph) const; - - /*! \brief Returns true if any (input/output) of node is (outside/inside) the sub-graph. */ - bool AnyInputOutside(const DataflowGraph::Node* node) const; - bool AnyInputInside(const DataflowGraph::Node* node) const; - bool AnyOutputOutside(const DataflowGraph::Node* node) const; - bool AnyOutputInside(const DataflowGraph::Node* node) const; - - public: - static constexpr const char* _type_key = "relay.collage.SubGraph"; - TVM_DECLARE_FINAL_OBJECT_INFO(SubGraphNode, Object); - - friend class SubGraph; -}; - -class SubGraph : public ObjectRef { - public: - /*! \brief Primitive constructor. The following constructors are generally more convenient. */ - SubGraph(const DataflowGraph& dataflow_graph, IndexSet inside, OpPatternKind kind = kOpaque, - String label = {}, std::vector nested_sub_graphs = {}); - - /*! \brief Constructs the empty sub-graph for \p dataflow_graph. */ - explicit SubGraph(const DataflowGraph& dataflow_graph); - - /*! \brief Returns true if this and that are disjoint. */ - bool AreDisjoint(const SubGraph& that) const; - - /*! - * \brief Returns true if: - * - \p this and \p that are disjoint, and - * - an output node of \p this coincides with an entry node of \p that, and - * - \p this and \p that are not obviously invalid after \p DisjointUnion - * (eg because such a sub-graph would produce a cycle). - * Note however that the \p DisjointUnion may not necessarily be valid even with the above - * checks. - */ - bool AreTouching(const DataflowGraph& dataflow_graph, const SubGraph& that) const; - - /*! - * \brief Returns true if: - * - all the outputs of \p this are entries for \p that, and - * - all the inputs of \p that are exits for \p this. - */ - bool AreSelfContained(const SubGraph& that) const; - - /*! - * \brief Returns disjoint union of this and \p that sub-graphs. The result may not be valid. - */ - SubGraph DisjointUnion(const DataflowGraph& dataflow_graph, const SubGraph& that) const; - - /*! - * \brief Returns copy of this sub-graph with all nodes placed inside a nested sub-graph with - * given attributes. - */ - SubGraph WithAttrs(const DataflowGraph& dataflow_graph, FunctionAttrsMap attrs) const; - - /*! - * \brief Returns copy of this sub-graph with all indexes substituted according to \p subst, - * whose range is w.r.t. \p new_dataflow_graph. - */ - SubGraph Subst(const DataflowGraph& new_dataflow_graph, - const std::unordered_map& subst) const; - - /*! - * \brief Returns the root expression of \p dataflow_graph rewritten according to all the - * given sub-graphs. The sub-graphs can be given in any order, but must be disjoint. - */ - static Expr ParallelRewrite(const DataflowGraph& dataflow_graph, - std::vector sub_graphs); - - TVM_DEFINE_OBJECT_REF_METHODS(SubGraph, ObjectRef, SubGraphNode); -}; - -struct SubGraphEqual { - bool operator()(const SubGraph& left, const SubGraph& right) const { - return *left.get() == *right.get(); - } -}; - -struct SubGraphHash { - size_t operator()(const SubGraph& sub_graph) const { return sub_graph->hash(); } -}; - -/*! - * \brief Pass to partition every global function according to the post-dfs indexes - * given in an array. Visible for testing from Python only, would never make sense to use - * as a generic pass! - */ -tvm::transform::Pass PartitionOnIndexesForTesting(Array indexes); - -} // namespace collage -} // namespace relay -} // namespace tvm - -#endif // TVM_RELAY_COLLAGE_SUB_GRAPH_H_ diff --git a/src/relay/collage/utils.cc b/src/relay/collage/utils.cc deleted file mode 100644 index 451e18c219d6..000000000000 --- a/src/relay/collage/utils.cc +++ /dev/null @@ -1,152 +0,0 @@ -/* - * 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 src/relay/collage/utils.cc - * \brief Misc helpers. - */ - -#include "./utils.h" - -#include "../../support/scalars.h" -#include "../op/memory/device_copy.h" - -namespace tvm { -namespace relay { -namespace collage { - -String GetSpecName(const Target& target) { - if (target.IsExternalCodegen()) { - return target->kind->name; - } else { - return std::string(kTVMSpecNamePrefix) + target->kind->name; - } -} - -String UnionLabels(String left, String right) { - if (left.empty()) { - return right; - } - if (right.empty()) { - return left; - } - return left + "+" + right; -} - -String NestLabels(String left, String right) { - if (left.empty()) { - return right; - } - if (right.empty()) { - return left; - } - if (right.size() > left.size()) { - std::string right_str = right; - if (right_str.substr(0, left.size()) == left) { - return right; - } - } - return left + "." + right; -} - -std::string KindToString(OpPatternKind kind) { - switch (kind) { - case kElemWise: - return "E"; - case kBroadcast: - return "B"; - case kInjective: - return "I"; - case kCommReduce: - return "R"; - case kOutEWiseFusable: - return "A"; - case kTuple: - return "T"; - case kOpaque: - return "O"; - } - return "?"; -} - -OpPatternKind CombineKinds(OpPatternKind left, OpPatternKind right) { - return std::max(left, right); -} - -bool CanInline(const Expr& expr) { - if (expr.as() || expr.as() || expr.as()) { - return true; - } - if (const auto* constant_node = expr.as()) { - return support::IsSimpleScalar(constant_node); - } - return false; -} - -bool IsSpecialOp(const OpNode* op_node) { - auto op = GetRef(op_node); - static auto fnoncomputational = Op::GetAttrMap("TNonComputational"); - if (fnoncomputational.count(op) && fnoncomputational[op]) { - // Operator has been marked as non-computational. - return true; - } - // TODO(mbs): This is incomplete. - static auto shape_of_op_ = Op::Get("shape_of"); - static auto vm_shape_of_op_ = Op::Get("vm.shape_of"); - if (op == DeviceCopyOp() || op == shape_of_op_ || op == vm_shape_of_op_) { - // Operator is compiled away by the VM compilation flow. - return true; - } - return false; -} - -bool MustBeLowered(const Expr& expr) { - if (const auto* call_node = expr.as()) { - if (const auto* function_node = call_node->op.as()) { - if (function_node->HasNonzeroAttr(attr::kPrimitive)) { - // We've already committed to this call being to one or more operators which must be - // lowered. - return true; - } - } else if (const auto* op_node = call_node->op.as()) { - if (!IsSpecialOp(op_node)) { - // The VM compilation path won't rewrite this call. - return true; - } - } - } - return false; -} - -std::vector SplitString(std::string stmt, const char* del) { - std::vector str_tokens; - int start = 0; - int end = stmt.find(del, 0); - str_tokens.emplace_back(stmt.substr(start, end)); - while (end != -1) { - stmt = stmt.substr(end + 1, stmt.size()); - end = stmt.find(del, 0); - str_tokens.emplace_back(stmt.substr(start, end)); - } - return str_tokens; -} - -} // namespace collage -} // namespace relay -} // namespace tvm diff --git a/src/relay/collage/utils.h b/src/relay/collage/utils.h deleted file mode 100644 index 630b3b22f199..000000000000 --- a/src/relay/collage/utils.h +++ /dev/null @@ -1,92 +0,0 @@ -/* - * 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 src/relay/collage/utils.h - * \brief Misc helpers. - */ - -#ifndef TVM_RELAY_COLLAGE_UTILS_H_ -#define TVM_RELAY_COLLAGE_UTILS_H_ - -#include -#include -#include -#include - -#include -#include - -namespace tvm { -namespace relay { -namespace collage { - -/*! - * \brief Distinguished partition spec names. - */ -constexpr const char* kTVMSpecNamePrefix = "tvm_"; -constexpr const char* kHostSpecName = "host"; - -/*! - * \brief Returns the partition spec name to use for \p target. For external codegen targets the - * spec name is just the target kind name. For TVM native targets the spec name is of the form - * "tvm_". - */ -String GetSpecName(const Target& target); - -/*! \brief Returns \p "+". */ -String UnionLabels(String left, String right); - -/*! \brief Returns \p ".". */ -String NestLabels(String outer, String inner); - -/*! \brief Returns abbreviation for \p kind. */ -std::string KindToString(OpPatternKind kind); - -/*! \brief Returns maximum of \p left and \p right. */ -OpPatternKind CombineKinds(OpPatternKind left, OpPatternKind right); - -/*! - * \brief Returns true if \p expr can be safely inlined in body of function extracted - * from sub-graph, even if \p expr was not technically matched by the pattern which produced - * the sub-graph. - */ -bool CanInline(const Expr& expr); - -/*! - * \brief Returns true if \p op_node can be directly handled by the VM. - */ -bool IsSpecialOp(const OpNode* op_node); - -/*! - * \brief Return true if the Relay expression node given by \p expr cannot be evaluated by - * the VM and must end up in a kernel. - */ -bool MustBeLowered(const Expr& expr); - -/*! - * \brief Returns the list of split strings of given statement with delimiter. - */ -std::vector SplitString(std::string stmt, const char* del); - -} // namespace collage -} // namespace relay -} // namespace tvm - -#endif // TVM_RELAY_COLLAGE_UTILS_H_ diff --git a/src/relay/ir/adt.cc b/src/relay/ir/adt.cc deleted file mode 100644 index 0389547a78f9..000000000000 --- a/src/relay/ir/adt.cc +++ /dev/null @@ -1,189 +0,0 @@ -/* - * 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 src/ir/adt.cc - * \brief AST nodes for Relay algebraic data types (ADTs). - */ -#include -#include - -namespace tvm { -namespace relay { - -PatternWildcard::PatternWildcard() { - ObjectPtr n = make_object(); - data_ = std::move(n); -} - -TVM_REGISTER_NODE_TYPE(PatternWildcardNode); - -TVM_REGISTER_GLOBAL("relay.ir.PatternWildcard").set_body_typed([]() { return PatternWildcard(); }); - -TVM_STATIC_IR_FUNCTOR(ReprPrinter, vtable) - .set_dispatch([](const ObjectRef& ref, ReprPrinter* p) { - p->stream << "PatternWildcardNode()"; - }); - -PatternVar::PatternVar(tvm::relay::Var var) { - ObjectPtr n = make_object(); - n->var = std::move(var); - data_ = std::move(n); -} - -TVM_REGISTER_NODE_TYPE(PatternVarNode); - -TVM_REGISTER_GLOBAL("relay.ir.PatternVar").set_body_typed([](tvm::relay::Var var) { - return PatternVar(var); -}); - -TVM_STATIC_IR_FUNCTOR(ReprPrinter, vtable) - .set_dispatch([](const ObjectRef& ref, ReprPrinter* p) { - auto* node = static_cast(ref.get()); - p->stream << "PatternVarNode(" << node->var << ")"; - }); - -PatternConstructor::PatternConstructor(Constructor constructor, tvm::Array patterns) { - ObjectPtr n = make_object(); - n->constructor = std::move(constructor); - n->patterns = std::move(patterns); - data_ = std::move(n); -} - -TVM_REGISTER_NODE_TYPE(PatternConstructorNode); - -TVM_REGISTER_GLOBAL("relay.ir.PatternConstructor") - .set_body_typed([](Constructor constructor, tvm::Array patterns) { - return PatternConstructor(constructor, patterns); - }); - -TVM_STATIC_IR_FUNCTOR(ReprPrinter, vtable) - .set_dispatch([](const ObjectRef& ref, ReprPrinter* p) { - auto* node = static_cast(ref.get()); - p->stream << "PatternConstructorNode(" << node->constructor << ", " << node->patterns << ")"; - }); - -PatternTuple::PatternTuple(tvm::Array patterns) { - ObjectPtr n = make_object(); - n->patterns = std::move(patterns); - data_ = std::move(n); -} - -TVM_REGISTER_NODE_TYPE(PatternTupleNode); - -TVM_REGISTER_GLOBAL("relay.ir.PatternTuple").set_body_typed([](tvm::Array patterns) { - return PatternTuple(patterns); -}); - -TVM_STATIC_IR_FUNCTOR(ReprPrinter, vtable) - .set_dispatch([](const ObjectRef& ref, ReprPrinter* p) { - auto* node = static_cast(ref.get()); - p->stream << "PatternTupleNode(" << node->patterns << ")"; - }); - -Clause::Clause(Pattern lhs, Expr rhs) { - ObjectPtr n = make_object(); - n->lhs = std::move(lhs); - n->rhs = std::move(rhs); - data_ = std::move(n); -} - -Clause WithFields(Clause clause, Optional opt_lhs, Optional opt_rhs) { - Pattern lhs = opt_lhs.value_or(clause->lhs); - Expr rhs = opt_rhs.value_or(clause->rhs); - - bool unchanged = lhs.same_as(clause->lhs) && rhs.same_as(clause->rhs); - - if (!unchanged) { - ClauseNode* cow_clause_node = clause.CopyOnWrite(); - cow_clause_node->lhs = lhs; - cow_clause_node->rhs = rhs; - } - return clause; -} - -TVM_REGISTER_NODE_TYPE(ClauseNode); - -TVM_REGISTER_GLOBAL("relay.ir.Clause").set_body_typed([](Pattern lhs, Expr rhs) { - return Clause(lhs, rhs); -}); - -TVM_STATIC_IR_FUNCTOR(ReprPrinter, vtable) - .set_dispatch([](const ObjectRef& ref, ReprPrinter* p) { - auto* node = static_cast(ref.get()); - p->stream << "ClauseNode(" << node->lhs << ", " << node->rhs << ")"; - }); - -Match::Match(Expr data, tvm::Array clauses, bool complete, Span span) { - ObjectPtr n = make_object(); - n->data = std::move(data); - n->clauses = std::move(clauses); - n->complete = complete; - n->span = std::move(span); - data_ = std::move(n); -} - -Match WithFields(Match match, Optional opt_data, Optional> opt_clauses, - Optional opt_complete, Optional opt_span) { - Expr data = opt_data.value_or(match->data); - Array clauses = opt_clauses.value_or(match->clauses); - Bool complete = opt_complete.value_or(Bool(match->complete)); - Span span = opt_span.value_or(match->span); - - bool unchanged = - data.same_as(match->data) && (complete == match->complete) && span.same_as(match->span); - - // Check that all clauses are unchanged - if (unchanged) { - bool all_clauses_unchanged = true; - if (clauses.size() == match->clauses.size()) { - for (size_t i = 0; i < clauses.size(); i++) { - all_clauses_unchanged &= clauses[i].same_as(match->clauses[i]); - } - } else { - all_clauses_unchanged = false; - } - unchanged &= all_clauses_unchanged; - } - if (!unchanged) { - MatchNode* cow_match_node = match.CopyOnWrite(); - cow_match_node->data = data; - cow_match_node->clauses = clauses; - cow_match_node->complete = complete; - cow_match_node->span = span; - } - return match; -} - -TVM_REGISTER_NODE_TYPE(MatchNode); - -TVM_REGISTER_GLOBAL("relay.ir.Match") - .set_body_typed([](Expr data, tvm::Array clauses, bool complete) { - return Match(data, clauses, complete); - }); - -TVM_STATIC_IR_FUNCTOR(ReprPrinter, vtable) - .set_dispatch([](const ObjectRef& ref, ReprPrinter* p) { - auto* node = static_cast(ref.get()); - p->stream << "MatchNode(" << node->data << ", " << node->clauses << ", " << node->complete - << ")"; - }); - -} // namespace relay -} // namespace tvm diff --git a/src/relay/ir/base.cc b/src/relay/ir/base.cc deleted file mode 100644 index deedd283c2ff..000000000000 --- a/src/relay/ir/base.cc +++ /dev/null @@ -1,43 +0,0 @@ -/* - * 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 base.cc - * \brief The core base types for Relay. - */ - -#include -#include -#include - -namespace tvm { -namespace relay { - -using namespace tvm::runtime; - -TVM_REGISTER_NODE_TYPE(IdNode); - -Id::Id(String name_hint) { - ObjectPtr n = make_object(); - n->name_hint = std::move(name_hint); - data_ = std::move(n); -} - -} // namespace relay -} // namespace tvm diff --git a/src/relay/ir/dataflow_matcher.cc b/src/relay/ir/dataflow_matcher.cc deleted file mode 100644 index 9d117adbbcaf..000000000000 --- a/src/relay/ir/dataflow_matcher.cc +++ /dev/null @@ -1,981 +0,0 @@ -/* - * 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 src/tvm/relay/dataflow_matcher.cc - * \brief The dataflow pattern matcher for Relay. - */ - -#include -#include -#include -#include -#include - -#include - -#include "dataflow_matcher_impl.h" - -namespace tvm { -namespace relay { - -// Pattern Matcher -bool DFPatternMatcher::Match(const DFPattern& pattern, const Expr& expr) { - VLOG(1) << "Match " << PrettyPrint(pattern) << " in:" << std::endl << PrettyPrint(expr); - memo_.clear(); - matched_nodes_.clear(); - return VisitDFPattern(pattern, expr); -} - -void DFPatternMatcher::ClearMap(size_t watermark) { - for (size_t i = watermark; i < matched_nodes_.size(); ++i) { - memo_.erase(matched_nodes_[i]); - } - matched_nodes_.erase(matched_nodes_.begin() + watermark, matched_nodes_.end()); -} - -bool DFPatternMatcher::VisitDFPattern(const DFPattern& pattern, const Expr& expr) { - if (memoize_ && memo_.count(pattern)) { - ICHECK_EQ(memo_[pattern].size(), 1); - return expr.same_as(memo_[pattern][0]); - } else { - auto watermark = matched_nodes_.size(); - auto out = DFPatternFunctor::VisitDFPattern(pattern, expr); - if (out) { - memo_[pattern].push_back(expr); - matched_nodes_.push_back(pattern); - VLOG(1) << "Matched " << PrettyPrint(pattern) << " at:" << std::endl << PrettyPrint(expr); - } else { - ClearMap(watermark); - } - return out; - } -} - -bool DFPatternMatcher::VisitDFPattern_(const AltPatternNode* op, const Expr& expr) { - return VisitDFPattern(op->left, expr) || VisitDFPattern(op->right, expr); -} - -bool MatchRetValue(const ObjectRef& lhs, const TVMRetValue& rhs) { - // Unwrapping arrays may find user-provided FFI types in the - // attributes (e.g. Defining pad_value as ((0,0), (0,0)) will result - // in runtime::Int. These need to be converted to compile-time IR - // types when encountered. - if (lhs->IsInstance() || - lhs->IsInstance() || - lhs->IsInstance()) { - TVMRetValue lhs_convert; - lhs_convert = lhs; - PrimExpr lhs_expr = lhs_convert; - return MatchRetValue(lhs_expr, rhs); - } - - // StructuralEqual doesn't check for conversions between FFI types - // and IR types, but the pattern-matcher should. Therefore, - // explicitly recurse into the array. - if (auto opt_lhs_array = lhs.as>()) { - if (Optional> opt_rhs_array = rhs) { - Array lhs_array = opt_lhs_array.value(); - Array rhs_array = opt_rhs_array.value(); - if (lhs_array.size() != rhs_array.size()) { - return false; - } - for (size_t i = 0; i < lhs_array.size(); i++) { - TVMRetValue rhs_item; - rhs_item = rhs_array[i]; - if (!MatchRetValue(lhs_array[i], rhs_item)) { - return false; - } - } - return true; - } else { - return false; - } - } - - switch (rhs.type_code()) { - case kDLInt: - if (auto* val = lhs.as()) { - return val->value == rhs.operator int64_t(); - } - break; - case kDLFloat: - if (auto* val = lhs.as()) { - return val->value == rhs.operator double(); - } - break; - case kTVMStr: - if (auto* val = lhs.as()) { - return val->value == rhs.operator std::string(); - } else if (auto* val = lhs.as()) { - return val->data == rhs.operator std::string(); - } - break; - case kTVMDataType: - if (auto* val = lhs.as()) { - return rhs.operator std::string() == val->value; - } else if (auto* val = lhs.as()) { - return rhs.operator std::string() == val->data; - } else { - ICHECK(false) << "PatternMatcher: Unsupported TVMDataType " << lhs; - } - break; - case kTVMObjectHandle: - if (rhs.IsObjectRef()) { - if (auto* val = lhs.as()) { - return rhs.operator String() == val->value; - } else if (auto* val = lhs.as()) { - return rhs.operator String() == val->data; - } - } else { - // Compare the objects for structural equality - static auto* structural_equal = runtime::Registry::Get("node.StructuralEqual"); - ICHECK(structural_equal) << "node.StructuralEqual is not registered."; - if ((*structural_equal)(lhs, GetRef(rhs.ptr()), false, true)) { - return true; - } - } - break; - default: - ICHECK(false) << "Unsupported type code in Pattern Node " << rhs.type_code(); - } - return false; -} - -bool DFPatternMatcher::VisitDFPattern_(const AttrPatternNode* attr_pattern, const Expr& expr) { - bool matches = VisitDFPattern(attr_pattern->pattern, expr); - if (!matches) { - return matches; - } - auto attributes = attr_pattern->attrs.as()->dict; - if (auto optional = expr.as()) { - Op op = optional.value(); - for (auto kv : attributes) { - auto attr_name = kv.first; - auto attr_value = kv.second; - if (Op::HasAttrMap(attr_name)) { - auto op_map = Op::GetAttrMap(attr_name); - if (op_map.count(op)) { - matches &= MatchRetValue(attr_value, op_map[op]); - } else { - matches = false; - } - } else { - matches = false; - } - } - } else if (auto* op = expr.as()) { - matches = true; - // TODO(mbrookhart): When OpNode Attrs move from TVMRetValue to the Object system, remove this - // and replace the whole thing with a Visitor-based approach - ReflectionVTable* reflection = ReflectionVTable::Global(); - auto attrs_node = const_cast(op->attrs.get()); - // attrs may be undefined on non-op calls so we check first - std::vector attr_names; - if (attrs_node) { - attr_names = reflection->ListAttrNames(attrs_node); - } - for (auto kv : attributes) { - std::string attr = kv.first; - if (matches && std::find(attr_names.begin(), attr_names.end(), attr) != attr_names.end()) { - matches &= MatchRetValue(kv.second, reflection->GetAttr(attrs_node, attr)); - } else { - matches = false; - break; - } - } - } else if (auto* op = expr.as()) { - matches = true; - for (auto kv : attributes) { - if (matches && op->attrs.defined() && op->attrs->dict.count(kv.first)) { - matches &= StructuralEqual()(kv.second, op->attrs->dict[kv.first]); - } else { - matches = false; - break; - } - } - } else { - matches = false; - } - return matches; -} - -Array reverse(const Array& args) { - Array new_args; - for (auto it = args.rbegin(); it != args.rend(); ++it) { - new_args.push_back(*it); - } - return new_args; -} - -bool DFPatternMatcher::VisitDFPattern_(const CallPatternNode* op, const Expr& expr) { - // utilities - auto get_op_node = [](const CallPatternNode* op) -> const tvm::OpNode* { - if (op) { - if (auto* expr_pattern = op->op.as()) { - return expr_pattern->expr.as(); - } - } - return nullptr; - }; - auto is_pattern_op = [&get_op_node](const CallPatternNode* op, std::string op_type) { - if (const auto* op_node = get_op_node(op)) { - if (op_node->name == op_type) { - return true; - } - } - return false; - }; - auto is_expr_op = [](const Expr& expr, std::string op_type) { - if (const auto* call_node = expr.as()) { - if (const auto* op_node = call_node->op.as()) { - if (op_node->name == op_type) { - return true; - } - } - } - return false; - }; - - // logic - auto watermark = matched_nodes_.size(); - if (const auto* call_node = expr.as()) { - auto matches_op = VisitDFPattern(op->op, call_node->op); - if (matches_op) { - auto watermark2 = matched_nodes_.size(); - - auto match_args = [this, &watermark2](const Array pattern_args, - const Array expr_args) { - bool matches = true; - size_t i = 0; - if (pattern_args.defined()) { - if (pattern_args.size() == expr_args.size()) { - while (matches && i < pattern_args.size()) { - matches &= VisitDFPattern(pattern_args[i], expr_args[i]); - ++i; - } - } else { - matches = false; - } - } - if (!matches) { - ClearMap(watermark2); - } - return matches; - }; - - // Standard case - if (match_args(op->args, call_node->args)) { - return true; - } - // Commutative Matching - if (const OpNode* op_node = get_op_node(op)) { - if ((op_node->name == "add") || (op_node->name == "multiply")) { - if (match_args(reverse(op->args), call_node->args)) { - return true; - } - } - } - } else { - ClearMap(watermark); - // associate divide/multiply - if (is_pattern_op(op, "divide")) { - if (const auto* arg_node = op->args[0].as()) { - if (is_pattern_op(arg_node, "multiply") && is_expr_op(expr, "multiply") && - (is_expr_op(call_node->args[0], "divide") || - is_expr_op(call_node->args[1], "divide"))) { - bool out = false; - for (size_t arg_id = 0; arg_id < 2; ++arg_id) { - auto div = CallPattern(op->op, {arg_node->args[arg_id], op->args[1]}); - auto mul = CallPattern(arg_node->op, {arg_node->args[(arg_id + 1) % 2], div}); - out = VisitDFPattern(mul, expr); - if (out) { - return true; - } else { - ClearMap(watermark); - } - } - return out; - } - } - } - if (is_pattern_op(op, "multiply")) { - // associate multiply/divide - for (size_t arg_id = 0; arg_id < 2; ++arg_id) { - if (auto* arg_node = op->args[arg_id].as()) { - if (is_pattern_op(arg_node, "divide") && is_expr_op(expr, "divide") && - (is_expr_op(call_node->args[0], "multiply") || - is_expr_op(call_node->args[1], "multiply"))) { - auto mul = CallPattern(op->op, {arg_node->args[0], op->args[(arg_id + 1) % 2]}); - auto div = CallPattern(arg_node->op, {mul, arg_node->args[1]}); - return VisitDFPattern(div, expr); - } - } - } - } - } - } - return false; -} - -// Recursively find the Dominator parent along all inputs paths. -bool DFPatternMatcher::MatchesPath(const DominatorPatternNode* op, const Expr& expr) { - // utilities - auto is_leaf_node = [](const Expr& expr) { - return expr.as() || expr.as(); - }; - - // logic - auto call_node = expr.as(); - auto index_node = expr_to_node(expr); - size_t arg_counter{0}; - for (auto node : index_node->inputs_) { - if (!(call_node && (node->ref() == call_node->op || is_leaf_node(node->ref())))) { - arg_counter += 1; - memoize_ = true; - if (!VisitDFPattern(op->parent, node->ref())) { - memoize_ = false; - if (!VisitDFPattern(op->path, node->ref())) { - return false; - } - if (!MatchesPath(op, node->ref())) { - return false; - } - } - } - } - if (!arg_counter) { - return false; - } - return true; -} - -// Iteratively ensure that the parent is dominated somewhere by the child or the path -bool DFPatternMatcher::DominatesParent(const DominatorPatternNode* op, const Expr& expr) { - std::stack stack; - std::unordered_set visited; - stack.push(expr); - while (!stack.empty()) { - Expr current = stack.top(); - stack.pop(); - for (auto node : expr_to_node(current)->dominator_children_) { - if (visited.count(node->node_ref_) == 0) { - if (VisitDFPattern(op->parent, node->ref())) { - return true; - } else { - stack.push(node->ref()); - } - visited.insert(node->node_ref_); - } - } - } - return false; -} - -bool DFPatternMatcher::VisitDFPattern_(const DominatorPatternNode* op, const Expr& expr) { - if (VisitDFPattern(op->child, expr)) { - bool matches_path = MatchesPath(op, expr); - memoize_ = true; - if (matches_path) { - return DominatesParent(op, expr); - } - } - return false; -} - -bool DFPatternMatcher::VisitDFPattern_(const ExprPatternNode* op, const Expr& expr) { - return StructuralEqual()(op->expr, expr); -} - -bool DFPatternMatcher::VisitDFPattern_(const FunctionPatternNode* op, const Expr& expr) { - bool matches = false; - if (const auto* func = expr.as()) { - matches = true; - if (op->params.defined()) { - size_t i = 0; - if (op->params.size() == func->params.size()) { - while (matches && i < op->params.size()) { - matches &= VisitDFPattern(op->params[i], func->params[i]); - ++i; - } - } else { - matches = false; - } - } - if (matches) { - matches &= VisitDFPattern(op->body, func->body); - } - } - return matches; -} - -bool DFPatternMatcher::VisitDFPattern_(const TupleGetItemPatternNode* op, const Expr& expr) { - bool matches = false; - if (const auto* tuple_get_item_node = expr.as()) { - matches = (op->index == -1 || op->index == tuple_get_item_node->index) && - VisitDFPattern(op->tuple, tuple_get_item_node->tuple); - } - return matches; -} - -bool DFPatternMatcher::VisitDFPattern_(const TuplePatternNode* op, const Expr& expr) { - bool matches = false; - if (const auto* tuple_node = expr.as()) { - matches = true; - if (op->fields.defined()) { - if (op->fields.size() == tuple_node->fields.size()) { - size_t i = 0; - while (matches && i < op->fields.size()) { - matches &= VisitDFPattern(op->fields[i], tuple_node->fields[i]); - ++i; - } - } else { - matches = false; - } - } - } - return matches; -} - -bool DFPatternMatcher::VisitDFPattern_(const IfPatternNode* op, const Expr& expr) { - if (const auto* if_node = expr.as()) { - auto cond = if_node->cond; - auto true_branch = if_node->true_branch; - auto false_branch = if_node->false_branch; - return VisitDFPattern(op->cond, cond) && VisitDFPattern(op->true_branch, true_branch) && - VisitDFPattern(op->false_branch, false_branch); - } - return false; -} - -bool DFPatternMatcher::VisitDFPattern_(const LetPatternNode* op, const Expr& expr) { - if (const auto* let_node = expr.as()) { - return VisitDFPattern(op->var, let_node->var) && VisitDFPattern(op->value, let_node->value) && - VisitDFPattern(op->body, let_node->body); - } - return false; -} - -Expr InferTypeWithModule(const Expr& expr, const IRModule& m) { - IRModule mod(m->functions, m->type_definitions, m->Imports()); - GlobalVarSupply global_var_supply = GlobalVarSupply(mod); - GlobalVar gvar = global_var_supply->FreshGlobal("_tmp", false); - BaseFunc func; - if (expr.as()) { - func = Downcast(expr); - } else { - func = relay::Function(relay::FreeVars(expr), expr, Type(), relay::FreeTypeVars(expr, mod)); - } - mod->Add(gvar, func); - mod = transform::InferType()(mod); - Expr ret; - if (expr.as()) { - ret = mod->Lookup(gvar); - } else { - ret = mod->Lookup(gvar).as()->body; - } - return ret; -} - -bool DFPatternMatcher::VisitDFPattern_(const TypePatternNode* op, const Expr& expr) { - auto expr_type = InferType(expr).as()->checked_type(); - return (StructuralEqual()(op->type, expr_type)) && VisitDFPattern(op->pattern, expr); -} - -bool DFPatternMatcher::VisitDFPattern_(const ShapePatternNode* op, const Expr& expr) { - auto expr_type = InferType(expr).as()->checked_type(); - if (const TensorTypeNode* tensor_type = expr_type.as()) { - return (StructuralEqual()(op->shape, tensor_type->shape)) && VisitDFPattern(op->pattern, expr); - } - return false; -} - -bool DFPatternMatcher::VisitDFPattern_(const DataTypePatternNode* op, const Expr& expr) { - auto expr_type = InferType(expr).as()->checked_type(); - if (const TensorTypeNode* tensor_type = expr_type.as()) { - return (StructuralEqual()(op->dtype, tensor_type->dtype)) && VisitDFPattern(op->pattern, expr); - } - return false; -} - -bool DFPatternMatcher::VisitDFPattern_(const VarPatternNode* op, const Expr& expr) { - bool matches = false; - if (const auto* var_node = expr.as()) { - matches = true; - if (op->name_hint() != "") { - matches &= op->name_hint() == var_node->name_hint(); - } - } - return matches; -} - -bool DFPatternMatcher::VisitDFPattern_(const ConstantPatternNode* op, const Expr& expr) { - return expr.as() != nullptr; -} - -bool DFPatternMatcher::VisitDFPattern_(const WildcardPatternNode* op, const Expr& expr) { - if (op->pattern) { - return VisitDFPattern(op->pattern.value(), expr); - } else { - return true; - } -} - -bool MatchPattern(DFPattern pattern, Expr expr) { - std::unique_ptr> expr_graph = CreateIndexedGraph(expr); - return DFPatternMatcher(expr_graph.get()).Match(pattern, expr); -} - -TVM_REGISTER_GLOBAL("relay.dataflow_pattern.match").set_body_typed(MatchPattern); - -/*! \brief Creates a new set of nodes based on Group inputs, used to create functions and perform - * group overlap analysis */ -class MatchExtractor : public ExprMutator { - public: - explicit MatchExtractor( - const std::unordered_map& inputs) - : inputs_(inputs) {} - const std::unordered_map& GetMemo() { - return this->memo_; - } - const std::string& GetName() { return name_; } - - protected: - Expr VisitExpr(const Expr& pre) override { - if (inputs_.count(pre)) { - return inputs_.at(pre); - } - return ExprMutator::VisitExpr(pre); - } - Expr VisitExpr_(const TupleNode* op) override { - auto out = ExprMutator::VisitExpr_(op); - name_ += "Tuple_"; - return out; - }; - Expr VisitExpr_(const FunctionNode* op) override { - auto out = ExprMutator::VisitExpr_(op); - name_ += "Function"; - return out; - }; - Expr VisitExpr_(const CallNode* call_node) override { - auto out = ExprMutator::VisitExpr_(call_node); - if (auto operation = call_node->op.as()) { - name_ += operation->name + "_"; - } else { - name_ += "Call_"; - } - return out; - }; - Expr VisitExpr_(const LetNode* op) override { - auto out = ExprMutator::VisitExpr_(op); - name_ += "Let_"; - return out; - }; - Expr VisitExpr_(const IfNode* op) override { - auto out = ExprMutator::VisitExpr_(op); - name_ += "If_"; - return out; - }; - Expr VisitExpr_(const TupleGetItemNode* op) override { - auto out = ExprMutator::VisitExpr_(op); - name_ += "TupleGetItem" + std::to_string(op->index) + "_"; - return out; - }; - Expr VisitExpr_(const MatchNode* op) override { - auto out = ExprMutator::VisitExpr_(op); - name_ += "Match_"; - return out; - }; - std::string name_; - const std::unordered_map inputs_; -}; - -/*! \brief Group expressions that match the pattern */ -const std::unordered_map& PatternGrouper::GroupMatches( - const DFPattern& pattern, const Expr& pre) { - groups_.clear(); - gid_assignments_.clear(); - - pattern_ = pattern; - pattern_graph_ = CreateIndexedGraph(pattern_); - std::unique_ptr> expr_graph = CreateIndexedGraph(pre); - DFPatternMatcher matcher(expr_graph.get()); - matcher_ = &matcher; - this->VisitExprs(); - return this->groups_; -} - -void PatternGrouper::VisitExprs() { - std::unordered_set pre_partitioned; - for (PostDfsIndex i = matcher_->size(); i != 0; --i) { - PostDfsIndex index = i - 1; - const auto current = matcher_->index_to_node(index)->ref(); - if (gid_assignments_.count(current) == 0) { // Don't visit nodes we've already grouped - if (auto op = current.as()) { - if (op->attrs.defined() && op->attrs->dict.count(attr::kPartitionedFromPattern) != 0) { - pre_partitioned.insert(current); - PostOrderVisit(op->body, - [&pre_partitioned](const Expr& expr) { pre_partitioned.insert(expr); }); - } - } - if (pre_partitioned.count(current) == 0 && matcher_->Match(pattern_, current)) { - CreateGroup(current); - } - } - } -} - -void PatternGrouper::CreateGroup(const Expr& expr) { - VLOG(1) << "Creating group for:" << std::endl << PrettyPrint(expr); - - int var_number = 0; - - auto node_map = matcher_->GetMemo(); - // Get fuzzy patterns - std::unordered_set fuzzy_matches; - for (PostDfsIndex index = 0; index < pattern_graph_->size(); ++index) { - auto node = pattern_graph_->index_to_node(index); - // Don't treat fuzzy Dominator patterns input variables for partition - if (auto op = node->ref().as()) { - for (auto fuzzy_op : {op->parent, op->path}) { - if (node_map.count(fuzzy_op)) { - for (auto match : node_map[fuzzy_op]) { - fuzzy_matches.insert(match); - } - } - } - } - // Don't treat Function params or body as input variables for partition - if (node->ref().as()) { - if (node_map.count(node->ref())) { - auto matches = node_map[node->ref()]; - for (auto match : matches) { - auto sub_graph = CreateIndexedGraph(match.as()->body); - for (PostDfsIndex sub_index = 0; sub_index < sub_graph->size(); ++sub_index) { - auto sub_node = sub_graph->index_to_node(sub_index); - fuzzy_matches.insert(sub_node->ref()); - } - } - } - } - } - - // Create input variables - Group group; - group.root_node = expr; - group.matched_nodes = node_map; - - std::unordered_map inputs; - Array params; - - for (PostDfsIndex index = 0; index < pattern_graph_->size(); ++index) { - auto node = pattern_graph_->index_to_node(index); - auto make_input = [&](const Expr& input) { - if (fuzzy_matches.count(input) == 0 && input.as() == nullptr && - input.as() == nullptr && !EmbedConst(input, node->ref())) { - // Avoid adding parameters repeatedly because multiple operatorss in the partition - // may use the same input. - if (inputs.find(input) != inputs.end()) { - return; - } - inputs[input] = - Var("FunctionVar_" + std::to_string(graph_number_) + "_" + std::to_string(var_number), - NullValue()); - group.args.push_back(input); - params.push_back(inputs[input]); - var_number++; - } - }; - auto tuple = node->ref().as(); - auto call = node->ref().as(); - if (tuple && !tuple->fields.defined()) { - if (node_map.count(node->ref())) { - auto matches = node_map[node->ref()]; - for (auto match : matches) { - for (auto input : match.as()->fields) { - make_input(input); - } - } - } - } else if (call && !call->args.defined()) { - if (node_map.count(node->ref())) { - auto matches = node_map[node->ref()]; - for (auto match : matches) { - for (auto input : match.as()->args) { - make_input(input); - } - } - } - } else if (node->inputs_.size() == 0) { - if (node_map.count(node->ref())) { - auto matches = node_map[node->ref()]; - for (auto match : matches) { - make_input(match); - } - } - } - } - - graph_number_++; - - // Extract a Function. Used in Partition directly, - // used to determine Group overlap in other passes - auto extractor = MatchExtractor(inputs); - auto body = extractor.Mutate(expr); - - group.function = Function(params, body, NullValue(), Array()); - VLOG(1) << "Candidate extracted function:" << std::endl << PrettyPrint(group.function); - group.name = extractor.GetName(); - // Check to make sure we aren't overlapping with another group or creating an invalid fusion - // The MatchExtractor will create a new graph by replacing nodes that match the inputs of the - // pattern with the input FunctionVar* Variables. The resulting memoization map will only - // contain nodes in the expression that matched the pattern. If a non-input node of the pattern - // (i.e., some piece of computation) overlaps with the nodes in a previous group, we'll have a - // situation where we try to rewrite the same node twice in the second rewriting or parition - // pass. This isn't valid, so we check for it here. We ignore Ops, functions, and constants - // because they exist more globally outside of the fusion. - // Similiarly, if interior nodes in a group are used outside of the group fusing to a single - // output would create an invalid graph tranformation, so we block the creation of such groups. - auto memo = extractor.GetMemo(); - for (auto kv : memo) { - VLOG(1) << "matched index " << matcher_->expr_to_node(kv.first)->index_; - } - - for (auto kv : memo) { - // Check to ensure that this node isn't an input or a global - if (inputs.count(kv.first) == 0 && kv.first.as() == nullptr && - kv.first.as() == nullptr && kv.first.as() == nullptr) { - if (gid_assignments_.count(kv.first) != 0) { - // check to see if the node is use in other groups - // Exit due to overlapping partitions - return; - } else if (kv.second != body) { - // if the node isn't the output of the group - auto node = matcher_->expr_to_node(kv.first); - for (auto* output : node->outputs_) { - if (memo.count(output->ref()) == 0) { - // A node inside the matched group contributes an output to nodes outside of the matched - // group... - auto root = matcher_->expr_to_node(expr); - if (!root->Dominates(output)) { - // ...and the outside dataflow does not come back to the root of the matched group. - // So reject the match since it would create a cycle. - VLOG(1) << "Rejecting group since would create a cycle with output " << output->index_ - << " for root " << root->index_ << " in graph:" << std::endl - << matcher_->expr_graph().ToString(); - return; - } - // else: We'll allow the output to be included in the matched group. - } - } - } - } - } - // Assign Group Ids - group.gid = ++gid_; - for (auto kv : extractor.GetMemo()) { - gid_assignments_[kv.first] = gid_; - } - - // Save Group - groups_[group.gid] = std::move(group); -} - -bool PatternGrouper::EmbedConst(const Expr& expr, const DFPattern pattern) { - bool embed = false; - if (expr.as()) { - if (pattern.as() != nullptr) { - embed = true; - } else if (auto expr_pat = pattern.as()) { - if (expr_pat->expr.as()) { - embed = true; - } - } else if (auto alt_pat = pattern.as()) { - if (matcher_->Match(alt_pat->left, expr)) { - embed = EmbedConst(expr, alt_pat->left); - } else { - embed = EmbedConst(expr, alt_pat->right); - } - } - } - return embed; -} - -// Rewrite - -DFPatternCallback::DFPatternCallback(DFPattern pattern, PackedFunc function, bool require_type, - bool rewrite_once) { - ObjectPtr n = make_object(); - n->pattern = std::move(pattern); - n->function = std::move(function); - n->require_type = require_type; - n->rewrite_once = rewrite_once; - data_ = std::move(n); -} - -TVM_REGISTER_NODE_TYPE(DFPatternCallbackNode); - -TVM_REGISTER_GLOBAL("relay.dataflow_pattern.DFPatternCallback") - .set_body_typed([](DFPattern pattern, PackedFunc function, bool require_type, - bool rewrite_once) { - return DFPatternCallback(pattern, function, require_type, rewrite_once); - }); - -Expr PatternRewriter::Rewrite(const Array& callbacks, const Expr& pre) { - VLOG_CONTEXT << "PatternRewriter"; - VLOG(1) << "rewriting:" << std::endl << PrettyPrint(pre); - auto post = pre; - auto last = post; - // rewrite the graph until it stops changing to make sure all rewrites are complete - int count = 0; - bool equal = true; - static auto* structural_equal = runtime::Registry::Get("node.StructuralEqual"); - ICHECK(structural_equal) << "node.StructuralEqual is not registered."; - // Keep track of callbacks that have finished rewriting - std::unordered_map done; - do { - last = post; - // We don't have to call InferType if previous pass has not modified anything - // We can just take previous typed state of the expression - bool types_invalidated = true; - for (auto callback : callbacks) { - if (!done[callback]) { - auto before = post; - auto post_typed = post; - callback_ = callback; - if (callback_->require_type && types_invalidated) { - post_typed = InferTypeWithModule(post, mod_); - } - auto grouper = PatternGrouper(); - groups_ = grouper.GroupMatches(callback_->pattern, post_typed); - gid_assignments_ = grouper.GetGIDAssignments(); - memo_.clear(); - VLOG(1) << "pre rewritten:" << std::endl << PrettyPrint(pre); - post = this->VisitExpr(post_typed); - VLOG(1) << "post rewritten:" << std::endl << PrettyPrint(post); - count++; - bool current_equal = (*structural_equal)(before, post, false, true); - if (callback_->require_type && current_equal) { - types_invalidated = false; - post = post_typed; - } else { - types_invalidated = true; - if (callback_->rewrite_once) { - done[callback] = true; - } - } - } - } - equal = (*structural_equal)(last, post, false, true); - } while (!equal && count < 100); - if (count >= 100) { - LOG(FATAL) << "Observed 100 rewrite passes, possible conflicting passes?"; - } - return post; -} - -Expr PatternRewriter::DispatchVisitExpr(const Expr& pre) { - auto post = MixedModeMutator::DispatchVisitExpr(pre); - if (gid_assignments_.count(pre) && pre == groups_[gid_assignments_[pre]].root_node) { - // Convert the pre-rewrite node map to a post-rewrite node map - auto group = groups_[gid_assignments_[pre]]; - std::unordered_map, ObjectPtrHash, ObjectPtrEqual> node_map; - for (auto kv : group.matched_nodes) { - Array tmp; - for (size_t i = 0; i < kv.second.size(); ++i) { - tmp.push_back(this->memo_[kv.second[i]]); - } - node_map.insert({kv.first, tmp}); - } - // run the user callback function - return callback_->function(pre, post, Map>(node_map)); - } - return post; -} - -Expr RewritePatterns(Array callbacks, Expr expr, IRModule mod) { - return PatternRewriter(mod).Rewrite(callbacks, expr); -} - -TVM_REGISTER_GLOBAL("relay.dataflow_pattern.rewrite").set_body_typed(RewritePatterns); - -/*! - * \brief PatternPartitioner replaces expressions that match a pattern with function call that - * perform the same computation but allow for further analysis and lowering. - * - * The class uses PatternGrouper to support the dominator pattern. - */ -class PatternPartitioner : protected MixedModeMutator { - public: - Expr Partition(const DFPattern& pattern, const Expr& pre, const Map& attrs, - PackedFunc check) { - if (pattern.as()) { - LOG(WARNING) << "Partioning a Function that isn't called doesn't make sense, skipping" - << pattern; - return pre; - } - auto grouper = PatternGrouper(); - groups_ = grouper.GroupMatches(pattern, pre); - gid_assignments_ = grouper.GetGIDAssignments(); - attrs_ = attrs; - check_ = check; - return this->VisitExpr(pre); - } - - protected: - Expr RewritePartition(const PatternGrouper::Group& group) { - Array args; - for (size_t i = 0; i < group.args.size(); ++i) { - args.push_back(memo_[group.args[i]]); - } - Function func = WithAttr(group.function, attr::kPartitionedFromPattern, String(group.name)); - if (!attrs_.empty()) { - for (auto kv : attrs_) { - func = WithAttr(std::move(func), kv.first, kv.second); - } - } - return Call(func, args); - } - - Expr DispatchVisitExpr(const Expr& pre) override { - auto post = MixedModeMutator::DispatchVisitExpr(pre); - if (gid_assignments_.count(pre) && pre == groups_[gid_assignments_[pre]].root_node && - static_cast(check_(pre))) { - post = RewritePartition(groups_[gid_assignments_[pre]]); - } - return post; - } - - Map attrs_; - std::unordered_map groups_; - std::unordered_map gid_assignments_; - PackedFunc check_; -}; - -Expr PartitionPattern(DFPattern pattern, Expr expr, Map attrs, - PackedFunc check) { - return PatternPartitioner().Partition(pattern, expr, attrs, check); -} - -TVM_REGISTER_GLOBAL("relay.dataflow_pattern.partition") - .set_body_typed([](DFPattern pattern, Expr expr, Map attrs, - PackedFunc check) { return PartitionPattern(pattern, expr, attrs, check); }); - -} // namespace relay -} // namespace tvm diff --git a/src/relay/ir/dataflow_matcher_impl.h b/src/relay/ir/dataflow_matcher_impl.h deleted file mode 100644 index a174d8e34eb7..000000000000 --- a/src/relay/ir/dataflow_matcher_impl.h +++ /dev/null @@ -1,178 +0,0 @@ -/* - * 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 src/tvm/relay/dataflow_matcher_impl.h - * \brief The auxiliary data structure for dataflow matcher. - */ -#ifndef TVM_RELAY_IR_DATAFLOW_MATCHER_IMPL_H_ -#define TVM_RELAY_IR_DATAFLOW_MATCHER_IMPL_H_ - -#include -#include -#include -#include - -#include -#include -#include -#include - -#include "indexed_graph.h" - -namespace tvm { -namespace relay { - -class DFPatternMatcher : public DFPatternFunctor { - public: - explicit DFPatternMatcher(const IndexedGraph* expr_graph) : expr_graph_(expr_graph) {} - bool Match(const DFPattern& pattern, const Expr& expr); - Map> GetMemo() { return Map>(memo_); } - - const IndexedGraph::Node* expr_to_node(const Expr& expr) const { - return expr_graph_->item_to_node(expr); - } - const IndexedGraph::Node* index_to_node(size_t index) const { - return expr_graph_->index_to_node(index); - } - size_t size() const { return expr_graph_->size(); } - const std::unordered_map, ObjectPtrHash, ObjectPtrEqual>& memo() const { - return memo_; - } - const IndexedGraph& expr_graph() const { return *expr_graph_; } - - protected: - bool VisitDFPattern(const DFPattern& pattern, const Expr& expr) override; - bool VisitDFPattern_(const AltPatternNode* op, const Expr& expr) override; - bool VisitDFPattern_(const AttrPatternNode* op, const Expr& expr) override; - bool VisitDFPattern_(const CallPatternNode* op, const Expr& expr) override; - bool VisitDFPattern_(const ConstantPatternNode* op, const Expr& expr) override; - bool VisitDFPattern_(const DataTypePatternNode* op, const Expr& expr) override; - bool VisitDFPattern_(const DominatorPatternNode* op, const Expr& expr) override; - bool VisitDFPattern_(const ExprPatternNode* op, const Expr& expr) override; - bool VisitDFPattern_(const FunctionPatternNode* op, const Expr& expr) override; - bool VisitDFPattern_(const IfPatternNode* op, const Expr& expr) override; - bool VisitDFPattern_(const LetPatternNode* op, const Expr& expr) override; - bool VisitDFPattern_(const ShapePatternNode* op, const Expr& expr) override; - bool VisitDFPattern_(const TupleGetItemPatternNode* op, const Expr& expr) override; - bool VisitDFPattern_(const TuplePatternNode* op, const Expr& expr) override; - bool VisitDFPattern_(const TypePatternNode* op, const Expr& expr) override; - bool VisitDFPattern_(const VarPatternNode* op, const Expr& expr) override; - bool VisitDFPattern_(const WildcardPatternNode* op, const Expr& expr) override; - - void ClearMap(size_t watermark); - bool MatchesPath(const DominatorPatternNode* op, const Expr& expr); - bool DominatesParent(const DominatorPatternNode* op, const Expr& expr); - - const IndexedGraph* expr_graph_; - std::unordered_map, ObjectPtrHash, ObjectPtrEqual> memo_; - std::vector matched_nodes_; - bool memoize_ = true; -}; - -/*! - * \brief PatternGrouper does pre-rewriting pattern matching and analysis - * - * This class creates a number of groups of matched expressions, ensures they don't overlap, and - * returns them to the caller for post-analysis rewriting. - * - * This is primarily needed to support the post-dominator analysis required for dominator pattern - * matching. - */ -class PatternGrouper { - public: - /*! \brief Internal Group class for storing analysis */ - struct Group { - Expr root_node; - int gid; - Map> matched_nodes; - std::string name; - Function function; - Array args; - }; - - /*! \brief Return the group assignments of expressions */ - inline const std::unordered_map& GetGIDAssignments() { - return gid_assignments_; - } - /*! \brief Group expressions that match the pattern */ - const std::unordered_map& GroupMatches(const DFPattern& pattern, const Expr& pre); - - protected: - /*! \brief Iteratively traverse the Expression in pre-order to find subgraphs - * - * If we traverse the graph in post-order, we can run into situtations where a small subgraph will - * match the pattern. Due to options like AltPattern, a larger subgraph with more nodes later in - * the graph may also match the pattern. With post-order traversal, we mark the smaller subgraph - * as matched and fail to catch the larger subgraph. This problem is fixed by using pre-order - * traversal. - */ - void VisitExprs(); - - /*! \brief Create a group based on a matched expression */ - void CreateGroup(const Expr& expr); - - /*! \brief EmbedConst implements rules for embedding constants into partitioned functions or - * lifting them into the function arguments. - * - * The rules depend on what pattern the ConstantNode matched. - * - * The basic rules are: - * If the constant matches ExprPattern(relay.const(*)) or a ConstantPattern(), embed the constant - * in the partitioned function. If the constant matched an AltPattern, recursively check the - * matched side of the pattern. For any other matching pattern (i.e, wildcard, VarPattern, etc), - * lift the constant into the arguments of the partitioned function. - */ - bool EmbedConst(const Expr& expr, const DFPattern pattern); - // Internal State - DFPattern pattern_; - std::unordered_map groups_; - std::unordered_map gid_assignments_; - DFPatternMatcher* matcher_ = nullptr; - std::unique_ptr> pattern_graph_; - int gid_ = 0; - int graph_number_ = 0; -}; - -/*! - * \brief PatternRewriter rewrites the expression by finding matches and allowing user callback - * function to rewrite those matches - * - * The class uses PatternGrouper to support the dominator pattern. - */ -class PatternRewriter : protected MixedModeMutator { - public: - explicit PatternRewriter(IRModule mod) : mod_(mod) {} - /*! \brief Rewrite can take a number of callbacks and will repeatedly rewrite the graph with the - * callbacks until it stops changing */ - virtual Expr Rewrite(const Array& callbacks, const Expr& pre); - - protected: - virtual Expr DispatchVisitExpr(const Expr& pre); - - IRModule mod_; - DFPatternCallback callback_; - std::unordered_map groups_; - std::unordered_map gid_assignments_; -}; - -} // namespace relay -} // namespace tvm - -#endif // TVM_RELAY_IR_DATAFLOW_MATCHER_IMPL_H_ diff --git a/src/relay/ir/dataflow_pattern.cc b/src/relay/ir/dataflow_pattern.cc deleted file mode 100644 index 637cb0665d38..000000000000 --- a/src/relay/ir/dataflow_pattern.cc +++ /dev/null @@ -1,558 +0,0 @@ -/* - * 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 src/tvm/relay/dataflow_pattern.cc - * \brief The dataflow pattern language for Relay. - */ -#include -#include - -namespace tvm { -namespace relay { - -DFPatternPrinter::FType& DFPatternPrinter::vtable() { - static FType inst; - return inst; -} - -String PrettyPrint(const DFPattern& pattern) { - std::stringstream string_stream{}; - string_stream << pattern; - return string_stream.str(); -} - -void DFPatternPrinter::Print(const ObjectRef& node) { - ICHECK(node.as()); - DFPattern pat = Downcast(node); - static const FType& f = vtable(); - string_stream.str(""); - if (!node.defined()) { - string_stream << "(nullptr)"; - } else if (memo_.find(pat) != memo_.end()) { - string_stream << "(invoke pattern id " << memo_[pat].first << ")"; - auxiliary_patterns.push_back(pat); - } else { - if (f.can_dispatch(node)) { - memo_.insert({pat, {memo_.size(), ""}}); - f(node, this); - memo_[pat].second = string_stream.str(); - } else { - // default value, output type key and addr. - string_stream << node->GetTypeKey() << "(" << node.get() << ")"; - } - } -} - -ExprPattern::ExprPattern(Expr expr) { - ObjectPtr n = make_object(); - n->expr = std::move(expr); - data_ = std::move(n); -} - -TVM_REGISTER_NODE_TYPE(ExprPatternNode); - -TVM_REGISTER_GLOBAL("relay.dataflow_pattern.ExprPattern").set_body_typed([](Expr e) { - return ExprPattern(e); -}); - -TVM_STATIC_IR_FUNCTOR(DFPatternPrinter, vtable) - .set_dispatch([](const ObjectRef& ref, DFPatternPrinter* p) { - ExprPattern pattern = Downcast(ref); - p->string_stream.str(""); - p->string_stream << pattern->expr; - }); - -VarPattern::VarPattern(String name_hint) { - ObjectPtr n = make_object(); - n->name = std::move(name_hint); - data_ = std::move(n); -} - -TVM_REGISTER_NODE_TYPE(VarPatternNode); - -TVM_REGISTER_GLOBAL("relay.dataflow_pattern.VarPattern").set_body_typed([](String name_hint) { - return VarPattern(name_hint); -}); - -TVM_STATIC_IR_FUNCTOR(DFPatternPrinter, vtable) - .set_dispatch([](const ObjectRef& ref, DFPatternPrinter* p) { - VarPattern pattern = Downcast(ref); - p->string_stream.str(""); - p->string_stream << "VarPattern(" << pattern->name_hint() << ")"; - }); - -TVM_REGISTER_NODE_TYPE(ConstantPatternNode); - -TVM_REGISTER_GLOBAL("relay.dataflow_pattern.ConstantPattern").set_body_typed([]() { - auto c = ConstantPattern(make_object()); - return c; -}); - -TVM_STATIC_IR_FUNCTOR(DFPatternPrinter, vtable) - .set_dispatch([](const ObjectRef& ref, DFPatternPrinter* p) { - p->string_stream.str(""); - p->string_stream << "ConstantPattern()"; - }); - -CallPattern::CallPattern(DFPattern op, Array args) { - ObjectPtr n = make_object(); - n->op = std::move(op); - n->args = std::move(args); - data_ = std::move(n); -} - -TVM_REGISTER_NODE_TYPE(CallPatternNode); - -TVM_REGISTER_GLOBAL("relay.dataflow_pattern.CallPattern") - .set_body_typed([](DFPattern op, Array args) { return CallPattern(op, args); }); - -TVM_STATIC_IR_FUNCTOR(DFPatternPrinter, vtable) - .set_dispatch([](const ObjectRef& ref, DFPatternPrinter* p) { - CallPattern pattern = Downcast(ref); - - p->Print(pattern->op); - std::string op_pattern_string{p->string_stream.str()}; - - std::vector args_pattern_string{}; - for (const DFPattern& arg : pattern->args) { - p->Print(arg); - args_pattern_string.push_back(p->string_stream.str()); - } - - p->string_stream.str(""); - p->string_stream << "(id " << p->memo_[pattern].first << "): "; - p->string_stream << "CallPatternNode(" << op_pattern_string << ", ["; - for (size_t i = 0; i < args_pattern_string.size(); ++i) { - if (i != 0) { - p->string_stream << ", "; - } - p->string_stream << args_pattern_string[i]; - } - p->string_stream << "])"; - }); - -FunctionPattern::FunctionPattern(Array params, DFPattern body) { - ObjectPtr n = make_object(); - n->params = std::move(params); - n->body = std::move(body); - data_ = std::move(n); -} -TVM_REGISTER_NODE_TYPE(FunctionPatternNode); - -TVM_REGISTER_GLOBAL("relay.dataflow_pattern.FunctionPattern") - .set_body_typed([](Array params, DFPattern body) { - return FunctionPattern(params, body); - }); - -TVM_STATIC_IR_FUNCTOR(DFPatternPrinter, vtable) - .set_dispatch([](const ObjectRef& ref, DFPatternPrinter* p) { - FunctionPattern pattern = Downcast(ref); - - std::vector params_pattern_string{}; - for (const DFPattern& param : pattern->params) { - p->Print(param); - params_pattern_string.push_back(p->string_stream.str()); - } - - p->Print(pattern->body); - std::string body_pattern_string{p->string_stream.str()}; - - p->string_stream.str(""); - p->string_stream << "(id " << p->memo_[pattern].first << "): "; - - p->string_stream << "FunctionPatternNode(["; - for (size_t i = 0; i < params_pattern_string.size(); ++i) { - if (i != 0) { - p->string_stream << ", "; - } - p->string_stream << params_pattern_string[i]; - } - p->string_stream << "]"; - p->string_stream << ", " << body_pattern_string << ")"; - }); - -LetPattern::LetPattern(DFPattern var, DFPattern value, DFPattern body) { - ObjectPtr n = make_object(); - n->var = std::move(var); - n->value = std::move(value); - n->body = std::move(body); - data_ = std::move(n); -} - -TVM_REGISTER_NODE_TYPE(LetPatternNode); - -TVM_REGISTER_GLOBAL("relay.dataflow_pattern.LetPattern") - .set_body_typed([](DFPattern var, DFPattern value, DFPattern body) { - return LetPattern(var, value, body); - }); - -TVM_STATIC_IR_FUNCTOR(DFPatternPrinter, vtable) - .set_dispatch([](const ObjectRef& ref, DFPatternPrinter* p) { - LetPattern pattern = Downcast(ref); - - p->Print(pattern->var); - std::string var_pattern_string{p->string_stream.str()}; - - p->Print(pattern->value); - std::string value_pattern_string{p->string_stream.str()}; - - p->Print(pattern->body); - std::string body_pattern_string{p->string_stream.str()}; - - p->string_stream.str(""); - p->string_stream << "(id " << p->memo_[pattern].first << "): "; - p->string_stream << "LetPatternNode(" << var_pattern_string << ", " << value_pattern_string - << ", " << body_pattern_string << ")"; - }); - -IfPattern::IfPattern(DFPattern cond, DFPattern true_branch, DFPattern false_branch) { - ObjectPtr n = make_object(); - n->cond = std::move(cond); - n->true_branch = std::move(true_branch); - n->false_branch = std::move(false_branch); - data_ = std::move(n); -} - -TVM_REGISTER_NODE_TYPE(IfPatternNode); - -TVM_REGISTER_GLOBAL("relay.dataflow_pattern.IfPattern") - .set_body_typed([](DFPattern cond, DFPattern true_branch, DFPattern false_branch) { - return IfPattern(cond, true_branch, false_branch); - }); - -TVM_STATIC_IR_FUNCTOR(DFPatternPrinter, vtable) - .set_dispatch([](const ObjectRef& ref, DFPatternPrinter* p) { - IfPattern pattern = Downcast(ref); - - p->Print(pattern->cond); - std::string cond_pattern_string{p->string_stream.str()}; - - p->Print(pattern->true_branch); - std::string true_branch_pattern_string{p->string_stream.str()}; - - p->Print(pattern->false_branch); - std::string false_branch_pattern_string{p->string_stream.str()}; - - p->string_stream.str(""); - p->string_stream << "(id " << p->memo_[pattern].first << "): "; - p->string_stream << "IfPattern(" << cond_pattern_string << ", " << true_branch_pattern_string - << ", " << false_branch_pattern_string << ")"; - }); - -TuplePattern::TuplePattern(tvm::Array fields) { - ObjectPtr n = make_object(); - n->fields = std::move(fields); - data_ = std::move(n); -} - -TVM_REGISTER_NODE_TYPE(TuplePatternNode); - -TVM_REGISTER_GLOBAL("relay.dataflow_pattern.TuplePattern") - .set_body_typed([](tvm::Array fields) { return TuplePattern(fields); }); - -TVM_STATIC_IR_FUNCTOR(DFPatternPrinter, vtable) - .set_dispatch([](const ObjectRef& ref, DFPatternPrinter* p) { - TuplePattern pattern = Downcast(ref); - - std::vector fields_pattern_string{}; - for (const DFPattern& field : pattern->fields) { - p->Print(field); - fields_pattern_string.push_back(p->string_stream.str()); - } - - p->string_stream.str(""); - p->string_stream << "(id " << p->memo_[pattern].first << "): "; - p->string_stream << "TuplePattern("; - p->string_stream << "["; - for (size_t i = 0; i < fields_pattern_string.size(); ++i) { - if (i != 0) { - p->string_stream << ", "; - } - p->string_stream << fields_pattern_string[i]; - } - p->string_stream << "]"; - p->string_stream << ")"; - }); - -TupleGetItemPattern::TupleGetItemPattern(DFPattern tuple, int index) { - ObjectPtr n = make_object(); - n->tuple = std::move(tuple); - n->index = index; - data_ = std::move(n); -} - -TVM_REGISTER_NODE_TYPE(TupleGetItemPatternNode); - -TVM_REGISTER_GLOBAL("relay.dataflow_pattern.TupleGetItemPattern") - .set_body_typed([](DFPattern tuple, int index) { return TupleGetItemPattern(tuple, index); }); - -TVM_STATIC_IR_FUNCTOR(DFPatternPrinter, vtable) - .set_dispatch([](const ObjectRef& ref, DFPatternPrinter* p) { - TupleGetItemPattern pattern = Downcast(ref); - - p->Print(pattern->tuple); - std::string tuple_pattern_string{p->string_stream.str()}; - - p->string_stream.str(""); - p->string_stream << "(id " << p->memo_[pattern].first << "): "; - p->string_stream << "TupleGetItemPatternNode("; - p->string_stream << tuple_pattern_string << ", " << pattern->index << ")"; - }); - -AltPattern::AltPattern(DFPattern left, DFPattern right) { - ObjectPtr n = make_object(); - n->left = std::move(left); - n->right = std::move(right); - data_ = std::move(n); -} - -TVM_REGISTER_NODE_TYPE(AltPatternNode); - -TVM_REGISTER_GLOBAL("relay.dataflow_pattern.AltPattern") - .set_body_typed([](DFPattern left, DFPattern right) { return AltPattern(left, right); }); - -TVM_STATIC_IR_FUNCTOR(DFPatternPrinter, vtable) - .set_dispatch([](const ObjectRef& ref, DFPatternPrinter* p) { - AltPattern pattern = Downcast(ref); - - p->Print(pattern->left); - std::string left_pattern_string{p->string_stream.str()}; - - p->Print(pattern->right); - std::string right_pattern_string{p->string_stream.str()}; - - p->string_stream.str(""); - p->string_stream << "(id " << p->memo_[pattern].first << "): "; - p->string_stream << "AltPattern(" << left_pattern_string << " | " << right_pattern_string - << ")"; - }); - -void WildcardPattern::redirect_to(DFPattern pat) const { - WildcardPatternNode* ptr = static_cast(get_mutable()); - ptr->pattern = pat; -} - -TVM_REGISTER_NODE_TYPE(WildcardPatternNode); - -TVM_REGISTER_GLOBAL("relay.dataflow_pattern.WildcardPattern_redirect_to") - .set_body_typed([](WildcardPattern wildcard, DFPattern pat) { - return wildcard.redirect_to(pat); - }); - -TVM_REGISTER_GLOBAL("relay.dataflow_pattern.WildcardPattern").set_body_typed([]() { - auto w = WildcardPattern(make_object()); - return w; -}); - -TVM_STATIC_IR_FUNCTOR(DFPatternPrinter, vtable) - .set_dispatch([](const ObjectRef& ref, DFPatternPrinter* p) { - p->string_stream.str(""); - p->string_stream << "*"; - }); - -TypePattern::TypePattern(DFPattern pattern, Type type) { - ObjectPtr n = make_object(); - n->pattern = std::move(pattern); - n->type = std::move(type); - data_ = std::move(n); -} - -TVM_REGISTER_NODE_TYPE(TypePatternNode); - -TVM_REGISTER_GLOBAL("relay.dataflow_pattern.TypePattern") - .set_body_typed([](DFPattern pattern, Type type) { return TypePattern(pattern, type); }); - -TVM_STATIC_IR_FUNCTOR(DFPatternPrinter, vtable) - .set_dispatch([](const ObjectRef& ref, DFPatternPrinter* p) { - TypePattern pattern = Downcast(ref); - - p->Print(pattern->pattern); - std::string pattern_pattern_string{p->string_stream.str()}; - - p->string_stream.str(""); - p->string_stream << "(id " << p->memo_[pattern].first << "): "; - p->string_stream << "TypePattern(" << pattern_pattern_string << " has type " << pattern->type - << ")"; - }); - -ShapePattern::ShapePattern(DFPattern pattern, Array shape) { - ObjectPtr n = make_object(); - n->pattern = std::move(pattern); - n->shape = std::move(shape); - data_ = std::move(n); -} - -TVM_REGISTER_NODE_TYPE(ShapePatternNode); - -TVM_REGISTER_GLOBAL("relay.dataflow_pattern.ShapePattern") - .set_body_typed([](DFPattern pattern, Array shape) { - return ShapePattern(pattern, shape); - }); - -TVM_STATIC_IR_FUNCTOR(DFPatternPrinter, vtable) - .set_dispatch([](const ObjectRef& ref, DFPatternPrinter* p) { - ShapePattern pattern = Downcast(ref); - - p->Print(pattern->pattern); - std::string pattern_pattern_string{p->string_stream.str()}; - - p->string_stream.str(""); - p->string_stream << "(id " << p->memo_[pattern].first << "): "; - }); - -DataTypePattern::DataTypePattern(DFPattern pattern, DataType dtype) { - ObjectPtr n = make_object(); - n->pattern = std::move(pattern); - n->dtype = std::move(dtype); - data_ = std::move(n); -} - -TVM_REGISTER_NODE_TYPE(DataTypePatternNode); - -TVM_REGISTER_GLOBAL("relay.dataflow_pattern.DataTypePattern") - .set_body_typed([](DFPattern pattern, DataType dtype) { - return DataTypePattern(pattern, dtype); - }); - -TVM_STATIC_IR_FUNCTOR(DFPatternPrinter, vtable) - .set_dispatch([](const ObjectRef& ref, DFPatternPrinter* p) { - DataTypePattern pattern = Downcast(ref); - - p->Print(pattern->pattern); - std::string pattern_pattern_string{p->string_stream.str()}; - - p->string_stream.str(""); - p->string_stream << "(id " << p->memo_[pattern].first << "): "; - p->string_stream << "DataTypePattern(" << pattern_pattern_string << " has dtype " - << pattern->dtype << ")"; - }); - -AttrPattern::AttrPattern(DFPattern pattern, DictAttrs attrs) { - ObjectPtr n = make_object(); - n->pattern = std::move(pattern); - n->attrs = std::move(attrs); - data_ = std::move(n); -} - -TVM_REGISTER_NODE_TYPE(AttrPatternNode); - -TVM_REGISTER_GLOBAL("relay.dataflow_pattern.AttrPattern") - .set_body_typed([](DFPattern pattern, DictAttrs attrs) { return AttrPattern(pattern, attrs); }); - -TVM_STATIC_IR_FUNCTOR(DFPatternPrinter, vtable) - .set_dispatch([](const ObjectRef& ref, DFPatternPrinter* p) { - AttrPattern pattern = Downcast(ref); - - p->Print(pattern->pattern); - std::string pattern_pattern_string{p->string_stream.str()}; - - p->string_stream.str(""); - p->string_stream << "(id " << p->memo_[pattern].first << "): "; - p->string_stream << "AttrPattern(" << pattern_pattern_string << " has attributes " - << pattern->attrs << ")"; - }); - -DominatorPattern::DominatorPattern(DFPattern parent, DFPattern path, DFPattern child) { - ObjectPtr n = make_object(); - n->parent = std::move(parent); - n->path = std::move(path); - - n->child = std::move(child); - data_ = std::move(n); -} - -TVM_REGISTER_NODE_TYPE(DominatorPatternNode); - -TVM_REGISTER_GLOBAL("relay.dataflow_pattern.DominatorPattern") - .set_body_typed([](DFPattern parent, DFPattern path, DFPattern child) { - return DominatorPattern(parent, path, child); - }); - -TVM_STATIC_IR_FUNCTOR(DFPatternPrinter, vtable) - .set_dispatch([](const ObjectRef& ref, DFPatternPrinter* p) { - DominatorPattern pattern = Downcast(ref); - - p->Print(pattern->parent); - std::string parent_pattern_string{p->string_stream.str()}; - - p->Print(pattern->path); - std::string path_pattern_string{p->string_stream.str()}; - - p->Print(pattern->child); - std::string child_pattern_string{p->string_stream.str()}; - - p->string_stream.str(""); - p->string_stream << "(id " << p->memo_[pattern].first << "): "; - p->string_stream << "DominatorPattern(" << parent_pattern_string << ", " - << path_pattern_string << ", " << child_pattern_string << ")"; - }); - -// Syntatic Sugar -DFPattern DFPattern::operator()(const std::vector& args) const { - return CallPattern(GetRef(this->get()), Array(args)); -} -DFPattern DFPattern::operator+(const DFPattern& other) const { - return IsOp("add")({GetRef(this->get()), other}); -} -DFPattern DFPattern::operator-(const DFPattern& other) const { - return IsOp("subtract")({GetRef(this->get()), other}); -} -DFPattern DFPattern::operator*(const DFPattern& other) const { - return IsOp("multiply")({GetRef(this->get()), other}); -} -DFPattern DFPattern::operator/(const DFPattern& other) const { - return IsOp("divide")({GetRef(this->get()), other}); -} -DFPattern DFPattern::operator||(const DFPattern& other) const { - return AltPattern(GetRef(this->get()), other); -} - -DFPattern DFPattern::Optional(const std::function& func) const { - DFPattern current = GetRef(this->get()); - return current || func(current); -} - -DFPattern DFPattern::HasAttr(const Map& attrs) const { - return AttrPattern(GetRef(this->get()), DictAttrs(attrs)); -} -DFPattern DFPattern::HasType(const Type& type) const { - return TypePattern(GetRef(this->get()), type); -} -DFPattern DFPattern::HasDtype(const DataType& dtype) const { - return DataTypePattern(GetRef(this->get()), dtype); -} -DFPattern DFPattern::HasDtype(const std::string& dtype) const { - return HasDtype(DataType(runtime::String2DLDataType(dtype))); -} -DFPattern DFPattern::HasShape(const Array shape) const { - return ShapePattern(GetRef(this->get()), shape); -} -DFPattern IsVar(const String& name) { return VarPattern(name); } -DFPattern IsConstant() { return ConstantPattern(make_object()); } -DFPattern IsWildcard() { return WildcardPattern(make_object()); } -DFPattern IsExpr(const Expr& expr) { return ExprPattern(expr); } -DFPattern IsOp(const String& op_name) { return IsExpr(Op::Get(op_name)); } -DFPattern IsTuple(const Array& fields) { return TuplePattern(fields); } -DFPattern IsTupleGetItem(const DFPattern tuple, int index) { - return TupleGetItemPattern(tuple, index); -} - -} // namespace relay -} // namespace tvm diff --git a/src/relay/ir/dataflow_pattern_functor.cc b/src/relay/ir/dataflow_pattern_functor.cc deleted file mode 100644 index 76b3fe068e45..000000000000 --- a/src/relay/ir/dataflow_pattern_functor.cc +++ /dev/null @@ -1,115 +0,0 @@ -/* - * 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 src/tvm/relay/dataflow_matcher.cc - * \brief The dataflow pattern matcher for Relay. - */ - -#include - -namespace tvm { -namespace relay { - -// DFPatternVisitor - -void DFPatternVisitor::VisitDFPattern(const DFPattern& pattern) { - if (this->visited_.count(pattern.get()) == 0) { - visited_.insert(pattern.get()); - DFPatternFunctor::VisitDFPattern(pattern); - } -} - -void DFPatternVisitor::VisitDFPattern_(const AltPatternNode* op) { - VisitDFPattern(op->left); - VisitDFPattern(op->right); -} - -void DFPatternVisitor::VisitDFPattern_(const AttrPatternNode* op) { VisitDFPattern(op->pattern); } - -void DFPatternVisitor::VisitDFPattern_(const CallPatternNode* op) { - VisitDFPattern(op->op); - if (op->args.defined()) { - for (auto arg : op->args) { - VisitDFPattern(arg); - } - } -} - -void DFPatternVisitor::VisitDFPattern_(const DataTypePatternNode* op) { - VisitDFPattern(op->pattern); -} - -void DFPatternVisitor::VisitDFPattern_(const DominatorPatternNode* op) { - VisitDFPattern(op->parent); - VisitDFPattern(op->path); - VisitDFPattern(op->child); -} - -void DFPatternVisitor::VisitDFPattern_(const ExprPatternNode* op) {} - -void DFPatternVisitor::VisitDFPattern_(const FunctionPatternNode* op) { - if (op->params.defined()) { - for (auto param : op->params) { - VisitDFPattern(param); - } - } - VisitDFPattern(op->body); -} - -void DFPatternVisitor::VisitDFPattern_(const ShapePatternNode* op) { VisitDFPattern(op->pattern); } - -void DFPatternVisitor::VisitDFPattern_(const TupleGetItemPatternNode* op) { - VisitDFPattern(op->tuple); -} - -void DFPatternVisitor::VisitDFPattern_(const TuplePatternNode* op) { - if (op->fields.defined()) { - for (auto field : op->fields) { - VisitDFPattern(field); - } - } -} - -void DFPatternVisitor::VisitDFPattern_(const IfPatternNode* op) { - VisitDFPattern(op->cond); - VisitDFPattern(op->true_branch); - VisitDFPattern(op->false_branch); -} - -void DFPatternVisitor::VisitDFPattern_(const LetPatternNode* op) { - VisitDFPattern(op->var); - VisitDFPattern(op->value); - VisitDFPattern(op->body); -} - -void DFPatternVisitor::VisitDFPattern_(const TypePatternNode* op) { VisitDFPattern(op->pattern); } - -void DFPatternVisitor::VisitDFPattern_(const VarPatternNode* op) {} - -void DFPatternVisitor::VisitDFPattern_(const ConstantPatternNode* op) {} - -void DFPatternVisitor::VisitDFPattern_(const WildcardPatternNode* op) { - if (op->pattern) { - VisitDFPattern(op->pattern.value()); - } -} - -} // namespace relay -} // namespace tvm diff --git a/src/relay/ir/error.cc b/src/relay/ir/error.cc deleted file mode 100644 index 940efd91aa52..000000000000 --- a/src/relay/ir/error.cc +++ /dev/null @@ -1,138 +0,0 @@ -/* - * 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 - -// clang-format off -#include -#include -#include -// clang-format on - -namespace tvm { -namespace relay { - -template -using NodeMap = std::unordered_map; - -void ErrorReporter::RenderErrors(const IRModule& module, bool use_color) { - // First we pick an error reporting strategy for each error. - // TODO(@jroesch): Spanned errors are currently not supported. - for (auto err : this->errors_) { - ICHECK(!err.span.defined()) << "attempting to use spanned errors, currently not supported"; - } - - NodeMap> error_maps; - - // Set control mode in order to produce colors; - if (use_color) { - rang::setControlMode(rang::control::Force); - } - - for (auto pair : this->node_to_gv_) { - auto node = pair.first; - auto global = Downcast(pair.second); - - auto has_errs = this->node_to_error_.find(node); - - ICHECK(has_errs != this->node_to_error_.end()); - - const auto& error_indices = has_errs->second; - - std::stringstream err_msg; - - err_msg << rang::fg::red; - err_msg << " "; - for (auto index : error_indices) { - err_msg << this->errors_[index].what() << "; "; - } - err_msg << rang::fg::reset; - - // Setup error map. - auto it = error_maps.find(global); - if (it != error_maps.end()) { - it->second.insert({node, err_msg.str()}); - } else { - error_maps.insert({global, {{node, err_msg.str()}}}); - } - } - - // Now we will construct the fully-annotated program to display to - // the user. - std::stringstream annotated_prog; - - // First we output a header for the errors. - annotated_prog << rang::style::bold << std::endl - << "Error(s) have occurred. The program has been annotated with them:" << std::endl - << std::endl - << rang::style::reset; - - // For each global function which contains errors, we will - // construct an annotated function. - for (auto pair : error_maps) { - auto global = pair.first; - auto err_map = pair.second; - auto func = module->Lookup(global); - - // We output the name of the function before displaying - // the annotated program. - annotated_prog << rang::style::bold << "In `" << global->name_hint << "`: " << std::endl - << rang::style::reset; - - // We then call into the Relay printer to generate the program. - // - // The annotation callback will annotate the error messages - // contained in the map. - annotated_prog << AsText(func, false, [&err_map](const ObjectRef& expr) { - auto it = err_map.find(expr); - if (it != err_map.end()) { - ICHECK_NE(it->second.size(), 0); - return it->second; - } else { - return std::string(""); - } - }); - } - - auto msg = annotated_prog.str(); - - if (use_color) { - rang::setControlMode(rang::control::Auto); - } - - // Finally we report the error, currently we do so to LOG(FATAL), - // it may be good to instead report it to std::cout. - LOG(FATAL) << annotated_prog.str() << std::endl; -} - -void ErrorReporter::ReportAt(const GlobalVar& global, const ObjectRef& node, - const CompileError& err) { - size_t index_to_insert = this->errors_.size(); - this->errors_.push_back(err); - auto it = this->node_to_error_.find(node); - if (it != this->node_to_error_.end()) { - it->second.push_back(index_to_insert); - } else { - this->node_to_error_.insert({node, {index_to_insert}}); - } - this->node_to_gv_.insert({node, global}); -} -} // namespace relay -} // namespace tvm diff --git a/src/relay/ir/expr.cc b/src/relay/ir/expr.cc deleted file mode 100644 index 062d9206cf92..000000000000 --- a/src/relay/ir/expr.cc +++ /dev/null @@ -1,710 +0,0 @@ -/* - * 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 src/relay/ir/expr.cc - * \brief The expression AST nodes of Relay. - */ -#include -#include -#include - -namespace tvm { - -GlobalVar WithFields(GlobalVar global_var, Optional opt_name_hint, Optional opt_type, - Optional opt_virtual_device, Optional opt_span) { - String name_hint = opt_name_hint.value_or(global_var->name_hint); - Type type = opt_type.value_or(global_var->checked_type()); - VirtualDevice virtual_device = opt_virtual_device.value_or(global_var->virtual_device()); - Span span = opt_span.value_or(global_var->span); - bool all_fields_unchanged = - name_hint.same_as(global_var->name_hint) && type.same_as(global_var->checked_type()) && - virtual_device.same_as(global_var->virtual_device()) && span.same_as(global_var->span); - if (!all_fields_unchanged) { - GlobalVarNode* cow_global_var_node = global_var.CopyOnWrite(); - cow_global_var_node->name_hint = name_hint; - cow_global_var_node->checked_type_ = type; - cow_global_var_node->virtual_device_ = virtual_device; - cow_global_var_node->span = span; - } - - return global_var; -} - -VirtualDevice RelayExprNode::virtual_device() const { - if (!this->virtual_device_.defined()) { - // virtual_device_ should always be defined, unless we imported this node from JSON using an old - // version of TVM, in which case we want to set it to the default, which is - // VirtualDevice::FullyUnconstrained(). - return VirtualDevice::FullyUnconstrained(); - } - return Downcast(this->virtual_device_); -} - -namespace relay { - -using tvm::ReprPrinter; -using namespace tvm::runtime; - -Constant::Constant(runtime::NDArray data, Span span) { - ObjectPtr n = make_object(); - n->data = std::move(data); - n->virtual_device_ = VirtualDevice::FullyUnconstrained(); - n->span = std::move(span); - data_ = std::move(n); -} - -TVM_REGISTER_NODE_TYPE(ConstantNode); - -TVM_REGISTER_GLOBAL("relay.ir.Constant").set_body_typed([](runtime::NDArray data, Span span) { - return Constant(data, span); -}); -TVM_REGISTER_GLOBAL("relay.ir.ConstantWithFields") - .set_body_typed([](Constant constant, Optional opt_data, - Optional opt_virtual_device, Optional opt_span) { - return WithFields(constant, opt_data, opt_virtual_device, opt_span); - }); - -TVM_STATIC_IR_FUNCTOR(ReprPrinter, vtable) - .set_dispatch([](const ObjectRef& ref, ReprPrinter* p) { - auto* node = static_cast(ref.get()); - const PackedFunc* fprint = Registry::Get("relay._constant_repr"); - ICHECK(fprint) << "unable to find printing function for constants"; - std::string data = (*fprint)(GetRef(node)); - p->stream << "Constant(" << data << ")"; - }); - -TensorType ConstantNode::tensor_type() const { - auto dtype = DataType(data->dtype); - Array shape; - for (int i = 0; i < data->ndim; i++) { - ICHECK_LE(data->shape[i], std::numeric_limits::max()); - ICHECK_GE(data->shape[i], std::numeric_limits::min()); - shape.push_back(tvm::IntImm(DataType::Int(32), data->shape[i])); - } - - return TensorType(shape, dtype); -} - -Constant WithFields(Constant constant, Optional opt_data, - Optional opt_virtual_device, Optional opt_span) { - runtime::NDArray data = opt_data.value_or(constant->data); - VirtualDevice virtual_device = opt_virtual_device.value_or(constant->virtual_device()); - Span span = opt_span.value_or(constant->span); - - bool all_fields_unchanged = data.same_as(constant->data) && - virtual_device.same_as(constant->virtual_device()) && - span.same_as(constant->span); - - if (!all_fields_unchanged) { - ConstantNode* cow_constant_node = constant.CopyOnWrite(); - cow_constant_node->data = data; - cow_constant_node->virtual_device_ = virtual_device; - cow_constant_node->span = span; - } - return constant; -} - -Tuple::Tuple(tvm::Array fields, Span span) { - ObjectPtr n = make_object(); - n->fields = std::move(fields); - n->virtual_device_ = VirtualDevice::FullyUnconstrained(); - n->span = std::move(span); - data_ = std::move(n); -} - -TVM_REGISTER_NODE_TYPE(TupleNode); - -TVM_REGISTER_GLOBAL("relay.ir.Tuple").set_body_typed([](tvm::Array fields, Span span) { - return Tuple(fields, span); -}); -TVM_REGISTER_GLOBAL("relay.ir.TupleWithFields") - .set_body_typed([](Tuple tuple, Optional> opt_fields, - Optional opt_virtual_device, Optional opt_span) { - return WithFields(tuple, opt_fields, opt_virtual_device, opt_span); - }); - -Tuple WithFields(Tuple tuple, Optional> opt_fields, - Optional opt_virtual_device, Optional opt_span) { - Array fields = opt_fields.value_or(tuple->fields); - VirtualDevice virtual_device = opt_virtual_device.value_or(tuple->virtual_device()); - Span span = opt_span.value_or(tuple->span); - - bool all_fields_unchanged = true; - if (fields.size() == tuple->fields.size()) { - for (size_t i = 0; i < fields.size(); i++) { - all_fields_unchanged &= fields[i].same_as(tuple->fields[i]); - } - } else { - all_fields_unchanged = false; - } - - all_fields_unchanged = all_fields_unchanged && virtual_device.same_as(tuple->virtual_device()) && - span.same_as(tuple->span); - if (!all_fields_unchanged) { - TupleNode* cow_tuple_node = tuple.CopyOnWrite(); - cow_tuple_node->fields = fields; - cow_tuple_node->virtual_device_ = virtual_device; - cow_tuple_node->span = span; - } - return tuple; -} - -TVM_STATIC_IR_FUNCTOR(ReprPrinter, vtable) - .set_dispatch([](const ObjectRef& ref, ReprPrinter* p) { - auto* node = static_cast(ref.get()); - p->stream << "Tuple(" << node->fields << ")"; - }); - -Var::Var(Id vid, Type type_annotation, Span span) { - ObjectPtr n = make_object(); - n->vid = std::move(vid); - n->type_annotation = std::move(type_annotation); - n->virtual_device_ = VirtualDevice::FullyUnconstrained(); - n->span = std::move(span); - data_ = std::move(n); -} - -/* static */ Var Var::GenSym(Type type_annotation, Span span) { - static size_t next_id = std::atomic(0); - std::ostringstream os; - os << "x_" << next_id++; - return Var(os.str(), std::move(type_annotation), std::move(span)); -} - -Var WithFields(Var var, Optional opt_vid, Optional opt_type_annotation, - Optional opt_virtual_device, Optional opt_span) { - Id vid = opt_vid.value_or(var->vid); - Type type_annotation = opt_type_annotation.value_or(var->type_annotation); - VirtualDevice virtual_device = opt_virtual_device.value_or(var->virtual_device()); - Span span = opt_span.value_or(var->span); - - bool unchanged = vid.same_as(var->vid) && type_annotation.same_as(var->type_annotation) && - virtual_device.same_as(var->virtual_device()) && span.same_as(var->span); - - if (!unchanged) { - VarNode* cow_var_node = var.CopyOnWrite(); - cow_var_node->vid = vid; - cow_var_node->type_annotation = type_annotation; - cow_var_node->virtual_device_ = virtual_device; - cow_var_node->span = span; - } - return var; -} - -TVM_REGISTER_NODE_TYPE(VarNode); - -TVM_REGISTER_GLOBAL("relay.ir.Var").set_body_typed([](String str, Type type_annotation, Span span) { - return Var(str, type_annotation, span); -}); -TVM_REGISTER_GLOBAL("relay.ir.VarWithFields") - .set_body_typed([](Var var, Optional opt_vid, Optional opt_type_annotation, - Optional opt_virtual_device, Optional opt_span) { - return WithFields(var, opt_vid, opt_type_annotation, opt_virtual_device, opt_span); - }); - -TVM_STATIC_IR_FUNCTOR(ReprPrinter, vtable) - .set_dispatch([](const ObjectRef& ref, ReprPrinter* p) { - auto* node = static_cast(ref.get()); - p->stream << "Var(" << node->name_hint(); - if (node->type_annotation.defined()) { - p->stream << ", ty="; - p->Print(node->type_annotation); - } - p->stream << ")"; - }); - -Call::Call(Expr op, Array args, Attrs attrs, Array type_args, Span span) { - ObjectPtr n = make_object(); - n->op = std::move(op); - n->args = std::move(args); - n->attrs = std::move(attrs); - n->type_args = std::move(type_args); - n->virtual_device_ = VirtualDevice::FullyUnconstrained(); - n->span = std::move(span); - data_ = std::move(n); -} - -Call WithFields(Call call, Optional opt_op, Optional> opt_args, - Optional opt_attrs, Optional> opt_type_args, - Optional opt_virtual_device, Optional opt_span) { - // Collect new values for fields. - Expr op = opt_op.value_or(call->op); - Array args = opt_args.value_or(call->args); - Attrs attrs = opt_attrs.value_or(call->attrs); - Array type_args = opt_type_args.value_or(call->type_args); - VirtualDevice virtual_device = opt_virtual_device.value_or(call->virtual_device()); - Span span = opt_span.value_or(call->span); - - // Check if anything changed. - bool unchanged = op.same_as(call->op) && attrs.same_as(call->attrs) && - virtual_device.same_as(call->virtual_device()) && span.same_as(call->span); - if (unchanged) { - if (args.size() == call->args.size()) { - for (size_t i = 0; i < args.size(); i++) { - unchanged &= args[i].same_as(call->args[i]); - } - } else { - unchanged = false; - } - } - if (unchanged) { - if (type_args.size() == call->type_args.size()) { - for (size_t i = 0; i < type_args.size(); i++) { - unchanged &= type_args[i].same_as(call->type_args[i]); - } - } else { - unchanged = false; - } - } - - if (!unchanged) { - // If call is only references, update it in place. Otherwise copy and update. - CallNode* cow_call_node = call.CopyOnWrite(); - cow_call_node->op = op; - cow_call_node->args = args; - cow_call_node->attrs = attrs; - cow_call_node->type_args = type_args; - cow_call_node->virtual_device_ = virtual_device; - cow_call_node->span = span; - } - return call; -} - -TVM_REGISTER_NODE_TYPE(CallNode); - -TVM_REGISTER_GLOBAL("relay.ir.Call") - .set_body_typed([](Expr op, Array args, Attrs attrs, Array type_args, Span span) { - return Call(op, args, attrs, type_args, span); - }); -TVM_REGISTER_GLOBAL("relay.ir.CallWithFields") - .set_body_typed([](Call call, Optional opt_op, Optional> opt_args, - Optional opt_attrs, Optional> opt_type_args, - Optional opt_virtual_device, Optional opt_span) { - return WithFields(call, opt_op, opt_args, opt_attrs, opt_type_args, opt_virtual_device, - opt_span); - }); - -TVM_STATIC_IR_FUNCTOR(ReprPrinter, vtable) - .set_dispatch([](const ObjectRef& ref, ReprPrinter* p) { - auto* node = static_cast(ref.get()); - p->stream << "CallNode(" << node->op << ", " << node->args << ", " << node->attrs << ", " - << node->type_args << ")"; - }); - -Let::Let(Var var, Expr value, Expr body, Span span) { - ObjectPtr n = make_object(); - n->var = std::move(var); - n->value = std::move(value); - n->body = std::move(body); - n->virtual_device_ = VirtualDevice::FullyUnconstrained(); - n->span = std::move(span); - data_ = std::move(n); -} - -Let WithFields(Let let, Optional opt_var, Optional opt_value, Optional opt_body, - Optional opt_virtual_device, Optional opt_span) { - Var var = opt_var.value_or(let->var); - Expr value = opt_value.value_or(let->value); - Expr body = opt_body.value_or(let->body); - VirtualDevice virtual_device = opt_virtual_device.value_or(let->virtual_device()); - Span span = opt_span.value_or(let->span); - - bool unchanged = var.same_as(let->var) && value.same_as(let->value) && body.same_as(let->body) && - virtual_device.same_as(let->virtual_device()) && span.same_as(let->span); - - if (!unchanged) { - LetNode* cow_let_node = let.CopyOnWrite(); - cow_let_node->var = var; - cow_let_node->value = value; - cow_let_node->body = body; - cow_let_node->virtual_device_ = virtual_device; - cow_let_node->span = span; - } - return let; -} - -TVM_REGISTER_NODE_TYPE(LetNode); - -TVM_REGISTER_GLOBAL("relay.ir.Let").set_body_typed([](Var var, Expr value, Expr body, Span span) { - return Let(var, value, body, span); -}); -TVM_REGISTER_GLOBAL("relay.ir.LetWithFields") - .set_body_typed([](Let let, Optional opt_var, Optional opt_value, - Optional opt_body, Optional opt_virtual_device, - Optional opt_span) { - return WithFields(let, opt_var, opt_value, opt_body, opt_virtual_device, opt_span); - }); - -TVM_STATIC_IR_FUNCTOR(ReprPrinter, vtable) - .set_dispatch([](const ObjectRef& ref, ReprPrinter* p) { - auto* node = static_cast(ref.get()); - p->stream << "LetNode(" << node->var << ", " << node->value << ", " << node->body << ")"; - }); - -If::If(Expr cond, Expr true_branch, Expr false_branch, Span span) { - ObjectPtr n = make_object(); - n->cond = std::move(cond); - n->true_branch = std::move(true_branch); - n->false_branch = std::move(false_branch); - n->virtual_device_ = VirtualDevice::FullyUnconstrained(); - n->span = std::move(span); - data_ = std::move(n); -} - -If WithFields(If if_expr, Optional opt_cond, Optional opt_true_branch, - Optional opt_false_branch, Optional opt_virtual_device, - Optional opt_span) { - Expr cond = opt_cond.value_or(if_expr->cond); - Expr true_branch = opt_true_branch.value_or(if_expr->true_branch); - Expr false_branch = opt_false_branch.value_or(if_expr->false_branch); - VirtualDevice virtual_device = opt_virtual_device.value_or(if_expr->virtual_device()); - Span span = opt_span.value_or(if_expr->span); - - bool unchanged = cond.same_as(if_expr->cond) && true_branch.same_as(if_expr->true_branch) && - false_branch.same_as(if_expr->false_branch) && - virtual_device.same_as(if_expr->virtual_device()) && span.same_as(if_expr->span); - - if (!unchanged) { - IfNode* cow_if_node = if_expr.CopyOnWrite(); - cow_if_node->cond = cond; - cow_if_node->true_branch = true_branch; - cow_if_node->false_branch = false_branch; - cow_if_node->virtual_device_ = virtual_device; - cow_if_node->span = span; - } - return if_expr; -} - -TVM_REGISTER_NODE_TYPE(IfNode); - -TVM_REGISTER_GLOBAL("relay.ir.If") - .set_body_typed([](Expr cond, Expr true_branch, Expr false_branch, Span span) { - return If(cond, true_branch, false_branch, span); - }); -TVM_REGISTER_GLOBAL("relay.ir.IfWithFields") - .set_body_typed([](If if_expr, Optional opt_cond, Optional opt_true_branch, - Optional opt_false_branch, Optional opt_virtual_device, - Optional opt_span) { - return WithFields(if_expr, opt_cond, opt_true_branch, opt_false_branch, opt_virtual_device, - opt_span); - }); - -TVM_STATIC_IR_FUNCTOR(ReprPrinter, vtable) - .set_dispatch([](const ObjectRef& ref, ReprPrinter* p) { - auto* node = static_cast(ref.get()); - p->stream << "IfNode(" << node->cond << ", " << node->true_branch << ", " - << node->false_branch << ")"; - }); - -TupleGetItem::TupleGetItem(Expr tuple, int index, Span span) { - ObjectPtr n = make_object(); - n->tuple = std::move(tuple); - n->index = index; - n->virtual_device_ = VirtualDevice::FullyUnconstrained(); - n->span = std::move(span); - data_ = std::move(n); -} - -TupleGetItem WithFields(TupleGetItem tuple_get_item, Optional opt_tuple, - Optional opt_index, Optional opt_virtual_device, - Optional opt_span) { - Expr tuple = opt_tuple.value_or(tuple_get_item->tuple); - Integer index = opt_index.value_or(tuple_get_item->index); - VirtualDevice virtual_device = opt_virtual_device.value_or(tuple->virtual_device()); - Span span = opt_span.value_or(tuple_get_item->span); - - bool unchanged = tuple.same_as(tuple_get_item->tuple) && (index == tuple_get_item->index) && - virtual_device.same_as(tuple_get_item->virtual_device()) && - span.same_as(tuple_get_item->span); - if (!unchanged) { - TupleGetItemNode* cow_tuple_get_item_node = tuple_get_item.CopyOnWrite(); - cow_tuple_get_item_node->tuple = tuple; - cow_tuple_get_item_node->index = index.IntValue(); - cow_tuple_get_item_node->span = span; - cow_tuple_get_item_node->virtual_device_ = virtual_device; - } - return tuple_get_item; -} - -TVM_REGISTER_NODE_TYPE(TupleGetItemNode); - -TVM_REGISTER_GLOBAL("relay.ir.TupleGetItem").set_body_typed([](Expr tuple, int index, Span span) { - return TupleGetItem(tuple, index, span); -}); -TVM_REGISTER_GLOBAL("relay.ir.TupleGetItemWithFields") - .set_body_typed([](TupleGetItem tuple_get_item, Optional opt_tuple, - Optional opt_index, Optional opt_virtual_device, - Optional opt_span) { - return WithFields(tuple_get_item, opt_tuple, opt_index, opt_virtual_device, opt_span); - }); - -TVM_STATIC_IR_FUNCTOR(ReprPrinter, vtable) - .set_dispatch([](const ObjectRef& ref, ReprPrinter* p) { - auto* node = static_cast(ref.get()); - p->stream << "TupleGetItemNode(" << node->tuple << ", " << node->index << ")"; - }); - -RefCreate::RefCreate(Expr value, Span span) { - ObjectPtr n = make_object(); - n->value = std::move(value); - n->virtual_device_ = VirtualDevice::FullyUnconstrained(); - n->span = std::move(span); - data_ = std::move(n); -} - -RefCreate WithFields(RefCreate ref_create, Optional opt_value, - Optional opt_virtual_device, Optional opt_span) { - Expr value = opt_value.value_or(ref_create->value); - VirtualDevice virtual_device = opt_virtual_device.value_or(ref_create->virtual_device()); - Span span = opt_span.value_or(ref_create->span); - - bool unchanged = value.same_as(ref_create->value) && - virtual_device.same_as(ref_create->virtual_device()) && - span.same_as(ref_create->span); - if (!unchanged) { - RefCreateNode* cow_ref_create_node = ref_create.CopyOnWrite(); - cow_ref_create_node->value = value; - cow_ref_create_node->virtual_device_ = virtual_device; - cow_ref_create_node->span = span; - } - return ref_create; -} - -TVM_REGISTER_NODE_TYPE(RefCreateNode); - -TVM_REGISTER_GLOBAL("relay.ir.RefCreate").set_body_typed([](Expr value, Span span) { - return RefCreate(value, span); -}); -TVM_REGISTER_GLOBAL("relay.ir.RefCreateWithFields") - .set_body_typed([](RefCreate ref_create, Optional opt_value, - Optional opt_virtual_device, Optional opt_span) { - return WithFields(ref_create, opt_value, opt_virtual_device, opt_span); - }); - -TVM_STATIC_IR_FUNCTOR(ReprPrinter, vtable) - .set_dispatch([](const ObjectRef& ref, ReprPrinter* p) { - auto* node = static_cast(ref.get()); - p->stream << "RefCreateNode(" << node->value << ")"; - }); - -RefRead::RefRead(Expr ref, Span span) { - ObjectPtr n = make_object(); - n->ref = std::move(ref); - n->virtual_device_ = VirtualDevice::FullyUnconstrained(); - n->span = std::move(span); - data_ = std::move(n); -} - -RefRead WithFields(RefRead ref_read, Optional opt_ref, - Optional opt_virtual_device, Optional opt_span) { - Expr ref = opt_ref.value_or(ref_read->ref); - VirtualDevice virtual_device = opt_virtual_device.value_or(ref_read->virtual_device()); - Span span = opt_span.value_or(ref_read->span); - - bool unchanged = ref.same_as(ref_read->ref) && - virtual_device.same_as(ref_read->virtual_device()) && - span.same_as(ref_read->span); - if (!unchanged) { - RefReadNode* cow_ref_read_node = ref_read.CopyOnWrite(); - cow_ref_read_node->ref = ref; - cow_ref_read_node->virtual_device_ = virtual_device; - cow_ref_read_node->span = span; - } - return ref_read; -} - -TVM_REGISTER_NODE_TYPE(RefReadNode); - -TVM_REGISTER_GLOBAL("relay.ir.RefRead").set_body_typed([](Expr ref, Span span) { - return RefRead(ref, span); -}); -TVM_REGISTER_GLOBAL("relay.ir.RefReadWithFields") - .set_body_typed([](RefRead ref_read, Optional opt_ref, - Optional opt_virtual_device, Optional opt_span) { - return WithFields(ref_read, opt_ref, opt_virtual_device, opt_span); - }); - -TVM_STATIC_IR_FUNCTOR(ReprPrinter, vtable) - .set_dispatch([](const ObjectRef& ref, ReprPrinter* p) { - auto* node = static_cast(ref.get()); - p->stream << "RefReadNode(" << node->ref << ")"; - }); - -RefWrite::RefWrite(Expr ref, Expr value, Span span) { - ObjectPtr n = make_object(); - n->ref = std::move(ref); - n->value = std::move(value); - n->virtual_device_ = VirtualDevice::FullyUnconstrained(); - n->span = std::move(span); - data_ = std::move(n); -} - -RefWrite WithFields(RefWrite ref_write, Optional opt_ref, Optional opt_value, - Optional opt_virtual_device, Optional opt_span) { - Expr ref = opt_ref.value_or(ref_write->ref); - Expr value = opt_value.value_or(ref_write->value); - VirtualDevice virtual_device = opt_virtual_device.value_or(ref_write->virtual_device()); - Span span = opt_span.value_or(ref_write->span); - - bool unchanged = ref.same_as(ref_write->ref) && value.same_as(ref_write->value) && - virtual_device.same_as(ref_write->virtual_device()) && - span.same_as(ref_write->span); - if (!unchanged) { - RefWriteNode* cow_ref_write_node = ref_write.CopyOnWrite(); - cow_ref_write_node->ref = ref; - cow_ref_write_node->value = value; - cow_ref_write_node->virtual_device_ = virtual_device; - cow_ref_write_node->span = span; - } - return ref_write; -} - -TVM_REGISTER_NODE_TYPE(RefWriteNode); - -TVM_REGISTER_GLOBAL("relay.ir.RefWrite").set_body_typed([](Expr ref, Expr value, Span span) { - return RefWrite(ref, value, span); -}); -TVM_REGISTER_GLOBAL("relay.ir.RefWriteWithFields") - .set_body_typed([](RefWrite ref_write, Optional opt_ref, Optional opt_value, - Optional opt_virtual_device, Optional opt_span) { - return WithFields(ref_write, opt_ref, opt_value, opt_virtual_device, opt_span); - }); - -TVM_STATIC_IR_FUNCTOR(ReprPrinter, vtable) - .set_dispatch([](const ObjectRef& ref, ReprPrinter* p) { - auto* node = static_cast(ref.get()); - p->stream << "RefWriteNode(" << node->ref << ", " << node->value << ")"; - }); - -TVM_REGISTER_GLOBAL("relay.ir.TempExprRealize").set_body_typed([](TempExpr temp) { - return temp->Realize(); -}); - -TVM_REGISTER_GLOBAL("relay.ir.Any").set_body_typed([]() { return Any(); }); - -/* - * Non-recursive traversal with dismantling unused call nodes, - * a derivative from ExpandDataflow method - */ -inline void Dismantle(const Expr& expr) { - std::stack> stack; - auto fpush_to_stack = [&stack](const Expr& expr) { - // do not visit nodes with more than 2 refs (one can be in stack) - if (expr.use_count() < 3) { - stack.push({expr, false}); - } - }; - fpush_to_stack(expr); - while (stack.size() > 0) { - const auto& node = stack.top().first; - if (stack.top().second) { - // dismantle node - // +1 ref in stack/deque; - if (node.use_count() < 3) { - if (auto* op = const_cast(node.as())) { - op->args = Array(); - } - if (auto* op = const_cast(node.as())) { - op->body = Expr(); - } - } - // eject - stack.pop(); - } else { - stack.top().second = true; - - // special handling - if (const auto* call_node = node.as()) { - // do not process args if used elsewhere - if (call_node->args.use_count() < 2) { - for (auto it = call_node->args.rbegin(); it != call_node->args.rend(); ++it) { - fpush_to_stack(*it); - } - } - } else if (const auto* tuple_node = node.as()) { - // do not process fields if used elsewhere - if (tuple_node->fields.use_count() < 2) { - for (auto it = tuple_node->fields.rbegin(); it != tuple_node->fields.rend(); ++it) { - fpush_to_stack(*it); - } - } - } else if (const auto* tuple_get_item_node = node.as()) { - // do not process tuple if used elsewhere - if (tuple_get_item_node->tuple.use_count() < 2) { - fpush_to_stack(tuple_get_item_node->tuple); - } - } else if (const auto* let_node = node.as()) { - // do not process let if used elsewhere - if (let_node->body.use_count() < 2) { - fpush_to_stack(let_node->body); - } - } - } - } -} - -/* - * Non-recursive destructor - */ -Call::~Call() { - // attempt to dismantle if referenced one or zero times - if (this->use_count() < 2) { - if (this->as() && this->as()->args.size()) { - Dismantle(*this); - } - } -} - -/* - * CallNode's deleter - */ -void CallNode::Deleter_(Object* ptr) { - auto p = reinterpret_cast(ptr); - // resore original deleter - p->deleter_ = p->saved_deleter_; - // create Call reference in order to invoke ~Call - auto c = GetRef(p); -} - -/* - * Non-recursive destructor - */ -Let::~Let() { - // attempt to dismantle if referenced one or zero times - if (this->use_count() < 2) { - if (this->as() && this->as()->body.defined()) { - Dismantle(*this); - } - } -} - -/* - * LetNode's deleter - */ -void LetNode::Deleter_(Object* ptr) { - auto p = reinterpret_cast(ptr); - // resore original deleter - p->deleter_ = p->saved_deleter_; - // create Let reference in order to invoke ~Let - auto c = GetRef(p); -} - -} // namespace relay -} // namespace tvm diff --git a/src/relay/ir/expr_functor.cc b/src/relay/ir/expr_functor.cc deleted file mode 100644 index 49ef3864aca8..000000000000 --- a/src/relay/ir/expr_functor.cc +++ /dev/null @@ -1,571 +0,0 @@ -/* - * 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 src/relay/expr_functor.cc - * \brief A wrapper around ExprFunctor which functionally updates the AST. - * - * ExprMutator uses memoization and self return in order to amortize - * the cost of using functional updates. - */ -#include -#include -#include -#include -#include - -#include - -#include "../op/annotation/annotation.h" -#include "../op/memory/on_device.h" - -namespace tvm { -namespace relay { -MixedModeVisitor::MixedModeVisitor(int visit_limit) { - ICHECK(visit_limit > 0) << "Dataflow visit limit must be greater than 0"; - ICHECK(visit_limit < 10) << "Dataflow visit limit must be less than 10"; - visit_limit_ = visit_limit; -} - -void MixedModeVisitor::VisitLeaf(const Expr& expr) { - if (visit_counter_[expr.get()] < visit_limit_) { - ExprFunctor::VisitExpr(expr); - } - visit_counter_[expr.get()]++; -} - -bool MixedModeVisitor::CheckVisited(const Expr& expr) { - if (visit_counter_[expr.get()] < visit_limit_) { - return false; - } else { - visit_counter_[expr.get()]++; - return true; - } -} - -void MixedModeVisitor::VisitExpr(const Expr& expr) { - auto fcheck_visited = [this](const Expr& expr) { return this->CheckVisited(expr); }; - auto fvisit_leaf = [this](const Expr& expr) { return this->VisitLeaf(expr); }; - if (visit_counter_[expr.get()] < visit_limit_) { - ExpandDataflow(expr, fcheck_visited, fvisit_leaf); - } -} - -// Overwrite the VisitExpr so we don't recurse for dataflow nodes -void MixedModeVisitor::VisitExpr_(const CallNode* op) {} - -// Overwrite the VisitExpr so we don't recurse for dataflow nodes -void MixedModeVisitor::VisitExpr_(const TupleNode* op) {} - -// Overwrite the VisitExpr so we don't recurse for dataflow nodes -void MixedModeVisitor::VisitExpr_(const TupleGetItemNode* op) {} - -void MixedModeMutator::VisitLeaf(const Expr& expr) { - if (!memo_.count(expr)) { - Expr ret = this->DispatchVisitExpr(expr); - memo_[expr] = ret; - } -} - -bool MixedModeMutator::CheckVisited(const Expr& expr) { - if (memo_.count(expr)) { - return true; - } else { - return false; - } -} - -Expr MixedModeMutator::DispatchVisitExpr(const Expr& expr) { return ExprMutator::VisitExpr(expr); } - -Expr MixedModeMutator::VisitExpr(const Expr& expr) { - auto fcheck_visited = [this](const Expr& expr) { return this->CheckVisited(expr); }; - auto fvisit_leaf = [this](const Expr& expr) { return this->VisitLeaf(expr); }; - if (memo_.count(expr)) { - return memo_[expr]; - } else { - ExpandDataflow(expr, fcheck_visited, fvisit_leaf); - return memo_[expr]; - } -} - -class PostOrderRewriter : public MixedModeMutator { - public: - explicit PostOrderRewriter(ExprRewriter* rewriter) : rewriter_(rewriter) {} - - Expr DispatchVisitExpr(const Expr& expr) final { - auto post = ExprFunctor::VisitExpr(expr); - return rewriter_->Rewrite(expr, post); - } - - using MixedModeMutator::VisitExpr_; - - Expr VisitExpr_(const LetNode* node) final { - auto pre_visit = [this](const LetNode* op) { - Expr var = this->Mutate(op->var); - Expr value = this->Mutate(op->value); - }; - auto post_visit = [this, node](const LetNode* op) { - Var var = Downcast(this->Mutate(op->var)); - Expr value = this->Mutate(op->value); - Expr body = this->Mutate(op->body); - Expr expr = GetRef(op); - Expr post; - if (var.same_as(op->var) && value.same_as(op->value) && body.same_as(op->body)) { - post = expr; - } else { - post = Let(var, value, body); - } - // avoid rewriting the first LetNode twice - if (op == node) { - this->memo_[expr] = post; - } else { - this->memo_[expr] = this->rewriter_->Rewrite(expr, post); - } - }; - ExpandANormalForm(node, pre_visit, post_visit); - return memo_[GetRef(node)]; - } - - protected: - ExprRewriter* rewriter_; -}; - -Expr PostOrderRewrite(const Expr& expr, ExprRewriter* rewriter) { - return PostOrderRewriter(rewriter).VisitExpr(expr); -} - -Expr ExprMutator::VisitExpr(const Expr& expr) { - auto it = this->memo_.find(expr); - if (it != this->memo_.end()) { - return it->second; - } else { - Expr new_expr = ExprFunctor::VisitExpr(expr); - memo_[expr] = new_expr; - return new_expr; - } -} - -Expr ExprMutator::VisitExpr_(const VarNode* var_node) { - Type type_annotation = var_node->type_annotation; - if (var_node->type_annotation.defined()) { - type_annotation = this->VisitType(var_node->type_annotation); - } - return WithFields(GetRef(var_node), var_node->vid, type_annotation); -} - -Expr ExprMutator::VisitExpr_(const ConstantNode* op) { return GetRef(op); } - -Expr ExprMutator::VisitExpr_(const GlobalVarNode* op) { return GetRef(op); } - -Expr ExprMutator::VisitExpr_(const OpNode* op) { return GetRef(op); } - -Expr ExprMutator::VisitExpr_(const TupleNode* tuple_node) { - tvm::Array fields; - fields.reserve(tuple_node->fields.size()); - - for (auto field : tuple_node->fields) { - auto new_field = this->Mutate(field); - fields.push_back(new_field); - } - return WithFields(GetRef(tuple_node), fields); -} - -Expr ExprMutator::VisitExpr_(const FunctionNode* func_node) { - tvm::Array ty_params; - - for (auto ty_param : func_node->type_params) { - TypeVar new_ty_param = Downcast(VisitType(ty_param)); - ty_params.push_back(new_ty_param); - } - - tvm::Array params; - for (auto param : func_node->params) { - Var new_param = Downcast(this->Mutate(param)); - params.push_back(new_param); - } - - auto ret_type = this->VisitType(func_node->ret_type); - auto body = this->Mutate(func_node->body); - - return WithFields(GetRef(func_node), params, body, ret_type, ty_params); -} - -Expr ExprMutator::VisitExpr_(const CallNode* call_node) { - auto new_op = this->Mutate(call_node->op); - - tvm::Array ty_args; - ty_args.reserve(call_node->type_args.size()); - - for (auto ty_arg : call_node->type_args) { - auto new_ty_arg = this->VisitType(ty_arg); - ty_args.push_back(new_ty_arg); - } - - tvm::Array call_args; - call_args.reserve(call_node->args.size()); - for (auto arg : call_node->args) { - auto new_arg = this->Mutate(arg); - call_args.push_back(new_arg); - } - - return WithFields(GetRef(call_node), new_op, call_args, {}, ty_args); -} - -Expr ExprMutator::VisitExpr_(const LetNode* let_node) { - Var var = Downcast(this->Mutate(let_node->var)); - auto value = this->Mutate(let_node->value); - auto body = this->Mutate(let_node->body); - - return WithFields(GetRef(let_node), var, value, body); -} - -Expr ExprMutator::VisitExpr_(const IfNode* if_node) { - auto cond = this->Mutate(if_node->cond); - auto true_b = this->Mutate(if_node->true_branch); - auto false_b = this->Mutate(if_node->false_branch); - - return WithFields(GetRef(if_node), cond, true_b, false_b); -} - -Expr ExprMutator::VisitExpr_(const TupleGetItemNode* get_item) { - Expr tuple = this->Mutate(get_item->tuple); - return WithFields(GetRef(get_item), tuple); -} - -Expr ExprMutator::VisitExpr_(const RefCreateNode* ref_create) { - Expr value = this->Mutate(ref_create->value); - return WithFields(GetRef(ref_create), value); -} - -Expr ExprMutator::VisitExpr_(const RefReadNode* ref_read) { - Expr ref = this->Mutate(ref_read->ref); - return WithFields(GetRef(ref_read), ref); -} - -Expr ExprMutator::VisitExpr_(const RefWriteNode* ref_write) { - Expr ref = this->Mutate(ref_write->ref); - Expr value = this->Mutate(ref_write->value); - return WithFields(GetRef(ref_write), ref, value); -} - -Expr ExprMutator::VisitExpr_(const ConstructorNode* c) { return GetRef(c); } - -Expr ExprMutator::VisitExpr_(const MatchNode* match_node) { - Array clauses; - for (const Clause& p : match_node->clauses) { - clauses.push_back(VisitClause(p)); - } - Expr data = Mutate(match_node->data); - - return WithFields(GetRef(match_node), data, clauses); -} - -Clause ExprMutator::VisitClause(const Clause& clause) { - Pattern lhs = VisitPattern(clause->lhs); - Expr rhs = Mutate(clause->rhs); - return WithFields(clause, lhs, rhs); -} - -Pattern ExprMutator::VisitPattern(const Pattern& p) { return p; } - -Type ExprMutator::VisitType(const Type& t) { return t; } - -void ExprVisitor::VisitExpr(const Expr& expr) { - auto it = visit_counter_.find(expr.get()); - if (it != visit_counter_.end()) { - ++it->second; - } else { - using TParent = ExprFunctor; - TParent::VisitExpr(expr); - visit_counter_.insert({expr.get(), 1}); - } -} - -void ExprVisitor::VisitExpr_(const VarNode* op) { - this->VisitSpan(op->span); - if (op->type_annotation.defined()) { - this->VisitType(op->type_annotation); - } -} - -void ExprVisitor::VisitExpr_(const GlobalVarNode* op) { this->VisitSpan(op->span); } - -void ExprVisitor::VisitExpr_(const ConstantNode* op) { this->VisitSpan(op->span); } - -void ExprVisitor::VisitExpr_(const TupleNode* op) { - this->VisitSpan(op->span); - for (auto field : op->fields) { - this->VisitExpr(field); - } -} - -void ExprVisitor::VisitExpr_(const FunctionNode* op) { - this->VisitSpan(op->span); - for (auto param : op->params) { - this->VisitExpr(param); - } - - this->VisitExpr(op->body); -} - -void ExprVisitor::VisitExpr_(const CallNode* op) { - this->VisitSpan(op->span); - this->VisitExpr(op->op); - - for (auto ty_arg : op->type_args) { - this->VisitType(ty_arg); - } - - for (auto arg : op->args) { - this->VisitExpr(arg); - } -} - -void ExprVisitor::VisitExpr_(const LetNode* op) { - this->VisitSpan(op->span); - this->VisitExpr(op->value); - this->VisitExpr(op->var); - this->VisitExpr(op->body); -} - -void ExprVisitor::VisitExpr_(const IfNode* op) { - this->VisitSpan(op->span); - this->VisitExpr(op->cond); - this->VisitExpr(op->true_branch); - this->VisitExpr(op->false_branch); -} - -void ExprVisitor::VisitExpr_(const OpNode* op) { return; } - -void ExprVisitor::VisitExpr_(const TupleGetItemNode* op) { - this->VisitSpan(op->span); - this->VisitExpr(op->tuple); -} - -void ExprVisitor::VisitExpr_(const RefCreateNode* op) { - this->VisitSpan(op->span); - this->VisitExpr(op->value); -} - -void ExprVisitor::VisitExpr_(const RefReadNode* op) { - this->VisitSpan(op->span); - this->VisitExpr(op->ref); -} - -void ExprVisitor::VisitExpr_(const RefWriteNode* op) { - this->VisitSpan(op->span); - this->VisitExpr(op->ref); - this->VisitExpr(op->value); -} - -void ExprVisitor::VisitExpr_(const ConstructorNode* op) { - // TODO(@jroesch): visit spans - for (const Type& t : op->inputs) { - this->VisitType(t); - } - this->VisitType(op->belong_to); -} - -void ExprVisitor::VisitExpr_(const MatchNode* op) { - this->VisitSpan(op->span); - this->VisitExpr(op->data); - for (const Clause& c : op->clauses) { - this->VisitClause(c); - } -} - -void ExprVisitor::VisitClause(const Clause& op) { - // TODO(@jroesch): visit spans - this->VisitPattern(op->lhs); - this->VisitExpr(op->rhs); -} - -void ExprVisitor::VisitPattern(const Pattern& p) { return; } - -void ExprVisitor::VisitType(const Type& t) { return; } - -void ExprVisitor::VisitSpan(const Span& span) { return; } - -// visitor to implement apply -class ExprApplyVisit : public ExprVisitor { - public: - explicit ExprApplyVisit(std::function f) : f_(f) {} - - void VisitExpr(const Expr& e) final { - if (visited_.count(e.get()) != 0) return; - visited_.insert(e.get()); - ExprVisitor::VisitExpr(e); - f_(e); - } - - private: - std::function f_; - std::unordered_set visited_; -}; - -void PostOrderVisit(const Expr& e, std::function fvisit) { - ExprApplyVisit(fvisit).VisitExpr(e); -} - -TVM_REGISTER_GLOBAL("relay.analysis.post_order_visit").set_body_typed([](Expr expr, PackedFunc f) { - PostOrderVisit(expr, [f](const Expr& n) { f(n); }); -}); - -// Implement bind. -class ExprBinder : public MixedModeMutator, PatternMutator { - public: - explicit ExprBinder(const tvm::Map& args_map) : args_map_(args_map) {} - - using MixedModeMutator::VisitExpr_; - - Expr VisitExpr_(const LetNode* op) final { - ICHECK(!args_map_.count(op->var)) << "Cannot bind an internel variable in let"; - return ExprMutator::VisitExpr_(op); - } - - Expr VisitExpr_(const FunctionNode* op) final { - for (Var param : op->params) { - ICHECK(!args_map_.count(param)) << "Cannnot bind an internal function parameter"; - } - return ExprMutator::VisitExpr_(op); - } - - Expr VisitExpr_(const VarNode* op) final { - auto id = GetRef(op); - auto it = args_map_.find(id); - if (it != args_map_.end()) { - return (*it).second; - } else { - return std::move(id); - } - } - - Pattern VisitPattern(const Pattern& p) final { return PatternMutator::VisitPattern(p); } - - Clause VisitClause(const Clause& clause) final { - Pattern lhs = VisitPattern(clause->lhs); - return WithFields(clause, lhs, VisitExpr(clause->rhs)); - } - - Var VisitVar(const Var& v) final { - ICHECK(!args_map_.count(v)) << "Cannnot bind an internal pattern variable"; - return v; - } - - private: - const tvm::Map& args_map_; -}; - -// This function should be called SubstAndBind, since it assumes any variables introduced -// in the substitution right hand side should be implicitly bound in the function. -Expr Bind(const Expr& expr, const tvm::Map& args_map) { - if (const FunctionNode* func = expr.as()) { - Expr new_body = ExprBinder(args_map).VisitExpr(func->body); - Array new_params; - for (size_t i = 0; i < func->params.size(); ++i) { - if (!args_map.count(func->params[i])) { - new_params.push_back(func->params[i]); - } - } - if (new_body.same_as(func->body) && new_params.size() == func->params.size()) { - return expr; - } - - auto ret = - Function(new_params, new_body, func->ret_type, func->type_params, func->attrs, func->span); - ret->virtual_device_ = func->virtual_device(); - - std::unordered_set set; - for (const auto& v : FreeVars(expr)) { - set.insert(v); - } - for (const auto& v : FreeVars(ret)) { - if (set.count(v) == 0) { - new_params.push_back(v); - } - } - - ret = - Function(new_params, new_body, func->ret_type, func->type_params, func->attrs, func->span); - ret->virtual_device_ = func->virtual_device(); - - VLOG(4) << "Expr:\n" << expr; - VLOG(4) << "Ret:\n" << ret; - - ICHECK_EQ(FreeVars(expr).size(), FreeVars(ret).size()); - return std::move(ret); - } else { - return ExprBinder(args_map).VisitExpr(expr); - } -} - -TVM_REGISTER_GLOBAL("relay.ir.Bind").set_body([](TVMArgs args, TVMRetValue* ret) { - ObjectRef input = args[0]; - if (input->IsInstance()) { - *ret = Bind(Downcast(input), args[1]); - } else { - ICHECK(input->IsInstance()); - *ret = Bind(Downcast(input), args[1]); - } -}); - -Function SubstituteBoundVars(const Function& func, const tvm::Map& args_map) { - Expr new_body = ExprBinder(args_map).VisitExpr(func->body); - Array new_params; - for (size_t i = 0; i < func->params.size(); i++) { - if (!args_map.count(func->params[i])) { - new_params.push_back(func->params[i]); - } else { - if (auto var = args_map[func->params[i]].as()) { - new_params.push_back(var.value()); - } else { - ICHECK(false) << "Expected all values in args_map to be vars, but found " - << args_map[func->params[i]]->GetTypeKey(); - } - } - } - auto ret = - Function(new_params, new_body, func->ret_type, func->type_params, func->attrs, func->span); - ret->virtual_device_ = func->virtual_device(); - return ret; -} - -void ExpandANormalForm(const LetNode* op, std::function pre_visit, - std::function post_visit) { - std::stack stack; - stack.push(op); - bool is_anormal = true; - while (is_anormal) { - const LetNode* current_op = stack.top(); - pre_visit(current_op); - if (const LetNode* new_op = current_op->body.as()) { - stack.push(new_op); - } else { - is_anormal = false; - } - } - while (stack.size()) { - const LetNode* current_op = stack.top(); - stack.pop(); - post_visit(current_op); - } -} - -} // namespace relay -} // namespace tvm diff --git a/src/relay/ir/function.cc b/src/relay/ir/function.cc deleted file mode 100644 index b5414b27cf22..000000000000 --- a/src/relay/ir/function.cc +++ /dev/null @@ -1,307 +0,0 @@ -/* - * 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 src/relay/ir/function.cc - * \brief Function in relay. - */ -#include -#include -#include -#include -#include - -namespace tvm { -namespace relay { - -Function::Function(tvm::Array params, Expr body, Type ret_type, - tvm::Array type_params, DictAttrs attrs, Span span) { - CHECK(attrs.defined()); - - ObjectPtr n = make_object(); - ICHECK(params.defined()); - ICHECK(type_params.defined()); - n->params = std::move(params); - n->body = std::move(body); - n->ret_type = std::move(ret_type); - n->type_params = std::move(type_params); - n->attrs = std::move(attrs); - n->virtual_device_ = VirtualDevice::FullyUnconstrained(); - n->span = std::move(span); - data_ = std::move(n); -} - -Function WithFields(Function function, Optional> opt_params, Optional opt_body, - Optional opt_ret_type, Optional> opt_ty_params, - Optional opt_attrs, Optional opt_virtual_device, - Optional opt_span) { - Array params = opt_params.value_or(function->params); - Expr body = opt_body.value_or(function->body); - Type ret_type = opt_ret_type.value_or(function->ret_type); - Array ty_params = opt_ty_params.value_or(function->type_params); - DictAttrs attrs = opt_attrs.value_or(function->attrs); - VirtualDevice virtual_device = opt_virtual_device.value_or(function->virtual_device()); - Span span = opt_span.value_or(function->span); - - bool unchanged = body.same_as(function->body) && ret_type.same_as(function->ret_type) && - attrs.same_as(function->attrs) && - virtual_device.same_as(function->virtual_device()) && - span.same_as(function->span); - - // Check that all the type params are unchanged - if (unchanged) { - bool all_ty_params_unchanged = true; - if (ty_params.size() == function->type_params.size()) { - for (size_t i = 0; i < ty_params.size(); i++) { - all_ty_params_unchanged &= ty_params[i].same_as(function->type_params[i]); - } - } else { - all_ty_params_unchanged = false; - } - unchanged &= all_ty_params_unchanged; - } - - // Check that all the params are unchanged - if (unchanged) { - bool all_params_unchanged = true; - if (params.size() == function->params.size()) { - for (size_t i = 0; i < params.size(); i++) { - all_params_unchanged &= params[i].same_as(function->params[i]); - } - } else { - all_params_unchanged = false; - } - unchanged &= all_params_unchanged; - } - - if (!unchanged) { - FunctionNode* cow_function_node = function.CopyOnWrite(); - cow_function_node->params = params; - cow_function_node->body = body; - cow_function_node->ret_type = ret_type; - cow_function_node->type_params = ty_params; - cow_function_node->attrs = attrs; - cow_function_node->virtual_device_ = virtual_device; - cow_function_node->span = span; - } - return function; -} - -FuncType FunctionNode::func_type_annotation() const { - Array param_types; - for (auto param : this->params) { - Type param_type = - (param->type_annotation.defined()) ? param->type_annotation : IncompleteType(Kind::kType); - param_types.push_back(param_type); - } - - Type ret_type = (this->ret_type.defined()) ? this->ret_type : IncompleteType(Kind::kType); - return FuncType(param_types, ret_type, this->type_params, {}); -} - -const FunctionNode* AsOptimizableFunctionNode(const BaseFunc& base_func) { - if (const auto* function_node = base_func.as()) { - if (!function_node->GetAttr(attr::kCompiler).defined() && - !function_node->HasNonzeroAttr(attr::kExtern) && - !function_node->HasNonzeroAttr(attr::kSkipOptimization)) { - return function_node; - } - } - return nullptr; -} - -TVM_REGISTER_GLOBAL("relay.ir.PrintRelayModule") - .set_body_typed([](IRModule mod) -> Optional { - for (const auto& it : mod->functions) { - if (it.second->IsInstance()) { - return PrettyPrint(mod); - } - } - return NullOpt; - }); - -TVM_REGISTER_GLOBAL("relay.ir.PrintIR") - .set_body_typed([](IRModule mod, String header, bool show_metadata) -> bool { - for (const auto& it : mod->functions) { - if (it.second->IsInstance()) { - LOG(INFO) << "PrintIR(" << header << "):\n" << AsText(mod, show_metadata); - return true; - } - } - return false; - }); - -TVM_REGISTER_GLOBAL("relay.ir.WarnIfMalformed") - .set_body_typed([](const IRModule& mod, const BaseFunc& base_func) -> void { - if (auto relay_func = base_func.as()) { - Function func = Downcast(relay::DeDup(relay_func.value())); - // Type check the item before we add it to the module. - auto fv = relay::FreeVars(func); - auto ftv = relay::FreeTypeVars(func, mod); - // TODO(@jroesch): refactor to use diagnostic context - ICHECK_EQ(fv.size(), 0) << "Function:" << std::endl - << PrettyPrint(func) << std::endl - << "contains free variables: " << fv; - ICHECK_EQ(ftv.size(), 0) << "Function:" << std::endl - << PrettyPrint(func) << std::endl - << "contains free type variables: " << fv; - } - }); -TVM_REGISTER_GLOBAL("relay.ir.IRModuleAdd") - .set_body_typed([](IRModule mod, GlobalVar var, ObjectRef val, bool update) -> IRModule { - if (val->IsInstance()) { - mod->Add(var, Downcast(val), update); - } else if (val->IsInstance()) { - GlobalVar gv = Downcast(val); - IRModule mod_copy(make_object(*mod.operator->())); - mod_copy = relay::transform::EtaExpand( - /* expand_constructor */ false, - /* expand_global_var */ true)(mod_copy); - auto func = mod_copy->Lookup(gv->name_hint); - mod->Add(var, Downcast(func), update); - } else { - auto func = relay::Function({}, Downcast(val), Type(nullptr), {}); - mod->Add(var, func, update); - } - return mod; - }); - -TVM_REGISTER_GLOBAL("relay.ir.IRModuleUpdateWithRenamer") - .set_body_typed([](IRModule self, IRModule mod) -> void { - struct Renamer : relay::ExprMutator, TypeMutator { - Map defs; - Map types; - std::unordered_map ctors; - - Renamer(Map defs_one, Map defs_two, - Map types_one, Map types_two, - std::unordered_map ctors_one, - std::unordered_map ctor_two) { - for (auto pair : defs_one) { - defs.Set(pair.first, pair.second); - } - - for (auto pair : defs_two) { - auto it = defs.find(pair.first); - if (it == defs.end()) { - defs.Set(pair.first, pair.second); - } - } - - for (auto pair : types_one) { - types.Set(pair.first, pair.second); - } - - for (auto pair : types_two) { - auto it = types.find(pair.first); - if (it == types.end()) { - types.Set(pair.first, pair.second); - } - } - } - - relay::Expr VisitExpr_(const GlobalVarNode* node) override { - return defs.at(node->name_hint); - } - - Type VisitType_(const GlobalTypeVarNode* node) override { - return types.at(node->name_hint); - } - }; - - Renamer renamer(self->global_var_map_, mod->global_var_map_, self->global_type_var_map_, - mod->global_type_var_map_, self->constructor_tag_map_, - mod->constructor_tag_map_); - - self->global_var_map_ = renamer.defs; - self->global_type_var_map_ = renamer.types; - self->constructor_tag_map_ = renamer.ctors; - - for (auto pair : mod->type_definitions) { - auto tvar = renamer.types.at(pair.first->name_hint); - auto ty = renamer.ExprMutator::VisitType(pair.second); - self->AddTypeDefUnchecked(tvar, Downcast(ty), true); - } - - for (auto pair : mod->functions) { - if (auto rfn = pair.second.as()) { - auto gvar = renamer.defs.at(pair.first->name_hint); - auto fn = renamer.VisitExpr(GetRef(rfn)); - self->AddUnchecked(gvar, Downcast(fn)); - } else { - // TODO(@jroesch): rename into IRModule. - self->AddUnchecked(pair.first, pair.second); - } - } - }); - -TVM_REGISTER_GLOBAL("relay.ir.FunctionFromExprInContext") - .set_body_typed([](RelayExpr expr, IRModule mod) -> Function { - return Function(relay::FreeVars(expr), expr, Type(), relay::FreeTypeVars(expr, mod)); - }); - -TVM_REGISTER_GLOBAL("relay.ir.FuncWithAttr") - .set_body_typed([](BaseFunc func, String key, ObjectRef value) -> Optional { - if (func->IsInstance()) { - return WithAttr(Downcast(std::move(func)), key, value); - } - return NullOpt; - }); - -TVM_REGISTER_GLOBAL("relay.ir.FuncWithoutAttr") - .set_body_typed([](BaseFunc func, String key) -> Optional { - if (func->IsInstance()) { - return WithoutAttr(Downcast(std::move(func)), key); - } - return NullOpt; - }); - -TVM_REGISTER_NODE_TYPE(FunctionNode); - -TVM_REGISTER_GLOBAL("relay.ir.Function") - .set_body_typed([](tvm::Array params, Expr body, Type ret_type, - tvm::Array ty_params, tvm::DictAttrs attrs, Span span) { - return Function(params, body, ret_type, ty_params, attrs, span); - }); -TVM_REGISTER_GLOBAL("relay.ir.FunctionWithFields") - .set_body_typed([](Function function, Optional> opt_params, Optional opt_body, - Optional opt_ret_type, Optional> opt_ty_params, - Optional opt_attrs, Optional opt_virtual_device, - Optional opt_span) { - return WithFields(function, opt_params, opt_body, opt_ret_type, opt_ty_params, opt_attrs, - opt_virtual_device, opt_span); - }); - -TVM_STATIC_IR_FUNCTOR(ReprPrinter, vtable) - .set_dispatch([](const ObjectRef& ref, ReprPrinter* p) { - // TODO(@jroesch): previously this had a debug printer, the debug printer - // can cause exponential behavior and is currently dangerous, for these - // cases we need some kind of de-duping. - // - // See old implementation: - // - // auto* node = static_cast(ref.get()); - // p->stream << "FunctionNode(" << node->params << ", " << node->ret_type << ", " << - // node->body - // << ", " << node->type_params << ", " << node->attrs << ")"; - p->stream << PrettyPrint(ref); - }); - -} // namespace relay -} // namespace tvm diff --git a/src/relay/ir/indexed_graph.cc b/src/relay/ir/indexed_graph.cc deleted file mode 100644 index f10920769d1f..000000000000 --- a/src/relay/ir/indexed_graph.cc +++ /dev/null @@ -1,554 +0,0 @@ -/* - * 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 src/relay/ir/indexed_graph.cc - * \brief A graph representation of the dataflow in a Relay expression or Relay (dataflow) - * pattern. - */ -#include "indexed_graph.h" - -#include -#include -#include -#include - -#include - -namespace tvm { -namespace relay { - -std::string RefToSummary(const Expr& expr) { - class Visitor : public ExprFunctor { - std::string VisitExpr_(const VarNode* op) final { return "%" + op->name_hint(); } - std::string VisitExpr_(const GlobalVarNode* op) final { return "@" + op->name_hint; } - std::string VisitExpr_(const ConstantNode* op) final { return "const"; } - std::string VisitExpr_(const TupleNode* op) final { - return "tuple(" + std::to_string(op->fields.size()) + ")"; - } - std::string VisitExpr_(const FunctionNode* op) final { return "fn"; } - std::string VisitExpr_(const CallNode* op) final { - return VisitExpr(op->op) + "(" + std::to_string(op->args.size()) + ")"; - } - std::string VisitExpr_(const LetNode* op) final { return "let"; } - std::string VisitExpr_(const IfNode* op) final { return "if"; } - std::string VisitExpr_(const OpNode* op) final { return op->name; } - std::string VisitExpr_(const TupleGetItemNode* op) final { - return "." + std::to_string(op->index); - } - std::string VisitExpr_(const RefCreateNode* op) final { return "ref_create"; } - std::string VisitExpr_(const RefReadNode* op) final { return "ref_read"; } - std::string VisitExpr_(const RefWriteNode* op) final { return "ref_write"; } - std::string VisitExpr_(const ConstructorNode* op) final { return "ctor"; } - std::string VisitExpr_(const MatchNode* op) final { return "match"; } - }; - return Visitor().VisitExpr(expr); -} - -std::string RefToSummary(const DFPattern& pattern) { - // TODO(mbs): Implement as debugging requires. - return ""; -} - -std::unique_ptr> CreateIndexedGraph(const Expr& expr) { - /*! - * \brief Adds indexed graph nodes in post-dfs order, and discovers which let-bound vars are to - * recursive functions. - */ - class Creator : public MixedModeVisitor { - public: - std::pair>, - std::unique_ptr>> - CreateGraph(const Expr& expr) { - VisitExpr(expr); - // Last visited node is implicitly used 'externally'. - graph_->item_to_node(expr)->is_external_ = true; - return {std::move(graph_), std::move(rec_calls_)}; - } - - protected: - using MixedModeVisitor::VisitExpr_; - - // By the default the MixedModeVisitor will place - // - callee and arguments before a call - // - tuple fields before a tuple - // - tuple before a tuple projection - void VisitLeaf(const Expr& expr) override { - if (const auto* var_node = expr.as()) { - if (var_node == current_let_bound_var_) { - // Don't visit occurrences of let-rec bound vars in the recursive function body. - // Instead, wait for them to be visited at call sites outside of the function. - VLOG(1) << "Ignore let-rec var '" << var_node->name_hint() << "'"; - return; - } - } - - MixedModeVisitor::VisitLeaf(expr); - graph_->AddNode(expr); - - if (const auto* call_node = expr.as()) { - if (const auto* var_node = call_node->op.as()) { - if (var_node == current_let_bound_var_) { - // Remember this is a recursive call to the let-rec bound function. - // The Annotator functor below will not record any dependency from the let-rec bound - // var to the expression so that the indexed graph is always a DAG. - VLOG(1) << "Remembering recursive call to '" << var_node->name_hint() << "'"; - rec_calls_->emplace(call_node); - } - } - } - } - - void VisitExpr_(const LetNode* let_node) override { - auto pre_visit = [&](const LetNode* op) { - // Let-bound values come before their let-bound variable. - const VarNode* prev_let_bound_var = current_let_bound_var_; - current_let_bound_var_ = op->var.get(); - VisitExpr(op->value); - current_let_bound_var_ = prev_let_bound_var; - VisitExpr(op->var); - }; - auto post_visit = [&](const LetNode* op) { - VisitExpr(op->body); - if (let_node != op) { - // Replicate VisitLeaf, which we are effectively bypassing. - visit_counter_[op]++; - graph_->AddNode(GetRef(op)); - } - }; - ExpandANormalForm(let_node, pre_visit, post_visit); - } - - class PatternCreator : public PatternVisitor { - public: - explicit PatternCreator(Creator* creator) : creator_(creator) {} - - private: - void VisitPattern_(const PatternVarNode* pattern_var_node) final { - creator_->VisitLeaf(pattern_var_node->var); - } - - Creator* creator_; - }; - - void VisitExpr_(const MatchNode* match_node) override { - // Matched data comes before match-bound vars then match rhs, in match order. - VisitExpr(match_node->data); - for (const Clause& c : match_node->clauses) { - PatternCreator pattern_creator(this); - pattern_creator.VisitPattern(c->lhs); - VisitExpr(c->rhs); - } - } - - /*! \brief Graph we are accumulated nodes into. */ - std::unique_ptr> graph_ = std::make_unique>(); - /*! \brief Variable the currently visited expression is to be let-bound to, if any. */ - const VarNode* current_let_bound_var_ = nullptr; - /*! \brief Accumulated calls to recursive functions. */ - std::unique_ptr> rec_calls_ = - std::make_unique>(); - }; - - /*! - * \brief Fills in the inputs and outputs for all nodes, then does dominator analysis. - * - * Thought we use the ExprFunctor to visit nodes, we never recurse and instead just inspect - * each sub-expression's immediate sub-sub-expressions to accumulate inputs and outputs. - */ - class Annotator : public ExprFunctor { - public: - explicit Annotator(std::pair>, - std::unique_ptr>> - args) - : graph_(std::move(args.first)), rec_calls_(std::move(args.second)) {} - - std::unique_ptr> Annotate() { - // Visit all of the nodes in topological order to get forward outputs - for (PostDfsIndex index = 0; index < graph_->size(); ++index) { - VisitExpr(graph_->index_to_node(index)->ref()); - } - // do the dominator analysis - graph_->PostDom(); - return std::move(graph_); - } - - /*! - * \brief Add \p parent as a possible output of the node corresponding to \p expr. - */ - void AddOutput(const Expr& expr, IndexedGraph::Node* parent) { - auto current = graph_->item_to_node(expr); - current->outputs_.push_back(parent); - parent->inputs_.push_back(current); - } - - protected: - void VisitExpr_(const VarNode* var_node) override {} - - void VisitExpr_(const GlobalVarNode* global_var_node) override {} - - void VisitExpr_(const ConstantNode* constant_node) override {} - - void VisitExpr_(const TupleNode* tuple_node) override { - auto node = graph_->item_to_node(GetRef(tuple_node)); - for (auto field : tuple_node->fields) { - AddOutput(field, node); - } - } - - void VisitExpr_(const FunctionNode* function_node) override { - auto node = graph_->item_to_node(GetRef(function_node)); - // Nothing to do for parameters -- each use of a parameter will contribute to its outputs. - AddOutput(function_node->body, node); - } - - void VisitExpr_(const CallNode* call_node) override { - auto node = graph_->item_to_node(GetRef(call_node)); - if (rec_calls_->count(call_node)) { - // We want the indexed graph to be a DAG, so don't consider a call to a let-rec bound - // function from inside the function to depend on the let-rec bound var. - VLOG(1) << "Ignoring op in call " << RefToSummary(GetRef(call_node)); - } else { - AddOutput(call_node->op, node); - } - for (auto arg : call_node->args) { - AddOutput(arg, node); - } - } - - void VisitExpr_(const LetNode* let_node) override { - auto node = graph_->item_to_node(GetRef(let_node)); - auto let_var_node = graph_->item_to_node(let_node->var); - AddOutput(let_node->value, let_var_node); - // Nothing to do for the let-bound variable -- each use of that variable in the let-body - // will contribute to its outputs. - AddOutput(let_node->body, node); - } - - void VisitExpr_(const IfNode* if_node) override { - auto node = graph_->item_to_node(GetRef(if_node)); - AddOutput(if_node->cond, node); - AddOutput(if_node->true_branch, node); - AddOutput(if_node->false_branch, node); - } - - void VisitExpr_(const OpNode* op_node) override {} - - void VisitExpr_(const TupleGetItemNode* tuple_get_item_node) override { - auto node = graph_->item_to_node(GetRef(tuple_get_item_node)); - AddOutput(tuple_get_item_node->tuple, node); - } - - void VisitExpr_(const RefCreateNode* ref_create_node) override { - auto node = graph_->item_to_node(GetRef(ref_create_node)); - AddOutput(ref_create_node->value, node); - } - - void VisitExpr_(const RefReadNode* ref_read_node) override { - auto node = graph_->item_to_node(GetRef(ref_read_node)); - AddOutput(ref_read_node->ref, node); - } - - void VisitExpr_(const RefWriteNode* ref_write_node) override { - auto node = graph_->item_to_node(GetRef(ref_write_node)); - AddOutput(ref_write_node->ref, node); - AddOutput(ref_write_node->value, node); - } - - void VisitExpr_(const ConstructorNode* constructor_node) override {} - - class PatternAnnotator : public PatternVisitor { - public: - PatternAnnotator(Annotator* annotator, const ExprNode* adt_node) - : annotator_(annotator), adt_node_(adt_node) {} - - private: - void VisitPattern_(const PatternVarNode* pattern_var_node) final { - auto node = annotator_->graph_->item_to_node(pattern_var_node->var); - annotator_->AddOutput(GetRef(adt_node_), node); - } - - Annotator* annotator_; - const ExprNode* adt_node_; - }; - - void VisitExpr_(const MatchNode* match_node) override { - // Data flows from the match data to pattern vars into match arms and out into overall - // match. - auto node = graph_->item_to_node(GetRef(match_node)); - for (const Clause& c : match_node->clauses) { - PatternAnnotator pattern_annotator(this, match_node->data.get()); - pattern_annotator.VisitPattern(c->lhs); - AddOutput(c->rhs, node); - } - } - - std::unique_ptr> graph_; - /*! \brief Accumulated calls to recursive functions. */ - std::unique_ptr> rec_calls_; - }; - - /*! \brief Fills in the basic blocks for all nodes. */ - class Blocker : public MixedModeVisitor { - public: - explicit Blocker(std::unique_ptr> graph) : graph_(std::move(graph)) {} - - std::unique_ptr> Scope(const Expr& expr) { - VisitExpr(expr); - return std::move(graph_); - } - - private: - using MixedModeVisitor::VisitExpr_; - - void VisitLeaf(const Expr& expr) override { - MixedModeVisitor::VisitLeaf(expr); - SetScope(expr); - } - - void VisitExpr_(const FunctionNode* function_node) override { - auto node = graph_->item_to_node(GetRef(function_node)); - basic_block_stack_.push_back(node); - ExprVisitor::VisitExpr_(function_node); - basic_block_stack_.pop_back(); - } - - void VisitExpr_(const IfNode* if_node) override { - VisitExpr(if_node->cond); - auto node = graph_->item_to_node(GetRef(if_node)); - basic_block_stack_.push_back(node); - VisitExpr(if_node->true_branch); - VisitExpr(if_node->false_branch); - basic_block_stack_.pop_back(); - } - - void VisitExpr_(const LetNode* let_node) override { - auto pre_visit = [&](const LetNode* op) { - VisitExpr(op->value); - VisitExpr(op->var); - }; - auto post_visit = [&](const LetNode* op) { - VisitExpr(op->body); - if (let_node != op) { - visit_counter_[op]++; - SetScope(GetRef(op)); - } - }; - ExpandANormalForm(let_node, pre_visit, post_visit); - } - - class PatternBlocker : public PatternVisitor { - public: - explicit PatternBlocker(Blocker* scoper) : scoper_(scoper) {} - - private: - void VisitPattern_(const PatternVarNode* pattern_var_node) final { - scoper_->SetScope(pattern_var_node->var); - } - - Blocker* scoper_; - }; - - void VisitExpr_(const MatchNode* match_node) override { - VisitExpr(match_node->data); - auto node = graph_->item_to_node(GetRef(match_node)); - basic_block_stack_.push_back(node); - for (const Clause& c : match_node->clauses) { - PatternBlocker pattern_scoper(this); - pattern_scoper.VisitPattern(c->lhs); - VisitExpr(c->rhs); - } - basic_block_stack_.pop_back(); - } - - void SetScope(const Expr& expr) { - auto node = graph_->item_to_node(expr); - if (!basic_block_stack_.empty()) { - node->basic_block_ = basic_block_stack_.back(); - } - } - - std::unique_ptr> graph_; - std::vector::Node*> basic_block_stack_; - }; - - VLOG(1) << "CreateIndexedGraph:" << std::endl << PrettyPrint(expr); - std::unique_ptr> graph = - Blocker(Annotator(Creator().CreateGraph(expr)).Annotate()).Scope(expr); - VLOG(1) << "graph:" << std::endl << graph->ToString(); -#if TVM_LOG_DEBUG - graph->CheckValid(); -#endif - return graph; -} - -std::unique_ptr> CreateIndexedGraph(const DFPattern& pattern) { - /*! \brief Creates an IndexedGraph and determines topological order */ - class Creator : public DFPatternVisitor { - public: - std::unique_ptr> CreateGraph(const DFPattern& pattern) { - graph_ = std::make_unique>(); - VisitDFPattern(pattern); - graph_->item_to_node(pattern)->is_external_ = true; - return std::move(graph_); - } - - protected: - void VisitDFPattern(const DFPattern& pattern) override { - if (this->visited_.count(pattern.get()) == 0) { - DFPatternVisitor::VisitDFPattern(pattern); - graph_->AddNode(pattern); - } - } - - std::unique_ptr> graph_; - }; - - /*! \brief Annotator takes an IndexedGraph, fills it's forward outputs, and does dominator tree - * analysis. - * - * Annotator use ExprFunctor to visit nodes, but iterates over them in pre-determined - * topological order instead of recursing. - */ - class Annotator : public DFPatternFunctor { - public: - Annotator(std::unique_ptr> graph) : graph_(std::move(graph)) {} - - std::unique_ptr> Annotate() { - // Visit all of the nodes in topological order to get forward outputs - for (PostDfsIndex index = 0; index < graph_->size(); ++index) { - VisitDFPattern(graph_->index_to_node(index)->ref()); - } - // do the dominator analysis - graph_->PostDom(); - return std::move(graph_); - } - - /*! Default visitation pushes the parent to the child's outputs */ - void AddOutput(const DFPattern& pattern, IndexedGraph::Node* parent) { - auto current = graph_->item_to_node(pattern); - if (parent) { - current->outputs_.push_back(parent); - parent->inputs_.push_back(current); - } - } - - protected: - void VisitDFPattern_(const AltPatternNode* op) override { - auto node = graph_->item_to_node(GetRef(op)); - AddOutput(op->left, node); - AddOutput(op->right, node); - } - - void VisitDFPattern_(const AttrPatternNode* op) override { - auto node = graph_->item_to_node(GetRef(op)); - AddOutput(op->pattern, node); - } - - void VisitDFPattern_(const CallPatternNode* op) override { - auto node = graph_->item_to_node(GetRef(op)); - AddOutput(op->op, node); - if (op->args.defined()) { - for (auto arg : op->args) { - AddOutput(arg, node); - } - } - } - - void VisitDFPattern_(const ConstantPatternNode* op) override {} - - void VisitDFPattern_(const DataTypePatternNode* op) override { - auto node = graph_->item_to_node(GetRef(op)); - AddOutput(op->pattern, node); - } - - void VisitDFPattern_(const DominatorPatternNode* op) override { - auto node = graph_->item_to_node(GetRef(op)); - AddOutput(op->parent, node); - AddOutput(op->path, node); - AddOutput(op->child, node); - } - - void VisitDFPattern_(const ExprPatternNode* op) override {} - - void VisitDFPattern_(const FunctionPatternNode* op) override { - auto node = graph_->item_to_node(GetRef(op)); - if (op->params.defined()) { - for (auto param : op->params) { - AddOutput(param, node); - } - } - AddOutput(op->body, node); - } - - void VisitDFPattern_(const ShapePatternNode* op) override { - auto node = graph_->item_to_node(GetRef(op)); - AddOutput(op->pattern, node); - } - - void VisitDFPattern_(const TupleGetItemPatternNode* op) override { - auto node = graph_->item_to_node(GetRef(op)); - AddOutput(op->tuple, node); - } - - void VisitDFPattern_(const TuplePatternNode* op) override { - auto node = graph_->item_to_node(GetRef(op)); - if (op->fields.defined()) { - for (auto field : op->fields) { - AddOutput(field, node); - } - } - } - - void VisitDFPattern_(const IfPatternNode* op) override { - auto node = graph_->item_to_node(GetRef(op)); - AddOutput(op->cond, node); - AddOutput(op->true_branch, node); - AddOutput(op->false_branch, node); - } - - void VisitDFPattern_(const LetPatternNode* op) override { - auto node = graph_->item_to_node(GetRef(op)); - AddOutput(op->var, node); - AddOutput(op->value, node); - AddOutput(op->body, node); - } - - void VisitDFPattern_(const TypePatternNode* op) override { - auto node = graph_->item_to_node(GetRef(op)); - AddOutput(op->pattern, node); - } - - void VisitDFPattern_(const VarPatternNode* op) override {} - - void VisitDFPattern_(const WildcardPatternNode* op) override { - if (op->pattern) { - auto node = graph_->item_to_node(GetRef(op)); - AddOutput(op->pattern.value(), node); - } - } - - std::unique_ptr> graph_; - }; - - return Annotator(Creator().CreateGraph(pattern)).Annotate(); -} - -} // namespace relay -} // namespace tvm diff --git a/src/relay/ir/indexed_graph.h b/src/relay/ir/indexed_graph.h deleted file mode 100644 index c1ce53f40da3..000000000000 --- a/src/relay/ir/indexed_graph.h +++ /dev/null @@ -1,371 +0,0 @@ -/* - * 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 src/relay/ir/indexed_graph.h - * \brief A graph representation of the dataflow in a Relay expression or Relay (dataflow) - * pattern. Each 'indexed graph' node is 1:1 with an expression/pattern 'node', hence the - * term 'IndexedGraph'. Dataflow is captured in a generic representation which is convenient - * for analysis, particularly pattern matching and partitioning. - * - * TODO(mbs): Copied from fuse_ops.cc, consider refactoring to share implementation. - */ -#ifndef TVM_RELAY_IR_INDEXED_GRAPH_H_ -#define TVM_RELAY_IR_INDEXED_GRAPH_H_ - -#include - -#include -#include -#include -#include -#include -#include -#include - -namespace tvm { -namespace relay { - -/*! \brief The index of a node in the post-dfs traversal of overall expression. */ -using PostDfsIndex = size_t; - -/*! - * \brief Returns a brief summary of the 'reference' expression or pattern. Only used by - * IndexedGraph::ToString() for debugging. - */ -std::string RefToSummary(const Expr& expr); -std::string RefToSummary(const DFPattern& pattern); - -/*! - * \brief Represents the implied dataflow of an expression or (dataflow) pattern as a DAG who's - * nodes are 1:1 with those in the underlying expression/pattern. - * - * Each indexed graph node captures: - * - Dataflow inputs. - * - Dataflow outputs (or a flag indicating the node is an implied output). - * - Dominator parent (ie closest node at which all outputs of the current node re-combine). - * - Dominator children (inverse of above). - * - Basic block (ie node representing the body of a function, arm of an if, etc). - * - * This class is templated so we can analyze both DFPatterns and Exprs with the same infrastructure. - * - * IndexedGraph should be instantiated through the CreateIndexedGraph utilities below. - */ -template -class IndexedGraph { - public: - using TNode = typename T::ContainerType; - - /*! \brief A Node in the graph. */ - struct Node { - /*! \brief Node Constructor - * \param ref The expression or dataflow pattern node this indexed graph node is augmenting. - * \param index The index of this node in the topological order - */ - Node(const TNode* ref, PostDfsIndex index) : node_ref_(ref), index_(index) {} - - /*! \brief The underlying expression or pattern node. */ - const TNode* node_ref_; - - T ref() const { - ICHECK(node_ref_ != nullptr); - return GetRef(node_ref_); - } - - /*! - * \brief The index of this node in post-dfs order. If left.index_ > right.index_ then - * left does not flow into right. If left.index_ = right.index_ then left and right are - * the same node. - */ - const PostDfsIndex index_; - - /*! \brief If true this node has implicit outputs, for example as the result of a function. */ - bool is_external_ = false; - /*! \brief Immediate dataflow inputs to this node. */ - std::vector inputs_; - /*! \brief Immediate dataflow outputs of this node -- may be empty if is_external_ is true. */ - std::vector outputs_; - - /*! - * \brief The node representing the 'basic block' containing this node: - * - Function bodies start a new basic block for their bodies. - * - The true and false branches of an if start their own blocks. - * - The arms of a match each have their own blocks. - */ - Node* basic_block_ = nullptr; - - /*! \brief The depth of this node in the dominator tree */ - size_t depth_ = 0; - /*! - * \brief The dominator parent of this node. This is the node N with least index such that - * all possible dataflows from this node pass through N. - */ - Node* dominator_parent_ = nullptr; - /*! \brief The nodes this node dominates. */ - std::vector dominator_children_; - - /*! - * Add to \p nodes all the nodes which are strictly downstream of \p this, ie can be - * reached by following output paths. - */ - void AccumulateDownstreamNodes(std::unordered_set* nodes) const { - std::stack stack; - stack.push(this); - while (!stack.empty()) { - const Node* current = stack.top(); - stack.pop(); - for (auto node : current->outputs_) { - if (nodes->count(node) == 0) { - stack.push(node); - nodes->insert(node); - } - } - } - } - - /*! - * \brief Returns true if \p this is a dominator of \p other. Ie all dataflow paths from \p - * other pass through \p this. - */ - bool Dominates(const Node* other) const { - std::stack stack; - std::unordered_set visited; - stack.push(this); - while (!stack.empty()) { - const Node* current = stack.top(); - stack.pop(); - for (auto node : current->dominator_children_) { - if (visited.count(node) == 0) { - if (other == node) { - return true; - } else { - stack.push(node); - } - visited.insert(node); - } - } - } - return false; - } - }; - - PostDfsIndex size() const { return topological_order_.size(); } - - Node* item_to_node(const T& item) { return item_to_node(item.get()); } - const Node* item_to_node(const T& item) const { return item_to_node(item.get()); } - - Node* item_to_node(const TNode* item) { - auto itr = node_map_.find(item); - ICHECK(itr != node_map_.end()) << PrettyPrint(GetRef(item)); - return itr->second; - } - - const Node* item_to_node(const TNode* item) const { - auto itr = node_map_.find(item); - ICHECK(itr != node_map_.end()) << PrettyPrint(GetRef(item)); - return itr->second; - } - - Node* index_to_node(PostDfsIndex index) { - ICHECK_LT(index, topological_order_.size()) << index; - return topological_order_[index].get(); - } - - const Node* index_to_node(PostDfsIndex index) const { - ICHECK_LT(index, topological_order_.size()) << index; - return topological_order_[index].get(); - } - - /*! - * \brief (For debugging only) Returns description of indexed graph with hints as to the - * sub-expressions or sub-patterns corresponding to each indexed graph node. - */ - std::string ToString() const { - std::ostringstream os; - os << "IndexedGraph(size = " << topological_order_.size() << ") {" << std::endl; - for (PostDfsIndex index = 0; index < topological_order_.size(); ++index) { - const Node* node = topological_order_[index].get(); - ICHECK_EQ(index, node->index_); - os << " " << index << " (" << RefToSummary(node->ref()) << "): inputs=["; - for (const auto* sub_node : node->inputs_) { - os << sub_node->index_ << ","; - } - os << "], outputs=["; - for (const auto* sub_node : node->outputs_) { - os << sub_node->index_ << ","; - } - os << "]"; - if (node->is_external_) { - os << ", external"; - } - if (node->basic_block_) { - os << ", basic_block=" << node->basic_block_->index_; - } - if (node->depth_ > 0) { - os << ", depth=" << node->depth_; - } - if (node->dominator_parent_) { - os << ", dom_parent=" << node->dominator_parent_->index_; - } - os << ", dom_children=["; - for (const auto* sub_node : node->dominator_children_) { - os << sub_node->index_ << ","; - } - os << "]" << std::endl; - } - os << "}"; - return os.str(); - } - - /*! - * Check-fails if the graph is ill-formed. For debugging only. - */ - void CheckValid() const { - ICHECK_GT(topological_order_.size(), 0); - for (PostDfsIndex index = 0; index < topological_order_.size(); ++index) { - const Node* node = topological_order_[index].get(); - // We have a node. - ICHECK(node); - // Bijections with post-dfs indexes and expressions/patterns are correct. - ICHECK_EQ(node->index_, index); - ICHECK(node->node_ref_); - auto itr = node_map_.find(node->node_ref_); - ICHECK(itr != node_map_.end()); - ICHECK_EQ(itr->second, node) << "at index " << index << " in:" << std::endl << ToString(); - // Inputs come before. - for (size_t i = 0; i < node->inputs_.size(); ++i) { - const Node* input = node->inputs_[i]; - ICHECK(input); - ICHECK_LT(input->index_, index); - ICHECK(std::find(input->outputs_.begin(), input->outputs_.end(), node) != - input->outputs_.end()); - } - // Outputs come after. - for (size_t i = 0; i < node->outputs_.size(); ++i) { - const Node* output = node->outputs_[i]; - ICHECK(output); - ICHECK_GT(output->index_, index); - ICHECK(std::find(output->inputs_.begin(), output->inputs_.end(), node) != - output->inputs_.end()); - } - ICHECK_GT(node->depth_, 0); - // Dominator children come before. - for (size_t i = 0; i < node->dominator_children_.size(); ++i) { - const Node* child = node->dominator_children_[i]; - ICHECK(child); - ICHECK_LT(child->index_, index); - } - if (node->dominator_parent_) { - // Dominator comes after. - ICHECK_GT(node->dominator_parent_->index_, index); - } - } - } - - private: - /*! \brief Construct the domination tree inside IndexedGraph */ - void PostDom() { - for (PostDfsIndex i = topological_order_.size(); i != 0; --i) { - PostDfsIndex index = i - 1; - auto* current = topological_order_[index].get(); - if (current->is_external_) { - current->depth_ = 1; - current->dominator_parent_ = nullptr; - } else { - auto parent = LeastCommonAncestor(current->outputs_); - current->depth_ = parent ? parent->depth_ + 1 : 1; - current->dominator_parent_ = parent; - if (parent) { - parent->dominator_children_.push_back(current); - } - } - } - } - - /*! \brief Find the least common ancestor of all outputs of a node */ - Node* LeastCommonAncestor(const std::vector& outputs) { - if (outputs.size() == 0) { - return nullptr; - } - auto parent = outputs.at(0); - for (size_t i = 1; i < outputs.size(); ++i) { - parent = LeastCommonAncestor(parent, outputs.at(i)); - } - return parent; - } - - /*! \brief Find the least common ancestor of two nodes */ - Node* LeastCommonAncestor(Node* lhs, Node* rhs) { - if (lhs == nullptr || rhs == nullptr) { - return nullptr; - } - PostDfsIndex lhs_index = lhs->index_; - PostDfsIndex rhs_index = rhs->index_; - while (lhs != rhs) { - ICHECK(lhs && rhs) << "LCA(" << lhs_index << ", " << rhs_index << ") on graph:" << std::endl - << ToString(); - if (lhs->depth_ < rhs->depth_) { - rhs = rhs->dominator_parent_; - } else if (lhs->depth_ > rhs->depth_) { - lhs = lhs->dominator_parent_; - } else { - rhs = rhs->dominator_parent_; - lhs = lhs->dominator_parent_; - } - } - return lhs; - } - - /*! - * \brief Appends a node corresponding to \p ref, and maintains the sub-expression/sub-pattern to - * node bijection. The insertion index will be the node's PostDfsIndex. All other node properties - * are accumulated in-place. - */ - void AddNode(const T& ref) { - PostDfsIndex index = topological_order_.size(); - auto node = std::make_unique(ref.get(), index); - node_map_[ref.get()] = node.get(); - topological_order_.emplace_back(std::move(node)); - } - - /*! - * \brief Map from underlying sub-expression or sub-pattern nodes to their indexed graph nodes. - */ - std::unordered_map node_map_; - /*! \brief All nodes in increasing post-dfs index order. This vector owns all the nodes. */ - std::vector> topological_order_; - - friend std::unique_ptr> CreateIndexedGraph(const Expr& expr); - friend std::unique_ptr> CreateIndexedGraph(const DFPattern& pattern); -}; - -/*! \brief Returns an Indexed Graph for \p expr, which much outlive the result. */ -std::unique_ptr> CreateIndexedGraph(const Expr& expr); - -/*! - * \brief Returns an Indexed Graph for \p pattern, which must outlive the result. - * The dataflow for a pattern mimics the dataflow for the expression which would match - * that pattern. - */ -std::unique_ptr> CreateIndexedGraph(const DFPattern& pattern); - -} // namespace relay -} // namespace tvm - -#endif // TVM_RELAY_IR_INDEXED_GRAPH_H_ diff --git a/src/relay/ir/op_strategy.cc b/src/relay/ir/op_strategy.cc deleted file mode 100644 index c675b2970536..000000000000 --- a/src/relay/ir/op_strategy.cc +++ /dev/null @@ -1,122 +0,0 @@ -/* - * 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 src/relay/ir/op_strategy.cc - * \brief The Relay operator Strategy and related data structure. - */ - -#include - -namespace tvm { -namespace relay { - -TVM_REGISTER_NODE_TYPE(OpImplementationNode); -TVM_REGISTER_NODE_TYPE(OpSpecializationNode); -TVM_REGISTER_NODE_TYPE(OpStrategyNode); - -Array OpImplementation::Compute(const Attrs& attrs, const Array& inputs, - const Type& out_type) { - return (*this)->fcompute(attrs, inputs, out_type); -} - -te::Schedule OpImplementation::Schedule(const Attrs& attrs, const Array& outs, - const Target& target) { - return (*this)->fschedule(attrs, outs, target); -} - -void OpSpecialization::AddImplementation(tvm::relay::FTVMCompute fcompute, - tvm::relay::FTVMSchedule fschedule, String name, - int plevel) { - auto n = make_object(); - n->fcompute = fcompute; - n->fschedule = fschedule; - n->name = std::move(name); - n->plevel = plevel; - (*this)->implementations.push_back(OpImplementation(n)); -} - -void OpStrategy::AddImplementation(FTVMCompute fcompute, FTVMSchedule fschedule, String name, - int plevel) { - auto curr_cond = te::SpecializedCondition::Current(); - auto self = this->operator->(); - Array specializations = self->specializations; - OpSpecialization op_spec; - for (OpSpecialization op_spec : specializations) { - if (op_spec->condition == curr_cond) { - op_spec.AddImplementation(fcompute, fschedule, std::move(name), plevel); - return; - } - } - ObjectPtr n = make_object(); - n->condition = curr_cond; - op_spec = OpSpecialization(n); - op_spec.AddImplementation(fcompute, fschedule, std::move(name), plevel); - self->specializations.push_back(op_spec); -} - -TVM_REGISTER_GLOBAL("relay.op._OpImplementationCompute") - .set_body([](TVMArgs args, TVMRetValue* rv) { - OpImplementation imp = args[0]; - Attrs attrs = args[1]; - Array inputs = args[2]; - Type out_type = args[3]; - *rv = imp.Compute(attrs, inputs, out_type); - }); - -TVM_REGISTER_GLOBAL("relay.op._OpImplementationSchedule") - .set_body([](TVMArgs args, TVMRetValue* rv) { - OpImplementation imp = args[0]; - Attrs attrs = args[1]; - Array outs = args[2]; - Target target = args[3]; - *rv = imp.Schedule(attrs, outs, target); - }); - -TVM_REGISTER_GLOBAL("relay.op._make.OpStrategy").set_body([](TVMArgs args, TVMRetValue* rv) { - ObjectPtr n = make_object(); - *rv = OpStrategy(n); -}); - -TVM_REGISTER_GLOBAL("relay.op._OpStrategyAddImplementation") - .set_body([](TVMArgs args, TVMRetValue* rv) { - OpStrategy strategy = args[0]; - FTVMCompute compute = args[1]; - FTVMSchedule schedule = args[2]; - std::string name = args[3]; - int plevel = args[4]; - strategy.AddImplementation(compute, schedule, name, plevel); - }); - -TVM_STATIC_IR_FUNCTOR(ReprPrinter, vtable) - .set_dispatch([](const ObjectRef& node, ReprPrinter* p) { - auto* op = static_cast(node.get()); - p->stream << "op_strategy(" << op->specializations << ")"; - }) - .set_dispatch([](const ObjectRef& node, ReprPrinter* p) { - auto* op = static_cast(node.get()); - p->stream << "op_spec(" << op->condition << ", " << op->implementations << ")"; - }) - .set_dispatch([](const ObjectRef& node, ReprPrinter* p) { - auto* op = static_cast(node.get()); - p->stream << "op_impl(name=" << op->name << ", level=" << op->plevel << ")"; - }); - -} // namespace relay -} // namespace tvm diff --git a/src/relay/ir/pattern_functor.cc b/src/relay/ir/pattern_functor.cc deleted file mode 100644 index 8c366bad641a..000000000000 --- a/src/relay/ir/pattern_functor.cc +++ /dev/null @@ -1,93 +0,0 @@ -/* - * 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 src/relay/ir/pattern_functor.cc - * \brief Implementations of visitors and mutators for ADT patterns. - */ - -#include - -namespace tvm { -namespace relay { - -Pattern PatternMutator::Mutate(const Pattern& pat) { return (*this)(pat); } - -Pattern PatternMutator::VisitPattern_(const PatternWildcardNode* op) { return GetRef(op); } - -Pattern PatternMutator::VisitPattern_(const PatternVarNode* op) { - return PatternVar(VisitVar(op->var)); -} - -Pattern PatternMutator::VisitPattern_(const PatternConstructorNode* op) { - std::vector pat; - for (const auto& p : op->patterns) { - pat.push_back(VisitPattern(p)); - } - return PatternConstructor(VisitConstructor(op->constructor), pat); -} - -Pattern PatternMutator::VisitPattern_(const PatternTupleNode* op) { - std::vector pat; - for (const auto& p : op->patterns) { - pat.push_back(VisitPattern(p)); - } - return PatternTuple(pat); -} - -Type PatternMutator::VisitType(const Type& t) { return t; } - -Var PatternMutator::VisitVar(const Var& v) { - if (var_map_.count(v) == 0) { - var_map_.insert(std::pair(v, Var(v->name_hint(), VisitType(v->type_annotation)))); - } - return var_map_.at(v); -} - -Constructor PatternMutator::VisitConstructor(const Constructor& v) { return v; } - -void PatternVisitor::VisitPattern_(const PatternWildcardNode* op) {} - -void PatternVisitor::VisitPattern_(const PatternVarNode* op) { VisitVar(op->var); } - -void PatternVisitor::VisitPattern_(const PatternConstructorNode* op) { - VisitConstructor(op->constructor); - for (const auto& p : op->patterns) { - VisitPattern(p); - } -} - -void PatternVisitor::VisitPattern_(const PatternTupleNode* op) { - for (const auto& p : op->patterns) { - VisitPattern(p); - } -} - -void PatternVisitor::VisitType(const Type& t) {} - -void PatternVisitor::VisitVar(const Var& v) { VisitType(v->type_annotation); } - -void PatternVisitor::VisitConstructor(const Constructor& c) { - for (const auto& inp : c->inputs) { - VisitType(inp); - } -} - -} // namespace relay -} // namespace tvm diff --git a/src/relay/ir/transform.cc b/src/relay/ir/transform.cc deleted file mode 100644 index dd31a1f7367d..000000000000 --- a/src/relay/ir/transform.cc +++ /dev/null @@ -1,179 +0,0 @@ -/* - * 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 relay/ir/transform.cc - * \brief Relay specific transformation passes. - */ -#include -#include -#include -#include - -namespace tvm { -namespace relay { -namespace transform { - -TVM_REGISTER_PASS_CONFIG_OPTION("relay.fallback_device_type", IntImm); - -class FunctionPass; - -/*! - * \brief Function-level passes are used to implement various global - * optimizations for a given Relay module. It fetches one function at a time - * from the function list in the module for optimization. - * - * Note that the scope of passes at this level is a Relay function. Therefore, - * we cannot add or delete a function through these passes as they are not aware - * of the global information. - */ -class FunctionPassNode : public PassNode { - public: - /* \brief The pass meta data.*/ - PassInfo pass_info; - - /*! \brief The packed pass function sketches the real optimization. For - * instance, we can implement a pass that works on a Relay function as a - * `pass_func` and let it run on a given module. The same `pass_func` will - * then be applied on each function in the module. - */ - runtime::TypedPackedFunc pass_func; - - FunctionPassNode() = default; - - void VisitAttrs(tvm::AttrVisitor* v) { v->Visit("pass_info", &pass_info); } - - /*! - * \brief Run a function pass on given pass context. - * - * \param mod The module that an optimization pass is applied on. - * \param mod The context that an optimization pass executes on. - * - * \return Return the updated module. - */ - IRModule operator()(IRModule mod, const PassContext& pass_ctx) const final; - - /*! - * \brief Get the pass information/meta data. - */ - PassInfo Info() const override { return pass_info; } - - static constexpr const char* _type_key = "relay.FunctionPass"; - TVM_DECLARE_FINAL_OBJECT_INFO(FunctionPassNode, PassNode); -}; - -class FunctionPass : public Pass { - public: - /*! - * \brief The constructor - * \param pass_func The packed function which implements a pass. - * \param pass_info The pass info. - */ - TVM_DLL FunctionPass( - runtime::TypedPackedFunc pass_func, - PassInfo pass_info); - - TVM_DEFINE_OBJECT_REF_METHODS(FunctionPass, Pass, FunctionPassNode); -}; - -FunctionPass::FunctionPass( - runtime::TypedPackedFunc pass_func, - PassInfo pass_info) { - auto n = make_object(); - n->pass_func = std::move(pass_func); - n->pass_info = std::move(pass_info); - data_ = std::move(n); -} - -// Perform Module -> Module optimizations at the Function level. -IRModule FunctionPassNode::operator()(IRModule mod, const PassContext& pass_ctx) const { - DiagnosticContext previous = DiagnosticContext::Default(mod); - - if (pass_ctx->diag_ctx) { - DiagnosticContext tmp = pass_ctx->diag_ctx.value(); - pass_ctx->diag_ctx = previous; - previous = tmp; - } else { - pass_ctx->diag_ctx = previous; - } - - ICHECK(pass_ctx->diag_ctx) - << "The diagnostic context was set at the top of this block this is a bug."; - - const PassInfo& pass_info = Info(); - - ICHECK(mod.defined()); - - VLOG_CONTEXT << pass_info->name; - VLOG(0) << "Executing function pass with opt level: " << pass_info->opt_level; - VLOG(1) << "Input module:" << std::endl << PrettyPrint(mod); - - IRModule updated_mod = mod->ShallowCopy(); - - std::vector> updates; - for (const auto& kv : mod->functions) { - // only process optimizable Relay Functions - if (const auto* function_node = AsOptimizableFunctionNode(kv.second)) { - Function updated_func = pass_func(GetRef(function_node), updated_mod, pass_ctx); - updates.push_back({kv.first, std::move(updated_func)}); - } - } - - for (const auto& pair : updates) { - updated_mod->Add(pair.first, pair.second, true); - } - - ICHECK(pass_ctx->diag_ctx) - << "The diagnostic context was set at the top of this block this is a bug."; - - pass_ctx->diag_ctx.value().Render(); - pass_ctx->diag_ctx = previous; - - VLOG(1) << "Output module:" << std::endl << PrettyPrint(updated_mod); - - // TODO(@jroesch): move away from eager type checking for performance reasons - // make issue. - return transform::InferType()(updated_mod); -} - -Pass CreateFunctionPass( - const runtime::TypedPackedFunc& pass_func, - int opt_level, String name, tvm::Array required, bool traceable) { - PassInfo pass_info = PassInfo(opt_level, name, required, traceable); - return FunctionPass(pass_func, pass_info); -} - -TVM_REGISTER_NODE_TYPE(FunctionPassNode); - -TVM_REGISTER_GLOBAL("relay._transform.MakeFunctionPass") - .set_body_typed( - [](runtime::TypedPackedFunc pass_func, - PassInfo pass_info) { return FunctionPass(pass_func, pass_info); }); - -TVM_STATIC_IR_FUNCTOR(ReprPrinter, vtable) - .set_dispatch([](const ObjectRef& ref, ReprPrinter* p) { - auto* node = static_cast(ref.get()); - const PassInfo info = node->Info(); - p->stream << "Run Function pass: " << info->name << " at the optimization level " - << info->opt_level; - }); - -} // namespace transform -} // namespace relay -} // namespace tvm diff --git a/src/relay/op/algorithm/argsort.cc b/src/relay/op/algorithm/argsort.cc deleted file mode 100644 index 455d413c2746..000000000000 --- a/src/relay/op/algorithm/argsort.cc +++ /dev/null @@ -1,69 +0,0 @@ -/* - * 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 argsort.cc - * \brief Argsort operators - */ -#include -#include - -namespace tvm { -namespace relay { - -TVM_REGISTER_NODE_TYPE(ArgsortAttrs); - -bool ArgsortRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - // `types` contains: [data, result] - const ArgsortAttrs* param = attrs.as(); - ICHECK_EQ(types.size(), 2); - const auto* data = types[0].as(); - if (data == nullptr) { - ICHECK(types[0].as()) - << "Argsort: expect input type to be TensorType but get " << types[0]; - return false; - } - reporter->Assign(types[1], TensorType(data->shape, param->dtype)); - return true; -} - -Expr MakeArgsort(Expr data, int axis, bool is_ascend, DataType dtype) { - auto attrs = make_object(); - attrs->axis = axis; - attrs->is_ascend = is_ascend; - attrs->dtype = dtype; - static const Op& op = Op::Get("argsort"); - return Call(op, {data}, Attrs(attrs), {}); -} - -TVM_REGISTER_GLOBAL("relay.op._make.argsort").set_body_typed(MakeArgsort); - -RELAY_REGISTER_OP("argsort") - .describe(R"doc(Returns the indices that would sort an -input array along the given axis. -)doc" TVM_ADD_FILELINE) - .set_num_inputs(1) - .set_attrs_type() - .add_argument("data", "Tensor", "Input data.") - .set_support_level(6) - .add_type_rel("Argsort", ArgsortRel); - -} // namespace relay -} // namespace tvm diff --git a/src/relay/op/algorithm/searchsorted.cc b/src/relay/op/algorithm/searchsorted.cc deleted file mode 100644 index be5921311660..000000000000 --- a/src/relay/op/algorithm/searchsorted.cc +++ /dev/null @@ -1,86 +0,0 @@ -/* - * 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 searchsorted.cc - * \brief SearchSorted operators - */ -#include -#include -#include - -namespace tvm { -namespace relay { - -TVM_REGISTER_NODE_TYPE(SearchSortedAttrs); - -bool SearchSortedRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - const SearchSortedAttrs* param = attrs.as(); - ICHECK_EQ(types.size(), 3); - const auto* sorted_sequence = types[0].as(); - const auto* values = types[1].as(); - ICHECK(sorted_sequence) << "Expects TensorType in the first input"; - ICHECK(values) << "Expects TensorType in the second input"; - ICHECK_GT(values->shape.size(), 0) << "The rank of `values` must be greater than one"; - - if (sorted_sequence->shape.size() > 1) { - ICHECK_EQ(sorted_sequence->shape.size(), values->shape.size()) - << "Ranks of `sorted_sequence` and values must be the same if `sorted_sequence` is " - "multi-dimensional."; - - for (size_t i = 0; i < values->shape.size() - 1; ++i) { - if (sorted_sequence->shape[i].as() && values->shape[i].as()) { - ICHECK_EQ(sorted_sequence->shape[i].as()->value, - values->shape[i].as()->value) - << "`sorted_sequence and `values` do not have the same shape along outer axes"; - } - } - } - - reporter->Assign(types[2], TensorType(values->shape, param->dtype)); - return true; -} - -Expr MakeSearchSorted(Expr sorted_sequence, Expr values, Bool right, DataType dtype) { - auto attrs = make_object(); - static const Op& op = Op::Get("searchsorted"); - attrs->dtype = dtype; - attrs->right = right; - return Call(op, {sorted_sequence, values}, Attrs(attrs), {}); -} - -TVM_REGISTER_GLOBAL("relay.op._make.searchsorted").set_body_typed(MakeSearchSorted); - -RELAY_REGISTER_OP("searchsorted") - .describe( - R"doc(Find indices where elements should be inserted to maintain order. -If `sorted_sequence` is N-dimensional, the innermost dimension of -`values` are searched in the corresponding dimension of `sorted_sequence`. -)doc" TVM_ADD_FILELINE) - .set_num_inputs(2) - .set_attrs_type() - .add_argument("sorted_sequence", "Tensor", - "Monotonically increasing sequence on the innermost dimension.") - .add_argument("values", "Tensor", "Values to search for.") - .set_support_level(6) - .add_type_rel("SearchSorted", SearchSortedRel); - -} // namespace relay -} // namespace tvm diff --git a/src/relay/op/algorithm/sort.cc b/src/relay/op/algorithm/sort.cc deleted file mode 100644 index 69a6ae55c71d..000000000000 --- a/src/relay/op/algorithm/sort.cc +++ /dev/null @@ -1,65 +0,0 @@ -/* - * 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 sort.cc - * \brief Sort operators - */ -#include -#include - -namespace tvm { -namespace relay { - -bool SortRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - // `types` contains: [data, result] - ICHECK_EQ(types.size(), 2); - const auto* data = types[0].as(); - if (data == nullptr) { - ICHECK(types[0].as()) - << "Sort: expect input type to be TensorType but get " << types[0]; - return false; - } - reporter->Assign(types[1], TensorType(data->shape, data->dtype)); - return true; -} - -Expr MakeSort(Expr data, int axis, bool is_ascend) { - auto attrs = make_object(); - attrs->axis = axis; - attrs->is_ascend = is_ascend; - static const Op& op = Op::Get("sort"); - return Call(op, {data}, Attrs(attrs), {}); -} - -TVM_REGISTER_GLOBAL("relay.op._make.sort").set_body_typed(MakeSort); - -RELAY_REGISTER_OP("sort") - .describe(R"doc(Returns the indices that would sort an -input array along the given axis. -)doc" TVM_ADD_FILELINE) - .set_num_inputs(1) - .set_attrs_type() - .add_argument("data", "Tensor", "Input data.") - .set_support_level(6) - .add_type_rel("Sort", SortRel); - -} // namespace relay -} // namespace tvm diff --git a/src/relay/op/algorithm/topk.cc b/src/relay/op/algorithm/topk.cc deleted file mode 100644 index c9f0a4396b06..000000000000 --- a/src/relay/op/algorithm/topk.cc +++ /dev/null @@ -1,133 +0,0 @@ -/* - * 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 topk.cc - * \brief TopK operators - */ -#include -#include -#include -#include - -#include "../../transforms/infer_layout_utils.h" - -namespace tvm { -namespace relay { - -TVM_REGISTER_NODE_TYPE(TopKAttrs); - -InferCorrectLayoutOutput TopKInferCorrectLayout(const Attrs& attrs, - const Array& new_in_layouts, - const Array& old_in_layouts, - const Array& old_in_types) { - const auto* attrs_ptr = attrs.as(); - ICHECK(attrs_ptr); - ObjectPtr param = make_object(*attrs_ptr); - - Array> old_in_shapes; - for (auto old_in_t : old_in_types) { - ICHECK(old_in_t.as()); - old_in_shapes.push_back(old_in_t.as()->shape); - } - - size_t axis = - param->axis < 0 ? param->axis + old_in_shapes[0].size() : static_cast(param->axis); - - Layout ret = Layout::Undef(); - - // If new_in_layouts are defined, this code tries to modify the layout. - if (new_in_layouts.defined() && old_in_layouts.defined()) { - const auto& sp_dim = old_in_layouts[0][axis]; - auto new_index = new_in_layouts[0].IndexOf(sp_dim); - param->axis = new_index; - ret = new_in_layouts[0]; - } else if (old_in_layouts.defined()) { - ret = old_in_layouts[0]; - } - - // TopK has 2 outputs, Values and Indices - return InferCorrectLayoutOutput({ret}, {ret, ret}, Attrs(param)); -} - -bool TopKRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - // `types` contains: [data, result] - const TopKAttrs* param = attrs.as(); - ICHECK_EQ(types.size(), 2); - const auto* data = types[0].as(); - if (data == nullptr) return false; - int ndim = data->shape.size(); - int axis = param->axis; - if (axis < 0) { - axis += ndim; - } - ICHECK(axis >= 0 && axis < ndim); - Array out_shape; - for (int i = 0; i < ndim; ++i) { - if (i != axis) { - out_shape.push_back(data->shape[i]); - } else { - const Integer& ck = param->k.value(); - if (ck->value < 1) { - out_shape.push_back(data->shape[i]); - } else { - out_shape.push_back(ck); - } - } - } - auto values_ty = TensorType(out_shape, data->dtype); - auto indices_ty = TensorType(out_shape, param->dtype); - if (param->ret_type == "both") { - reporter->Assign(types[1], TupleType({values_ty, indices_ty})); - } else if (param->ret_type == "values") { - reporter->Assign(types[1], values_ty); - } else if (param->ret_type == "indices") { - reporter->Assign(types[1], indices_ty); - } else { - LOG(FATAL) << "Unsupported ret type: " << param->ret_type; - } - return true; -} - -Expr MakeTopK(Expr data, int k, int axis, String ret_type, bool is_ascend, DataType dtype) { - auto attrs = make_object(); - attrs->k = Integer(k); - attrs->axis = axis; - attrs->ret_type = ret_type; - attrs->is_ascend = is_ascend; - attrs->dtype = dtype; - static const Op& op = Op::Get("topk"); - return Call(op, {data}, Attrs(attrs), {}); -} - -TVM_REGISTER_GLOBAL("relay.op._make.topk").set_body_typed(MakeTopK); - -RELAY_REGISTER_OP("topk") - .describe(R"doc(Get the top k elements in an input tensor along the given axis. -)doc" TVM_ADD_FILELINE) - .set_num_inputs(1) - .set_attrs_type() - .add_argument("data", "Tensor", "Input data.") - .set_attr("FInferCorrectLayout", TopKInferCorrectLayout) - .set_support_level(6) - .add_type_rel("TopK", TopKRel); - -} // namespace relay -} // namespace tvm diff --git a/src/relay/op/annotation/annotation.cc b/src/relay/op/annotation/annotation.cc deleted file mode 100644 index bd3162dfde86..000000000000 --- a/src/relay/op/annotation/annotation.cc +++ /dev/null @@ -1,205 +0,0 @@ -/* - * 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 src/relay/op/annotation/annotation.cc - * \brief Helpers for working with various 'annotations' attributes. - */ - -#include "./annotation.h" - -#include -#include -#include -#include -#include -#include - -#include "../../transforms/infer_layout_utils.h" -#include "../type_relations.h" - -namespace tvm { -namespace relay { - -Expr StopFusion(Expr data) { - static const Op& op = Op::Get("annotation.stop_fusion"); - return Call(op, {data}, Attrs{}, {}); -} - -TVM_REGISTER_GLOBAL("relay.op.annotation._make.stop_fusion").set_body_typed([](Expr data) { - return StopFusion(data); -}); - -RELAY_REGISTER_OP("annotation.stop_fusion") - .describe( - R"code(Annotate an expression to prevent it being fused with following expressions.)code" TVM_ADD_FILELINE) - .set_num_inputs(1) - .add_argument("data", "Tensor", "The input data.") - .add_type_rel("Identity", IdentityRel) - .set_support_level(10) - .set_attr("TOpPattern", kOpaque) - .set_attr("TOpIsStateful", false) - .set_attr("FInferCorrectLayout", ElemwiseArbitraryLayout) - .set_attr("FTVMCompute", - [](const Attrs& attrs, const Array& inputs, - const Type& out_dtype) -> Array { - return {topi::identity(inputs[0])}; - }); - -// relay.annotation.cast_hint -TVM_REGISTER_NODE_TYPE(CastHintAttrs); - -Expr CastHint(Expr data, DataType dtype) { - auto attrs = make_object(); - attrs->dtype = dtype; - static const Op& op = Op::Get("annotation.cast_hint"); - return Call(op, {data}, Attrs{attrs}, {}); -} - -RELAY_REGISTER_OP("annotation.cast_hint") - .describe( - R"code(Annotate an expression to be cast into specific data type.)code" TVM_ADD_FILELINE) - .set_num_inputs(1) - .add_argument("data", "Tensor", "The input data.") - .add_type_rel("Identity", IdentityRel) - .set_support_level(10) - .set_attr("TOpPattern", kOpaque) - .set_attr("TOpIsStateful", false) - .set_attr("FInferCorrectLayout", ElemwiseArbitraryLayout) - .set_attr("FTVMCompute", - [](const Attrs& attrs, const Array& inputs, - const Type& out_dtype) -> Array { - return {topi::identity(inputs[0])}; - }); - -RELAY_REGISTER_OP("annotation.bitpack_start") - .describe(R"code( -Mark the start of bitpacking. -)code" TVM_ADD_FILELINE) - .set_num_inputs(1) - .add_argument("data", "Tensor", "The input data.") - .set_support_level(10) - .add_type_rel("Identity", IdentityRel) - .set_attr("TOpPattern", kOpaque) - .set_attr("TOpIsStateful", false) - .set_attr("FInferCorrectLayout", ElemwiseArbitraryLayout) - .set_attr("FTVMCompute", - [](const Attrs& attrs, const Array& inputs, - const Type& out_dtype) -> Array { - return {topi::identity(inputs[0])}; - }); - -RELAY_REGISTER_OP("annotation.bitpack_end") - .describe(R"code( -Mark the end of bitpacking. -)code" TVM_ADD_FILELINE) - .set_num_inputs(1) - .add_argument("data", "Tensor", "The input data.") - .set_support_level(10) - .add_type_rel("Identity", IdentityRel) - .set_attr("TOpPattern", kOpaque) - .set_attr("TOpIsStateful", false) - .set_attr("FInferCorrectLayout", ElemwiseArbitraryLayout) - .set_attr("FTVMCompute", - [](const Attrs& attrs, const Array& inputs, - const Type& out_dtype) -> Array { - return {topi::identity(inputs[0])}; - }); - -TVM_REGISTER_GLOBAL("relay.op.annotation._make.checkpoint").set_body_typed([](Expr data) { - static const Op& op = Op::Get("annotation.checkpoint"); - return Call(op, {data}, Attrs{}, {}); -}); - -RELAY_REGISTER_OP("annotation.checkpoint") - .describe(R"code( -Mark a checkpoint for checkpointing memory optimization. -)code" TVM_ADD_FILELINE) - .set_num_inputs(1) - .set_support_level(10) - .add_argument("data", "Tensor", "The input data.") - .add_type_rel("Identity", IdentityRel) - .set_attr("TOpPattern", kOpaque) - .set_attr("TOpIsStateful", false) - .set_attr("FInferCorrectLayout", ElemwiseArbitraryLayout) - .set_attr("FTVMCompute", - [](const Attrs& attrs, const Array& inputs, - const Type& out_dtype) -> Array { - Array outputs; - for (size_t i = 0; i < inputs.size(); ++i) { - outputs.push_back(topi::identity(inputs[i])); - } - return outputs; - }); - -TVM_REGISTER_NODE_TYPE(CompilerAttrs); - -RELAY_REGISTER_OP("annotation.compiler_begin") - .describe(R"code( -Beginning of a region that is handled by a given compiler. -)code" TVM_ADD_FILELINE) - .set_num_inputs(1) - .add_argument("data", "Tensor", "The input data.") - .set_support_level(10) - .add_type_rel("Identity", IdentityRel) - .set_attr("TOpPattern", kOpaque) - .set_attr("TOpIsStateful", false) - .set_attr("FInferCorrectLayout", ElemwiseArbitraryLayout) - .set_attr("FTVMCompute", - [](const Attrs& attrs, const Array& inputs, - const Type& out_dtype) -> Array { - return {topi::identity(inputs[0])}; - }); - -TVM_REGISTER_GLOBAL("relay.op.annotation._make.compiler_begin") - .set_body_typed([](Expr expr, String compiler) { - auto attrs = make_object(); - attrs->compiler = compiler; - static const Op& op = Op::Get("annotation.compiler_begin"); - return Call(op, {expr}, Attrs(attrs), {}); - }); - -RELAY_REGISTER_OP("annotation.compiler_end") - .describe(R"code( -End of a region that is handled by a given compiler. -)code" TVM_ADD_FILELINE) - .set_num_inputs(1) - .add_argument("data", "Tensor", "The input data.") - .set_support_level(10) - .add_type_rel("Identity", IdentityRel) - .set_attr("TOpPattern", kOpaque) - .set_attr("TOpIsStateful", false) - .set_attr("FInferCorrectLayout", ElemwiseArbitraryLayout) - .set_attr("FTVMCompute", - [](const Attrs& attrs, const Array& inputs, - const Type& out_dtype) -> Array { - return {topi::identity(inputs[0])}; - }); - -TVM_REGISTER_GLOBAL("relay.op.annotation._make.compiler_end") - .set_body_typed([](Expr expr, String compiler) { - auto attrs = make_object(); - attrs->compiler = compiler; - static const Op& op = Op::Get("annotation.compiler_end"); - return Call(op, {expr}, Attrs(attrs), {}); - }); - -} // namespace relay -} // namespace tvm diff --git a/src/relay/op/annotation/annotation.h b/src/relay/op/annotation/annotation.h deleted file mode 100644 index 1675b7281ebb..000000000000 --- a/src/relay/op/annotation/annotation.h +++ /dev/null @@ -1,46 +0,0 @@ -/* - * 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 relay/op/annotation/annotation.h - * \brief Helpers for working with various 'annotation' attributes. - */ -#ifndef TVM_RELAY_OP_ANNOTATION_ANNOTATION_H_ -#define TVM_RELAY_OP_ANNOTATION_ANNOTATION_H_ - -#include -#include -#include -#include - -#include - -namespace tvm { -namespace relay { - -/*! \brief Wraps \p data in a "stop_fusion" annotation. */ -Expr StopFusion(Expr data); - -/*! \brief Wraps \p data in a "cast_hint" annotation for \p dtype. */ -Expr CastHint(Expr data, DataType dtype); - -} // namespace relay -} // namespace tvm - -#endif // TVM_RELAY_OP_ANNOTATION_ANNOTATION_H_ diff --git a/src/relay/op/call/call.cc b/src/relay/op/call/call.cc deleted file mode 100644 index ab8e7e12d213..000000000000 --- a/src/relay/op/call/call.cc +++ /dev/null @@ -1,139 +0,0 @@ -/* - * 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 src/relay/op/call/call.cc - * \brief Operators for calling lowered functions. - */ - -#include "./call.h" - -#include -#include -#include -#include - -#include "../../transforms/infer_layout_utils.h" - -namespace tvm { -namespace relay { - -TVM_REGISTER_NODE_TYPE(CallLoweredAttrs); - -// call_lowered -bool CallLoweredRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - // Types = [func, call_args, ret_type] - if (types.size() != 3u) { - return false; - } - const auto* func_type = types[0].as(); - if (!func_type) { - return false; - } - - const auto* tuple_type_node = types[1].as(); - if (!tuple_type_node) { - return false; - } - - // Constraint to ensure function arguments are the same type as the inputs to the function (modulo - // the Tuple wrapper) - reporter->Assign(GetRef(tuple_type_node), TupleType(func_type->arg_types, {})); - // Constraint to ensure the output of call_lowered is the same as the function's return type - reporter->Assign(types[2], func_type->ret_type); - return true; -} - -const Op& CallLoweredOp() { return Op::Get("call_lowered"); } - -Call CallLowered(GlobalVar lowered_func, Array args, CallLoweredAttrs call_lowered_attrs, - Span span) { - auto attrs = make_object(std::move(call_lowered_attrs)); - return Call(CallLoweredOp(), {std::move(lowered_func), Tuple(std::move(args))}, - Attrs(std::move(attrs)), /*type_args=*/{}, std::move(span)); -} - -TVM_REGISTER_GLOBAL("relay.op.call_lowered") - .set_body_typed([](Expr lowered_func, Array args, Attrs attrs, Span span) { - const auto* lowered_func_node = lowered_func.as(); - ICHECK(lowered_func_node) << "Function to call should be GlobalVarNode, but got:" << std::endl - << PrettyPrint(lowered_func); - const auto* call_lowered_attrs = attrs.as(); - ICHECK(call_lowered_attrs) << "Expected attributes to be CallLoweredAttrs, but got " - << attrs->GetTypeKey(); - return CallLowered(GetRef(lowered_func_node), std::move(args), *call_lowered_attrs, - std::move(span)); - }); - -RELAY_REGISTER_OP("call_lowered") - .describe(R"code(Invoke an operation compiled by TVM.)code" TVM_ADD_FILELINE) - .set_num_inputs(2) - .set_attrs_type() - .add_argument("func", "Function", "The lowered function to call.") - .add_argument("call_args", "Tuple", "The input tensors.") - .add_type_rel("CallLoweredRel", CallLoweredRel) - .set_support_level(10) - .set_attr("TOpPattern", kOpaque) - .set_attr("TOpIsStateful", false) - .set_attr("TNonComputational", true) - .set_attr("FInferCorrectLayout", ElemwiseArbitraryLayout); - -CallLoweredProps GetCallLoweredProps(const CallNode* call_node) { - if (call_node->op == CallLoweredOp()) { - ICHECK(call_node->args.size() == 2) << "Expected call_lowered to have 2 arguments."; - const auto* function_node = call_node->args[0].as(); - ICHECK(function_node) << "Expected first arg to call_lowered to be a GlobalVar. "; - - const auto* tuple_args = call_node->args[1].as(); - ICHECK(tuple_args) << "Expected second arg to call_lowered to be a Tuple of input arguments."; - - ICHECK(call_node->attrs.defined()) << "Expecting call_lowered to have attributes."; - const auto* call_lowered_attrs = call_node->attrs.as(); - ICHECK(call_lowered_attrs) << "Expected call_lowered op to have CallLoweredAttrs, but found " - << call_node->attrs->GetTypeKey(); - // If the call_node has type_args then they are for the polymorphic 'call_lowered' operator - // itself which expects the function type and argument type as parameters. - return {GetRef(function_node), tuple_args->fields, *call_lowered_attrs}; - } - return {}; -} - -Call GetAnyCall(const CallNode* call_node) { - CallLoweredProps props = GetCallLoweredProps(call_node); - if (props.lowered_func.defined()) { - auto call_lowered_attrs = make_object(props.attrs); - return Call(std::move(props.lowered_func), std::move(props.arguments), - Attrs(std::move(call_lowered_attrs)), - /*type_args=*/{}, call_node->span); - } else { - return GetRef(call_node); - } -} - -bool IsReshapeOnly(const CallLoweredProps& props) { - if (props.attrs.metadata.count("relay_attrs")) { - auto dict_attrs = Downcast(props.attrs.metadata["relay_attrs"]); - return dict_attrs.HasNonzeroAttr(attr::kReshapeOnly); - } - return false; -} - -} // namespace relay -} // namespace tvm diff --git a/src/relay/op/call/call.h b/src/relay/op/call/call.h deleted file mode 100644 index 6193c9249ee2..000000000000 --- a/src/relay/op/call/call.h +++ /dev/null @@ -1,96 +0,0 @@ -/* - * 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 src/relay/op/call/call.h - * \brief Operators for calling lowered functions. - */ -#ifndef TVM_RELAY_OP_CALL_CALL_H_ -#define TVM_RELAY_OP_CALL_CALL_H_ - -#include -#include - -#include - -namespace tvm { -namespace relay { - -/*! - * \brief Returns the Relay call_lowered op. Use this helper to avoid extraneous calls to - * Registry::Get. - */ -const Op& CallLoweredOp(); - -/*! - * \brief Helper to construct a Relay call with the "call_lowered" op. - * - * The callee must: - * - Be a global bound to a PrimFunc or an externally defined functions. - * - Accept only tensor arguments and return tensor results. - * - Arguments and results correspond to the flattened form (see FlattenTupleType) of the - * Relay Function type. - * - Return results by output pointer, ie use DPS. - * The arguments remain in Relay form (ie not flattened). - * The result remains in Relay form (ie returned from the call and not flattened). - * - * \param lowered_func Lowered function to call with call_lowered. - * \param args Arguments to be passed to the function. - * \param call_lowered_attrs Function attributes. - * \param span TVM span for propagating debugging info. - * \return - */ -Call CallLowered(GlobalVar lowered_func, Array args, CallLoweredAttrs call_lowered_attrs, - Span span); - -/*! - * \brief Lowered function and the arguments to call it with. - */ -struct CallLoweredProps { - /*! \brief Global variable pointing to the lowered function. */ - GlobalVar lowered_func; - /*! \brief Array of the arguments to call lowered_func with. */ - Array arguments; - /*! \brief Attributes from the call_lowered op. */ - CallLoweredAttrs attrs; -}; - -/*! - * \brief Helper to extract the lowered function and its arguments from a Call("call_lowered", ...). - * Returns the null/empty \p CallLoweredProps if \p call_node is not in that form. - */ -CallLoweredProps GetCallLoweredProps(const CallNode* call_node); - -/*! - * \brief Returns \p call_node in 'standard' Relay form. Ie if \p call_node is a call_lowered - * then returns it in un-lowered form, otherwise returns \p call_node directly. - * - * Useful for passes which can act uniformly on calls irrespective of their form. - */ -Call GetAnyCall(const CallNode* call_node); - -/*! - * \brief Returns true if lowered call described by \p props is to a reshape primitive. - */ -bool IsReshapeOnly(const CallLoweredProps& props); - -} // namespace relay -} // namespace tvm - -#endif // TVM_RELAY_OP_CALL_CALL_H_ diff --git a/src/relay/op/debug.cc b/src/relay/op/debug.cc deleted file mode 100644 index 4b5e7d97f87d..000000000000 --- a/src/relay/op/debug.cc +++ /dev/null @@ -1,71 +0,0 @@ -/* - * 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 nn.cc - * \brief Property def of nn operators. - */ - -#include -#include -#include -#include - -#include - -#include "./op_common.h" -#include "./type_relations.h" - -namespace tvm { -namespace relay { - -TVM_REGISTER_NODE_TYPE(DebugAttrs); - -Array DebugCompute(const Attrs& attrs, const Array& inputs, - const Type& out_type) { - return Array{topi::identity(inputs[0])}; -} - -RELAY_REGISTER_OP("debug") - .describe(R"code(Enter the interpreter's debugger. - -)code" TVM_ADD_FILELINE) - .set_num_inputs(1) - .add_argument("program", "Tuple", "The program to execute before debugging.") - .set_support_level(1) - .set_attrs_type() - .add_type_rel("Debug", IdentityRel) - .set_attr("TOpPattern", kOpaque) - .set_attr("FTVMCompute", DebugCompute); - -Expr MakeDebug(Expr expr, String name) { - auto dattrs = make_object(); - if (name.size() > 0) { - dattrs->debug_func = EnvFunc::Get(name); - } else { - dattrs->debug_func = EnvFunc(); - } - static const Op& op = Op::Get("debug"); - return Call(op, {expr}, Attrs(dattrs), {}); -} - -TVM_REGISTER_GLOBAL("relay.op._make.debug").set_body_typed(MakeDebug); - -} // namespace relay -} // namespace tvm diff --git a/src/relay/op/dyn/algorithm/topk.cc b/src/relay/op/dyn/algorithm/topk.cc deleted file mode 100644 index 0ce0a18b2170..000000000000 --- a/src/relay/op/dyn/algorithm/topk.cc +++ /dev/null @@ -1,107 +0,0 @@ -/* - * 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 topk.cc - * \brief TopK operators - */ -#include -#include -#include - -namespace tvm { -namespace relay { -namespace dyn { - -bool TopKRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - // `types` contains: [data, k, result] - const TopKAttrs* param = attrs.as(); - ICHECK_EQ(types.size(), 3); - const auto* data = types[0].as(); - const auto* k = types[1].as(); - if (data == nullptr) { - ICHECK(types[0].as()) - << "tile: expect input type to be TensorType but get " << types[0]; - return false; - } - if (k == nullptr) { - ICHECK(types[1].as()) - << "tile: expect input type to be TensorType but get " << types[1]; - return false; - } - ICHECK(k->shape.size() <= 1) << "Parameter k must be a Scalar or a Tensor of shape (1, )"; - if (k->shape.size() == 1) { - const IntImmNode* k_shape = k->shape[0].as(); - ICHECK(k_shape) << "Parameter k must have static shape"; - ICHECK_EQ(k_shape->value, 1) << "Parameter k must be a Scalar or a Tensor of shape (1, )"; - } - int ndim = data->shape.size(); - int axis = param->axis; - if (axis < 0) { - axis += ndim; - } - ICHECK(axis >= 0 && axis < ndim); - Array out_shape; - for (int i = 0; i < ndim; ++i) { - if (i != axis) { - out_shape.push_back(data->shape[i]); - } else { - out_shape.push_back(Any()); - } - } - auto values_ty = TensorType(out_shape, data->dtype); - auto indices_ty = TensorType(out_shape, param->dtype); - if (param->ret_type == "both") { - reporter->Assign(types[2], TupleType({values_ty, indices_ty})); - } else if (param->ret_type == "values") { - reporter->Assign(types[2], values_ty); - } else if (param->ret_type == "indices") { - reporter->Assign(types[2], indices_ty); - } else { - LOG(FATAL) << "Unsupported ret type: " << param->ret_type; - } - return true; -} - -Expr MakeTopK(Expr data, Expr k, int axis, String ret_type, bool is_ascend, DataType dtype) { - auto attrs = make_object(); - attrs->axis = axis; - attrs->ret_type = ret_type; - attrs->is_ascend = is_ascend; - attrs->dtype = dtype; - static const Op& op = Op::Get("dyn.topk"); - return Call(op, {data, k}, Attrs(attrs), {}); -} - -TVM_REGISTER_GLOBAL("relay.op.dyn._make.topk").set_body_typed(MakeTopK); - -RELAY_REGISTER_OP("dyn.topk") - .describe(R"doc(Get the top k elements in an input tensor along the given axis. -)doc" TVM_ADD_FILELINE) - .set_num_inputs(2) - .set_attrs_type() - .add_argument("data", "Tensor", "Input data.") - .add_argument("k", "Tensor", "Number of top elements.") - .set_support_level(6) - .add_type_rel("DynTopK", TopKRel); - -} // namespace dyn -} // namespace relay -} // namespace tvm diff --git a/src/relay/op/dyn/image/resize.cc b/src/relay/op/dyn/image/resize.cc deleted file mode 100644 index 1f5f6b43763f..000000000000 --- a/src/relay/op/dyn/image/resize.cc +++ /dev/null @@ -1,115 +0,0 @@ -/* - * 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 resize.cc - * \brief Image resize operators - */ -#include -#include -#include - -#include "../../op_common.h" - -namespace tvm { -namespace relay { -namespace dyn { - -TVM_REGISTER_NODE_TYPE(Resize2DAttrs); - -bool Resize2DRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - // {data, size, roi, out} - ICHECK_EQ(types.size(), 4); - const auto* data = types[0].as(); - if (data == nullptr) return false; - - static const Layout kNCHW("NCHW"); - - const Resize2DAttrs* param = attrs.as(); - ICHECK(param != nullptr); - const Layout in_layout(param->layout); - auto layout_converter = tir::BijectiveLayout(in_layout, kNCHW); - ICHECK(layout_converter.defined()) - << "Resize only support input layouts that are convertible from NCHW." - << " But got " << in_layout; - - auto oshape = layout_converter.ForwardShape(data->shape); - oshape.Set(2, Any()); - oshape.Set(3, Any()); - - DataType out_dtype = param->out_dtype; - if (out_dtype.bits() == 0) { - out_dtype = data->dtype; - } - - // assign output type - reporter->Assign(types[3], TensorType(layout_converter.BackwardShape(oshape), out_dtype)); - return true; -} - -// Positional relay function to create image operator -// used by frontend FFI. -Expr MakeResize2D(Expr data, Expr size, Expr roi, String layout, String method, - String coordinate_transformation_mode, String rounding_method, double cubic_alpha, - double cubic_exclude, double extrapolation_value, DataType out_dtype) { - auto attrs = make_object(); - attrs->layout = std::move(layout); - attrs->method = std::move(method); - attrs->coordinate_transformation_mode = coordinate_transformation_mode; - attrs->rounding_method = rounding_method; - attrs->cubic_alpha = cubic_alpha; - attrs->cubic_exclude = cubic_exclude; - attrs->extrapolation_value = extrapolation_value; - attrs->out_dtype = out_dtype; - static const Op& op = Op::Get("dyn.image.resize2d"); - return Call(op, {data, size, roi}, Attrs(attrs), {}); -} - -TVM_REGISTER_GLOBAL("relay.op.dyn.image._make.resize2d").set_body_typed(MakeResize2D); - -RELAY_REGISTER_OP("dyn.image.resize2d") - .describe(R"code(Perform resize to input array with nearest neighbour or bilinear interpolation. - -- **data**: data is 4D array of shape - (batch_size, channels, in_height, in_width) for NCHW - (batch_size, in_height, in_width, channels) for NHWC - -- **size**: data is 2D array of shape (2,) with values - (new_height, new_width) - -- **out**: Output is 4D array of shape - for layout NCHW - (batch_size, channels, size[0], size[1]) - - for layout NHWC - (batch_size, size[0], size[1], channels) -)code" TVM_ADD_FILELINE) - .set_attrs_type() - .set_num_inputs(3) - .add_argument("data", "Tensor", "The input tensor.") - .add_argument("size", "Tensor", "The output size tensor.") - .add_argument("roi", "Tensor", "The region of interest for tf_crop_and_resize.") - .set_support_level(5) - .add_type_rel("DynResize2D", Resize2DRel) - .set_attr("TOpPattern", kInjective); - -} // namespace dyn -} // namespace relay -} // namespace tvm diff --git a/src/relay/op/dyn/nn/pad.cc b/src/relay/op/dyn/nn/pad.cc deleted file mode 100644 index 101ad5de7f57..000000000000 --- a/src/relay/op/dyn/nn/pad.cc +++ /dev/null @@ -1,125 +0,0 @@ -/* - * 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 pad.cc - * \brief Implementation of dynamic pad - */ -#include -#include -#include -#include -#include - -#include - -#include "../../make_op.h" -#include "../../op_common.h" - -namespace tvm { -namespace relay { -namespace dyn { - -// relay.dyn.nn.pad - -bool PadRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - // types = [data_type, pad_width_type, pad_value_type, ret_type] - ICHECK_EQ(types.size(), 4); - const auto* data = types[0].as(); - if (data == nullptr) return false; - - const auto* pad_width = types[1].as(); - if (pad_width == nullptr) return false; - - const auto* pad_value = types[2].as(); - if (pad_value == nullptr) return false; - - int data_rank = data->shape.size(); - ICHECK(data_rank) << "Data shape must have static rank"; - - int pad_width_rank = pad_width->shape.size(); - ICHECK_EQ(pad_width_rank, 2) << "Pad width must be 2D"; - - const PadAttrs* param = attrs.as(); - ICHECK(param != nullptr); - - std::vector oshape; - for (int i = 0; i < data_rank; i++) { - oshape.push_back(Any()); - } - - reporter->Assign(types[3], TensorType(oshape, data->dtype)); - return true; -} - -Array PadCompute(const Attrs& attrs, const Array& inputs, - const Type& out_type) { - const auto* param = attrs.as(); - ICHECK(param); - - auto data = inputs[0]; - auto pad_width = inputs[1]; - - te::Tensor cast_pad_value = topi::cast(inputs[2], inputs[0]->dtype); - const PrimExpr& pad_value = cast_pad_value(Array()); - - Array pad_before; - Array pad_after; - - for (int i = 0; i < pad_width->shape[0].as()->value; ++i) { - pad_before.push_back(pad_width[i][0]); - pad_after.push_back(pad_width[i][1]); - } - - const auto* out_ttype = out_type.as(); - ICHECK(out_ttype != nullptr); - - return Array{topi::pad(inputs[0], pad_before, pad_after, pad_value, "T_pad", - topi::kElementWise, param->pad_mode, - &out_type.as()->shape)}; -} - -// Handler to create a call to the padding op used by front-end FFI -Expr MakePad(Expr data, Expr pad_width, Expr pad_value, String pad_mode) { - auto attrs = make_object(); - attrs->pad_mode = std::move(pad_mode); - static const Op& op = Op::Get("dyn.nn.pad"); - return Call(op, {data, pad_width, pad_value}, Attrs(attrs), {}); -} - -TVM_REGISTER_GLOBAL("relay.op.dyn.nn._make.pad").set_body_typed(MakePad); - -RELAY_REGISTER_OP("dyn.nn.pad") - .describe(R"code(Pad for n-D tensor. - -)code" TVM_ADD_FILELINE) - .set_attrs_type() - .set_num_inputs(3) - .add_argument("data", "Tensor", "Tensor that will be padded") - .add_argument("pad_width", "Tensor", "Tensor of how much to pad by") - .add_argument("pad_val", "double", "The value to fill the padded area with") - .set_support_level(2) - .add_type_rel("DynamicPad", PadRel) - .set_attr("TOpPattern", kInjective) - .set_attr("FTVMCompute", PadCompute); - -} // namespace dyn -} // namespace relay -} // namespace tvm diff --git a/src/relay/op/dyn/nn/upsampling.cc b/src/relay/op/dyn/nn/upsampling.cc deleted file mode 100644 index 93869757e96f..000000000000 --- a/src/relay/op/dyn/nn/upsampling.cc +++ /dev/null @@ -1,201 +0,0 @@ -/* - * 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 upsampling.cc - * \brief upsampling operator - */ - -#include "upsampling.h" - -#include -#include -#include -#include - -#include -#include - -#include "../../op_common.h" - -namespace tvm { -namespace relay { -namespace dyn { - -bool UpSamplingRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - // types = [data_type, scale_h_type, scale_w_type, ret_type] - ICHECK_EQ(types.size(), 4); - const auto* data = types[0].as(); - const auto* scale_h = types[1].as(); - const auto* scale_w = types[2].as(); - if (data == nullptr) return false; - if (scale_h == nullptr) return false; - if (scale_w == nullptr) return false; - - ICHECK_EQ(scale_h->shape.size(), 0); - ICHECK_EQ(scale_w->shape.size(), 0); - static const Layout kNCHW("NCHW"); - - const UpSamplingAttrs* param = attrs.as(); - ICHECK(param); - const Layout in_layout(param->layout); - - auto layout_converter = tir::BijectiveLayout(in_layout, kNCHW); - ICHECK(layout_converter.defined()) - << "UpSampling only supports input layouts that are convertible from NCHW." - << " But got " << in_layout; - - auto nchw_oshape = layout_converter.ForwardShape(data->shape); - - nchw_oshape.Set(2, Any()); - nchw_oshape.Set(3, Any()); - auto oshape = layout_converter.BackwardShape(nchw_oshape); - - reporter->Assign(types[3], TensorType(oshape, data->dtype)); - return true; -} - -// Positional relay function to create upsampling operator -// used by frontend FFI. -Expr MakeUpSampling(Expr data, Expr scale_h, Expr scale_w, String layout, String method, - bool align_corners) { - auto attrs = make_object(); - attrs->layout = std::move(layout); - attrs->method = std::move(method); - attrs->align_corners = align_corners; - - static const Op& op = Op::Get("dyn.nn.upsampling"); - return Call(op, {data, scale_h, scale_w}, Attrs(attrs), {}); -} - -TVM_REGISTER_GLOBAL("relay.op.dyn.nn._make.upsampling").set_body_typed(MakeUpSampling); - -RELAY_REGISTER_OP("dyn.nn.upsampling") - .describe( - R"code(Perform upsampling on input array with nearest neighbour or bilinear interpolation. - -- **data**: data is 4D array of shape - (batch_size, channels, in_height, in_width) for NCHW - (batch_size, in_height, in_width, channels) for NHWC - -- **scale_h**: scale_h is a double of the amount to scale height by - -- **scale_w**: scale_w is a double of the amount to scale width by - -- **out**: Output is 4D array of shape - for layout NCHW - (batch_size, channels, in_height*scale_h, in_width*scale_w) - - for layout NHWC - (batch_size, in_height*scale_h, in_width*scale_w, channels) - -)code" TVM_ADD_FILELINE) - .set_attrs_type() - .set_num_inputs(3) - .add_argument("data", "Tensor", "The input tensor.") - .add_argument("scale_h", "double", "The scale for the height.") - .add_argument("scale_w", "double", "The scale for the width.") - .set_support_level(2) - .add_type_rel("DynamicUpSampling", UpSamplingRel) - .set_attr("FInferCorrectLayout", - UpsamplingInferCorrectLayout) - .set_attr("TOpPattern", kInjective); - -// UpSampling3D -bool UpSampling3DRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - // types = [data_type, scale_d_type, scale_h_type, scale_w_type, ret_type] - ICHECK_EQ(types.size(), 5); - const auto* data = types[0].as(); - if (data == nullptr) return false; - - static const Layout kNCDHW("NCDHW"); - - const UpSampling3DAttrs* param = attrs.as(); - ICHECK(param != nullptr); - const Layout in_layout(param->layout); - - auto layout_converter = tir::BijectiveLayout(in_layout, kNCDHW); - ICHECK(layout_converter.defined()) - << "UpSampling3D only support input layouts that are convertible from NCDHW." - << " But got " << in_layout; - - auto ncdhw_oshape = layout_converter.ForwardShape(data->shape); - - ncdhw_oshape.Set(2, Any()); - ncdhw_oshape.Set(3, Any()); - ncdhw_oshape.Set(4, Any()); - - auto oshape = layout_converter.BackwardShape(ncdhw_oshape); - - reporter->Assign(types[4], TensorType(oshape, data->dtype)); - return true; -} - -Expr MakeUpSampling3D(Expr data, Expr scale_d, Expr scale_h, Expr scale_w, String layout, - String method, String coordinate_transformation_mode) { - auto attrs = make_object(); - attrs->layout = std::move(layout); - attrs->method = std::move(method); - attrs->coordinate_transformation_mode = coordinate_transformation_mode; - - static const Op& op = Op::Get("dyn.nn.upsampling3d"); - return Call(op, {data, scale_d, scale_h, scale_w}, Attrs(attrs), {}); -} - -TVM_REGISTER_GLOBAL("relay.op.dyn.nn._make.upsampling3d").set_body_typed(MakeUpSampling3D); - -RELAY_REGISTER_OP("dyn.nn.upsampling3d") - .describe(R"code(Perform upsampling on input array with nearest neighbour or -bilinear interpolation. - -- **data**: data is 5D array of shape - (batch_size, channels, in_depth, in_height, in_width) for NCDHW - (batch_size, in_depth, in_height, in_width, channels) for NDHWC - -- **scale_d**: scale_d is a double of the amount to scale depth by - -- **scale_h**: scale_h is a double of the amount to scale height by - -- **scale_w**: scale_w is a double of the amount to scale width by - -- **out**: Output is 5D array of shape - for layout NCDHW - (batch_size, channels, in_depth*scale_d, in_height*scale_h, in_width*scale_w) - - for layout NDHWC - (batch_size, in_depth*scale_d, in_height*scale_h, in_width*scale_w, channels) - -)code" TVM_ADD_FILELINE) - .set_attrs_type() - .set_num_inputs(4) - .add_argument("data", "Tensor", "The input tensor.") - .add_argument("scale_d", "double", "The scale for the depth.") - .add_argument("scale_h", "double", "The scale for the height.") - .add_argument("scale_w", "double", "The scale for the width.") - .set_support_level(2) - .add_type_rel("DynamicUpSampling3D", UpSampling3DRel) - .set_attr("FInferCorrectLayout", - UpsamplingInferCorrectLayout) - .set_attr("TOpPattern", kInjective); - -} // namespace dyn -} // namespace relay -} // namespace tvm diff --git a/src/relay/op/dyn/nn/upsampling.h b/src/relay/op/dyn/nn/upsampling.h deleted file mode 100644 index f7b37abbc6d8..000000000000 --- a/src/relay/op/dyn/nn/upsampling.h +++ /dev/null @@ -1,72 +0,0 @@ -/* - * 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 src/relay/op/dyn/nn/upsampling.h - * \brief implementation of the InferCorrectLayout pass for dynamic upsampling - */ - -#ifndef TVM_RELAY_OP_DYN_NN_UPSAMPLING_H_ -#define TVM_RELAY_OP_DYN_NN_UPSAMPLING_H_ - -#include -#include - -#include "../../op_common.h" - -namespace tvm { -namespace relay { -namespace dyn { - -template -InferCorrectLayoutOutput UpsamplingInferCorrectLayout(const Attrs& attrs, - const Array& new_in_layouts, - const Array& old_in_layouts, - const Array& old_in_types) { - const auto* attrs_ptr = attrs.as(); - ICHECK(attrs_ptr); - ObjectPtr params = make_object(*attrs_ptr); - - if (new_in_layouts.defined()) { - ICHECK_GT(new_in_layouts.size(), 0); - - Layout raw_layout(params->layout); - Layout input = new_in_layouts[0]; - if (input.IndexOf(LayoutAxis::Get('W')) == raw_layout.IndexOf(LayoutAxis::Get('W')) && - input.IndexOf(LayoutAxis::Get('H')) == raw_layout.IndexOf(LayoutAxis::Get('H')) && - !input.Contains(LayoutAxis::Get('w')) && !input.Contains(LayoutAxis::Get('h')) && - (input.IndexOf(LayoutAxis::Get('D')) == -1 || - (input.IndexOf(LayoutAxis::Get('D')) == raw_layout.IndexOf(LayoutAxis::Get('D')) && - !input.Contains(LayoutAxis::Get('d'))))) { - params->layout = input.name(); // modify self to follow the input layout - } - } - - Layout inferred_layout(params->layout); - Layout param_layout("NCHW"); - return InferCorrectLayoutOutput({inferred_layout, param_layout, param_layout}, {inferred_layout}, - Attrs(params)); -} - -} // namespace dyn -} // namespace relay -} // namespace tvm - -#endif // TVM_RELAY_OP_DYN_NN_UPSAMPLING_H_ diff --git a/src/relay/op/dyn/tensor/transform.cc b/src/relay/op/dyn/tensor/transform.cc deleted file mode 100644 index d5cc6608662b..000000000000 --- a/src/relay/op/dyn/tensor/transform.cc +++ /dev/null @@ -1,759 +0,0 @@ -/* - * 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 transform.cc - * \brief Dynamic Transform operators. - */ -#include "transform.h" - -#include -#include -#include -#include -#include -#include -#include -#include - -#include -#include -#include - -#include "../../../transforms/infer_layout_utils.h" - -namespace tvm { -namespace relay { -namespace dyn { - -/* relay.dyn.reshape */ - -bool ReshapeRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - // types: [data, newshape, result] - ICHECK_EQ(types.size(), 3); - - const auto* data = types[0].as(); - if (data == nullptr) { - ICHECK(types[0].as()) - << "reshape: expect input type to be TensorType but get " << types[0]; - return false; - } - - Array oshape; - const auto* newshape = types[1].as(); - if (newshape == nullptr) { - ICHECK(types[1].as()) - << "reshape: expect input type to be TensorType but get " << types[1]; - return false; - } - - const IntImmNode* rank = newshape->shape[0].as(); - ICHECK(rank != nullptr) << "Dynamic Reshape doesn't support Dynamic Rank"; - for (int i = 0; i < rank->value; i++) { - oshape.push_back(Any()); - } - - reporter->Assign(types[2], TensorType(oshape, data->dtype)); - return true; -} - -Array ReshapeCompute(const Attrs& attrs, const Array& inputs, - const Type& out_type) { - const auto* out_ttype = out_type.as(); - ICHECK(out_ttype != nullptr); - Array newshape; - for (auto val : out_ttype->shape) { - if (val->IsInstance()) { - newshape.push_back(val.as()->ToVar()); - } else { - newshape.push_back(val); - } - } - return {topi::reshape(inputs[0], newshape)}; -} - -Expr MakeReshape(Expr data, Expr newshape, bool allowzero = false) { - auto attrs = make_object(); - attrs->allowzero = allowzero; - static const Op& op = Op::Get("dyn.reshape"); - return Call(op, {data, newshape}, Attrs(attrs), {}); -} - -TVM_REGISTER_GLOBAL("relay.op.dyn._make.reshape").set_body_typed(MakeReshape); - -RELAY_REGISTER_OP("dyn.reshape") - .describe(R"code(Reshapes the input array based on the values in the newshape array. - - To give user more convenience in without doing manual shape inference, - some dimensions of the shape can take special values from the set {0, -1, -3}. - The significance of each is explained below: - - ``0`` copy this dimension from the input to the output shape. - - .. code-block:: python - - data.shape = (2,3,4), newshape = (4,0,2), result.shape = (4,3,2) - data.shape = (2,3,4), newshape = (2,0,0), result.shape = (2,3,4) - - ``-1`` infers the dimension of the output shape by using the remainder of - the input dimensions keeping the size of the new array same as that of the input array. - At most one dimension of shape can be -1. - - .. code-block:: python - - data.shape = (2,3,4), newshape = (6,1,-1), result.shape = (6,1,4) - data.shape = (2,3,4), newshape = (3,-1,8), result.shape = (3,1,8) - data.shape = (2,3,4), newshape = (-1,), result.shape = (24,) - - ``-3`` use the product of two consecutive dimensions of the input shape - as the output dimension. - - .. code-block:: python - - data.shape = (2,3,4), newshape = (-3,4), result.shape = (6,4) - data.shape = (2,3,4,5), newshape = (-3,-3), result.shape = (6,20) - data.shape = (2,3,4), newshape = (0,-3), result.shape = (2,12) - - Special values -2 and -4 from the standard reshape op would introduce dynamic rank - in this op. Thus, they are not permitted. - - )code" TVM_ADD_FILELINE) - .set_num_inputs(2) - .set_attrs_type() - .add_argument("data", "Tensor", "The input tensor.") - .add_argument("newshape", "Tensor", "The shape of output tensor.") - .set_support_level(3) - .add_type_rel("DynamicReshape", ReshapeRel) - .set_attr("FTVMCompute", ReshapeCompute) - .set_attr("TOpPattern", kInjective) - .set_attr("TReshapeOp", true); - -// tile operator -// TVM_REGISTER_NODE_TYPE(TileAttrs); - -bool TileRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - // `types` contains: [data, reps, result] - ICHECK_EQ(types.size(), 3); - const auto* data = types[0].as(); - const auto* reps = types[1].as(); - if (data == nullptr) { - ICHECK(types[0].as()) - << "tile: expect input type to be TensorType but get " << types[0]; - return false; - } - if (reps == nullptr) { - ICHECK(types[1].as()) - << "tile: expect input type to be TensorType but get " << types[1]; - return false; - } - const IntImmNode* reps_shape = reps->shape[0].as(); - ICHECK(reps_shape) << "Parameter reps must have static shape"; - const size_t ndim = data->shape.size(); - const size_t rndim = reps_shape->value; - size_t tndim = (ndim > rndim) ? ndim : rndim; - std::vector oshape; - oshape.reserve(tndim); - for (size_t i = 0; i < tndim; ++i) { - oshape.emplace_back(Any()); - } - reporter->Assign(types[2], TensorType(oshape, data->dtype)); - return true; -} - -Array TileCompute(const Attrs& attrs, const Array& inputs, - const Type& out_type) { - ICHECK_EQ(inputs.size(), 2); - const auto* out_ttype = out_type.as(); - size_t rndim = inputs[1]->shape[0].as()->value; - return {topi::dyn_tile(inputs[0], out_ttype->shape, rndim)}; -} - -Expr MakeTile(Expr data, Expr reps) { - auto attrs = make_object(); - static const Op& op = Op::Get("dyn.tile"); - return Call(op, {data, reps}, Attrs(attrs), {}); -} - -TVM_REGISTER_GLOBAL("relay.op.dyn._make.tile").set_body_typed(MakeTile); - -RELAY_REGISTER_OP("dyn.tile") - .describe(R"code(Repeat the whole array multiple times. - -- **data**: The input data to the operator. -- **reps**: The number of times to repeat the operator. - -)code" TVM_ADD_FILELINE) - .set_num_inputs(2) - .set_attrs_type() - .add_argument("data", "Tensor", "The input tensor.") - .add_argument("reps", "Tensor", "The number of times to repeat the input on each axis.") - .set_support_level(3) - .add_type_rel("DynamicTile", TileRel) - .set_attr("FTVMCompute", TileCompute) - .set_attr("TOpPattern", kInjective); - -// broadcast_to operator -bool BroadCastToRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - // types = [data_type, broadcast_shape_type, ret_type] - ICHECK_EQ(types.size(), 3); - - const auto* input_type = types[0].as(); - const auto* target_type = types[1].as(); - if (target_type == nullptr) { - return false; - } - if (input_type == nullptr) { - return false; - } - auto out_dtype = input_type->dtype; - // rank must be static - const IntImmNode* rank = target_type->shape[0].as(); - ICHECK(rank) - << "Target shape must have static rank"; // rank must be static even in dyn pass - // could add support for dyn rank in futures - - std::vector oshape; - for (int i = 0; i < rank->value; ++i) { - oshape.push_back(Any()); - } - - reporter->Assign(types[2], TensorType(oshape, out_dtype)); - return true; -} - -Expr MakeBroadCastTo(Expr data, Expr shape) { - static const Op& op = Op::Get("dyn.broadcast_to"); - auto attrs = make_object(); - return Call(op, {data, shape}, Attrs(attrs), {}); -} - -Array BroadCastToCompute(const Attrs& attrs, const Array& inputs, - const Type& out_type) { - const auto* out_ttype = out_type.as(); - return {topi::broadcast_to(inputs[0], out_ttype->shape)}; -} - -TVM_REGISTER_GLOBAL("relay.op.dyn._make.broadcast_to").set_body_typed(MakeBroadCastTo); - -RELAY_REGISTER_OP("dyn.broadcast_to") - .describe(R"code(Broadcast the first input to match the shape argument. -)code" TVM_ADD_FILELINE) - .set_num_inputs(2) - .set_attrs_type() - .add_argument("data", "Tensor", "The input tensor.") - .add_argument("shape", "Tensor", "Target shape.") - .set_support_level(4) - .add_type_rel("DynamicBroadCastTo", BroadCastToRel) - .set_attr("FTVMCompute", BroadCastToCompute) - .set_attr("TOpPattern", kBroadcast); - -// zeros and ones operator -bool InitOpRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - // types = [zeros_shape, ret_type] - ICHECK_EQ(types.size(), 2); - const InitOpAttrs* param = attrs.as(); - const auto* fill_shape = types[0].as(); - DataType out_dtype = param->dtype; - - const IntImmNode* shape_shape = fill_shape->shape[0].as(); - ICHECK(shape_shape) << "Parameter shape must have static rank"; - - std::vector oshape; - for (int i = 0; i < shape_shape->value; ++i) { - oshape.push_back(Any()); - } - - reporter->Assign(types[1], TensorType(oshape, out_dtype)); - return true; -} - -Expr MakeZeros(Expr shape, DataType dtype) { - auto attrs = make_object(); - attrs->dtype = std::move(dtype); - static const Op& op = Op::Get("dyn.zeros"); - return Call(op, {shape}, Attrs(attrs), {}); -} - -TVM_REGISTER_GLOBAL("relay.op.dyn._make.zeros").set_body_typed(MakeZeros); - -RELAY_REGISTER_OP("dyn.zeros") - .describe(R"code(Fill array with zeros. - -)code" TVM_ADD_FILELINE) - .set_attrs_type() - .set_num_inputs(1) - .add_argument("shape", "Tensor", "Target shape.") - .set_support_level(3) - .add_type_rel("DynamicInitOp", InitOpRel); - -Expr MakeOnes(Expr shape, DataType dtype) { - auto attrs = make_object(); - attrs->dtype = std::move(dtype); - static const Op& op = Op::Get("dyn.ones"); - return Call(op, {shape}, Attrs(attrs), {}); -} - -TVM_REGISTER_GLOBAL("relay.op.dyn._make.ones").set_body_typed(MakeOnes); - -RELAY_REGISTER_OP("dyn.ones") - .describe(R"code(Fill array with ones. - -)code" TVM_ADD_FILELINE) - .set_attrs_type() - .set_num_inputs(1) - .add_argument("shape", "Tensor", "Target shape.") - .set_support_level(3) - .add_type_rel("DynamicInitOp", InitOpRel); - -bool OneHotRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - // `types` contains: [indices, on_value, off_value, result] - ICHECK_EQ(types.size(), 5); - const auto* indices = types[0].as(); - ICHECK(indices); - - const auto param = attrs.as(); - - Array oshape; - int ndim = indices->shape.size() + 1; - int indices_index = 0; - int true_axis = (param->axis == -1) ? indices->shape.size() : param->axis; - for (int i = 0; i < ndim; i++) { - if (i == true_axis) { - oshape.push_back(Any()); - } else { - oshape.push_back(indices->shape[indices_index++]); - } - } - - reporter->Assign(types[4], TensorType(oshape, param->dtype)); - return true; -} - -Array OneHotCompute(const Attrs& attrs, const Array& inputs, - const Type& out_type) { - const auto* param = attrs.as(); - ICHECK(param != nullptr); - const auto* out_ttype = out_type.as(); - return Array{topi::one_hot(inputs[0], inputs[1](), inputs[2](), -1, param->axis, - param->dtype, out_ttype->shape)}; -} - -Expr MakeOneHot(Expr indices, Expr on_value, Expr off_value, Expr depth, int axis, DataType dtype) { - auto attrs = make_object(); - attrs->axis = axis; - attrs->dtype = dtype; - static const Op& op = Op::Get("dyn.one_hot"); - return Call(op, {indices, on_value, off_value, depth}, Attrs(attrs), {}); -} - -TVM_REGISTER_GLOBAL("relay.op.dyn._make.one_hot").set_body_typed(MakeOneHot); - -RELAY_REGISTER_OP("dyn.one_hot") - .describe(R"code(Returns a one-hot tensor where the locations repsented by indices take value 1, - other locations take value 0. Final dimension is x depth. - - **indices** Locations to set to 1. - - **on_value** Value to fill at indices. - - **off_value** Value to fill at all other positions besides indices. - - **depth** Depth of the one-hot dimension. - - **axis** Axis to fill. - - **dtype**)code" TVM_ADD_FILELINE) - .set_attrs_type() - .set_num_inputs(4) - .add_argument("indices", "Tensor", "Locations to set to on_value.") - .add_argument("on_value", "Expr", "Value to fill at indices.") - .add_argument("off_value", "Expr", "Value to fill at all other positions besides indices.") - .add_argument("depth", "Expr", "Value to fill at all other positions besides indices.") - .set_support_level(10) - .add_type_rel("DynOneHot", OneHotRel) - .set_attr("FTVMCompute", OneHotCompute) - .set_attr("TOpPattern", kOutEWiseFusable); - -bool FullRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - ICHECK_EQ(types.size(), 3); - const InitOpAttrs* param = attrs.as(); - const auto* fill_value = types[0].as(); - const auto* fill_shape = types[1].as(); - if (fill_value == nullptr) { - return false; - } - if (fill_shape == nullptr) { - return false; - } - - DataType out_dtype = param->dtype; - if (out_dtype.bits() == 0) { - out_dtype = fill_value->dtype; - } - - ICHECK_EQ(fill_value->shape.size(), 0) - << "Fill value should be a scalar but has dimension " << fill_value->shape.size() << "."; - - const IntImmNode* rank = fill_shape->shape[0].as(); - ICHECK(rank) << "Parameter shape must have static rank"; - - std::vector oshape; - for (int i = 0; i < rank->value; ++i) { - oshape.push_back(Any()); - } - reporter->Assign(types[2], TensorType(oshape, out_dtype)); - return true; -} - -Expr MakeFull(Expr fill_value, Expr shape, DataType dtype) { - auto attrs = make_object(); - attrs->dtype = std::move(dtype); - static const Op& op = Op::Get("dyn.full"); - return Call(op, {fill_value, shape}, Attrs(attrs), {}); -} -Array FullCompute(const Attrs& attrs, const Array& inputs, - const Type& out_type) { - const auto* out_ttype = out_type.as(); - return {topi::full(out_ttype->shape, out_ttype->dtype, inputs[0]())}; -} -TVM_REGISTER_GLOBAL("relay.op.dyn._make.full").set_body_typed(MakeFull); - -RELAY_REGISTER_OP("dyn.full") - .describe(R"code(Fill array with scalar value. - -)code" TVM_ADD_FILELINE) - .set_attrs_type() - .set_num_inputs(2) - .add_argument("fill_value", "double", "The value to fill.") - .add_argument("shape", "Tensor", "Target shape.") - .set_support_level(3) - .add_type_rel("DynamicFull", FullRel) - .set_attr("FTVMCompute", FullCompute) - .set_attr("TOpPattern", kElemWise); - -bool StridedSliceRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - // [data, begin, end, strides, out] - ICHECK_EQ(types.size(), 5); - const StridedSliceAttrs* param = attrs.as(); - if (param == nullptr) { - return false; - } - const auto* data = types[0].as(); - if (data == nullptr) { - return false; - } - auto dshape = data->shape; - int64_t num_axis = dshape.size(); - - const auto* begin = types[1].as(); - if (begin == nullptr) { - return false; - } - ICHECK(begin); - - // calculate output shape - std::vector oshape(num_axis); - int64_t num_dynamic_axes = begin->shape[0].as()->value; - for (int64_t i = 0; i < num_dynamic_axes; ++i) { - oshape[i] = Any(); - } - - for (int64_t i = num_dynamic_axes; i < num_axis; ++i) { - oshape[i] = dshape[i]; - } - - reporter->Assign(types[4], TensorType(oshape, data->dtype)); - return true; -} - -Array StridedSliceCompute(const Attrs& attrs, const Array& inputs, - const Type& out_type) { - te::Tensor data = inputs[0]; - te::Tensor begin = inputs[1]; - te::Tensor end = inputs[2]; - te::Tensor strides = inputs[3]; - // Dynamic computation - int64_t data_rank = data->shape.size(); - int64_t num_dynamic_axes = begin->shape[0].as()->value; - ICHECK(end->shape[0].as()->value == num_dynamic_axes && - strides->shape[0].as()->value == num_dynamic_axes) - << "begin, end, strides should have the same length if they are dynamic variables"; - ICHECK(num_dynamic_axes <= data_rank) - << "the number of dynamic axes to slice should be less than or equal to the data rank"; - return Array{topi::dynamic_strided_slice(data, begin, end, strides)}; -} - -Expr MakeStridedSlice(Expr data, Expr begin, Expr end, Expr strides, String slice_mode) { - auto attrs = make_object(); - attrs->slice_mode = slice_mode; - static const Op& op = Op::Get("dyn.strided_slice"); - return Call(op, {data, begin, end, strides}, Attrs(attrs), {}); -} - -TVM_REGISTER_GLOBAL("relay.op.dyn._make.strided_slice").set_body_typed(MakeStridedSlice); - -RELAY_REGISTER_OP("dyn.strided_slice") - .describe(R"code(Strided slice of an array. - -Examples:: - - x = [[ 1., 4., 7., 10.], - [ 2., 5., 8., 11.], - [ 3., 6., 9., 12.]] - - strided_slice(x, begin=[0, 1], end=[2, 4], stride=[1, 1]) = [[ 4., 7., 10.], - [ 5., 8., 11.]] - - x = [[[ 1., 2.], - [ 3., 4.]], - - [[ 5., 6.], - [ 7., 8.]]] - - strided_slice(x, begin=[0, 0], end=[2, 2]) = [[[ 1., 2.], - [ 3., 4.]], - - [[ 5., 6.], - [ 7., 8.]]] -)code" TVM_ADD_FILELINE) - .set_num_inputs(4) - .add_argument("data", "Tensor", "The input tensor.") - .add_argument("begin", "Tensor", "The indices to begin with in the slicing.") - .add_argument("end", "Tensor", "Indices indicating end of the slice.") - .add_argument("strides", "Tensor", "The stride values.") - .add_argument("slice_mode", "Tensor", "The slice mode.") - .set_support_level(4) - .set_attrs_type() - .add_type_rel("DynStridedSlice", StridedSliceRel) - .set_attr("FTVMCompute", StridedSliceCompute) - .set_attr("TOpPattern", kInjective) - .set_attr("AnyCodegenStrategy", kVariableDimensions); - -bool SparseToDenseRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - ICHECK_EQ(num_inputs, 4); - auto sparse_indices = types[0].as(); - auto sparse_values = types[1].as(); - auto default_value = types[2].as(); - auto output_shape = types[3].as(); - - if (sparse_indices == nullptr || sparse_values == nullptr || default_value == nullptr || - output_shape == nullptr) { - return false; - } - - CHECK(sparse_indices->dtype.is_int()) << "sparse_indices must be tensor of integers"; - - CHECK_LE(sparse_indices->shape.size(), 3) - << "sparse_indices must be a tensor of either 0D, 1D or 2D"; - - CHECK_LE(sparse_values->shape.size(), 2) << "sparse_values must be a tensor of either 0D, 1D"; - - CHECK_EQ(default_value->shape.size(), 0) << "default_value should be a scalar"; - - Array oshape; - for (int i = 0; i < output_shape->shape[0].as()->value; i++) { - oshape.push_back(Any()); - } - reporter->Assign(types[4], TensorType(oshape, sparse_values->dtype)); - return true; -} - -Array SparseToDenseCompute(const Attrs& attrs, const Array& inputs, - const Type& out_type) { - ICHECK_EQ(inputs.size(), 4); - const auto* out_ttype = out_type.as(); - ICHECK(out_ttype); - return {topi::sparse_to_dense(inputs[0], out_ttype->shape, inputs[1], inputs[2]())}; -} - -TVM_REGISTER_GLOBAL("relay.op.dyn._make.sparse_to_dense") - .set_body_typed([](Expr indices, Expr output_shape, Expr values, Expr default_value) { - static const Op& op = Op::Get("dyn.sparse_to_dense"); - return Call(op, {indices, values, default_value, output_shape}); - }); - -RELAY_REGISTER_OP("dyn.sparse_to_dense") - .describe(R"code(A dense tensor from a sparse representation. - - - **sparse_indices**: A 0-D, 1-D, or 2-D tensor of integers containing location of sparse values - - - **output_shape**: A list of integers. Shape of the dense output tensor. - - - **sparse_values**: A 0-D or 1-D tensor containing the sparse values for the sparse indices. - - - **default_value**: A 0-D tensor containing the default value for the remaining locations. Defaults to 0. - - Example:: - - sparse_to_dense([0, 0], [1, 2]], [3, 4], [1, 2], 0) = [[1, 0, 0, 0], [0, 0, 2, 0], [0, 0, 0, 0]] - - )code" TVM_ADD_FILELINE) - .set_num_inputs(4) - .set_support_level(3) - .add_argument("sparse_indices", "Tensor", "Contains sparse indices.") - .add_argument("sparse_values", "Tensor", "Contains values for sparse indices.") - .add_argument("default_value", "Tensor", "Value to set for non-sparse indices. Defaults to 0.") - .add_argument("output_shape", "Tensor", "Shape of the dense output tensor") - .add_type_rel("DynSparseToDense", SparseToDenseRel) - .set_attr("TOpIsStateful", false) - .set_attr("TOpPattern", kOpaque) - .set_attr("FInferCorrectLayout", ElemwiseArbitraryLayout) - .set_attr("FTVMCompute", SparseToDenseCompute); - -/* relay.dyn.unsqueeze */ -TVM_REGISTER_NODE_TYPE(DynExpandDimsAttrs); - -bool ExpandDimsRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - ICHECK_EQ(num_inputs, 2); - const auto* data_type = types[0].as(); - if (data_type == nullptr) { - ICHECK(types[0].as()) - << "expand_dims: expect input type to be TensorType but get " << types[0]; - return false; - } - - const auto* param = attrs.as(); - - // We don't know the output shape until we see the value of the axis input - int ndim = data_type->shape.size(); - Array oshape(ndim + param->num_newaxis, Any()); - - const auto* axis_type = types[1].as(); - ICHECK(axis_type->shape.size() == 0) << "Axis should be a scalar got shape " << axis_type->shape; - - // Set output shape - reporter->Assign(types[2], TensorType(oshape, data_type->dtype)); - return true; -} - -Array ExpandDimsCompute(const Attrs& attrs, const Array& inputs, - const Type& out_type) { - // inputs = [Input tensor, axis to expand] - ICHECK_EQ(inputs.size(), 2); - - const auto* param = attrs.as(); - - Array ishape = inputs[0]->shape; - const TensorTypeNode* out_ttype = out_type.as(); - int ndim_out = out_ttype->shape.size(); - int ndim_in = ishape.size(); - ICHECK_EQ(ndim_in + param->num_newaxis, ndim_out); - - Array newshape; - for (auto val : out_ttype->shape) { - // These vars will be populated by the VM executor with the results - // of the shape_func for the op. - newshape.push_back(val.as()->ToVar()); - } - - return {topi::reshape(inputs[0], newshape)}; -} - -Expr MakeExpandDims(Expr data, Expr axis_tensor, int num_newaxis) { - auto attrs = make_object(); - attrs->num_newaxis = num_newaxis; - static const Op& op = Op::Get("dyn.expand_dims"); - return Call(op, {data, axis_tensor}, Attrs(attrs), {}); -} - -TVM_REGISTER_GLOBAL("relay.op.dyn._make.expand_dims").set_body_typed(MakeExpandDims); - -RELAY_REGISTER_OP("dyn.expand_dims") - .describe(R"code(Insert one new axis at the position given by `axis` - -- **data**: The input data to the operator. -- **axis**: The axis to insert a new dimension - -)code" TVM_ADD_FILELINE) - .set_num_inputs(2) - .add_argument("data", "Tensor", "The input tensor.") - .add_argument("axis", "Tensor", "The axis to insert at a dimension.") - .set_support_level(3) - .add_type_rel("DynamicExpandDims", ExpandDimsRel) - .set_attr("FTVMCompute", ExpandDimsCompute) - .set_attr("TOpPattern", kInjective); - -bool DynSqueezeRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - // [data, axes, output] - ICHECK_EQ(types.size(), 3); - const auto* data = types[0].as(); - if (data == nullptr) { - return false; - } - const auto* axes = types[1].as(); - if (axes == nullptr) { - return false; - } - - ICHECK_EQ(axes->shape.size(), 1) << "Got" << axes->shape.size() << "expected 1"; - ICHECK(axes->shape[0].as()) << "axes expected to be static rank"; - size_t output_rank = data->shape.size() - axes->shape[0].as()->value; - std::vector result_shape(output_rank, Any()); - reporter->Assign(types[2], TensorType(result_shape, data->dtype)); - return true; -} - -Array SqueezeCompute(const Attrs& attrs, const Array& inputs, - const Type& out_type) { - const auto* out_ttype = out_type.as(); - ICHECK(out_ttype != nullptr); - Array newshape; - for (auto val : out_ttype->shape) { - newshape.push_back(val.as()->ToVar()); - } - return {topi::reshape(inputs[0], newshape)}; -} - -Expr MakeDynSqueeze(Expr data, Expr axes) { - auto attrs = make_object(); - static const Op& op = Op::Get("dyn.squeeze"); - return Call(op, {data, axes}, Attrs(attrs), {}); -} - -TVM_REGISTER_GLOBAL("relay.op.dyn._make.squeeze").set_body_typed(MakeDynSqueeze); - -RELAY_REGISTER_OP("dyn.squeeze") - .describe(R"code(Remove axes of value 1 in input tensor at the dimensions given by axes - -- **data**: The input data to the operator. -- **axes**: The axes to squeeze. - -)code" TVM_ADD_FILELINE) - .set_num_inputs(2) - .set_attrs_type() - .add_argument("data", "Tensor", "The input tensor.") - .add_argument("axes", "Tensor", "The axes to squeeze.") - .set_support_level(3) - .add_type_rel("DynSqueeze", DynSqueezeRel) - .set_attr("FTVMCompute", SqueezeCompute) - .set_attr("TOpPattern", kInjective) - .set_attr("TReshapeOp", true); - -} // namespace dyn -} // namespace relay -} // namespace tvm diff --git a/src/relay/op/dyn/tensor/transform.h b/src/relay/op/dyn/tensor/transform.h deleted file mode 100644 index 98b0474a7e2b..000000000000 --- a/src/relay/op/dyn/tensor/transform.h +++ /dev/null @@ -1,32 +0,0 @@ -/* - * 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 src/relay/op/tensor/transform.h - * \brief Transform op attributes that can be shared among Relay and its dialects. - */ -#ifndef TVM_RELAY_OP_DYN_TENSOR_TRANSFORM_H_ -#define TVM_RELAY_OP_DYN_TENSOR_TRANSFORM_H_ - -namespace tvm { -namespace relay { -namespace dyn {} // namespace dyn -} // namespace relay -} // namespace tvm -#endif // TVM_RELAY_OP_DYN_TENSOR_TRANSFORM_H_ diff --git a/src/relay/op/image/dilation2d.cc b/src/relay/op/image/dilation2d.cc deleted file mode 100644 index ef3e2592e3fc..000000000000 --- a/src/relay/op/image/dilation2d.cc +++ /dev/null @@ -1,151 +0,0 @@ -/* - * 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 dilation2d.cc - * \brief Morphological dilation operator - */ -#include -#include -#include - -#include "../op_common.h" - -namespace tvm { -namespace relay { - -// relay.image.dilation2d -TVM_REGISTER_NODE_TYPE(Dilation2DAttrs); - -template -InferCorrectLayoutOutput Dilation2DInferCorrectLayout(const Attrs& attrs, - const Array& new_in_layouts, - const Array& old_in_layouts, - const Array& old_in_types) { - const T* params = attrs.as(); - return InferCorrectLayoutOutput({params->data_layout, params->kernel_layout}, - {params->data_layout}, attrs); -} - -// Positional relay function to create dilation2d operator -// used by frontend FFI. -Expr MakeDilation2D(Expr data, Expr weight, Array strides, Array padding, - Array dilations, String data_layout, String kernel_layout, - DataType out_dtype) { - auto attrs = make_object(); - attrs->strides = std::move(strides); - attrs->padding = std::move(padding); - attrs->dilations = std::move(dilations); - attrs->data_layout = std::move(data_layout); - attrs->kernel_layout = std::move(kernel_layout); - attrs->out_dtype = std::move(out_dtype); - static const Op& op = Op::Get("image.dilation2d"); - return Call(op, {data, weight}, Attrs(attrs), {}); -} - -template -bool Dilation2DRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - ICHECK_EQ(types.size(), 3); - const auto* data = types[0].as(); - const auto* weight = types[1].as(); - if (data == nullptr) return false; - static const Layout kNCHW("NCHW"); - static const Layout kOIHW("IHW"); - - const AttrType* param = attrs.as(); - ICHECK(param != nullptr); - const Layout in_layout(param->data_layout); - const Layout kernel_layout(param->kernel_layout); - - const auto trans_in_layout = tir::BijectiveLayout(in_layout, kNCHW); - ICHECK(trans_in_layout.defined()) - << "Dilation2D only support input layouts that are convertible from NCHW." - << " But got " << in_layout; - - const auto trans_kernel_layout = tir::BijectiveLayout(kernel_layout, kOIHW); - ICHECK(trans_kernel_layout.defined()) - << "Dilation2D only support kernel layouts that are convertible from OIHW." - << " But got " << kernel_layout; - - Layout out_layout(param->data_layout); - const auto trans_out_layout = tir::BijectiveLayout(out_layout, kNCHW); - ICHECK(trans_out_layout.defined()) - << "Dilation2D only support output layouts that are convertible from NCHW." - << " But got " << out_layout; - - Array dshape_nchw = trans_in_layout.ForwardShape(data->shape); - - IndexExpr channels, dilated_ksize_y, dilated_ksize_x; - - // use weight to infer the conv shape. - if (weight == nullptr) return false; - auto wshape = trans_kernel_layout.ForwardShape(weight->shape); - channels = wshape[0]; - - dilated_ksize_y = 1 + (wshape[1] - 1) * param->dilations[0]; - dilated_ksize_x = 1 + (wshape[2] - 1) * param->dilations[1]; - - // dilation - Array oshape({dshape_nchw[0], channels, 0, 0}); - IndexExpr pad_h, pad_w; - GetPaddingHeightWidth(param->padding, &pad_h, &pad_w); - if (!dshape_nchw[2].as()) { - oshape.Set(2, indexdiv(dshape_nchw[2] + pad_h - dilated_ksize_y, param->strides[0]) + 1); - } else { - oshape.Set(2, dshape_nchw[2]); - } - - if (!dshape_nchw[3].as()) { - oshape.Set(3, indexdiv(dshape_nchw[3] + pad_w - dilated_ksize_x, param->strides[1]) + 1); - } else { - oshape.Set(3, dshape_nchw[3]); - } - - DataType out_dtype = param->out_dtype; - if (out_dtype.bits() == 0) { - out_dtype = data->dtype; - } - oshape = trans_out_layout.BackwardShape(oshape); - // assign output type - reporter->Assign(types[2], TensorType(oshape, out_dtype)); - return true; -} - -TVM_REGISTER_GLOBAL("relay.op.image._make.dilation2d").set_body_typed(MakeDilation2D); - -RELAY_REGISTER_OP("image.dilation2d") - .describe(R"code(Computes grayscale dilation of 4D input and 3D filter. -- **data**: This depends on the `layout` parameter. Input is 4D array of shape - (batch_size, in_channels, height, width) if `layout` is `NCHW`. -- **weight**: (in_channels, height, width) -- **out**: This depends on the `layout` parameter. Output is 4D array of shape - (batch_size, channels, out_height, out_width) if `layout` is `NCHW`. -)code" TVM_ADD_FILELINE) - .set_attrs_type() - .set_num_inputs(2) - .add_argument("data", "Tensor", "The input tensor.") - .add_argument("weight", "Tensor", "The weight tensor.") - .set_support_level(2) - .add_type_rel("Dilation2D", Dilation2DRel) - .set_attr("FInferCorrectLayout", - Dilation2DInferCorrectLayout); - -} // namespace relay -} // namespace tvm diff --git a/src/relay/op/image/grid_sample.cc b/src/relay/op/image/grid_sample.cc deleted file mode 100644 index 46b9714aeeb8..000000000000 --- a/src/relay/op/image/grid_sample.cc +++ /dev/null @@ -1,212 +0,0 @@ -/* - * 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 grid_sample.cc - * \brief affine_grid and grid_sample operator - */ -#include -#include -#include - -#include "../op_common.h" - -namespace tvm { -namespace relay { - -// relay.image.affine_grid -TVM_REGISTER_NODE_TYPE(AffineGridAttrs); - -bool AffineGridRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - ICHECK_EQ(types.size(), 2); - const auto* data = types[0].as(); - if (data == nullptr) return false; - auto batch_size = data->shape[0]; - - const AffineGridAttrs* param = attrs.as(); - ICHECK(param != nullptr); - - Array oshape; - - ICHECK(data->shape.size() == 3U && reporter->AssertEQ(data->shape[1], 2) && - reporter->AssertEQ(data->shape[2], 3)) - << "data should be an" - "affine matrix with shape [batch_size, 2, 3]"; - ICHECK(param->target_shape.defined() && param->target_shape.size() == 2) - << "target_shape should be 2D"; - oshape.push_back(batch_size); - oshape.push_back(2); - oshape.push_back(param->target_shape[0]); - oshape.push_back(param->target_shape[1]); - - // assign output type - reporter->Assign(types[1], TensorType(oshape, data->dtype)); - return true; -} - -// Positional relay function to create affine_grid operator -// used by frontend FFI. -Expr MakeAffineGrid(Expr data, Array target_shape) { - auto attrs = make_object(); - attrs->target_shape = std::move(target_shape); - static const Op& op = Op::Get("image.affine_grid"); - return Call(op, {data}, Attrs(attrs), {}); -} - -TVM_REGISTER_GLOBAL("relay.op.image._make.affine_grid").set_body_typed(MakeAffineGrid); - -RELAY_REGISTER_OP("image.affine_grid") - .describe(R"code(affine_grid operator that generates 2D sampling grid. - -This operation is described in https://arxiv.org/pdf/1506.02025.pdf. It generates a uniform -sampling grid within the target shape and normalizes it to [-1, 1]. The provided affine -transformation is then applied on the sampling grid. - -- **data**: data is 3D array of shape [batch, 2, 3], which defines an affine transformation. - -- **out**: out is 4D array of shape [batch, 2, height, width], where each vector - :math:`out[b, :, h, w]` represents the coordinate :math:`(x, y)` - -)code" TVM_ADD_FILELINE) - .set_attrs_type() - .set_num_inputs(1) - .add_argument("data", "Tensor", "The affine matrix.") - .set_support_level(5) - .add_type_rel("AffineGrid", AffineGridRel) - .set_attr("TOpPattern", kInjective); - -// relay.image.grid_sample -TVM_REGISTER_NODE_TYPE(GridSampleAttrs); - -bool GridSampleRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - ICHECK_EQ(types.size(), 3); - const auto* data = types[0].as(); - const auto* grid = types[1].as(); - if (!data || !grid) return false; - const auto* param = attrs.as(); - ICHECK(param); - const Layout in_layout(param->layout); - - if (data->shape.size() == 4) { - static const Layout kNCHW("NCHW"); - auto layout_converter = tir::BijectiveLayout(in_layout, kNCHW); - auto oshape = layout_converter.ForwardShape(data->shape); - oshape.Set(2, grid->shape[2]); - oshape.Set(3, grid->shape[3]); - - // assign output type - reporter->Assign(types[2], TensorType(layout_converter.BackwardShape(oshape), data->dtype)); - return true; - } else if (data->shape.size() == 5) { - static const Layout kNDCHW("NCDHW"); - auto layout_converter = tir::BijectiveLayout(in_layout, kNDCHW); - auto oshape = layout_converter.ForwardShape(data->shape); - oshape.Set(2, grid->shape[2]); - oshape.Set(3, grid->shape[3]); - oshape.Set(4, grid->shape[4]); - - // assign output type - reporter->Assign(types[2], TensorType(layout_converter.BackwardShape(oshape), data->dtype)); - return true; - } - - return false; -} - -// Positional relay function to create affine_grid operator -// used by frontend FFI. -Expr MakeGridSample(Expr data, Expr grid, String method, String layout, String padding_mode, - bool align_corners) { - auto attrs = make_object(); - attrs->method = std::move(method); - attrs->layout = std::move(layout); - attrs->padding_mode = std::move(padding_mode); - attrs->align_corners = std::move(align_corners); - - static const Op& op = Op::Get("image.grid_sample"); - return Call(op, {data, grid}, Attrs(attrs), {}); -} - -TVM_REGISTER_GLOBAL("relay.op.image._make.grid_sample").set_body_typed(MakeGridSample); - -RELAY_REGISTER_OP("image.grid_sample") - .describe(R"code(Applies grid sampling to input feature map. - -Given :math:`data` and :math:`grid`, then the output is computed by - -.. math:: - - x_{src} = grid[batch, 0, y_{dst}, x_{dst}] \\ - y_{src} = grid[batch, 1, y_{dst}, x_{dst}] \\ - output[batch, channel, y_{dst}, x_{dst}] = G(data[batch, channel, y_{src}, x_{src}]) - -For 5-D, the output is computed by - -.. math:: - - x_{src} = grid[batch, 0, z_{dst}, y_{dst}, x_{dst}] \\ - y_{src} = grid[batch, 1, z_{dst}, y_{dst}, x_{dst}] \\ - z_{src} = grid[batch, 2, z_{dst}, y_{dst}, x_{dst}] \\ - output[batch, channel, z_{src}, y_{dst}, x_{dst}] - = G(data[batch, channel, z_{src}, y_{src}, x_{src}]) - -:math:`x_{dst}`, :math:`y_{dst}` enumerate all spatial locations in :math:`output`, and -:math:`G()` denotes the interpolation function. - -The out-boundary points will be padded with zeros if padding_mode is "zeros", or -border pixel value if padding_mode is "border", or -inner pixel value if padding_mode is "reflection". - -The left-top corner (-1, -1) and right-bottom corner (1, 1) in grid will be map to -(0, 0) and (h - 1, w - 1) of data if align_corners is "True", or -(-0.5, -0.5) and (h - 0.5, w - 0.5) of data if align_corners is "False". - -The shape of the output will be -4-D (data.shape[0], data.shape[1], grid.shape[2], grid.shape[3]), or -5-D (data.shape[0], data.shape[1], grid.shape[2], grid.shape[3], grid.shape[4]). - -The operator assumes that :math:`data` and :math:`grid` has been normalized to [-1, 1]. - -grid_sample often cooperates with affine_grid which generates sampling grids for grid_sample. - -- **data**: data is of 4-D shape (batch_size, channels, in_height, in_width), or - of 5-D shape (batch_size, channels, in_depth, in_height, in_width) - -- **grid**: grid is of 4-D shape [batch, 2, out_height, out_width] - where each vector :math:`out[b, :, h, w]` represents the coordinate :math:`(x, y)`, - or of 5-D of shape [batch, 3, out_depth, out_height, out_width] - where each vector :math:`out[b, :, d, h, w]` represents the coordinate - :math:`(x, y, z)` - -- **out**: out is of 4-D shape (batch, in_channel, out_height, out_width), or - of 5-D shape [batch, channel, out_depth, out_height, out_width] - -)code" TVM_ADD_FILELINE) - .set_num_inputs(2) - .set_attrs_type() - .add_argument("data", "Tensor", "The input tensor.") - .add_argument("grid", "Tensor", "The grid tensor.") - .set_support_level(5) - .add_type_rel("GridSample", GridSampleRel) - .set_attr("TOpPattern", kInjective); - -} // namespace relay -} // namespace tvm diff --git a/src/relay/op/image/resize.cc b/src/relay/op/image/resize.cc deleted file mode 100644 index ca05a4bdce43..000000000000 --- a/src/relay/op/image/resize.cc +++ /dev/null @@ -1,369 +0,0 @@ -/* - * 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 resize.cc - * \brief Image resize operators - */ -#include -#include -#include - -#include "../make_op.h" -#include "../op_common.h" - -namespace tvm { -namespace relay { - -template -InferCorrectLayoutOutput ResizeInferCorrectLayout(const Attrs& attrs, - const Array& new_in_layouts, - const Array& old_in_layouts, - const Array& old_in_types) { - const auto* attrs_ptr = attrs.as(); - CHECK(attrs_ptr); - ObjectPtr params = make_object(*attrs_ptr); - - if (new_in_layouts.defined()) { - ICHECK_EQ(new_in_layouts.size(), 1); - - Layout raw_layout(params->layout); - Layout new_layout = new_in_layouts[0]; - Layout old_layout = old_in_layouts[0]; - if (!new_layout.Equals(old_layout) && raw_layout.Equals(old_layout) && - new_layout->axes.size() == old_layout->axes.size()) { - // Follow input layout - params->layout = new_layout.name(); - } - } - - return InferCorrectLayoutOutput({params->layout}, {params->layout}, Attrs(params)); -} - -TVM_REGISTER_NODE_TYPE(Resize1DAttrs); - -bool Resize1DRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - ICHECK_EQ(types.size(), 2); - const auto* data = types[0].as(); - if (data == nullptr) return false; - - static const Layout kNCW("NCW"); - - const Resize1DAttrs* param = attrs.as(); - ICHECK(param != nullptr); - ICHECK(param->size.size() == 1); - ICHECK(param->roi.size() == 2); - const Layout in_layout(param->layout); - auto layout_converter = tir::BijectiveLayout(in_layout, kNCW); - ICHECK(layout_converter.defined()) - << "Resize only support input layouts that are convertible from NCW." - << " But got " << in_layout; - - auto oshape = layout_converter.ForwardShape(data->shape); - oshape.Set(2, param->size[0]); - - DataType out_dtype = param->out_dtype; - if (out_dtype.bits() == 0) { - out_dtype = data->dtype; - } - - // assign output type - reporter->Assign(types[1], TensorType(layout_converter.BackwardShape(oshape), out_dtype)); - return true; -} - -// Positional relay function to create image operator -// used by frontend FFI. -Expr MakeResize1D(Expr data, Array size, Array roi, String layout, - String method, String coordinate_transformation_mode, String rounding_method, - double cubic_alpha, int cubic_exclude, double extrapolation_value, - DataType out_dtype) { - auto attrs = make_object(); - attrs->size = std::move(size); - attrs->roi = std::move(roi); - attrs->layout = std::move(layout); - attrs->method = std::move(method); - attrs->coordinate_transformation_mode = coordinate_transformation_mode; - attrs->rounding_method = rounding_method; - attrs->cubic_alpha = cubic_alpha; - attrs->cubic_exclude = cubic_exclude; - attrs->extrapolation_value = extrapolation_value; - attrs->out_dtype = out_dtype; - static const Op& op = Op::Get("image.resize1d"); - return Call(op, {data}, Attrs(attrs), {}); -} - -TVM_REGISTER_GLOBAL("relay.op.image._make.resize1d").set_body_typed(MakeResize1D); - -RELAY_REGISTER_OP("image.resize1d") - .describe(R"code(Perform resize to input array with nearest neighbour or bilinear interpolation. - -- **data**: data is 3D array of shape - (batch_size, channels, in_width) for NCW - (batch_size, in_width, channels) for NWC - -- **out**: Output is 3D array of shape - for layout NCW - (batch_size, channels, size[0]) - - for layout NWC - (batch_size, size[0], channels) -)code" TVM_ADD_FILELINE) - .set_attrs_type() - .set_num_inputs(1) - .add_argument("data", "Tensor", "The input tensor.") - .set_support_level(5) - .add_type_rel("Resize1D", Resize1DRel) - .set_attr("FInferCorrectLayout", ResizeInferCorrectLayout) - .set_attr("TOpPattern", kInjective); - -TVM_REGISTER_NODE_TYPE(Resize2DAttrs); - -bool Resize2DRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - ICHECK_EQ(types.size(), 2); - const auto* data = types[0].as(); - if (data == nullptr) return false; - - static const Layout kNCHW("NCHW"); - - const Resize2DAttrs* param = attrs.as(); - ICHECK(param != nullptr); - ICHECK(param->size.size() == 2); - ICHECK(param->roi.size() == 4); - const Layout in_layout(param->layout); - auto layout_converter = tir::BijectiveLayout(in_layout, kNCHW); - ICHECK(layout_converter.defined()) - << "Resize only support input layouts that are convertible from NCHW." - << " But got " << in_layout; - - auto oshape = layout_converter.ForwardShape(data->shape); - oshape.Set(2, param->size[0]); - oshape.Set(3, param->size[1]); - - DataType out_dtype = param->out_dtype; - if (out_dtype.bits() == 0) { - out_dtype = data->dtype; - } - - // assign output type - reporter->Assign(types[1], TensorType(layout_converter.BackwardShape(oshape), out_dtype)); - return true; -} - -// Positional relay function to create image operator -// used by frontend FFI. -Expr MakeResize2D(Expr data, Array size, Array roi, String layout, - String method, String coordinate_transformation_mode, String rounding_method, - double cubic_alpha, int cubic_exclude, double extrapolation_value, - DataType out_dtype) { - auto attrs = make_object(); - attrs->size = std::move(size); - attrs->roi = std::move(roi); - attrs->layout = std::move(layout); - attrs->method = std::move(method); - attrs->coordinate_transformation_mode = coordinate_transformation_mode; - attrs->rounding_method = rounding_method; - attrs->cubic_alpha = cubic_alpha; - attrs->cubic_exclude = cubic_exclude; - attrs->extrapolation_value = extrapolation_value; - attrs->out_dtype = out_dtype; - static const Op& op = Op::Get("image.resize2d"); - return Call(op, {data}, Attrs(attrs), {}); -} - -TVM_REGISTER_GLOBAL("relay.op.image._make.resize2d").set_body_typed(MakeResize2D); - -RELAY_REGISTER_OP("image.resize2d") - .describe(R"code(Perform resize to input array with nearest neighbour or bilinear interpolation. - -- **data**: data is 4D array of shape - (batch_size, channels, in_height, in_width) for NCHW - (batch_size, in_height, in_width, channels) for NHWC - -- **out**: Output is 4D array of shape - for layout NCHW - (batch_size, channels, size[0], size[1]) - - for layout NHWC - (batch_size, size[0], size[1], channels) -)code" TVM_ADD_FILELINE) - .set_attrs_type() - .set_num_inputs(1) - .add_argument("data", "Tensor", "The input tensor.") - .set_support_level(5) - .add_type_rel("Resize2D", Resize2DRel) - .set_attr("FInferCorrectLayout", ResizeInferCorrectLayout) - .set_attr("TOpPattern", kInjective); - -TVM_REGISTER_NODE_TYPE(Resize3DAttrs); - -bool Resize3DRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - ICHECK_EQ(types.size(), 2); - const auto* data = types[0].as(); - if (data == nullptr) return false; - - static const Layout kNCDHW("NCDHW"); - - const Resize3DAttrs* param = attrs.as(); - ICHECK(param != nullptr); - ICHECK(param->size.size() == 3); - ICHECK(param->roi.size() == 6); - const Layout in_layout(param->layout); - auto layout_converter = tir::BijectiveLayout(in_layout, kNCDHW); - ICHECK(layout_converter.defined()) - << "Resize3d only support input layouts that are convertible from NCDHW." - << " But got " << in_layout; - - auto oshape = layout_converter.ForwardShape(data->shape); - oshape.Set(2, param->size[0]); - oshape.Set(3, param->size[1]); - oshape.Set(4, param->size[2]); - - DataType out_dtype = param->out_dtype; - if (out_dtype.bits() == 0) { - out_dtype = data->dtype; - } - - // assign output type - reporter->Assign(types[1], TensorType(layout_converter.BackwardShape(oshape), out_dtype)); - return true; -} - -// Positional relay function to create image operator -// used by frontend FFI. -Expr MakeResize3D(Expr data, Array size, Array roi, String layout, - String method, String coordinate_transformation_mode, String rounding_method, - double cubic_alpha, int cubic_exclude, double extrapolation_value, - DataType out_dtype) { - auto attrs = make_object(); - attrs->size = std::move(size); - attrs->roi = std::move(roi); - attrs->layout = std::move(layout); - attrs->method = std::move(method); - attrs->coordinate_transformation_mode = coordinate_transformation_mode; - attrs->rounding_method = rounding_method; - attrs->cubic_alpha = cubic_alpha; - attrs->cubic_exclude = cubic_exclude; - attrs->extrapolation_value = extrapolation_value; - attrs->out_dtype = out_dtype; - static const Op& op = Op::Get("image.resize3d"); - return Call(op, {data}, Attrs(attrs), {}); -} - -TVM_REGISTER_GLOBAL("relay.op.image._make.resize3d").set_body_typed(MakeResize3D); - -RELAY_REGISTER_OP("image.resize3d") - .describe(R"code( -Perform resize3d to input array with nearest neighbour or bilinear interpolation. - -- **data**: data is 5D array of shape - (batch_size, channels, in_depth, in_height, in_width) for NCDHW - (batch_size, in_depth, in_height, in_width, channels) for NDHWC - -- **out**: Output is 5D array of shape - for layout NCDHW - (batch_size, channels, size[0], size[1], size[2]) - - for layout NDHWC - (batch_size, size[0], size[1], size[2], channels) -)code" TVM_ADD_FILELINE) - .set_attrs_type() - .set_num_inputs(1) - .add_argument("data", "Tensor", "The input tensor.") - .set_support_level(5) - .add_type_rel("Resize3d", Resize3DRel) - .set_attr("TOpPattern", kInjective); - -TVM_REGISTER_NODE_TYPE(CropAndResizeAttrs); - -bool CropAndResizeRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - ICHECK_EQ(types.size(), 4); - const auto* data = types[0].as(); - const auto* boxes = types[1].as(); - const auto* box_indices = types[2].as(); - if (data == nullptr || boxes == nullptr || box_indices == nullptr) return false; - - const CropAndResizeAttrs* param = attrs.as(); - ICHECK(param != nullptr); - auto crop_size = param->crop_size; - - DataType out_dtype = param->out_dtype; - if (out_dtype.bits() == 0) { - out_dtype = data->dtype; - } - - // 4-D tensor of shape [num_boxes, crop_height, crop_width, depth] - static const Layout kNCHW("NCHW"); - const Layout in_layout(param->layout); - auto layout_converter = tir::BijectiveLayout(in_layout, kNCHW); - auto oshape = layout_converter.ForwardShape(data->shape); - oshape.Set(0, boxes->shape[0]); - oshape.Set(2, crop_size[0]); - oshape.Set(3, crop_size[1]); - auto bshape = layout_converter.BackwardShape(oshape); - // assign output type - reporter->Assign(types[3], TensorType(bshape, out_dtype)); - return true; -} - -Expr MakeCropAndResize(Expr data, Expr boxes, Expr box_indices, Array crop_size, - String layout, String method, double extrapolation_value, - DataType out_dtype) { - auto attrs = make_object(); - attrs->crop_size = std::move(crop_size); - attrs->layout = std::move(layout); - attrs->method = std::move(method); - attrs->extrapolation_value = std::move(extrapolation_value); - attrs->out_dtype = out_dtype; - static const Op& op = Op::Get("image.crop_and_resize"); - return Call(op, {data, boxes, box_indices}, Attrs(attrs), {}); -} - -TVM_REGISTER_GLOBAL("relay.op.image._make.crop_and_resize").set_body_typed(MakeCropAndResize); - -RELAY_REGISTER_OP("image.crop_and_resize") - .describe( - R"code(Perform crop and resize to input array with nearest neighbour or bilinear interpolation. - -- **data**: data is 4D array of shape - (batch_size, channels, in_height, in_width) for NCHW - (batch_size, in_height, in_width, channels) for NHWC - -- **out**: Output is 4D array of shape - for layout NCHW - (batch_size, channels, crop_size[0], crop_size[1]) - - for layout NHWC - (batch_size, crop_size[0], crop_size[1], channels) -)code" TVM_ADD_FILELINE) - .set_num_inputs(3) - .add_argument("data", "Tensor", "The input tensor.") - .add_argument("boxes", "Tensor", "The boxes tensor.") - .add_argument("box_indices", "Tensor", "The box indices tensor.") - .set_attrs_type() - .set_support_level(5) - .add_type_rel("CropAndResize", CropAndResizeRel) - .set_attr("TOpPattern", kInjective); - -} // namespace relay -} // namespace tvm diff --git a/src/relay/op/make_op.h b/src/relay/op/make_op.h deleted file mode 100644 index 222aba4bd25b..000000000000 --- a/src/relay/op/make_op.h +++ /dev/null @@ -1,128 +0,0 @@ -/* - * 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 tvm/relay/op/make_op.h - * \brief Header of internal operator functions - * to assist in creating ops in C++ - */ -#ifndef TVM_RELAY_OP_MAKE_OP_H_ -#define TVM_RELAY_OP_MAKE_OP_H_ - -#include -#include -#include - -// Include Templated Make Functions -#include "nn/convolution_make.h" -#include "nn/pooling.h" - -namespace tvm { -namespace relay { - -Expr MakeBroadCastTo(Expr data, Array shape); - -Expr MakeCast(Expr data, DataType dtype); - -Expr MakeClip(Expr a, double a_min, double a_max); - -Expr MakeConcatenate(Expr data, int axis); - -Expr MakeMatmul(Expr tensor_a, Expr tensor_b, IndexExpr units, DataType out_dtype, bool transpose_a, - bool transpose_b); - -Expr MakeDense(Expr data, Expr weight, IndexExpr units, DataType out_dtype); - -Expr MakeBatchMatmul(Expr lhs, Expr rhs, DataType out_dtype, bool transpose_a, bool transpose_b); - -Expr MakeExpandDims(Expr data, int axis, int num_newaxis); - -Expr MakeFixedPointMultiplyPerAxis(Expr x, Expr m, Expr lshift, Expr rshift, - bool is_lshift_required, bool is_rshift_required, - Array axis); - -Expr MakeFull(Expr fill_value, Array shape, DataType dtype); - -Expr MakeLayoutTransform(Expr data, String src_layout, String dst_layout); - -Expr MakeMetaScheduleLayoutTransform(Expr data, tir::IndexMap index_map); - -Expr MakeAutoSchedulerLayoutTransform(Expr data, String src_layout, String dst_layout); - -Expr MakeOnes(Array shape, DataType dtype); - -Expr MakePad(Expr data, Array> pad_width, Expr pad_value, String pad_mode); - -Expr MakeReduce(Expr data, Array axis, bool keepdims, bool exclude, String op_name); - -Expr MakeRepeat(Expr data, int repeats, int axis); - -Expr MakeReshape(Expr data, Array newshape, bool allowzero = false); - -Expr MakeReshapeLike(Expr lhs, Expr rhs, int lhs_begin, Integer lhs_end, int rhs_begin, - Integer rhs_end); - -Expr MakeSplit(Expr data, Variant> indices_or_sections, int axis); - -Expr MakeSqueeze(Expr data, Array axis); - -Expr MakeStack(Expr data, int axis); - -Expr MakeTranspose(Expr data, Array axes); - -Expr MakeStridedSlice(Expr data, Array begin, Array end, Array strides, - String slice_mode, - Optional> axes = NullValue>()); - -Expr MakeTile(Expr data, Array reps); - -Expr MakeTopK(Expr data, int k, int axis, String ret_type, bool is_ascend, DataType dtype); - -Expr MakeUpSampling(Expr data, double scale_h, double scale_w, String layout, String method, - bool align_corners); - -Expr MakeUpSampling3D(Expr data, double scale_d, double scale_h, double scale_w, String layout, - String method, String coordinate_transformation_mode); - -Expr MakeVariance(Expr data, Expr mean, Array axis, bool keepdims, bool exclude, - bool unbiased); - -Expr MakeZeros(Array shape, DataType dtype); - -Expr MakeOneHot(Expr indices, Expr on_value, Expr off_value, int depth, int axis, DataType dtype); - -Expr MakeResize2D(Expr data, Array size, Array roi, String layout, - String method, String coordinate_transformation_mode, String rounding_method, - double cubic_alpha, int cubic_exclude, double extrapolation_value, - DataType out_dtype); - -Expr MakeSparseToDense(Expr indices, Array output_shape, Expr values, Expr default_value); - -Expr MakeArange(Expr start, Expr stop, Expr step, DataType dtype); - -Expr MakeShapeOf(Expr data, DataType dtype); - -Expr MakeTake(Expr data, Expr indices, Integer batch_dims, Integer axis, String mode); - -Expr MakeBiasAdd(Expr data, Expr bias, int axis); - -} // namespace relay -} // namespace tvm -#endif // TVM_RELAY_OP_MAKE_OP_H_ diff --git a/src/relay/op/memory/device_copy.cc b/src/relay/op/memory/device_copy.cc deleted file mode 100644 index a59e25ce1e13..000000000000 --- a/src/relay/op/memory/device_copy.cc +++ /dev/null @@ -1,123 +0,0 @@ -/* - * 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 relay/op/memory/device_copy.cc - * \brief Helpers for working with "device_copy" attributes. - */ - -#include "./device_copy.h" - -#include -#include -#include -#include -#include -#include -#include - -#include - -#include "../../transforms/infer_layout_utils.h" -#include "../annotation/annotation.h" -#include "../call/call.h" -#include "../type_relations.h" - -namespace tvm { -namespace relay { - -// relay.device_copy -TVM_REGISTER_NODE_TYPE(DeviceCopyAttrs); - -const Op& DeviceCopyOp() { - static const Op& op = Op::Get("device_copy"); - return op; -} - -Expr DeviceCopy(Expr expr, VirtualDevice src_virtual_device, VirtualDevice dst_virtual_device) { - ICHECK(!src_virtual_device->IsFullyUnconstrained()); - ICHECK(!dst_virtual_device->IsFullyUnconstrained()); - auto attrs = make_object(); - attrs->src_virtual_device = std::move(src_virtual_device); - attrs->dst_virtual_device = std::move(dst_virtual_device); - Span span = expr->span; - return Call(DeviceCopyOp(), {std::move(expr)}, Attrs(std::move(attrs)), /*type_args=*/{}, - std::move(span)); -} - -TVM_REGISTER_GLOBAL("relay.op._make.DeviceCopy").set_body_typed(DeviceCopy); - -Expr MaybeDeviceCopy(Expr expr, VirtualDevice src_virtual_device, - VirtualDevice dst_virtual_device) { - if (src_virtual_device == dst_virtual_device) { - // No copy needed. - return expr; - } - return DeviceCopy(std::move(expr), std::move(src_virtual_device), std::move(dst_virtual_device)); -} - -RELAY_REGISTER_OP("device_copy") - .describe(R"code( -Copy data from one tensor to another. The source and destination might be -on different devices. -)code" TVM_ADD_FILELINE) - .set_num_inputs(1) - .add_argument("data", "Tensor", "The input data.") - .set_support_level(10) - .add_type_rel("Identity", IdentityRel) - .set_attrs_type_key("relay.attrs.DeviceCopyAttrs") - .set_attr("TOpPattern", kOpaque) - .set_attr("TOpIsStateful", false) - .set_attr("FInferCorrectLayout", ElemwiseArbitraryLayout) - .set_attr("FTVMCompute", - [](const Attrs& attrs, const Array& inputs, - const Type& out_dtype) -> Array { - return {topi::identity(inputs[0])}; - }); - -// Get device copy props for original device copy op -DeviceCopyProps GetDeviceCopyProps(const CallNode* call_node) { - if (call_node->op == DeviceCopyOp()) { - ICHECK_EQ(call_node->args.size(), 1) << "device_copy expects one argument"; - ICHECK(call_node->attrs.defined()) << "device_copy requires attributes"; - const auto* device_copy_attrs = call_node->attrs.as(); - ICHECK(device_copy_attrs != nullptr) << "device_copy requires DeviceCopyAttrs"; - // Follow nesting: - // device_copy(device_copy(expr, src_virtual_device=S, dst_virtual_device=T), - // src_virtual_device=T, dst_virtual_device=U) ==> {expr, S, U} - auto inner = GetDeviceCopyProps(call_node->args[0]); - if (inner.body.defined()) { - return {inner.body, inner.src_virtual_device, device_copy_attrs->dst_virtual_device}; - } else { - return {call_node->args[0], device_copy_attrs->src_virtual_device, - device_copy_attrs->dst_virtual_device}; - } - } - return {}; -} - -DeviceCopyProps GetDeviceCopyProps(const Expr& expr) { - if (const auto* call_node = expr.as()) { - return GetDeviceCopyProps(call_node); - } - return {}; -} - -} // namespace relay -} // namespace tvm diff --git a/src/relay/op/memory/device_copy.h b/src/relay/op/memory/device_copy.h deleted file mode 100644 index bb74324d5444..000000000000 --- a/src/relay/op/memory/device_copy.h +++ /dev/null @@ -1,84 +0,0 @@ -/* - * 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 relay/op/memory/device_copy.h - * \brief Helpers for working with "device_copy" attributes. - */ - -#ifndef TVM_RELAY_OP_MEMORY_DEVICE_COPY_H_ -#define TVM_RELAY_OP_MEMORY_DEVICE_COPY_H_ - -#include -#include - -#include - -#include "../call/call.h" - -namespace tvm { -namespace relay { - -/*! \brief Returns the "device_copy" operator. */ -const Op& DeviceCopyOp(); - -/*! - * \brief Wraps \p expr in a "device_copy" CallNode indicating it should be evaluated and - * stored at \p src_virtual_device but then copied to \p dst_virtual_device. - */ -Expr DeviceCopy(Expr expr, VirtualDevice src_virtual_device, VirtualDevice dst_virtual_device); - -/*! - * \brief Wraps \p expr in a "device_copy" CallNode indicating it should be evaluated and - * stored at \p src_virtual_device but then copied to \p dst_virtual_device.However, return \p expr - * directly if \p src_virtual_device and \p dst_virtual_device are (structurally) the same. - */ -Expr MaybeDeviceCopy(Expr expr, VirtualDevice src_virtual_device, VirtualDevice dst_virtual_device); - -/*! \brief Result of \p GetDeviceCopyProps. */ -struct DeviceCopyProps { - Expr body; // = null - VirtualDevice src_virtual_device = VirtualDevice::FullyUnconstrained(); - VirtualDevice dst_virtual_device = VirtualDevice::FullyUnconstrained(); - - DeviceCopyProps() = default; - - DeviceCopyProps(Expr body, VirtualDevice src_virtual_device, VirtualDevice dst_virtual_device) - : body(std::move(body)), - src_virtual_device(std::move(src_virtual_device)), - dst_virtual_device(std::move(dst_virtual_device)) {} -}; - -/*! - * \brief Returns the body expression, source, and destination \p VirtualDevices for \p call_node - * if it is a "device_copy" CallNode. Otherwise returns the null expression and unconstrained - * virtual device. - */ -DeviceCopyProps GetDeviceCopyProps(const CallNode* call_node); - -/*! - * \brief Returns the body expression, source, and destination \p VirtualDevices for \p expr if it - * is a "device_copy" Call. Otherwise returns the null expression and unconstrained virtual device. - */ -DeviceCopyProps GetDeviceCopyProps(const Expr& expr); - -} // namespace relay -} // namespace tvm - -#endif // TVM_RELAY_OP_MEMORY_DEVICE_COPY_H_ diff --git a/src/relay/op/memory/memory.cc b/src/relay/op/memory/memory.cc deleted file mode 100644 index 008dbff841f3..000000000000 --- a/src/relay/op/memory/memory.cc +++ /dev/null @@ -1,314 +0,0 @@ -/* - * 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 src/relay/op/memory/memory.cc - * \brief Operators for manifest shape-aware memory allocation in Relay. - */ - -#include "memory.h" - -#include -#include -#include -#include -#include -#include -#include - -#include -#include - -#include "../../transforms/infer_layout_utils.h" -#include "../annotation/annotation.h" -#include "../op_common.h" -#include "../type_relations.h" -#include "on_device.h" - -namespace tvm { -namespace relay { - -TVM_REGISTER_NODE_TYPE(AllocStorageAttrs); -TVM_REGISTER_NODE_TYPE(AllocTensorAttrs); - -// The passing value in attrs and args doesn't seem super great. -// We should consider a better solution, i.e the type relation -// being able to see the arguments as well? -Expr AllocStorage(Expr size, Expr shape, Expr alignment, VirtualDevice virtual_device, - DataType dtype_hint) { - auto attrs = make_object(); - attrs->dtype = dtype_hint; - attrs->virtual_device = std::move(virtual_device); - static const Op& op = Op::Get("memory.alloc_storage"); - return Call(op, {std::move(size), std::move(shape), std::move(alignment)}, - Attrs(std::move(attrs)), {}); -} - -TVM_REGISTER_GLOBAL("relay.op.memory._make.alloc_storage").set_body_typed(AllocStorage); - -bool AllocStorageRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - ICHECK_EQ(types.size(), 4u); - auto size_type = types[0]; - auto tensor_type = size_type.as(); - ICHECK(tensor_type != nullptr); - ICHECK_EQ(tensor_type->dtype, DataType::Int(64)); - ICHECK_EQ(tensor_type->shape.size(), 0); - - // Tensor shape - auto tt = types[1].as(); - ICHECK(tt != nullptr) << "must be tensor type"; - - auto align_type = types[2]; - auto align_ttype = align_type.as(); - ICHECK(align_ttype != nullptr); - ICHECK_EQ(align_ttype->dtype, DataType::Int(64)); - ICHECK_EQ(align_ttype->shape.size(), 0); - auto mod = reporter->GetModule(); - ICHECK(mod.defined()); - auto storage_name = mod->GetGlobalTypeVar("Storage"); - auto storage = TypeCall(storage_name, {}); - reporter->Assign(types[3], storage); - return true; -} - -RELAY_REGISTER_OP("memory.alloc_storage") - .describe(R"code(Explicitly allocate storage to be used by tensors.)code" TVM_ADD_FILELINE) - .set_num_inputs(3) - .add_argument("size", "Tensor", "The size of the storage to allocate.") - .add_argument("shape", "Tensor", "The shape of the storage to allocate.") - .add_argument("alignment", "Tensor", "The alignment of the storage.") - .add_type_rel("AllocStorage", AllocStorageRel) - .set_attrs_type_key("relay.attrs.AllocStorageAttrs") - .set_support_level(10) - .set_attr("TOpPattern", kOpaque) - .set_attr("TOpIsStateful", false) - .set_attr("TNonComputational", true) - .set_attr("FInferCorrectLayout", ElemwiseArbitraryLayout); - -const Op& MemoryAllocTensorOp() { - static const Op& op = Op::Get("memory.alloc_tensor"); - return op; -} - -Expr AllocTensor(Expr storage, Expr offset, Expr shape, DataType dtype, - Array assert_shape) { - auto attrs = make_object(); - attrs->dtype = dtype; - if (assert_shape.defined()) { - attrs->assert_shape = assert_shape; - } else { - // Look through any on_device for the shape argument expression. - const auto* constant_node = AsIgnoringOnDevice(shape); - ICHECK(constant_node); - attrs->const_shape = GetRef(constant_node); - } - return Call(MemoryAllocTensorOp(), {storage, offset, shape}, Attrs(attrs), {}); -} - -TVM_REGISTER_GLOBAL("relay.op.memory._make.alloc_tensor").set_body_typed(AllocTensor); - -std::vector FromConstShape(Constant konst) { - runtime::NDArray shape = konst->data; - std::vector raw_shape; - ICHECK_EQ(shape->ndim, 1u); - ICHECK_EQ(shape->dtype.code, 0U) << "The dtype of constant shape must be int32 or int64, but got " - << runtime::DLDataType2String(shape->dtype); - ICHECK(shape->dtype.bits == 64 || shape->dtype.bits == 32) - << "The dtype of constant shape must be int32 or int64, but got" - << runtime::DLDataType2String(shape->dtype); - - if (shape->dtype.bits == 32) { - const int32_t* int_ptr = reinterpret_cast(shape->data); - for (auto i = 0; i < shape->shape[0]; i++) { - raw_shape.push_back(int_ptr[i]); - } - } else if (shape->dtype.bits == 64) { - const int64_t* int_ptr = reinterpret_cast(shape->data); - for (auto i = 0; i < shape->shape[0]; i++) { - raw_shape.push_back(int_ptr[i]); - } - } - - return raw_shape; -} - -bool AllocTensorRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - ICHECK_EQ(types.size(), 4u); - auto alloc_attrs = attrs.as(); - ICHECK(alloc_attrs != nullptr) << "must be alloc_tensor attributes"; - // First argument should be storage. - auto mod = reporter->GetModule(); - ICHECK(mod.defined()); - auto storage_name = mod->GetGlobalTypeVar("Storage"); - auto storage = relay::TypeCall(storage_name, {}); - reporter->Assign(types[0], storage); - // Second argument should be the offset. - auto offset_type = types[1].as(); - ICHECK(offset_type != nullptr) << "must be a scalar type"; - - // Third argument should be shape tensor. - auto tt = types[2].as(); - ICHECK(tt != nullptr) << "must be tensor type"; - - // Be careful about having to allocate scalars. - int64_t dims = 0; - if (tt->shape.size() != 0) { - auto rank = tt->shape[0].as(); - ICHECK(rank != nullptr); - dims = rank->value; - } - - // Constant node case. - Type alloc_type; - if (alloc_attrs->const_shape.defined()) { - auto con = alloc_attrs->const_shape; - auto sh = FromConstShape(con); - ICHECK_EQ(sh.size(), dims); - Array out_shape; - for (auto i = 0u; i < dims; i++) { - out_shape.push_back(tvm::Integer(sh[i])); - } - alloc_type = TensorType(out_shape, alloc_attrs->dtype); - } else { - ICHECK(alloc_attrs->assert_shape.defined()) - << "the assert_shape must be set when const_shape is not"; - alloc_type = TensorType(alloc_attrs->assert_shape, alloc_attrs->dtype); - return true; - } - - reporter->Assign(types[3], alloc_type); - return true; -} - -RELAY_REGISTER_OP("memory.alloc_tensor") - .describe(R"code(Explicitly allocate storage to be used by tensors.)code" TVM_ADD_FILELINE) - .set_num_inputs(3) - .add_argument("storage", "Storage", "The storage to allocate from.") - .add_argument("offset", "Tensor", "The offset into the backing storage.") - .add_argument("shape", "Tensor", "The shape of the tensor to allocate.") - .add_type_rel("AllocTensor", AllocTensorRel) - .set_attrs_type_key("relay.attrs.AllocTensorAttrs") - .set_support_level(10) - .set_attr("TOpPattern", kOpaque) - .set_attr("TOpIsStateful", false) - .set_attr("TNonComputational", true) - .set_attr("FInferCorrectLayout", ElemwiseArbitraryLayout); - -bool KillRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - ICHECK_EQ(types.size(), 2u); - // TODO(@jroesch): should only support tensors. - reporter->Assign(types[1], TupleType::Empty()); - return true; -} - -RELAY_REGISTER_OP("memory.kill") - .describe(R"code(Mark a variable for release to the allocator.)code" TVM_ADD_FILELINE) - .set_num_inputs(1) - .add_argument("to_free", "Variable", "The variable to free.") - .add_type_rel("Kill", KillRel) - .set_support_level(10) - .set_attr("TOpPattern", kOpaque) - .set_attr("TOpIsStateful", true) - .set_attr("TNonComputational", true) - .set_attr("FInferCorrectLayout", ElemwiseArbitraryLayout); - -static void FlattenTupleTypeAux(const Type& type, std::vector* out) { - if (auto tt = type.as()) { - out->push_back(tt.value()); - } else if (auto tuple_ty = type.as()) { - for (auto field : tuple_ty->fields) { - FlattenTupleTypeAux(field, out); - } - } else { - LOG(FATAL) << "unsupported " << type; - } -} - -std::vector FlattenTupleType(const Type& type) { - std::vector out; - FlattenTupleTypeAux(type, &out); - return out; -} - -static void FromTupleTypeAux(const Type& type, const Expr& expr, std::vector* out) { - if (type.as()) { - out->push_back(expr); - } else if (auto tuple_ty = type.as()) { - for (size_t i = 0; i < tuple_ty->fields.size(); i++) { - FromTupleTypeAux(tuple_ty->fields[i], TupleGetItem(expr, i), out); - } - } else { - LOG(FATAL) << "unsupported " << type; - } -} - -std::vector FromTupleType(const Type& type, const Expr& expr) { - std::vector out; - FromTupleTypeAux(type, expr, &out); - return out; -} - -static void ToTupleTypeAux(const Type& type, const std::vector& exprs, int* index, - std::vector* out) { - if (type.as()) { - out->push_back(exprs[*index]); - *index += 1; - } else if (auto tuple_ty = type.as()) { - std::vector tuple_out; - for (size_t i = 0; i < tuple_ty->fields.size(); i++) { - ToTupleTypeAux(tuple_ty->fields[i], exprs, index, &tuple_out); - } - out->push_back(Tuple(tuple_out)); - } else { - LOG(FATAL) << "unsupported " << type; - } -} - -// Pack the sequence of expressions according to the provided TupleType. -Expr ToTupleType(const Type& t, const std::vector& exprs) { - if (t.as() && exprs.size() == 1) { - return exprs[0]; - } else { - std::vector out; - int index = 0; - ToTupleTypeAux(t, exprs, &index, &out); - return out[0]; - } -} - -TVM_REGISTER_GLOBAL("relay.op.memory._make.FlattenTupleType").set_body_typed([](Type type) { - auto types = FlattenTupleType(type); - return Array(types.begin(), types.end()); -}); - -TVM_REGISTER_GLOBAL("relay.op.memory._make.FromTupleType").set_body_typed([](Type type, Expr expr) { - auto exprs = FromTupleType(type, expr); - return Array(exprs.begin(), exprs.end()); -}); - -TVM_REGISTER_GLOBAL("relay.op.memory._make.ToTupleType") - .set_body_typed([](Type t, Array array) { - return ToTupleType(t, std::vector(array.begin(), array.end())); - }); - -} // namespace relay -} // namespace tvm diff --git a/src/relay/op/memory/memory.h b/src/relay/op/memory/memory.h deleted file mode 100644 index 5533553393ec..000000000000 --- a/src/relay/op/memory/memory.h +++ /dev/null @@ -1,50 +0,0 @@ -/* - * 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 src/relay/op/memory/memory.h - * \brief Operators for memory related operations in Relay. - */ - -#ifndef TVM_RELAY_OP_MEMORY_MEMORY_H_ -#define TVM_RELAY_OP_MEMORY_MEMORY_H_ - -#include - -#include - -#include "tvm/relay/expr.h" - -namespace tvm { -namespace relay { - -Expr AllocStorage(Expr size, Expr shape, Expr alignment, VirtualDevice virtual_device, - DataType dtype_hint); -/*! \brief Returns the "memory.alloc_tensor" operator. */ -const Op& MemoryAllocTensorOp(); -Expr AllocTensor(Expr storage, Expr offset, Expr shape, DataType dtype, - Array assert_shape); -Expr ToTupleType(const Type& ty, const std::vector& exprs); -std::vector FromTupleType(const Type& type, const Expr& expr); -std::vector FlattenTupleType(const Type& type); - -} // namespace relay -} // namespace tvm - -#endif // TVM_RELAY_OP_MEMORY_MEMORY_H_ diff --git a/src/relay/op/memory/on_device.cc b/src/relay/op/memory/on_device.cc deleted file mode 100644 index 155b6daf0848..000000000000 --- a/src/relay/op/memory/on_device.cc +++ /dev/null @@ -1,147 +0,0 @@ -/* - * 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 src/relay/op/memory/on_device.cc - * \brief Helpers for working with the "on_device" 'annotation' call. - */ - -#include "./on_device.h" - -#include -#include -#include -#include -#include - -#include "../../transforms/infer_layout_utils.h" -#include "../type_relations.h" - -namespace tvm { -namespace relay { - -TVM_REGISTER_NODE_TYPE(OnDeviceAttrs); - -const Op& OnDeviceOp() { - static const Op& op = Op::Get("on_device"); - return op; -} - -Call OnDevice(Expr body, VirtualDevice virtual_device, bool constrain_result, bool constrain_body) { - ICHECK((!constrain_result && !constrain_body) || !virtual_device->IsFullyUnconstrained()); - auto attrs = make_object(); - attrs->virtual_device = (constrain_result || constrain_body) - ? std::move(virtual_device) - : VirtualDevice::FullyUnconstrained(); - attrs->constrain_result = constrain_result; - attrs->constrain_body = constrain_body; - Span span = body->span; // about to be moved - return Call(OnDeviceOp(), {std::move(body)}, Attrs(std::move(attrs)), /*type_args=*/{}, - std::move(span)); -} - -TVM_REGISTER_GLOBAL("relay.op.annotation._make.OnDevice").set_body_typed(OnDevice); - -Expr MaybeOnDevice(Expr body, VirtualDevice virtual_device, bool constrain_result, - bool constrain_body) { - if (virtual_device->IsFullyUnconstrained()) { - // Nothing to annotate with. - return body; - } - if (body->IsInstance() || body->IsInstance()) { - // These operators are device polymorphic so no annotation is required. - return body; - } - if (body->IsInstance() || body->IsInstance()) { - // The device can be recovered from the binding site of the global or local variable. - return body; - } - if (body->IsInstance()) { - // If a primitive function then it is device polymorphic. Otherwise the device is captured - // by the function's "result_virtual_device" attribute. - return body; - } - OnDeviceProps props = GetOnDeviceProps(body); - if (props.body.defined()) { - // The user is asking for - // on_device(on_device(body, virtual_device=inner), virtual_device=outer) - // ^ ^ ^ - // outer middle inner - // First recover the implied constraints (if any) for outer and inner, and check they don't - // contradict. - const VirtualDevice& inner = props.virtual_device; - const VirtualDevice& outer = virtual_device; - bool constrain_outer = constrain_result; - bool constrain_inner = props.constrain_body; - if (constrain_outer && constrain_inner) { - ICHECK(inner == outer) << "Cannot constrain result and body of nested on_device calls to " - "different virtual devices"; - } - // There are two possible ways the middle sub-expression may be constrained, check they don't - // contradict. - bool constrain_middle_via_outer = constrain_body; - bool constrain_middle_via_inner = props.constrain_result; - if (constrain_middle_via_outer && constrain_middle_via_inner) { - ICHECK(inner == outer) << "Cannot constrain intermediate result of nested on_device calls to " - "different virtual devices"; - } - // We can now ignore the middle constraint. - // If the outer on_device has any constraint then use virtual_device given for it. - // Otherwise we can use the existing inner virtual_device. - return OnDevice(props.body, (constrain_inner || constrain_outer) ? outer : inner, - constrain_outer, constrain_inner); - } else { - return OnDevice(body, std::move(virtual_device), constrain_result, constrain_body); - } -} - -RELAY_REGISTER_OP("on_device") - .describe(R"code(Annotate an expression with device type)code" TVM_ADD_FILELINE) - .set_num_inputs(1) - .add_argument("body", "Expr", "The sub-expression to be annotated.") - .set_support_level(10) - .add_type_rel("Identity", IdentityRel) - .set_attrs_type_key("relay.attrs.OnDeviceAttrs") - .set_attr("TOpPattern", kOpaque) - .set_attr("TOpIsStateful", false) - .set_attr("FInferCorrectLayout", ElemwiseArbitraryLayout) - .set_attr("TNonComputational", true); - -OnDeviceProps GetOnDeviceProps(const CallNode* call_node) { - if (call_node->op == OnDeviceOp()) { - ICHECK_EQ(call_node->args.size(), 1) << "on_device expects one argument"; - ICHECK(call_node->attrs.defined()) << "on_device requires attributes"; - const auto* on_device_attrs = call_node->attrs.as(); - ICHECK(on_device_attrs != nullptr) << "on_device requires OnDeviceAttrs"; - return {call_node->args[0], on_device_attrs->virtual_device, on_device_attrs->constrain_result, - on_device_attrs->constrain_body}; - } - return {}; -} - -OnDeviceProps GetOnDeviceProps(const Expr& expr) { - if (const auto* call_node = expr.as()) { - return GetOnDeviceProps(call_node); - } - return {}; -} - -} // namespace relay -} // namespace tvm diff --git a/src/relay/op/memory/on_device.h b/src/relay/op/memory/on_device.h deleted file mode 100644 index b597af8fc7fa..000000000000 --- a/src/relay/op/memory/on_device.h +++ /dev/null @@ -1,160 +0,0 @@ -/* - * 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 relay/op/memory/on_device.h - * \brief Helpers for working with the "on_device" 'annotation' call. - */ -#ifndef TVM_RELAY_OP_MEMORY_ON_DEVICE_H_ -#define TVM_RELAY_OP_MEMORY_ON_DEVICE_H_ - -#include -#include -#include -#include - -#include -#include - -namespace tvm { -namespace relay { - -/*! \brief Returns the "on_device" operator. */ -const Op& OnDeviceOp(); - -/*! - * \brief Wraps \p body in an "on_device" CallNode for \p virtual_device. - * - * See \p OnDeviceAttrs for an overview. - */ -Call OnDevice(Expr body, VirtualDevice virtual_device, bool constrain_result = false, - bool constrain_body = true); - -/*! \brief Result of \p GetOnDeviceProps. */ -struct OnDeviceProps { - Expr body; // = null - VirtualDevice virtual_device = VirtualDevice::FullyUnconstrained(); - bool constrain_result = false; - bool constrain_body = false; - - OnDeviceProps() = default; - - OnDeviceProps(Expr body, VirtualDevice virtual_device, bool constrain_result, bool constrain_body) - : body(std::move(body)), - virtual_device(std::move(virtual_device)), - constrain_result(constrain_result), - constrain_body(constrain_body) {} - - bool is_fixed() const { return constrain_result && constrain_body; } - bool is_normal() const { return !constrain_result && constrain_body; } -}; - -/*! - * \brief Wraps \p body in an "on_device" CallNode, taking all fields other than \p body from \p - * props. - */ -inline Call OnDeviceWithProps(Expr body, const OnDeviceProps& props) { - return OnDevice(std::move(body), props.virtual_device, props.constrain_result, - props.constrain_body); -} - -/*! - * \brief Wraps \p body in an "on_device" CallNode, but don't constrain the body or result to - * any particular virtual device. This allows a "device_copy" to be inserted by PlanDevices - * where required, while at the same time not introducing unnecessary freedom in the device - * choices. - */ -inline Call OnDeviceCopyOk(Expr body) { - return OnDevice(std::move(body), VirtualDevice::FullyUnconstrained(), - /*constrain_result=*/false, /*constrain_body=*/false); -} - -/*! - * \brief Wraps \p expr in an "on_device" CallNode for \p virtual_device and \p constraint if the - * \p VirtualDevice for \p expr cannot otherwise be recovered by the lexical scoping convention. - * This means we will NOT wrap if: - * - \p virtual_device is full unconstrained, which signals there are no device annotations - * already in play. - * - \p expr is an operator or primitive function literal. These are device polymorphic. - * - \p expr is a non-primitive function literal. The device is captured by the - * "result_virtual_device" attribute on the function itself. - * - \p expr is a global var. The device is on the function attributes the global is bound to. - * - \p expr is a local var. The device is tracked by the device aware visitors for us. - * - \p expr is a constructor. These are device polymorphic. - * Nested on_device calls will never be constructed, they are instead merged on-the-fly. - */ -Expr MaybeOnDevice(Expr body, VirtualDevice virtual_device, bool constrain_result = false, - bool constrain_body = true); - -/*! \brief As for MaybeOnDevice, but with both body and result constrained. */ -inline Expr MaybeOnDeviceFixed(Expr body, VirtualDevice virtual_device) { - return MaybeOnDevice(std::move(body), std::move(virtual_device), /*constrain_result=*/true, - /*constrain_body=*/true); -} - -/*! \brief As for MaybeOnDevice, but with fields other than body taken from \p props. */ -inline Expr MaybeOnDeviceWithProps(Expr body, const OnDeviceProps& props) { - return MaybeOnDevice(std::move(body), props.virtual_device, props.constrain_result, - props.constrain_body); -} - -/*! - * \brief Returns the body expression, \p VirtualDevice, and constraint field for \p call_node if it - * is an "on_device" CallNode. Otherwise returns the null expression, the unconstrained - * \p VirtualDevice, and \p kBody. - */ -OnDeviceProps GetOnDeviceProps(const CallNode* call_node); - -/*! - * \brief Returns the body expression, \p VirtualDevice, and constraint field for \p expr if it is - * an "on_device" CallNode. Otherwise returns the null expression, the unconstrained \p - * VirtualDevice, and \p kBody. - */ -OnDeviceProps GetOnDeviceProps(const Expr& expr); - -/*! - * \brief Returns the body of \p expr if it is an "on_device" annotation, otherwise returns - * \p expr directly. - */ -inline Expr IgnoreOnDevice(const Expr& expr) { - OnDeviceProps props = GetOnDeviceProps(expr); - return props.body.defined() ? props.body : expr; -} - -/*! - * \brief Returns \p expr as \p NodeType, or null if it is not of that type. Looks through - * any "on_device" annotations. - */ -template -const NodeType* AsIgnoringOnDevice(const Expr& expr) { - const auto* node = expr.as(); - if (node != nullptr) { - return node; - } - OnDeviceProps props = GetOnDeviceProps(expr); - if (!props.body.defined()) { - return nullptr; - } - return props.body.as(); -} - -} // namespace relay -} // namespace tvm - -#endif // TVM_RELAY_OP_MEMORY_ON_DEVICE_H_ diff --git a/src/relay/op/nn/bitserial.cc b/src/relay/op/nn/bitserial.cc deleted file mode 100644 index 496aa3514d88..000000000000 --- a/src/relay/op/nn/bitserial.cc +++ /dev/null @@ -1,260 +0,0 @@ -/* - * 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 bitserial.cc - * \brief Property def of bitserial operators. - */ - -#include -#include -#include - -#include "../../transforms/infer_layout_utils.h" -#include "../op_common.h" - -namespace tvm { -namespace relay { - -// relay.nn.bitpack -TVM_REGISTER_NODE_TYPE(BitPackAttrs); - -template -InferCorrectLayoutOutput BinaryConv2DInferCorrectLayout( - const Attrs& attrs, const Array& new_in_layouts, const Array& old_in_layouts, - const Array& old_in_types) { - const T* params = attrs.as(); - - // We always make other operators to fit the layouts of convolution layers - // So this inference ignores all inputs - return InferCorrectLayoutOutput({params->data_layout, params->kernel_layout}, - {params->data_layout}, attrs); -} - -bool BitPackRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - const BitPackAttrs* param = attrs.as(); - ICHECK_EQ(types.size(), 2); - const auto* data = types[0].as(); - ICHECK(data); - int ndim = data->shape.size(); - int bits = param->bits; - int pack_axis = param->pack_axis; - int bit_axis = param->bit_axis; - DataType pack_type = param->pack_type; - - int pack_bits = pack_type.bits(); - - Array out_shape; - for (int i = 0; i < ndim; ++i) { - if (i == bit_axis) { - out_shape.push_back(bits); - if (i == pack_axis) { - out_shape.push_back(indexdiv(data->shape[i], pack_bits)); - } else { - out_shape.push_back(data->shape[i]); - } - } else if (i == pack_axis) { - out_shape.push_back(indexdiv(data->shape[i], pack_bits)); - } else { - out_shape.push_back(data->shape[i]); - } - } - // Add extra check for last axis expansion. - if (bit_axis == ndim) { - out_shape.push_back(bits); - } - - reporter->Assign(types[1], TensorType(out_shape, pack_type)); - return true; -} - -Expr MakeBitPack(Expr data, int bits, int pack_axis, int bit_axis, DataType pack_type, - String name) { - auto attrs = make_object(); - attrs->bits = bits; - attrs->pack_axis = pack_axis; - attrs->bit_axis = bit_axis; - attrs->pack_type = pack_type; - attrs->name = name; - static const Op& op = Op::Get("nn.bitpack"); - return Call(op, {data}, Attrs(attrs), {}); -} - -TVM_REGISTER_GLOBAL("relay.op.nn._make.bitpack").set_body_typed(MakeBitPack); - -RELAY_REGISTER_OP("nn.bitpack") - .describe(R"code(Bitpack layer that prepares data for bitserial operations. - -This layer backs the bits of an input into a single datatype, allowing -efficient implementation of bitserial operations. - -- **data**: Input tensor of any shape, dimension that is to be - packed must be divisible by number of bits. -- **out**: Packed tensor with shape appropriately compressed. -)code" TVM_ADD_FILELINE) - .set_num_inputs(1) - .set_attrs_type() - .add_argument("data", "Tensor", "Input data.") - .set_support_level(2) - .add_type_rel("BitPack", BitPackRel) - .set_attr("TOpPattern", kInjective); - -// relay.nn.bitserial_conv2d -TVM_REGISTER_NODE_TYPE(BinaryConv2DAttrs); - -bool BinaryConv2DRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - ICHECK_EQ(types.size(), 3); - const auto* data = types[0].as(); - if (data == nullptr) return false; - - const BinaryConv2DAttrs* param = attrs.as(); - ICHECK(param != nullptr); - - static const Layout kNCHW("NCHW"); - - const Layout in_layout(param->data_layout); - const auto trans_in_layout = tir::BijectiveLayout(in_layout, kNCHW); - Array dshape_nchw = trans_in_layout.ForwardShape(data->shape); - ICHECK(param->channels.defined()); - ICHECK(param->kernel_size.defined()); - Array oshape({dshape_nchw[0], param->channels, 0, 0}); - IndexExpr pad_h, pad_w; - GetPaddingHeightWidth(param->padding, &pad_h, &pad_w); - oshape.Set(2, (dshape_nchw[2] + pad_h - param->kernel_size[0]) / param->strides[0] + 1); - oshape.Set(3, (dshape_nchw[3] + pad_w - param->kernel_size[1]) / param->strides[1] + 1); - DataType out_dtype = param->out_dtype; - oshape = trans_in_layout.BackwardShape(oshape); - // assign output type - reporter->Assign(types[2], TensorType(oshape, out_dtype)); - return true; -} - -// Positional relay function to create binaryconv2d operator -// used by frontend FFI. -Expr MakeBinaryConv2D(Expr data, Expr weight, Array strides, Array padding, - IndexExpr channels, Array kernel_size, int activation_bits, - int weight_bits, String data_layout, String kernel_layout, - DataType pack_dtype, DataType out_dtype, bool unipolar) { - auto attrs = make_object(); - attrs->strides = std::move(strides); - attrs->padding = std::move(padding); - attrs->channels = std::move(channels); - attrs->kernel_size = std::move(kernel_size); - attrs->activation_bits = activation_bits; - attrs->weight_bits = weight_bits; - attrs->data_layout = std::move(data_layout); - attrs->kernel_layout = std::move(kernel_layout); - attrs->pack_dtype = std::move(pack_dtype); - attrs->out_dtype = std::move(out_dtype); - attrs->unipolar = unipolar; - static const Op& op = Op::Get("nn.bitserial_conv2d"); - return Call(op, {data, weight}, Attrs(attrs), {}); -} - -TVM_REGISTER_GLOBAL("relay.op.nn._make.bitserial_conv2d").set_body_typed(MakeBinaryConv2D); - -RELAY_REGISTER_OP("nn.bitserial_conv2d") - .describe(R"code(2D convolution using packed binary computation. - -This layer creates a convolution kernel that is convolved with the -layer input using bitserial computation. This enables faster processing -on some platforms. - -- **data**: 4D input tensor that can be either `NCHW` or `NHWC` layout. - -- **weight**: Weight tensor that can either be prepacked (5D) or unpacked (4D). - When data is NCHW, weight is expected to be OIHW or OIHWi. - When data is NHWC weight is expected to be HWIO or HWIOi. - -- **out**: Output with same layout as input. -)code" TVM_ADD_FILELINE) - .set_attrs_type() - .set_num_inputs(2) - .add_argument("data", "Tensor", "The input tensor.") - .add_argument("weight", "Tensor", "The weight tensor.") - .set_support_level(2) - .add_type_rel("BinaryConv2D", BinaryConv2DRel) - .set_attr("FInferCorrectLayout", - BinaryConv2DInferCorrectLayout) - .set_attr("TOpPattern", kOutEWiseFusable); - -// relay.nn.bitserial_dense -TVM_REGISTER_NODE_TYPE(BinaryDenseAttrs); - -bool BinaryDenseRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - ICHECK_EQ(types.size(), 3); - const auto* data = types[0].as(); - if (data == nullptr) return false; - - const BinaryDenseAttrs* param = attrs.as(); - ICHECK(param != nullptr); - - ICHECK(static_cast(data->shape.size()) != 0); - ICHECK(param->units.defined()); - - Array oshape = data->shape; - oshape.Set((oshape.size() - 1), param->units); - - DataType out_dtype = param->out_dtype; - if (out_dtype.bits() == 0) { - out_dtype = data->dtype; - } - - // Assign output type. - reporter->Assign(types[2], TensorType(oshape, out_dtype)); - return true; -} - -// Positional relay function to create bitserial dense operator used by frontend FFI. -Expr MakeBinaryDense(Expr data, Expr weight, IndexExpr units, int data_bits, int weight_bits, - DataType pack_dtype, DataType out_dtype, bool unipolar) { - auto attrs = make_object(); - attrs->units = units; - attrs->data_bits = data_bits; - attrs->weight_bits = weight_bits; - attrs->pack_dtype = pack_dtype; - attrs->out_dtype = out_dtype; - attrs->unipolar = unipolar; - static const Op& op = Op::Get("nn.bitserial_dense"); - return Call(op, {data, weight}, Attrs(attrs), {}); -} - -TVM_REGISTER_GLOBAL("relay.op.nn._make.bitserial_dense").set_body_typed(MakeBinaryDense); - -RELAY_REGISTER_OP("nn.bitserial_dense") - .describe(R"code(Applies a quantized linear transformation: :math:`Y = XW^T`. - -- **data**: `(x1, x2, ..., xn, input_dim)` -- **weight**: `(units, input_dim)` -- **out**: `(x1, x2, ..., xn, units)`. - -)code" TVM_ADD_FILELINE) - .set_attrs_type() - .set_num_inputs(2) - .add_argument("data", "2D Tensor", "Input data.") - .add_argument("weight", "2D Tensor", "Weight matrix.") - .set_support_level(1) - .add_type_rel("BinaryDense", BinaryDenseRel) - .set_attr("TOpPattern", kOutEWiseFusable); - -} // namespace relay -} // namespace tvm diff --git a/src/relay/op/nn/convolution.cc b/src/relay/op/nn/convolution.cc deleted file mode 100644 index 547b533ccc9b..000000000000 --- a/src/relay/op/nn/convolution.cc +++ /dev/null @@ -1,1942 +0,0 @@ -/* - * 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 convolution.cc - * \brief Convolution operators - */ -#include "convolution.h" - -#include -#include -#include - -#include - -#include "../../transforms/infer_layout_utils.h" -#include "../op_common.h" -#include "convolution_make.h" - -namespace tvm { -namespace relay { - -Expr MakeConvWinogradWeightTransform(Expr weight, int tile_size, std::string op_name) { - auto attrs = make_object(); - attrs->tile_size = tile_size; - const Op& op = Op::Get(op_name); - return Call(op, {weight}, Attrs(attrs), {}); -} - -Expr MakeConvGemmWeightTransform(Expr weight, int tile_N, int tile_K, std::string op_name) { - auto attrs = make_object(); - attrs->tile_N = tile_N; - attrs->tile_K = tile_K; - const Op& op = Op::Get(op_name); - return Call(op, {weight}, Attrs(attrs), {}); -} - -// relay.nn.conv1d -TVM_REGISTER_NODE_TYPE(Conv1DAttrs); - -bool Conv1DRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - ICHECK_EQ(types.size(), 3); - const auto* data = types[0].as(); - const auto* weight = types[1].as(); - if (data == nullptr) return false; - static const Layout kNCW("NCW"); - static const Layout kOIW("OIW"); - - const auto* param = attrs.as(); - ICHECK(param != nullptr); - const Layout in_layout(param->data_layout); - const Layout kernel_layout(param->kernel_layout); - - const auto trans_in_layout = tir::BijectiveLayout(in_layout, kNCW); - ICHECK(trans_in_layout.defined()) - << "Conv only support input layouts that are convertible from NCW." - << " But got " << in_layout; - - const auto trans_kernel_layout = tir::BijectiveLayout(kernel_layout, kOIW); - ICHECK(trans_kernel_layout.defined()) - << "Conv only support kernel layouts that are convertible from OIW." - << " But got " << kernel_layout; - - Layout out_layout(param->out_layout == "" ? param->data_layout : param->out_layout); - const auto trans_out_layout = tir::BijectiveLayout(out_layout, kNCW); - ICHECK(trans_out_layout.defined()) - << "Conv only support output layouts that are convertible from NCW." - << " But got " << out_layout; - - Array dshape_ncw = trans_in_layout.ForwardShape(data->shape); - - IndexExpr channels, dilated_ksize; - // infer weight if the kernel_size and channels are defined - if (param->kernel_size.defined() && param->channels.defined()) { - Array wshape; - - wshape = {{param->channels, indexdiv(dshape_ncw[1], param->groups), param->kernel_size[0]}}; - - wshape = trans_kernel_layout.BackwardShape(wshape); - channels = param->channels; - dilated_ksize = 1 + (param->kernel_size[0] - 1) * param->dilation[0]; - DataType weight_dtype = data->dtype; - if (weight != nullptr) { - weight_dtype = weight->dtype; - } - // assign result to reporter - reporter->Assign(types[1], TensorType(wshape, weight_dtype)); - } else { - // use weight to infer the conv shape. - if (weight == nullptr) return false; - auto wshape = trans_kernel_layout.ForwardShape(weight->shape); - if (param->kernel_size.defined()) { - // check the size - ICHECK(reporter->AssertEQ(param->kernel_size[0], wshape[2])) - << "Conv1D: shape of weight is inconsistent with kernel_size, " - << " kernel_size=" << param->kernel_size << " wshape=" << wshape; - } - if (param->channels.defined()) { - ICHECK(reporter->AssertEQ(param->channels, wshape[0])) - << "Conv1D: shape of weight is inconsistent with channels, " - << " channels=" << param->channels << " wshape=" << wshape; - } - if (!dshape_ncw[1].as() && !wshape[1].as()) { - ICHECK(reporter->AssertEQ(dshape_ncw[1], wshape[1])); - } - channels = wshape[0]; - dilated_ksize = 1 + (wshape[2] - 1) * param->dilation[0]; - } - // dilation - Array oshape({dshape_ncw[0], channels, 0}); - - if (!dshape_ncw[2].as()) { - oshape.Set(2, indexdiv(dshape_ncw[2] + param->padding[0] + param->padding[1] - dilated_ksize, - param->strides[0]) + - 1); - } else { - oshape.Set(2, dshape_ncw[2]); - } - - DataType out_dtype = param->out_dtype; - if (out_dtype.bits() == 0) { - out_dtype = data->dtype; - } - oshape = trans_out_layout.BackwardShape(oshape); - // assign output type - reporter->Assign(types[2], TensorType(oshape, out_dtype)); - return true; -} - -TVM_REGISTER_GLOBAL("relay.op.nn._make.conv1d") - .set_body_typed([](Expr data, Expr weight, Array strides, Array padding, - Array dilation, int groups, IndexExpr channels, - Array kernel_size, String data_layout, String kernel_layout, - String out_layout, DataType out_dtype) { - return MakeConv(data, weight, strides, padding, dilation, groups, channels, - kernel_size, data_layout, kernel_layout, out_layout, out_dtype, - "nn.conv1d"); - }); - -RELAY_REGISTER_OP("nn.conv1d") - .describe(R"code(1D convolution layer (e.g. spatial convolution over sequences). - -This layer creates a convolution kernel that is convolved -with the layer input to produce a tensor of outputs. - -- **data**: This depends on the `layout` parameter. Input is 3D array of shape - (batch_size, in_channels, width) if `layout` is `NCW`. -- **weight**: (channels, in_channels, kernel_size) -- **out**: This depends on the `layout` parameter. Output is 3D array of shape - (batch_size, channels, out_width) if `layout` is `NCW`. - -)code" TVM_ADD_FILELINE) - .set_attrs_type() - .set_num_inputs(2) - .add_argument("data", "Tensor", "The input tensor.") - .add_argument("weight", "Tensor", "The weight tensor.") - .set_support_level(2) - .add_type_rel("Conv1D", Conv1DRel) - .set_attr("FInferCorrectLayout", ConvInferCorrectLayout) - .set_attr("TOpPattern", kOutEWiseFusable); - -// relay.nn.conv2d -TVM_REGISTER_NODE_TYPE(Conv2DAttrs); - -bool Conv2DRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - ICHECK_EQ(types.size(), 3); - const auto* data = types[0].as(); - const auto* weight = types[1].as(); - if (data == nullptr) return false; - static const Layout kNCHW("NCHW"); - Layout kOIHW("OIHW"); - - const auto* param = attrs.as(); - DataType out_dtype = param->out_dtype; - if (out_dtype.bits() == 0) { - out_dtype = data->dtype; - if (out_dtype.bits() == 0 && weight != nullptr) { - out_dtype = weight->dtype; - } - } - TensorType meta_schedule_weight{nullptr}; - if (param->meta_schedule_original_shape.size() != 0) { - meta_schedule_weight = TensorType(param->meta_schedule_original_shape, out_dtype); - weight = meta_schedule_weight.get(); - } - ICHECK(param != nullptr); - const Layout in_layout(param->data_layout); - const Layout kernel_layout(param->kernel_layout); - - bool is_dnnl_group_conv = false; - if (param->groups > 1 && kernel_layout.name().find("G") != std::string::npos) { - kOIHW = Layout("GOIHW"); - is_dnnl_group_conv = true; - } - - const auto trans_in_layout = tir::BijectiveLayout(in_layout, kNCHW); - if (!trans_in_layout.defined()) { - reporter->GetDiagCtx().Emit( - Diagnostic::Error(reporter->GetSpan()) - << "conv2d only support input layouts that are convertible from NCHW." - << " The provided layout is: " << in_layout); - return false; - } - - const auto trans_kernel_layout = tir::BijectiveLayout(kernel_layout, kOIHW); - if (!trans_kernel_layout.defined()) { - reporter->GetDiagCtx().Emit(Diagnostic::Error(reporter->GetSpan()) - << "conv2d only support kernel layouts that are convertible from " - << kOIHW << "." - << " The provided layout is: " << kernel_layout); - return false; - } - - Layout out_layout(param->out_layout == "" ? param->data_layout : param->out_layout); - const auto trans_out_layout = tir::BijectiveLayout(out_layout, kNCHW); - if (!trans_out_layout.defined()) { - reporter->GetDiagCtx().Emit( - Diagnostic::Error(reporter->GetSpan()) - << "conv2d only support output layouts that are convertible from NCHW." - << "The provided layout is: " << out_layout); - return false; - } - - Array dshape_nchw = trans_in_layout.ForwardShape(data->shape); - bool is_depthwise = false; - if (param->groups > 1) { - if (!(weight && weight->shape.defined())) { - reporter->GetDiagCtx().Emit( - Diagnostic::Error(reporter->GetSpan()) - << "Weight shape must be specified when groups is greater than 1."); - return false; - } - - Array wshape_oihw = trans_kernel_layout.ForwardShape(weight->shape); - if (tvm::tir::ExprDeepEqual()(param->groups, dshape_nchw[1]) && - tvm::tir::ExprDeepEqual()(param->groups, wshape_oihw[0])) { - is_depthwise = true; - } - } - - IndexExpr channels, dilated_ksize_y, dilated_ksize_x; - // infer weight if the kernel_size and channels are defined - if (param->kernel_size.defined() && param->channels.defined()) { - ICHECK_EQ(param->kernel_size.size(), 2); - ICHECK_EQ(param->dilation.size(), 2); - Array wshape; - - if (is_dnnl_group_conv) { - // infer weight's shape for group convolution - wshape = {{param->groups, indexdiv(param->channels, param->groups), - indexdiv(dshape_nchw[1], param->groups), param->kernel_size[0], - param->kernel_size[1]}}; - } else if (is_depthwise) { - // infer weight's shape for depthwise convolution - wshape = {{dshape_nchw[1], indexdiv(param->channels, dshape_nchw[1]), param->kernel_size[0], - param->kernel_size[1]}}; - } else { - wshape = {{param->channels, indexdiv(dshape_nchw[1], param->groups), param->kernel_size[0], - param->kernel_size[1]}}; - } - - wshape = trans_kernel_layout.BackwardShape(wshape); - channels = param->channels; - dilated_ksize_y = 1 + (param->kernel_size[0] - 1) * param->dilation[0]; - dilated_ksize_x = 1 + (param->kernel_size[1] - 1) * param->dilation[1]; - DataType weight_dtype = data->dtype; - if (weight != nullptr) { - weight_dtype = weight->dtype; - } - - if (param->auto_scheduler_rewritten_layout.size() != 0) { - // If the layout is rewritten by auto-scheduler, - // we just forcly apply the layout provided by auto-scheduler and - // skip the normal inference logic. - {} // do nothing - } else if (param->meta_schedule_original_shape.size() == 0) { - // Normal case: assign result to reporter - reporter->Assign(types[1], TensorType(wshape, weight_dtype)); - } - } else { - // use weight to infer the conv shape. - if (weight == nullptr) return false; - - Array wshape; - if (param->auto_scheduler_rewritten_layout.size() != 0) { - // works for the default kernel layout "HWIO" - ICHECK_EQ(param->kernel_layout, "HWIO"); - wshape = auto_scheduler::GetShapeFromRewrittenLayout(param->auto_scheduler_rewritten_layout, - {"ry", "rx", "rc", "ff"}); - } else { - wshape = weight->shape; - } - - wshape = trans_kernel_layout.ForwardShape(wshape); - if (param->kernel_size.defined()) { - ICHECK_EQ(param->kernel_size.size(), 2); - - if (!reporter->AssertEQ(param->kernel_size[0], wshape[2])) { - reporter->GetDiagCtx().Emit(Diagnostic::Error(reporter->GetSpan()) - << "Conv2D: shape of weight is inconsistent with kernel_size," - << " kernel_size=" << param->kernel_size - << " wshape=" << wshape); - } - - if (!reporter->AssertEQ(param->kernel_size[1], wshape[3])) { - reporter->GetDiagCtx().Emit(Diagnostic::Error(reporter->GetSpan()) - << "Conv2D: shape of weight is inconsistent with kernel_size," - << " kernel_size=" << param->kernel_size - << " wshape=" << wshape); - return false; - } - } - - if (param->channels.defined() && !reporter->AssertEQ(param->channels, wshape[0])) { - reporter->GetDiagCtx().Emit( - Diagnostic::Error(reporter->GetSpan()) - << "conv2D: the first dimensions of the weight tensor (" << wshape << ")" - << "does not match the number of channels (" << param->channels << ")."); - return false; - } - - if (!dshape_nchw[1].as() && !wshape[1].as()) { - if (!reporter->AssertEQ(indexdiv(dshape_nchw[1], param->groups), wshape[1])) { - reporter->GetDiagCtx().Emit(Diagnostic::Error(reporter->GetSpan()) - << "conv2d: requires that `" - << indexdiv(dshape_nchw[1], param->groups) << "`," - << " the input channels (" << dshape_nchw[1] << ")" - << " divided by groups (" << param->groups << ")" - << ",\n must match the input channels" - << " of the weight `" << wshape[1] - << "`, where the weight shape is (" << wshape << ")."); - return false; - } - } - channels = wshape[0]; - dilated_ksize_y = 1 + (wshape[2] - 1) * param->dilation[0]; - dilated_ksize_x = 1 + (wshape[3] - 1) * param->dilation[1]; - } - // dilation - Array oshape({dshape_nchw[0], channels, 0, 0}); - - IndexExpr pad_h, pad_w; - GetPaddingHeightWidth(param->padding, &pad_h, &pad_w); - if (!dshape_nchw[2].as()) { - oshape.Set(2, indexdiv(dshape_nchw[2] + pad_h - dilated_ksize_y, param->strides[0]) + 1); - } else { - oshape.Set(2, dshape_nchw[2]); - } - - if (!dshape_nchw[3].as()) { - oshape.Set(3, indexdiv(dshape_nchw[3] + pad_w - dilated_ksize_x, param->strides[1]) + 1); - } else { - oshape.Set(3, dshape_nchw[3]); - } - oshape = trans_out_layout.BackwardShape(oshape); - // assign output type - reporter->Assign(types[2], TensorType(oshape, out_dtype)); - return true; -} - -TVM_REGISTER_GLOBAL("relay.op.nn._make.conv2d") - .set_body_typed([](Expr data, Expr weight, Array strides, Array padding, - Array dilation, int groups, IndexExpr channels, - Array kernel_size, String data_layout, String kernel_layout, - String out_layout, DataType out_dtype) { - return MakeConv(data, weight, strides, padding, dilation, groups, channels, - kernel_size, data_layout, kernel_layout, out_layout, out_dtype, - "nn.conv2d"); - }); - -RELAY_REGISTER_OP("nn.conv2d") - .describe(R"code(2D convolution layer (e.g. spatial convolution over images). - -This layer creates a convolution kernel that is convolved -with the layer input to produce a tensor of outputs. - -- **data**: This depends on the `layout` parameter. Input is 4D array of shape - (batch_size, in_channels, height, width) if `layout` is `NCHW`. -- **weight**: (channels, in_channels, kernel_size[0], kernel_size[1]) -- **out**: This depends on the `layout` parameter. Output is 4D array of shape - (batch_size, channels, out_height, out_width) if `layout` is `NCHW`. - -)code" TVM_ADD_FILELINE) - .set_attrs_type() - .set_num_inputs(2) - .add_argument("data", "Tensor", "The input tensor.") - .add_argument("weight", "Tensor", "The weight tensor.") - .set_support_level(2) - .add_type_rel("Conv2D", Conv2DRel) - .set_attr("FInferCorrectLayout", ConvInferCorrectLayout) - .set_attr("TOpPattern", kOutEWiseFusable); - -// relay.nn.conv3d -TVM_REGISTER_NODE_TYPE(Conv3DAttrs); - -bool Conv3DRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - ICHECK_EQ(types.size(), 3); - const auto* data = types[0].as(); - const auto* weight = types[1].as(); - if (data == nullptr) return false; - static const Layout kNCDHW("NCDHW"); - static const Layout kOIDHW("OIDHW"); - - const auto* param = attrs.as(); - ICHECK(param != nullptr); - DataType out_dtype = param->out_dtype; - if (out_dtype.bits() == 0) { - out_dtype = data->dtype; - if (out_dtype.bits() == 0 && weight != nullptr) { - out_dtype = weight->dtype; - } - } - TensorType meta_schedule_weight{nullptr}; - if (param->meta_schedule_original_shape.size() != 0) { - meta_schedule_weight = TensorType(param->meta_schedule_original_shape, out_dtype); - weight = meta_schedule_weight.get(); - } - const Layout in_layout(param->data_layout); - const Layout kernel_layout(param->kernel_layout); - - const auto trans_in_layout = tir::BijectiveLayout(in_layout, kNCDHW); - ICHECK(trans_in_layout.defined()) - << "Conv only support input layouts that are convertible from NCDHW." - << " But got " << in_layout; - - const auto trans_kernel_layout = tir::BijectiveLayout(kernel_layout, kOIDHW); - ICHECK(trans_kernel_layout.defined()) - << "Conv only support kernel layouts that are convertible from OIDHW." - << " But got " << kernel_layout; - - Layout out_layout(param->out_layout == "" ? param->data_layout : param->out_layout); - const auto trans_out_layout = tir::BijectiveLayout(out_layout, kNCDHW); - ICHECK(trans_out_layout.defined()) - << "Conv only support output layouts that are convertible from NCDHW." - << " But got " << out_layout; - - Array dshape_ncdhw = trans_in_layout.ForwardShape(data->shape); - - IndexExpr channels, dilated_ksize_z, dilated_ksize_y, dilated_ksize_x; - // infer weight if the kernel_size and channels are defined - if (param->kernel_size.defined() && param->channels.defined()) { - ICHECK_EQ(param->kernel_size.size(), 3); - ICHECK_EQ(param->dilation.size(), 3); - - bool is_depthwise = false; - if (param->groups > 1) { - if (!(weight && weight->shape.defined())) { - reporter->GetDiagCtx().Emit( - Diagnostic::Error(reporter->GetSpan()) - << "Weight shape must be specified when groups is greater than 1."); - return false; - } - - Array wshape_oidhw = trans_kernel_layout.ForwardShape(weight->shape); - if (tvm::tir::ExprDeepEqual()(param->groups, dshape_ncdhw[1]) && - tvm::tir::ExprDeepEqual()(param->groups, wshape_oidhw[0])) { - is_depthwise = true; - } - } - - Array wshape; - if (is_depthwise) { - auto channel_multiplier = indexdiv(param->channels, dshape_ncdhw[1]); - wshape = {dshape_ncdhw[1], channel_multiplier, param->kernel_size[0], param->kernel_size[1], - param->kernel_size[2]}; - } else { - wshape = {param->channels, indexdiv(dshape_ncdhw[1], param->groups), param->kernel_size[0], - param->kernel_size[1], param->kernel_size[2]}; - } - - wshape = trans_kernel_layout.BackwardShape(wshape); - channels = param->channels; - dilated_ksize_z = 1 + (param->kernel_size[0] - 1) * param->dilation[0]; - dilated_ksize_y = 1 + (param->kernel_size[1] - 1) * param->dilation[1]; - dilated_ksize_x = 1 + (param->kernel_size[2] - 1) * param->dilation[2]; - DataType weight_dtype = data->dtype; - if (weight != nullptr) { - weight_dtype = weight->dtype; - } - - if (param->auto_scheduler_rewritten_layout.size() != 0) { - // If the layout is rewritten by auto-scheduler, - // we just forcly apply the layout provided by auto-scheduler and - // skip the normal inference logic. - {} // do nothing - } else if (param->meta_schedule_original_shape.size() == 0) { - // Normal case: assign result to reporter - reporter->Assign(types[1], TensorType(wshape, weight_dtype)); - } - - } else { - // use weight to infer the conv shape. - if (weight == nullptr) return false; - - Array wshape; - if (param->auto_scheduler_rewritten_layout.size() != 0) { - // works for the default kernel layout "DHWIO" - ICHECK_EQ(param->kernel_layout, "DHWIO"); - wshape = auto_scheduler::GetShapeFromRewrittenLayout(param->auto_scheduler_rewritten_layout, - {"rd", "rh", "rw", "rc", "cc"}); - } else { - wshape = weight->shape; - } - - wshape = trans_kernel_layout.ForwardShape(wshape); - if (param->kernel_size.defined()) { - ICHECK_EQ(param->kernel_size.size(), 3); - // check the size - ICHECK(reporter->AssertEQ(param->kernel_size[0], wshape[2]) && - reporter->AssertEQ(param->kernel_size[1], wshape[3]) && - reporter->AssertEQ(param->kernel_size[2], wshape[4])) - << "Conv3D: shape of weight is inconsistent with kernel_size, " - << " kernel_size=" << param->kernel_size << " wshape=" << wshape; - } - - if (param->channels.defined()) { - ICHECK(reporter->AssertEQ(param->channels, wshape[0])) - << "Conv3D: shape of weight is inconsistent with channels, " - << " channels=" << param->channels << " wshape=" << wshape; - } - - if (!dshape_ncdhw[1].as() && !wshape[1].as()) { - ICHECK(reporter->AssertEQ(indexdiv(dshape_ncdhw[1], param->groups), wshape[1])); - } - channels = wshape[0]; - dilated_ksize_z = 1 + (wshape[2] - 1) * param->dilation[0]; - dilated_ksize_y = 1 + (wshape[3] - 1) * param->dilation[1]; - dilated_ksize_x = 1 + (wshape[4] - 1) * param->dilation[2]; - } - // dilation - Array oshape({dshape_ncdhw[0], channels, 0, 0, 0}); - - IndexExpr pad_d, pad_h, pad_w; - GetPaddingDepthHeightWidth(param->padding, &pad_d, &pad_h, &pad_w); - if (!dshape_ncdhw[2].as()) { - oshape.Set(2, indexdiv(dshape_ncdhw[2] + pad_d - dilated_ksize_z, param->strides[0]) + 1); - } else { - oshape.Set(2, dshape_ncdhw[2]); - } - - if (!dshape_ncdhw[3].as()) { - oshape.Set(3, indexdiv(dshape_ncdhw[3] + pad_h - dilated_ksize_y, param->strides[1]) + 1); - } else { - oshape.Set(3, dshape_ncdhw[3]); - } - - if (!dshape_ncdhw[4].as()) { - oshape.Set(4, indexdiv(dshape_ncdhw[4] + pad_w - dilated_ksize_x, param->strides[2]) + 1); - } else { - oshape.Set(4, dshape_ncdhw[4]); - } - oshape = trans_out_layout.BackwardShape(oshape); - // assign output type - reporter->Assign(types[2], TensorType(oshape, out_dtype)); - return true; -} - -TVM_REGISTER_GLOBAL("relay.op.nn._make.conv3d") - .set_body_typed([](Expr data, Expr weight, Array strides, Array padding, - Array dilation, int groups, IndexExpr channels, - Array kernel_size, String data_layout, String kernel_layout, - String out_layout, DataType out_dtype) { - return MakeConv(data, weight, strides, padding, dilation, groups, channels, - kernel_size, data_layout, kernel_layout, out_layout, out_dtype, - "nn.conv3d"); - }); - -RELAY_REGISTER_OP("nn.conv3d") - .describe(R"code(3D convolution layer (e.g. convolution over 3D image data, -like Magnetic Resonance Imaging (MRI) data in medicine). - -This layer creates a convolution kernel that is convolved -with the layer input to produce a tensor of outputs. - -- **data**: This depends on the `layout` parameter. Input is 5D array of shape - (batch_size, in_channels, depth, height, width) if `layout` is `NCDHW`. -- **weight**: (channels, in_channels, kernel_size[0], kernel_size[1], kernel_size[2]) -- **out**: This depends on the `layout` parameter. Output is 5D array of shape - (batch_size, channels, out_depth, out_height, out_width) if `layout` is `NCDHW`. - -)code" TVM_ADD_FILELINE) - .set_attrs_type() - .set_num_inputs(2) - .add_argument("data", "Tensor", "The input tensor.") - .add_argument("weight", "Tensor", "The weight tensor.") - .set_support_level(2) - .add_type_rel("Conv3D", Conv3DRel) - .set_attr("FInferCorrectLayout", ConvInferCorrectLayout) - .set_attr("TOpPattern", kOutEWiseFusable); - -// relay.nn.conv3d_transpose -TVM_REGISTER_NODE_TYPE(Conv3DTransposeAttrs); - -template -bool Conv3DTransposeRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - ICHECK_EQ(types.size(), 3); - const auto* data = types[0].as(); - const auto* weight = types[1].as(); - if (data == nullptr) return false; - - static const Layout kNCDHW("NCDHW"); - static const Layout kIODHW("IODHW"); - - const Conv3DTransposeAttrs* param = attrs.as(); - ICHECK(param != nullptr); - const Layout in_layout(param->data_layout); - const Layout kernel_layout(param->kernel_layout); - - const auto trans_in_layout = tir::BijectiveLayout(in_layout, kNCDHW); - ICHECK(trans_in_layout.defined()) - << "Conv3d_transpose only support input layouts that are convertible from NCDHW." - << " But got " << in_layout; - - const auto trans_kernel_layout = tir::BijectiveLayout(kernel_layout, kIODHW); - ICHECK(trans_kernel_layout.defined()) - << "Conv3d_transpose only support kernel layouts that are convertible from IODHW." - << " But got " << kernel_layout; - - Layout out_layout(param->out_layout == "" ? param->data_layout : param->out_layout); - const auto trans_out_layout = tir::BijectiveLayout(out_layout, kNCDHW); - ICHECK(trans_out_layout.defined()) - << "Conv3d_transpose only support output layouts that are convertible from NCDHW." - << " But got " << out_layout; - - IndexExpr channels, dilated_ksize_d, dilated_ksize_y, dilated_ksize_x; - - auto dshape_ncdhw = trans_in_layout.ForwardShape(data->shape); - - // infer weight if the kernel_size and channels are defined - if (param->kernel_size.defined() && param->channels.defined()) { - ICHECK_EQ(param->kernel_size.size(), 3); - ICHECK_EQ(param->dilation.size(), 3); - - Array wshape({dshape_ncdhw[1], indexdiv(param->channels, param->groups), - param->kernel_size[0], param->kernel_size[1], param->kernel_size[2]}); - - wshape = trans_kernel_layout.BackwardShape(wshape); - dilated_ksize_d = 1 + (param->kernel_size[0] - 1) * param->dilation[0]; - dilated_ksize_y = 1 + (param->kernel_size[1] - 1) * param->dilation[1]; - dilated_ksize_x = 1 + (param->kernel_size[2] - 1) * param->dilation[2]; - channels = param->channels; - - DataType weight_dtype = data->dtype; - if (weight != nullptr) { - weight_dtype = weight->dtype; - } - // assign result to reporter - reporter->Assign(types[1], TensorType(wshape, weight_dtype)); - } else { - // use weight to infer the conv shape. - if (weight == nullptr) return false; - auto wshape = trans_kernel_layout.ForwardShape(weight->shape); - if (param->kernel_size.defined()) { - ICHECK_EQ(param->kernel_size.size(), 3); - // check the size - ICHECK(reporter->AssertEQ(param->kernel_size[0], wshape[2]) && - reporter->AssertEQ(param->kernel_size[1], wshape[3]) && - reporter->AssertEQ(param->kernel_size[2], wshape[4])) - << "Conv3DTransposed: shape of weight is inconsistent with kernel_size, " - << " kernel_size=" << param->kernel_size << " wshape=" << Array(wshape); - } - if (param->channels.defined()) { - ICHECK(reporter->AssertEQ(indexdiv(param->channels, param->groups), wshape[1])) - << "Conv3DTransposed: shape of weight is inconsistent out_channels, " - << " out_channels // groups != weight.shape[1] " - << " out_channels=" << param->channels << " groups=" << param->groups - << " wshape=" << Array(wshape); - } - if (!dshape_ncdhw[1].as() && !wshape[0].as()) { - ICHECK(reporter->AssertEQ(dshape_ncdhw[1], wshape[0])); - } - channels = wshape[1]; - dilated_ksize_d = 1 + (wshape[2] - 1) * param->dilation[0]; - dilated_ksize_x = 1 + (wshape[3] - 1) * param->dilation[1]; - dilated_ksize_y = 1 + (wshape[4] - 1) * param->dilation[2]; - } - - // dilation - Array oshape({dshape_ncdhw[0], channels, 0, 0, 0}); - IndexExpr pad_d, pad_h, pad_w; - GetPaddingDepthHeightWidth(param->padding, &pad_d, &pad_h, &pad_w); - - if (!dshape_ncdhw[2].as()) { - oshape.Set(2, (param->strides[0] * (dshape_ncdhw[2] - 1) + dilated_ksize_d - pad_d + - param->output_padding[0])); - } else { - oshape.Set(2, dshape_ncdhw[2]); - } - if (!dshape_ncdhw[3].as()) { - oshape.Set(3, (param->strides[1] * (dshape_ncdhw[3] - 1) + dilated_ksize_y - pad_h + - param->output_padding[1])); - } else { - oshape.Set(3, dshape_ncdhw[3]); - } - if (!dshape_ncdhw[4].as()) { - oshape.Set(4, (param->strides[2] * (dshape_ncdhw[4] - 1) + dilated_ksize_x - pad_w + - param->output_padding[2])); - } else { - oshape.Set(4, dshape_ncdhw[4]); - } - - DataType out_dtype = param->out_dtype; - if (out_dtype.bits() == 0) { - out_dtype = data->dtype; - } - oshape = trans_out_layout.BackwardShape(oshape); - reporter->Assign(types[2], TensorType(oshape, out_dtype)); - return true; -} - -TVM_REGISTER_GLOBAL("relay.op.nn._make.conv3d_transpose") - .set_body_typed([](Expr data, Expr weight, Array strides, Array padding, - Array dilation, int groups, IndexExpr channels, - Array kernel_size, String data_layout, String kernel_layout, - String out_layout, Array output_padding, DataType out_dtype) { - return MakeConvTranspose( - data, weight, strides, padding, dilation, groups, channels, kernel_size, data_layout, - kernel_layout, out_layout, output_padding, out_dtype, "nn.conv3d_transpose"); - }); - -RELAY_REGISTER_OP("nn.conv3d_transpose") - .describe(R"code(Transposed 3D convolution layer (sometimes called Deconvolution 3D). - -The need for transposed convolutions generally arises -from the desire to use a transformation going in the opposite direction -of a normal convolution, i.e., from something that has the shape of the -output of some convolution to something that has the shape of its input -while maintaining a connectivity pattern that is compatible with -said convolution. - -- **data**: This depends on the `layout` parameter. Input is 5D array of shape - (batch_size, in_channels, depth, height, width) if `layout` is `NCDHW`. -- **weight**: (in_channels, channels, kernel_size[0], kernel_size[1], kernel_size[2]) -- **bias**: (channels,) -- **out**: This depends on the `layout` parameter. Output is 5D array of shape - (batch_size, channels, out_depth, out_height, out_width) if `layout` is `NCDHW`. - - out_depth and out_height and out_width are calculated as:: - out_depth = (depth-1)*strides[0]-2*padding[0]+kernel_size[0]+output_padding[0] - out_height = (height-1)*strides[1]-2*padding[1]+kernel_size[1]+output_padding[1] - out_width = (width-1)*strides[2]-2*padding[2]+kernel_size[2]+output_padding[2] - -)code" TVM_ADD_FILELINE) - .set_attrs_type() - .set_num_inputs(2) - .add_argument("data", "Tensor", "The input tensor.") - .add_argument("weight", "Tensor", "The weight tensor.") - .set_support_level(2) - .set_attr("FInferCorrectLayout", - ConvInferCorrectLayout) - .add_type_rel("Conv3DTranspose", Conv3DTransposeRel) - .set_attr("TOpPattern", kOutEWiseFusable); - -// relay.nn.conv2d_transpose -TVM_REGISTER_NODE_TYPE(Conv2DTransposeAttrs); - -bool Conv2DTransposeRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - ICHECK_EQ(types.size(), 3); - const auto* data = types[0].as(); - const auto* weight = types[1].as(); - if (data == nullptr) return false; - - static const Layout kNCHW("NCHW"); - Layout kIOHW("IOHW"); - - const Conv2DTransposeAttrs* param = attrs.as(); - ICHECK(param != nullptr); - const Layout in_layout(param->data_layout); - const Layout kernel_layout(param->kernel_layout); - - bool is_dnnl_group_conv = false; - if (param->groups > 1 && kernel_layout.name().find("G") != std::string::npos) { - kIOHW = Layout("GIOHW"); - is_dnnl_group_conv = true; - } - - const auto trans_in_layout = tir::BijectiveLayout(in_layout, kNCHW); - ICHECK(trans_in_layout.defined()) - << "Conv2DTransposed only support input layouts that are convertible from NCHW." - << " But got " << in_layout; - - const auto trans_kernel_layout = tir::BijectiveLayout(kernel_layout, kIOHW); - ICHECK(trans_kernel_layout.defined()) - << "Conv2DTransposed only support kernel layouts that are convertible from " << kIOHW << "." - << " But got " << kernel_layout << " " << kIOHW; - - Layout out_layout(param->out_layout == "" ? param->data_layout : param->out_layout); - const auto trans_out_layout = tir::BijectiveLayout(out_layout, kNCHW); - ICHECK(trans_out_layout.defined()) - << "Conv2DTransposed only support output layouts that are convertible from NCHW." - << " But got " << out_layout; - - IndexExpr channels, dilated_ksize_y, dilated_ksize_x; - - auto dshape_nchw = trans_in_layout.ForwardShape(data->shape); - - // infer weight if the kernel_size and channels are defined - if (param->kernel_size.defined() && param->channels.defined()) { - ICHECK_EQ(param->kernel_size.size(), 2); - ICHECK_EQ(param->dilation.size(), 2); - - Array wshape; - if (is_dnnl_group_conv) { - // infer weight's shape for group convolution - wshape = {{param->groups, indexdiv(dshape_nchw[1], param->groups), - indexdiv(param->channels, param->groups), param->kernel_size[0], - param->kernel_size[1]}}; - } else { - // infer weight's shape for depthwise convolution - wshape = {{dshape_nchw[1], indexdiv(param->channels, param->groups), param->kernel_size[0], - param->kernel_size[1]}}; - } - - wshape = trans_kernel_layout.BackwardShape(wshape); - dilated_ksize_y = 1 + (param->kernel_size[0] - 1) * param->dilation[0]; - dilated_ksize_x = 1 + (param->kernel_size[1] - 1) * param->dilation[1]; - channels = param->channels; - - DataType weight_dtype = data->dtype; - if (weight != nullptr) { - weight_dtype = weight->dtype; - } - // assign result to reporter - reporter->Assign(types[1], TensorType(wshape, weight_dtype)); - } else { - // use weight to infer the conv shape. - if (weight == nullptr) return false; - auto wshape = trans_kernel_layout.ForwardShape(weight->shape); - if (param->kernel_size.defined()) { - ICHECK_EQ(param->kernel_size.size(), 2); - // check the size - ICHECK(reporter->AssertEQ(param->kernel_size[0], wshape[2]) && - reporter->AssertEQ(param->kernel_size[1], wshape[3])) - << "Conv2DTransposed: shape of weight is inconsistent with kernel_size, " - << " kernel_size=" << param->kernel_size << " wshape=" << Array(wshape); - } - if (param->channels.defined()) { - ICHECK(reporter->AssertEQ(indexdiv(param->channels, param->groups), wshape[1])) - << "Conv2DTransposed: shape of weight is inconsistent with out_channels, " - << " out_channels // groups != weight.shape[1] " - << " out_channels=" << param->channels << " groups=" << param->groups - << " weight.shape=" << Array(wshape); - } - if (!dshape_nchw[1].as() && !wshape[0].as()) { - ICHECK(reporter->AssertEQ(dshape_nchw[1], wshape[0])) - << "Conv2DTransposed: shape of weight is inconsistent with in_channels." - << " data.shape= " << Array(dshape_nchw) << " groups= " << param->groups - << " weight.shape= " << Array(wshape); - } - channels = wshape[1]; - dilated_ksize_y = 1 + (wshape[2] - 1) * param->dilation[0]; - dilated_ksize_x = 1 + (wshape[3] - 1) * param->dilation[1]; - } - // dilation - Array oshape({dshape_nchw[0], channels, 0, 0}); - IndexExpr pad_h, pad_w; - GetPaddingHeightWidth(param->padding, &pad_h, &pad_w); - if (!dshape_nchw[2].as()) { - oshape.Set(2, (param->strides[0] * (dshape_nchw[2] - 1) + dilated_ksize_y - pad_h + - param->output_padding[0])); - } else { - oshape.Set(2, dshape_nchw[2]); - } - if (!dshape_nchw[3].as()) { - oshape.Set(3, (param->strides[1] * (dshape_nchw[3] - 1) + dilated_ksize_x - pad_w + - param->output_padding[1])); - } else { - oshape.Set(3, dshape_nchw[3]); - } - - DataType out_dtype = param->out_dtype; - if (out_dtype.bits() == 0) { - out_dtype = data->dtype; - } - oshape = trans_out_layout.BackwardShape(oshape); - reporter->Assign(types[2], TensorType(oshape, out_dtype)); - return true; -} - -TVM_REGISTER_GLOBAL("relay.op.nn._make.conv2d_transpose") - .set_body_typed([](Expr data, Expr weight, Array strides, Array padding, - Array dilation, int groups, IndexExpr channels, - Array kernel_size, String data_layout, String kernel_layout, - String out_layout, Array output_padding, DataType out_dtype) { - return MakeConvTranspose( - data, weight, strides, padding, dilation, groups, channels, kernel_size, data_layout, - kernel_layout, out_layout, output_padding, out_dtype, "nn.conv2d_transpose"); - }); - -RELAY_REGISTER_OP("nn.conv2d_transpose") - .describe(R"code(Transposed 2D convolution layer (sometimes called Deconvolution). - -The need for transposed convolutions generally arises -from the desire to use a transformation going in the opposite direction -of a normal convolution, i.e., from something that has the shape of the -output of some convolution to something that has the shape of its input -while maintaining a connectivity pattern that is compatible with -said convolution. - -- **data**: This depends on the `layout` parameter. Input is 4D array of shape - (batch_size, in_channels, height, width) if `layout` is `NCHW`. -- **weight**: (in_channels, channels, kernel_size[0], kernel_size[1]) -- **bias**: (channels,) -- **out**: This depends on the `layout` parameter. Output is 4D array of shape -v (batch_size, channels, out_height, out_width) if `layout` is `NCHW`. - - out_height and out_width are calculated as:: - out_height = (height-1)*strides[0]-2*padding[0]+kernel_size[0]+output_padding[0] - out_width = (width-1)*strides[1]-2*padding[1]+kernel_size[1]+output_padding[1] - -)code" TVM_ADD_FILELINE) - .set_attrs_type() - .set_num_inputs(2) - .add_argument("data", "Tensor", "The input tensor.") - .add_argument("weight", "Tensor", "The weight tensor.") - .set_support_level(2) - .set_attr("FInferCorrectLayout", - ConvInferCorrectLayout) - .add_type_rel("Conv2DTranspose", Conv2DTransposeRel) - .set_attr("TOpPattern", kOutEWiseFusable); - -// relay.nn.conv1d_transpose -TVM_REGISTER_NODE_TYPE(Conv1DTransposeAttrs); - -bool Conv1DTransposeRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - ICHECK_EQ(types.size(), 3); - const auto* data = types[0].as(); - const auto* weight = types[1].as(); - if (data == nullptr) return false; - - static const Layout kNCW("NCW"); - static const Layout kIOW("IOW"); - - const Conv1DTransposeAttrs* param = attrs.as(); - ICHECK(param != nullptr); - const Layout in_layout(param->data_layout); - const Layout kernel_layout(param->kernel_layout); - - const auto trans_in_layout = tir::BijectiveLayout(in_layout, kNCW); - ICHECK(trans_in_layout.defined()) - << "Conv only support input layouts that are convertible from NCW." - << " But got " << in_layout; - - const auto trans_kernel_layout = tir::BijectiveLayout(kernel_layout, kIOW); - ICHECK(trans_kernel_layout.defined()) - << "Conv only support kernel layouts that are convertible from IOW." - << " But got " << kernel_layout; - - Layout out_layout(param->out_layout == "" ? param->data_layout : param->out_layout); - const auto trans_out_layout = tir::BijectiveLayout(out_layout, kNCW); - ICHECK(trans_out_layout.defined()) - << "Conv only support output layouts that are convertible from NCW." - << " But got " << out_layout; - - IndexExpr channels, dilated_ksize_y, dilated_ksize_x; - - auto dshape_ncw = trans_in_layout.ForwardShape(data->shape); - - // infer weight if the kernel_size and channels are defined - if (param->kernel_size.defined() && param->channels.defined()) { - ICHECK_EQ(param->kernel_size.size(), 1); - ICHECK_EQ(param->dilation.size(), 1); - - Array wshape( - {dshape_ncw[1], indexdiv(param->channels, param->groups), param->kernel_size[0]}); - - wshape = trans_kernel_layout.BackwardShape(wshape); - dilated_ksize_x = 1 + (param->kernel_size[0] - 1) * param->dilation[0]; - channels = param->channels; - - DataType weight_dtype = data->dtype; - if (weight != nullptr) { - weight_dtype = weight->dtype; - } - // assign result to reporter - reporter->Assign(types[1], TensorType(wshape, weight_dtype)); - } else { - // use weight to infer the conv shape. - if (weight == nullptr) return false; - auto wshape = trans_kernel_layout.ForwardShape(weight->shape); - if (param->kernel_size.defined()) { - ICHECK_EQ(param->kernel_size.size(), 1); - // check the size - ICHECK(reporter->AssertEQ(param->kernel_size[0], wshape[2])) - << "Conv1DTraspose: shape of weight is inconsistent with kernel_size, " - << " kernel_size=" << param->kernel_size << " wshape=" << Array(wshape); - } - if (param->channels.defined()) { - ICHECK(reporter->AssertEQ(indexdiv(param->channels, param->groups), wshape[1])) - << "Conv1DTraspose: shape of weight is inconsistent with channels, " - << " out_channels // groups != weight.shape[1] " - << " out_channels=" << param->channels << " groups=" << param->groups - << " wshape=" << Array(wshape); - } - if (!dshape_ncw[1].as() && !wshape[0].as()) { - ICHECK(reporter->AssertEQ(dshape_ncw[1], wshape[0])); - } - channels = wshape[1]; - dilated_ksize_x = 1 + (wshape[2] - 1) * param->dilation[0]; - } - // dilation - IndexExpr pad_w; - GetPaddingWidth(param->padding, &pad_w); - Array oshape({dshape_ncw[0], channels, 0}); - if (!dshape_ncw[2].as()) { - oshape.Set(2, (param->strides[0] * (dshape_ncw[2] - 1) + dilated_ksize_x - pad_w + - param->output_padding[0])); - } else { - oshape.Set(2, dshape_ncw[2]); - } - - DataType out_dtype = param->out_dtype; - if (out_dtype.bits() == 0) { - out_dtype = data->dtype; - } - oshape = trans_out_layout.BackwardShape(oshape); - reporter->Assign(types[2], TensorType(oshape, out_dtype)); - return true; -} - -TVM_REGISTER_GLOBAL("relay.op.nn._make.conv1d_transpose") - .set_body_typed([](Expr data, Expr weight, Array strides, Array padding, - Array dilation, int groups, IndexExpr channels, - Array kernel_size, String data_layout, String kernel_layout, - String out_layout, Array output_padding, DataType out_dtype) { - return MakeConvTranspose( - data, weight, strides, padding, dilation, groups, channels, kernel_size, data_layout, - kernel_layout, out_layout, output_padding, out_dtype, "nn.conv1d_transpose"); - }); - -RELAY_REGISTER_OP("nn.conv1d_transpose") - .describe(R"code(Transposed 1D convolution layer (sometimes called Deconvolution). - -The need for transposed convolutions generally arises -from the desire to use a transformation going in the opposite direction -of a normal convolution, i.e., from something that has the shape of the -output of some convolution to something that has the shape of its input -while maintaining a connectivity pattern that is compatible with -said convolution. - -- **data**: This depends on the `layout` parameter. Input is 3D array of shape - (batch_size, in_channels, width) if `layout` is `NCW`. -- **weight**: (in_channels, channels, kernel_size[0]) -- **bias**: (channels,) -- **out**: This depends on the `layout` parameter. Output is 3D array of shape - (batch_size, channels, out_width) if `layout` is `NCW`. - - out_width is calculated as:: - out_width = (width-1)*strides[0]-2*padding[0]+kernel_size[0]+output_padding[0] - -)code" TVM_ADD_FILELINE) - .set_attrs_type() - .set_num_inputs(2) - .add_argument("data", "Tensor", "The input tensor.") - .add_argument("weight", "Tensor", "The weight tensor.") - .set_support_level(2) - .add_type_rel("Conv1DTranspose", Conv1DTransposeRel) - .set_attr("TOpPattern", kOutEWiseFusable); - -// relay.nn.contrib_conv2d_winograd_without_weight_transform -TVM_REGISTER_NODE_TYPE(Conv2DWinogradAttrs); - -TVM_REGISTER_GLOBAL("relay.op.nn._make.contrib_conv2d_winograd_without_weight_transform") - .set_body_typed([](Expr data, Expr weight, int tile_size, Array strides, - Array padding, Array dilation, int groups, - IndexExpr channels, Array kernel_size, String data_layout, - String kernel_layout, String out_layout, DataType out_dtype) { - return MakeConvWinograd( - data, weight, tile_size, strides, padding, dilation, groups, channels, kernel_size, - data_layout, kernel_layout, out_layout, out_dtype, - "nn.contrib_conv2d_winograd_without_weight_transform"); - }); - -RELAY_REGISTER_OP("nn.contrib_conv2d_winograd_without_weight_transform") - .describe(R"code(Compute conv2d with winograd algorithm. Only supports NCHW layout. - This operator assumes the weight tensor is already pre-transformed by - nn.contrib_conv2d_winograd_weight_transform. - -- **data**: Input is 4D array of shape (batch_size, in_channels, height, width) -- **weight**: Any shape - We do not check the shape for this input tensor. Since different backend - has different layout strategy. - -- **out**: Output is 4D array of shape (batch_size, channels, out_height, out_width) -)code" TVM_ADD_FILELINE) - .set_attrs_type() - .set_num_inputs(2) - .add_argument("data", "Tensor", "The input tensor.") - .add_argument("weight", "Tensor", "The weight tensor.") - .set_support_level(10) - .add_type_rel("Conv2DWinograd", Conv2DWinogradRel) - .set_attr("FInferCorrectLayout", - ConvInferCorrectLayout) - .set_attr("TOpPattern", kOutEWiseFusable); - -// relay.nn.contrib_conv2d_winograd_weight_transform -TVM_REGISTER_NODE_TYPE(ConvWinogradWeightTransformAttrs); - -bool Conv2DWinogradWeightTransformRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - ICHECK_EQ(types.size(), 2); - const auto* data = types[0].as(); - if (data == nullptr) return false; - - const ConvWinogradWeightTransformAttrs* param = attrs.as(); - ICHECK(param != nullptr); - - ICHECK_EQ(data->shape.size(), 4) << "Only support NCHW normal kernel layout"; - - std::vector oshape{ - param->tile_size + data->shape[2] - 1, - param->tile_size + data->shape[3] - 1, - data->shape[0], - data->shape[1], - }; - - reporter->Assign(types[1], TensorType(Array(oshape), data->dtype)); - return true; -} - -TVM_REGISTER_GLOBAL("relay.op.nn._make.contrib_conv2d_winograd_weight_transform") - .set_body_typed([](Expr weight, int tile_size) { - return MakeConvWinogradWeightTransform(weight, tile_size, - "nn.contrib_conv2d_winograd_weight_transform"); - }); - -RELAY_REGISTER_OP("nn.contrib_conv2d_winograd_weight_transform") - .describe(R"code(Weight transformation of winograd fast convolution algorithm. - -Separate this into another operator in order to enable Precompute Pass to compute the -weight transformation in advance. - -- **weight**: (channels, in_channels, kernel_size[0], kernel_size[1]) -)code" TVM_ADD_FILELINE) - .set_attrs_type() - .set_num_inputs(1) - .add_argument("weight", "Tensor", "The weight tensor.") - .set_support_level(10) - .add_type_rel("Conv2DWinogradWeightTransform", Conv2DWinogradWeightTransformRel) - .set_attr("TOpPattern", kOutEWiseFusable); - -// relay.nn.contrib_conv3d_winograd_without_weight_transform -TVM_REGISTER_NODE_TYPE(Conv3DWinogradAttrs); - -bool Conv3DWinogradRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - ICHECK_EQ(types.size(), 3); - const auto* data = types[0].as(); - if (data == nullptr) return false; - static const Layout kNCDHW("NCDHW"); - static const Layout kOIDHW("OIDHW"); - - const auto* param = attrs.as(); - ICHECK(param != nullptr); - const Layout in_layout(param->data_layout); - const Layout kernel_layout(param->kernel_layout); - - const auto trans_in_layout = tir::BijectiveLayout(in_layout, kNCDHW); - ICHECK(trans_in_layout.defined()) - << "Conv only support input layouts that are convertible from NCDHW." - << " But got " << in_layout; - - const auto trans_kernel_layout = tir::BijectiveLayout(kernel_layout, kOIDHW); - ICHECK(trans_kernel_layout.defined()) - << "Conv only support kernel layouts that are convertible from OIDHW." - << " But got " << kernel_layout; - - Layout out_layout(param->out_layout == "" ? param->data_layout : param->out_layout); - const auto trans_out_layout = tir::BijectiveLayout(out_layout, kNCDHW); - ICHECK(trans_out_layout.defined()) - << "Conv only support output layouts that are convertible from NCDHW." - << " But got " << out_layout; - - Array dshape_ncdhw = trans_in_layout.ForwardShape(data->shape); - - IndexExpr channels, dilated_ksize_d, dilated_ksize_y, dilated_ksize_x; - - ICHECK(param->kernel_size.defined() && param->channels.defined()) - << "The kernel size and channels of a Conv must be set or inferred by previous pass"; - - ICHECK_EQ(param->kernel_size.size(), 3); - ICHECK_EQ(param->dilation.size(), 3); - - channels = param->channels; - dilated_ksize_d = 1 + (param->kernel_size[0] - 1) * param->dilation[0]; - dilated_ksize_y = 1 + (param->kernel_size[1] - 1) * param->dilation[1]; - dilated_ksize_x = 1 + (param->kernel_size[2] - 1) * param->dilation[2]; - - // NOTE: Do not check weight shape here! - // Different backend requires different layout to compute - // the batch gemm stage in winograd efficiently, but we want to - // make this op work for all backends. - // So we accept all weight shapes, and assume the TOPI developers - // can handle this correctly in alter_op_layout. - - // dilation - Array oshape({dshape_ncdhw[0], channels, 0, 0, 0}); - - IndexExpr pad_d, pad_h, pad_w; - GetPaddingDepthHeightWidth(param->padding, &pad_d, &pad_h, &pad_w); - if (!dshape_ncdhw[2].as()) { - oshape.Set(2, (dshape_ncdhw[2] + pad_d - dilated_ksize_d) / param->strides[0] + 1); - } else { - oshape.Set(2, dshape_ncdhw[2]); - } - if (!dshape_ncdhw[2].as()) { - oshape.Set(3, (dshape_ncdhw[3] + pad_h - dilated_ksize_y) / param->strides[1] + 1); - } else { - oshape.Set(3, dshape_ncdhw[3]); - } - if (!dshape_ncdhw[4].as()) { - oshape.Set(4, (dshape_ncdhw[4] + pad_w - dilated_ksize_x) / param->strides[2] + 1); - } else { - oshape.Set(4, dshape_ncdhw[4]); - } - - DataType out_dtype = param->out_dtype; - if (out_dtype.bits() == 0) { - out_dtype = data->dtype; - } - oshape = trans_out_layout.BackwardShape(oshape); - // assign output type - reporter->Assign(types[2], TensorType(oshape, out_dtype)); - return true; -} - -TVM_REGISTER_GLOBAL("relay.op.nn._make.contrib_conv3d_winograd_without_weight_transform") - .set_body_typed([](Expr data, Expr weight, int tile_size, Array strides, - Array padding, Array dilation, int groups, - IndexExpr channels, Array kernel_size, String data_layout, - String kernel_layout, String out_layout, DataType out_dtype) { - return MakeConvWinograd( - data, weight, tile_size, strides, padding, dilation, groups, channels, kernel_size, - data_layout, kernel_layout, out_layout, out_dtype, - "nn.contrib_conv3d_winograd_without_weight_transform"); - }); - -RELAY_REGISTER_OP("nn.contrib_conv3d_winograd_without_weight_transform") - .describe(R"code(Compute conv3d with winograd algorithm. Only supports NCDHW layout. - This operator assumes the weight tensor is already pre-transformed by - nn.contrib_conv3d_winograd_weight_transform. - -- **data**: Input is 5D array of shape (batch_size, in_channels, depth, height, width) -- **weight**: Any shape - We do not check the shape for this input tensor. Since different backend - has different layout strategy. - -- **out**: Output is 5D array of shape (batch_size, channels, depth, out_height, out_width) -)code" TVM_ADD_FILELINE) - .set_attrs_type() - .set_num_inputs(2) - .add_argument("data", "Tensor", "The input tensor.") - .add_argument("weight", "Tensor", "The weight tensor.") - .set_support_level(10) - .add_type_rel("Conv3DWinograd", Conv3DWinogradRel) - .set_attr("FInferCorrectLayout", - ConvInferCorrectLayout) - .set_attr("TOpPattern", kOutEWiseFusable); - -// relay.nn.contrib_conv3d_winograd_weight_transform -TVM_REGISTER_GLOBAL("relay.op.nn._make.contrib_conv3d_winograd_weight_transform") - .set_body_typed([](Expr weight, int tile_size) { - return MakeConvWinogradWeightTransform(weight, tile_size, - "nn.contrib_conv3d_winograd_weight_transform"); - }); - -bool Conv3DWinogradWeightTransformRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - ICHECK_EQ(types.size(), 2); - const auto* data = types[0].as(); - if (data == nullptr) return false; - - const ConvWinogradWeightTransformAttrs* param = attrs.as(); - ICHECK(param != nullptr); - - ICHECK_EQ(data->shape.size(), 5) << "Only support NCDHW normal kernel layout"; - - // Shape of packed weights depends on whether depth is being transformed or not. - Array oshape({0, 0, 0, data->shape[0], data->shape[1]}); - auto* depth_imm = data->shape[2].as(); - bool transform_depth = (depth_imm->value > 2) && (depth_imm->value < 8); - if (transform_depth) { - oshape.Set(0, param->tile_size + data->shape[2] - 1); - oshape.Set(1, param->tile_size + data->shape[3] - 1); - oshape.Set(2, param->tile_size + data->shape[4] - 1); - } else { - oshape.Set(0, param->tile_size + data->shape[3] - 1); - oshape.Set(1, param->tile_size + data->shape[4] - 1); - oshape.Set(2, data->shape[2]); - } - - reporter->Assign(types[1], TensorType(oshape, data->dtype)); - return true; -} - -RELAY_REGISTER_OP("nn.contrib_conv3d_winograd_weight_transform") - .describe(R"code(Weight transformation of winograd fast 3d convolution algorithm. - -Separate this into another operator in order to enable Precompute Pass to compute the -weight transformation in advance. - -- **weight**: (channels, in_channels, kernel_size[0], kernel_size[1], kernel_size[2]) -)code" TVM_ADD_FILELINE) - .set_attrs_type() - .set_num_inputs(1) - .add_argument("weight", "Tensor", "The weight tensor.") - .set_support_level(10) - .add_type_rel("Conv3DWinogradWeightTransform", Conv3DWinogradWeightTransformRel) - .set_attr("TOpPattern", kOutEWiseFusable); - -// relay.nn.contrib_conv2d_winograd_nnpack_weight_transform -TVM_REGISTER_NODE_TYPE(Conv2DWinogradNNPACKWeightTransformAttrs); - -bool Conv2DWinogradNNPACKWeightTransformRel(const Array& types, int num_inputs, - const Attrs& attrs, const TypeReporter& reporter) { - ICHECK_EQ(types.size(), 2); - const auto* data = types[0].as(); - if (data == nullptr) { - return false; - } - - const Conv2DWinogradNNPACKWeightTransformAttrs* param = - attrs.as(); - ICHECK(param != nullptr); - - ICHECK_EQ(data->shape.size(), 4) << "Only support NCHW normal kernel layout"; - - std::vector oshape{ - data->shape[0], - data->shape[1], - 8, - 8, - }; - - DataType out_dtype = param->out_dtype; - if (out_dtype.bits() == 0) { - out_dtype = data->dtype; - } - reporter->Assign(types[1], TensorType(Array(oshape), out_dtype)); - return true; -} - -Expr MakeConv2DWinogradNNPACKWeightTransform(Expr weight, int convolution_algorithm, - DataType out_dtype) { - auto attrs = make_object(); - attrs->convolution_algorithm = convolution_algorithm; - attrs->out_dtype = std::move(out_dtype); - static const Op& op = Op::Get("nn.contrib_conv2d_winograd_nnpack_weight_transform"); - return Call(op, {weight}, Attrs(attrs), {}); -} - -TVM_REGISTER_GLOBAL("relay.op.nn._make.contrib_conv2d_winograd_nnpack_weight_transform") - .set_body_typed(MakeConv2DWinogradNNPACKWeightTransform); - -RELAY_REGISTER_OP("nn.contrib_conv2d_winograd_nnpack_weight_transform") - .describe(R"code(Weight transformation of winograd fast convolution algorithm with NNPACK. -Separate this into another symbol in order to enable Precompute Pass to compute the -weight transformation in advance. - -- **weight**: (channels, in_channels, kernel_size[0], kernel_size[1]) - -)code" TVM_ADD_FILELINE) - .set_attrs_type() - .set_num_inputs(1) - .add_argument("weight", "Tensor", "The weight tensor.") - .set_support_level(10) - .add_type_rel("Conv2DWinogradNNPACKWeightTransform", Conv2DWinogradNNPACKWeightTransformRel) - .set_attr("TOpPattern", kOpaque); - -// relay.nn.contrib_conv2d_gemm_without_weight_transform -TVM_REGISTER_GLOBAL("relay.op.nn._make.contrib_conv2d_gemm_without_weight_transform") - .set_body_typed([](Expr data, Expr weight, Array strides, Array padding, - Array dilation, int groups, IndexExpr channels, - Array kernel_size, tvm::String data_layout, - tvm::String kernel_layout, tvm::String out_layout, DataType out_dtype) { - return MakeConvGemm( - data, weight, strides, padding, dilation, groups, channels, kernel_size, data_layout, - kernel_layout, out_layout, out_dtype, "nn.contrib_conv2d_gemm_without_weight_transform"); - }); - -bool Conv2DGemmRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - ICHECK_EQ(types.size(), 3); - const auto* data = types[0].as(); - if (data == nullptr) return false; - static const Layout kNHWC("NHWC"); - static const Layout kHWIO("HWIO"); - - const auto* param = attrs.as(); - ICHECK(param != nullptr); - const Layout in_layout(param->data_layout); - const Layout kernel_layout(param->kernel_layout); - - const auto trans_in_layout = tir::BijectiveLayout(in_layout, kNHWC); - ICHECK(trans_in_layout.defined()) - << "Conv only support input layouts that are convertible from NHWC." - << " But got " << in_layout; - - const auto trans_kernel_layout = tir::BijectiveLayout(kernel_layout, kHWIO); - ICHECK(trans_kernel_layout.defined()) - << "Conv only support kernel layouts that are convertible from HWIO." - << " But got " << kernel_layout; - - Layout out_layout(param->out_layout == "" ? param->data_layout : param->out_layout); - const auto trans_out_layout = tir::BijectiveLayout(out_layout, kNHWC); - ICHECK(trans_out_layout.defined()) - << "Conv only support output layouts that are convertible from NHWC." - << " But got " << out_layout; - - Array dshape_nhwc = trans_in_layout.ForwardShape(data->shape); - - IndexExpr channels, dilated_ksize_y, dilated_ksize_x; - - ICHECK(param->kernel_size.defined() && param->channels.defined()) - << "The kernel size and channels of a Conv must be set or inferred by previous pass"; - - ICHECK_EQ(param->kernel_size.size(), 2); - ICHECK_EQ(param->dilation.size(), 2); - - channels = param->channels; - dilated_ksize_y = 1 + (param->kernel_size[0] - 1) * param->dilation[0]; - dilated_ksize_x = 1 + (param->kernel_size[1] - 1) * param->dilation[1]; - - // NOTE: Do not check weight shape here! - - // dilation - Array oshape({dshape_nhwc[0], 0, 0, channels}); - - IndexExpr pad_h, pad_w; - GetPaddingHeightWidth(param->padding, &pad_h, &pad_w); - if (!dshape_nhwc[2].as()) { - oshape.Set(1, (dshape_nhwc[1] + pad_h - dilated_ksize_y) / param->strides[0] + 1); - } else { - oshape.Set(1, dshape_nhwc[1]); - } - if (!dshape_nhwc[3].as()) { - oshape.Set(2, (dshape_nhwc[2] + pad_w - dilated_ksize_x) / param->strides[1] + 1); - } else { - oshape.Set(2, dshape_nhwc[2]); - } - - DataType out_dtype = param->out_dtype; - if (out_dtype.bits() == 0) { - out_dtype = data->dtype; - } - oshape = trans_out_layout.BackwardShape(oshape); - // assign output type - reporter->Assign(types[2], TensorType(oshape, out_dtype)); - return true; -} - -RELAY_REGISTER_OP("nn.contrib_conv2d_gemm_without_weight_transform") - .describe(R"code(Compute conv2d with gemm algorithm. Only supports NHWC layout. - This operator assumes the weight tensor is already pre-transformed by - nn.contrib_conv2d_gemm_weight_transform. - -- **data**: Input is 4D array of shape (batch_size, height, width, in_channels) -- **weight**: Any shape - We do not check the shape for this input tensor. Since different backend - has different layout strategy. - -- **out**: Output is 4D array of shape (batch_size, channels, out_height, out_width) -)code" TVM_ADD_FILELINE) - .set_attrs_type() - .set_num_inputs(2) - .add_argument("data", "Tensor", "The input tensor.") - .add_argument("weight", "Tensor", "The weight tensor.") - .set_support_level(10) - .add_type_rel("Conv2DGemm", Conv2DGemmRel) - .set_attr("FInferCorrectLayout", ConvInferCorrectLayout) - .set_attr("TOpPattern", kOutEWiseFusable); - -// relay.nn.contrib_conv2d_gemm_weight_transform - -TVM_REGISTER_NODE_TYPE(ConvGemmWeightTransformAttrs); - -// Gemm convolution shape relations -// In order to run GEMM we need to transform the K x N weights matrix W. -// -// For integer datatypes, the high level idea is to subdivide W in tiles of tile_K x tile_N, and -// transpose and interleave them. The final output is a [N//tile_N, K//tile_K, tile_N, tile_K] -// matrix that we call W_interleaved_t. -// -// In the following picture, we show how the first [tile_K,tile_N] block of W is transformed -// for tile_N = 4 and tile_K = 16 -// -// W[0,0,:,:] W_interleaved_t[0,0,:,:] -// +-------------------------------+ +----------------------------------- + -// |W[0,0] W[0,1] W[0,2] W[0,3] | |W[0,0] W[1,0] W[2,0] ... W[15,0]| -// |W[1,0] W[1,1] W[1,2] W[1,3] | --\ |W[0,1] W[1,1] W[2,1] ... W[15,1]| -// |W[2,0] W[2,1] W[2,2] W[2,3] | --/ |W[0,2] W[1,2] W[2,2] ... W[15,2]| -// | ... ... ... ... | |W[0,3] W[1,3] W[2,3] ... W[15,3]| -// | ... ... ... ... | +------------------------------------+ -// |W[15,0] W[15,1] W[15,2] W[15,3]| -// +-------------------------------+ -// -// Alternatively, for floating point datatypes, we subdivide W in tiles of tile_K x tile_N size, -// then interleave these tiles, without transposing. The final output is a [N//tile_N, K//tile_K, -// tile_K, tile_N] matrix called W_interleaved. -// -// In the following illustration, we show how the tiles are interleaved. -// Note that the inside of each tile is kept unchanged during this tranformation. -// -// W[:,:,:,:] W_interleaved[:,:,:,:] -// +--------+--------+--------+ +--------+--------+ -// | | | | | | | -// | tile_1 | tile_2 | tile_3 | | tile_1 | tile_4 | -// | | | | --\ | | | -// +--------+--------+--------+ --/ +--------+--------+ -// | | | | | | | -// | tile_4 | tile_5 | tile_6 | | tile_2 | tile_5 | -// | | | | | | | -// +--------+--------+--------+ +--------+--------+ -// | | | -// | tile_3 | tile_6 | -// | | | -// +--------+--------+ -// -// Tile K is the direction of the reduction in both cases. So, if our target can reduce k elements -// at the time, we should set tile_K = k. -// Tile N is connected with the number of registers available for the given target. -// -bool Conv2DGemmWeightTransformRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - ICHECK_EQ(types.size(), 2); - const auto* weight = types[0].as(); - if (weight == nullptr) return false; - - const ConvGemmWeightTransformAttrs* param = attrs.as(); - ICHECK(param != nullptr); - int n = param->tile_N; - int k = param->tile_K; - - ICHECK_EQ(weight->shape.size(), 4) << "Only support HWIO kernel layout"; - - const auto K = weight->shape[0] * weight->shape[1] * weight->shape[2]; - const auto N = weight->shape[3]; - - auto K_mod_k = indexmod(K, k * 4); - auto N_mod_n = indexmod(N, n); - - auto pad_K = tvm::if_then_else(K_mod_k != 0, k * 4 - K_mod_k, tir::make_zero(DataType::Int(32))); - auto pad_N = tvm::if_then_else(N_mod_n != 0, n - N_mod_n, tir::make_zero(DataType::Int(32))); - - const auto N_padded = N + pad_N; - const auto K_padded = K + pad_K; - - Array oshape; - if (weight->dtype.bits() == 8 && (weight->dtype.is_int() || weight->dtype.is_uint())) - oshape = { - indexdiv(N_padded, n), - indexdiv(K_padded, k), - n, - k, - }; - else - oshape = { - indexdiv(N_padded, n), - indexdiv(K_padded, k), - k, - n, - }; - - reporter->Assign(types[1], TensorType(oshape, weight->dtype)); - return true; -} - -TVM_REGISTER_GLOBAL("relay.op.nn._make.contrib_conv2d_gemm_weight_transform") - .set_body_typed([](Expr weights, int tile_rows, int tile_cols) { - return MakeConvGemmWeightTransform(weights, tile_rows, tile_cols, - "nn.contrib_conv2d_gemm_weight_transform"); - }); - -RELAY_REGISTER_OP("nn.contrib_conv2d_gemm_weight_transform") - .describe(R"code(Weight transformation of GEMM convolution algorithm. - -Separate this into another operator in order to enable Precompute Pass to compute the -weight transformation in advance. - -)code" TVM_ADD_FILELINE) - .set_attrs_type() - .set_num_inputs(1) - .add_argument("weights", "Tensor", "The weights tensor.") - .set_support_level(10) - .add_type_rel("Conv2DGemmWeightTransform", Conv2DGemmWeightTransformRel) - .set_attr("TOpPattern", kOutEWiseFusable); - -// Positional relay function to create conv2d NCHWc operator -// used by frontend FFI. -TVM_REGISTER_GLOBAL("relay.op.nn._make.contrib_conv2d_NCHWc") - .set_body_typed([](Expr data, Expr weight, Array strides, Array padding, - Array dilation, int groups, IndexExpr channels, - Array kernel_size, String data_layout, String kernel_layout, - String out_layout, DataType out_dtype) { - return MakeConv(data, weight, strides, padding, dilation, groups, channels, - kernel_size, data_layout, kernel_layout, out_layout, out_dtype, - "nn.contrib_conv2d_NCHWc"); - }); - -RELAY_REGISTER_OP("nn.contrib_conv2d_NCHWc") - .describe(R"code(Compute conv2d with NCHWc data layout. Only supports NCHW layout. -- **data**: Input is 5D packed tensor. -- **weight**: 6D packed tensor. - -- **out**: Output is 5D packed tensor -)code" TVM_ADD_FILELINE) - .set_attrs_type() - .set_num_inputs(2) - .add_argument("data", "Tensor", "The input tensor.") - .add_argument("weight", "Tensor", "The weight tensor.") - .set_support_level(10) - .add_type_rel("Conv2DNCHWc", Conv2DWinogradRel) - .set_attr("FInferCorrectLayout", ConvInferCorrectLayout) - .set_attr("TOpPattern", kOutEWiseFusable); - -// Positional relay function to create depthwise conv2d NCHWc operator -// used by frontend FFI. -TVM_REGISTER_GLOBAL("relay.op.nn._make.contrib_depthwise_conv2d_NCHWc") - .set_body_typed([](Expr data, Expr weight, Array strides, Array padding, - Array dilation, int groups, IndexExpr channels, - Array kernel_size, String data_layout, String kernel_layout, - String out_layout, DataType out_dtype) { - return MakeConv(data, weight, strides, padding, dilation, groups, channels, - kernel_size, data_layout, kernel_layout, out_layout, out_dtype, - "nn.contrib_depthwise_conv2d_NCHWc"); - }); - -RELAY_REGISTER_OP("nn.contrib_depthwise_conv2d_NCHWc") - .describe(R"code(Compute conv2d with NCHWc data layout. Only supports NCHW layout. -- **data**: Input is 5D packed tensor. -- **weight**: 6D packed tensor. - -- **out**: Output is 5D packed tensor -)code" TVM_ADD_FILELINE) - .set_attrs_type() - .set_num_inputs(2) - .add_argument("data", "Tensor", "The input tensor.") - .add_argument("weight", "Tensor", "The weight tensor.") - .set_support_level(10) - .add_type_rel("Conv2D", Conv2DRel) - .set_attr("FInferCorrectLayout", ConvInferCorrectLayout) - .set_attr("TOpPattern", kOutEWiseFusable); - -TVM_REGISTER_NODE_TYPE(DeformableConv2DAttrs); - -// Deformable Convolution shape relations. -bool DeformableConv2DRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - ICHECK_EQ(types.size(), 4); - const auto* data = types[0].as(); - const auto* weight = types[2].as(); - - ICHECK(data); - static const Layout kNCHW("NCHW"); - static const Layout kOIHW("OIHW"); - - auto* param = attrs.as(); - ICHECK(param != nullptr); - const Layout in_layout(param->data_layout); - const Layout kernel_layout(param->kernel_layout); - - const auto trans_in_layout = tir::BijectiveLayout(in_layout, kNCHW); - if (!trans_in_layout.defined()) { - reporter->GetDiagCtx().Emit( - Diagnostic::Error(reporter->GetSpan()) - << "deformable_conv2d only support input layouts that are convertible from NCHW." - << " The provided layout is: " << in_layout); - return false; - } - - const auto trans_kernel_layout = tir::BijectiveLayout(kernel_layout, kOIHW); - if (!trans_kernel_layout.defined()) { - reporter->GetDiagCtx().Emit( - Diagnostic::Error(reporter->GetSpan()) - << "deformable_conv2d only support kernel layouts that are convertible from OIHW." - << " The provided layout is: " << kernel_layout); - return false; - } - - Layout out_layout(param->out_layout == "" ? param->data_layout : param->out_layout); - const auto trans_out_layout = tir::BijectiveLayout(out_layout, kNCHW); - if (!trans_out_layout.defined()) { - reporter->GetDiagCtx().Emit( - Diagnostic::Error(reporter->GetSpan()) - << "deformable_conv2d only support output layouts that are convertible from NCHW." - << "The provided layout is: " << out_layout); - return false; - } - - Array dshape_nchw = trans_in_layout.ForwardShape(data->shape); - - IndexExpr channels, dilated_ksize_y, dilated_ksize_x, ksize_y, ksize_x; - - // infer weight shape if kernel_size and channels are defiend - if (param->kernel_size.defined() && param->channels.defined()) { - ICHECK_EQ(param->kernel_size.size(), 2); - ICHECK_EQ(param->dilation.size(), 2); - Array wshape({param->channels, indexdiv(dshape_nchw[1], param->groups), - param->kernel_size[0], param->kernel_size[1]}); - - wshape = trans_kernel_layout.BackwardShape(wshape); - channels = param->channels; - ksize_y = param->kernel_size[0]; - ksize_x = param->kernel_size[1]; - dilated_ksize_y = 1 + (param->kernel_size[0] - 1) * param->dilation[0]; - dilated_ksize_x = 1 + (param->kernel_size[1] - 1) * param->dilation[1]; - // assign result to reporter - reporter->Assign(types[2], TensorType(wshape, data->dtype)); - } else { - // use weight to infer the conv shape. - if (weight == nullptr) return false; - auto wshape = trans_kernel_layout.ForwardShape(weight->shape); - - if (param->kernel_size.defined()) { - ICHECK_EQ(param->kernel_size.size(), 2); - // check the size - ICHECK(reporter->AssertEQ(param->kernel_size[0], wshape[2]) && - reporter->AssertEQ(param->kernel_size[1], wshape[3])) - << "DeformableConv2D: shape of weight is inconsistent with kernel_size, " - << " kernel_size=" << param->kernel_size << " wshape=" << wshape; - } - if (param->channels.defined()) { - ICHECK(reporter->AssertEQ(param->channels, wshape[0])) - << "DeformableConv2D: shape of weight is inconsistent with channels, " - << " channels=" << param->channels << " wshape=" << wshape; - } - if (!dshape_nchw[1].as() && !wshape[1].as()) { - ICHECK(reporter->AssertEQ(indexdiv(dshape_nchw[1], param->groups), wshape[1])); - } - channels = wshape[0]; - ksize_y = wshape[2]; - ksize_x = wshape[3]; - dilated_ksize_y = 1 + (wshape[2] - 1) * param->dilation[0]; - dilated_ksize_x = 1 + (wshape[3] - 1) * param->dilation[1]; - } - // dilation - Array oshape({dshape_nchw[0], channels, 0, 0}); - - IndexExpr pad_h, pad_w; - GetPaddingHeightWidth(param->padding, &pad_h, &pad_w); - oshape.Set(2, indexdiv(dshape_nchw[2] + pad_h - dilated_ksize_y, param->strides[0]) + 1); - oshape.Set(3, indexdiv(dshape_nchw[3] + pad_w - dilated_ksize_x, param->strides[1]) + 1); - DataType out_dtype = param->out_dtype; - - // infer offset shape - Array offset_shape( - {dshape_nchw[0], 2 * ksize_y * ksize_x * param->deformable_groups, oshape[2], oshape[3]}); - offset_shape = trans_in_layout.BackwardShape(offset_shape); - reporter->Assign(types[1], TensorType(offset_shape, data->dtype)); - if (out_dtype.bits() == 0) { - out_dtype = data->dtype; - } - - oshape = trans_out_layout.BackwardShape(oshape); - reporter->Assign(types[3], TensorType(oshape, out_dtype)); - return true; -} - -InferCorrectLayoutOutput DeformableConvInferCorrectLayout( - const Attrs& attrs, const Array& new_in_layouts, const Array& old_in_layouts, - const Array& old_in_types) { - const auto* params = attrs.as(); - return InferCorrectLayoutOutput( - {params->data_layout, params->data_layout, params->kernel_layout}, - {params->out_layout == "" ? params->data_layout : params->out_layout}, attrs); -} - -RELAY_REGISTER_OP("nn.deformable_conv2d") - .describe(R"code(Compute 2-D deformable convolution on 4-D input. -The deformable convolution operation is described in https://arxiv.org/abs/1703.06211 - -For 2-D deformable convolution, the shapes are -- **data**: (batch_size, channel, height, width) -- **offset**: (batch_size, deformable_groups * kernel[0] * kernel[1] * 2, out_height, out_width) -- **weight**: (num_filter, channel, kernel[0], kernel[1]) -- **out**: (batch_size, num_filter, out_height, out_width). - -If `deformable_groups` is larger than 1, denoted by *dg*, then split the -input `offset` evenly into *dg* parts along the channel axis, and also evenly split `out` -evenly into *dg* parts along the channel axis. Next compute the deformable convolution, apply the -*i*-th part of the offset part on the *i*-th out. - -If `groups` is larger than 1, denoted by *g*, then split the input `data` evenly into *g* parts -along the channel axis, and also evenly split `weight` along the first dimension. Next compute -the convolution on the *i*-th part of the data with the *i*-th weight part. The output is obtained -by concating all the *g* results. -)code" TVM_ADD_FILELINE) - .set_attrs_type() - .set_num_inputs(3) - .add_argument("data", "Tensor", "The input tensor.") - .add_argument("offset", "Tensor", "The offset tensor.") - .add_argument("weight", "Tensor", "The weight tensor.") - .set_support_level(5) - .add_type_rel("DeformableConv2D", DeformableConv2DRel) - .set_attr("FInferCorrectLayout", DeformableConvInferCorrectLayout) - .set_attr("TOpPattern", kOutEWiseFusable); - -// Positional relay function to create deformable_conv2d operator -// used by frontend FFI. -TVM_REGISTER_GLOBAL("relay.op.nn._make.deformable_conv2d") - .set_body_typed([](Expr data, Expr offset, Expr weight, Array strides, - Array padding, Array dilation, int deformable_groups, - int groups, int channels, Array kernel_size, String data_layout, - String kernel_layout, String out_layout, DataType out_dtype) { - return MakeDeformableConv( - data, offset, weight, strides, padding, dilation, deformable_groups, groups, channels, - kernel_size, data_layout, kernel_layout, out_layout, out_dtype, "nn.deformable_conv2d"); - }); - -inline Expr MakeConv2dBackwardWeight(Expr grad, Expr data, Array strides, - Array padding, Array dilation, - int groups, IndexExpr channels, Array kernel_size, - std::string grad_layout, std::string data_layout, - std::string kernel_layout, DataType out_dtype) { - auto attrs = make_object(); - attrs->strides = std::move(strides); - attrs->padding = std::move(padding); - attrs->dilation = std::move(dilation); - attrs->groups = groups; - attrs->channels = std::move(channels); - attrs->kernel_size = std::move(kernel_size); - attrs->out_dtype = std::move(out_dtype); - attrs->data_layout = std::move(grad_layout); - attrs->kernel_layout = std::move(data_layout); - attrs->out_layout = std::move(kernel_layout); - const Op& op = Op::Get("nn.conv2d_backward_weight"); - return Call(op, {grad, data}, Attrs(attrs), {}); -} - -TVM_REGISTER_GLOBAL("relay.op.nn._make.conv2d_backward_weight") - .set_body_typed([](Expr grad, Expr data, Array strides, Array padding, - Array dilation, int groups, IndexExpr channels, - Array kernel_size, String grad_layout, String data_layout, - String kernel_layout, DataType out_dtype) { - return MakeConv2dBackwardWeight(grad, data, strides, padding, dilation, groups, channels, - kernel_size, grad_layout, data_layout, kernel_layout, - out_dtype); - }); - -bool Conv2DBackwardWeightRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - ICHECK_EQ(types.size(), 3); - const auto* grad = types[0].as(); - const auto* data = types[1].as(); - if (data == nullptr) return false; - - static const Layout kNCHW("NCHW"); - static const Layout kOIHW("OIHW"); - - const auto* param = attrs.as(); - ICHECK(param != nullptr); - // Require kernel_size to be passed, to simplify the output shape determination. - ICHECK(param->kernel_size.defined()) << "kernel_size attribute needs to be specified"; - - // We repurpose Conv2dAttrs for Conv2DBackwardWeight, note the meanings of layouts. - const Layout grad_layout(param->data_layout); - const Layout in_layout(param->kernel_layout); - const Layout kernel_layout(param->out_layout); - - const auto trans_grad_layout = tir::BijectiveLayout(grad_layout, kNCHW); - const auto trans_in_layout = tir::BijectiveLayout(in_layout, kNCHW); - const auto trans_kernel_layout = tir::BijectiveLayout(kernel_layout, kOIHW); - - Array dshape_nchw = trans_in_layout.ForwardShape(data->shape); - Array grad_shape_nchw = trans_grad_layout.ForwardShape(grad->shape); - - auto in_channels = dshape_nchw[1]; - auto out_channels = grad_shape_nchw[1]; - - auto in_channels_intimm = in_channels.as(); - auto out_channels_intimm = out_channels.as(); - ICHECK(in_channels_intimm); - ICHECK(out_channels_intimm); - - IndexExpr weight_dim_i; - if (in_channels_intimm->value == out_channels_intimm->value && - in_channels_intimm->value == param->groups) { - // depthwise - ICHECK(param->channels.defined()) - << "out_channels attribute not specified for depth wise conv2d."; - weight_dim_i = indexdiv(param->channels, param->groups); - } else { - weight_dim_i = indexdiv(in_channels, param->groups); - } - - Array wshape_oihw{out_channels, weight_dim_i, param->kernel_size[0], - param->kernel_size[1]}; - auto wshape = trans_kernel_layout.BackwardShape(wshape_oihw); - - const auto dw_dtype = (param->out_dtype == DataType() || param->out_dtype.is_void()) - ? grad->dtype - : param->out_dtype; - - reporter->Assign(types[2], TensorType(wshape, dw_dtype)); - return true; -} - -RELAY_REGISTER_OP("nn.conv2d_backward_weight") - .describe(R"code(The gradient of the 2D convolution layer with respect to the weight. - -This layer computes the gradient of the conv2d op with respect to weight, -given the original input data and the output gradient. - -- **grad**: (batch, channels, out_height, out_width) if `layout` is `NCHW`. -- **data**: This depends on the `layout` parameter. Input is 4D array of shape - (batch_size, in_channels, height, width) if `layout` is `NCHW`. -- **out**: This depends on the `layout` parameter. Output is 4D array of shape - (channels, in_channels, kernel_size[0], kernel_size[1]) if `layout` is `NCHW`. -)code" TVM_ADD_FILELINE) - .set_attrs_type() - .set_num_inputs(2) - .add_argument("grad", "Tensor", "The gradient tensor.") - .add_argument("data", "Tensor", "The input tensor.") - .set_support_level(2) - .add_type_rel("Conv2DBackwardWeight", Conv2DBackwardWeightRel) - .set_attr("FInferCorrectLayout", ConvInferCorrectLayout) - .set_attr("TOpPattern", kOutEWiseFusable); - -} // namespace relay -} // namespace tvm diff --git a/src/relay/op/nn/convolution.h b/src/relay/op/nn/convolution.h deleted file mode 100644 index 62552ee4783e..000000000000 --- a/src/relay/op/nn/convolution.h +++ /dev/null @@ -1,138 +0,0 @@ -/* - * 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 src/relay/op/nn/convolution.h - * \brief Properties def of convlution operator for sharing. - */ -#ifndef TVM_RELAY_OP_NN_CONVOLUTION_H_ -#define TVM_RELAY_OP_NN_CONVOLUTION_H_ - -#include -#include -#include - -#include -#include -#include - -#include "../op_common.h" - -namespace tvm { -namespace relay { - -bool Conv2DRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter); - -bool Conv2DTransposeRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter); - -template -bool Conv2DWinogradRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - ICHECK_EQ(types.size(), 3); - const auto* data = types[0].as(); - if (data == nullptr) return false; - static const Layout kNCHW("NCHW"); - static const Layout kOIHW("OIHW"); - - const AttrType* param = attrs.as(); - ICHECK(param != nullptr); - const Layout in_layout(param->data_layout); - const Layout kernel_layout(param->kernel_layout); - - const auto trans_in_layout = tir::BijectiveLayout(in_layout, kNCHW); - ICHECK(trans_in_layout.defined()) - << "Conv only support input layouts that are convertible from NCHW." - << " But got " << in_layout; - - const auto trans_kernel_layout = tir::BijectiveLayout(kernel_layout, kOIHW); - ICHECK(trans_kernel_layout.defined()) - << "Conv only support kernel layouts that are convertible from OIHW." - << " But got " << kernel_layout; - - Layout out_layout(param->out_layout == "" ? param->data_layout : param->out_layout); - const auto trans_out_layout = tir::BijectiveLayout(out_layout, kNCHW); - ICHECK(trans_out_layout.defined()) - << "Conv only support output layouts that are convertible from NCHW." - << " But got " << out_layout; - - Array dshape_nchw = trans_in_layout.ForwardShape(data->shape); - - IndexExpr channels, dilated_ksize_y, dilated_ksize_x; - - ICHECK(param->kernel_size.defined() && param->channels.defined()) - << "The kernel size and channels of a Conv must be set or inferred by previous pass"; - - ICHECK_EQ(param->kernel_size.size(), 2); - ICHECK_EQ(param->dilation.size(), 2); - - channels = param->channels; - dilated_ksize_y = 1 + (param->kernel_size[0] - 1) * param->dilation[0]; - dilated_ksize_x = 1 + (param->kernel_size[1] - 1) * param->dilation[1]; - - // NOTE: Do not check weight shape here! - // Different backend requires different layout to compute - // the batch gemm stage in winograd efficiently, but we want to - // make this op work for all backends. - // So we accept all weight shapes, and assume the TOPI developers - // can handle this correctly in alter_op_layout. - - // dilation - Array oshape({dshape_nchw[0], channels, 0, 0}); - - IndexExpr pad_h, pad_w; - GetPaddingHeightWidth(param->padding, &pad_h, &pad_w); - if (!dshape_nchw[2].as()) { - oshape.Set(2, (dshape_nchw[2] + pad_h - dilated_ksize_y) / param->strides[0] + 1); - } else { - oshape.Set(2, dshape_nchw[2]); - } - if (!dshape_nchw[3].as()) { - oshape.Set(3, (dshape_nchw[3] + pad_w - dilated_ksize_x) / param->strides[1] + 1); - } else { - oshape.Set(3, dshape_nchw[3]); - } - - DataType out_dtype = param->out_dtype; - if (out_dtype.bits() == 0) { - out_dtype = data->dtype; - } - oshape = trans_out_layout.BackwardShape(oshape); - // assign output type - reporter->Assign(types[2], TensorType(oshape, out_dtype)); - return true; -} - -template -InferCorrectLayoutOutput ConvInferCorrectLayout(const Attrs& attrs, - const Array& new_in_layouts, - const Array& old_in_layouts, - const Array& old_in_types) { - const T* params = attrs.as(); - // We always make other operators to fit the layouts of convolution layers - // So this inference ignores all inputs - return InferCorrectLayoutOutput( - {params->data_layout, params->kernel_layout}, - {params->out_layout == "" ? params->data_layout : params->out_layout}, attrs); -} - -} // namespace relay -} // namespace tvm -#endif // TVM_RELAY_OP_NN_CONVOLUTION_H_ diff --git a/src/relay/op/nn/convolution_make.h b/src/relay/op/nn/convolution_make.h deleted file mode 100644 index d343940b9ca7..000000000000 --- a/src/relay/op/nn/convolution_make.h +++ /dev/null @@ -1,149 +0,0 @@ -/* - * 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 src/relay/op/nn/convolution_make.h - * \brief utilities for creating convolution ops - */ -#ifndef TVM_RELAY_OP_NN_CONVOLUTION_MAKE_H_ -#define TVM_RELAY_OP_NN_CONVOLUTION_MAKE_H_ - -#include -#include - -#include -#include -#include - -namespace tvm { -namespace relay { - -template -inline Expr MakeConv(Expr data, Expr weight, Array strides, Array padding, - Array dilation, int groups, IndexExpr channels, - Array kernel_size, std::string data_layout, - std::string kernel_layout, std::string out_layout, DataType out_dtype, - std::string op_name) { - auto attrs = make_object(); - attrs->strides = std::move(strides); - attrs->padding = std::move(padding); - attrs->dilation = std::move(dilation); - attrs->groups = groups; - attrs->channels = std::move(channels); - attrs->kernel_size = std::move(kernel_size); - attrs->data_layout = std::move(data_layout); - attrs->kernel_layout = std::move(kernel_layout); - attrs->out_layout = std::move(out_layout); - attrs->out_dtype = std::move(out_dtype); - const Op& op = Op::Get(op_name); - return Call(op, {data, weight}, Attrs(attrs), {}); -} - -template -inline Expr MakeConvWinograd(Expr data, Expr weight, int tile_size, Array strides, - Array padding, Array dilation, int groups, - IndexExpr channels, Array kernel_size, - std::string data_layout, std::string kernel_layout, - std::string out_layout, DataType out_dtype, std::string op_name) { - auto attrs = make_object(); - attrs->tile_size = tile_size; - attrs->strides = std::move(strides); - attrs->padding = std::move(padding); - attrs->dilation = std::move(dilation); - attrs->groups = groups; - attrs->channels = std::move(channels); - attrs->kernel_size = std::move(kernel_size); - attrs->data_layout = std::move(data_layout); - attrs->kernel_layout = std::move(kernel_layout); - attrs->out_layout = std::move(out_layout); - attrs->out_dtype = std::move(out_dtype); - const Op& op = Op::Get(op_name); - return Call(op, {data, weight}, Attrs(attrs), {}); -} - -template -inline Expr MakeConvGemm(Expr data, Expr weight, Array strides, Array padding, - Array dilation, int groups, IndexExpr channels, - Array kernel_size, std::string data_layout, - std::string kernel_layout, std::string out_layout, DataType out_dtype, - std::string op_name) { - auto attrs = make_object(); - attrs->strides = std::move(strides); - attrs->padding = std::move(padding); - attrs->dilation = std::move(dilation); - attrs->groups = groups; - attrs->channels = std::move(channels); - attrs->kernel_size = std::move(kernel_size); - attrs->data_layout = std::move(data_layout); - attrs->kernel_layout = std::move(kernel_layout); - attrs->out_layout = std::move(out_layout); - attrs->out_dtype = std::move(out_dtype); - const Op& op = Op::Get(op_name); - return Call(op, {data, weight}, Attrs(attrs), {}); -} - -template -inline Expr MakeConvTranspose(Expr data, Expr weight, Array strides, - Array padding, Array dilation, int groups, - IndexExpr channels, Array kernel_size, - std::string data_layout, std::string kernel_layout, - std::string out_layout, Array output_padding, - DataType out_dtype, std::string op_name) { - auto attrs = make_object(); - attrs->strides = std::move(strides); - attrs->padding = std::move(padding); - attrs->dilation = std::move(dilation); - attrs->groups = groups; - attrs->channels = std::move(channels); - attrs->kernel_size = std::move(kernel_size); - attrs->data_layout = std::move(data_layout); - attrs->kernel_layout = std::move(kernel_layout); - attrs->out_layout = std::move(out_layout); - attrs->output_padding = std::move(output_padding); - attrs->out_dtype = std::move(out_dtype); - const Op& op = Op::Get(op_name); - return Call(op, {data, weight}, Attrs(attrs), {}); -} - -template -inline Expr MakeDeformableConv(Expr data, Expr offset, Expr weight, Array strides, - Array padding, Array dilation, - int deformable_groups, int groups, int channels, - Array kernel_size, std::string data_layout, - std::string kernel_layout, std::string out_layout, - DataType out_dtype, std::string op_name) { - auto attrs = make_object(); - attrs->strides = strides; - attrs->padding = padding; - attrs->dilation = dilation; - attrs->deformable_groups = deformable_groups; - attrs->groups = groups; - attrs->channels = channels; - attrs->kernel_size = kernel_size; - attrs->data_layout = data_layout; - attrs->kernel_layout = kernel_layout; - attrs->out_layout = out_layout; - attrs->out_dtype = out_dtype; - const Op& op = Op::Get(op_name); - return Call(op, {data, offset, weight}, Attrs{attrs}, {}); -} - -} // namespace relay -} // namespace tvm -#endif // TVM_RELAY_OP_NN_CONVOLUTION_MAKE_H_ diff --git a/src/relay/op/nn/correlation.cc b/src/relay/op/nn/correlation.cc deleted file mode 100644 index 8abc9909e83c..000000000000 --- a/src/relay/op/nn/correlation.cc +++ /dev/null @@ -1,136 +0,0 @@ -/* - * 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 correlation.cc - * \brief Correlation operators - */ -#include -#include -#include -#include -#include - -#include - -#include "../op_common.h" - -namespace tvm { -namespace relay { - -// relay.nn.correlation -TVM_REGISTER_NODE_TYPE(CorrelationAttrs); - -InferCorrectLayoutOutput CorrelationInferCorrectLayout( - const Attrs& attrs, const Array& new_in_layouts, const Array& old_in_layouts, - const Array& old_in_types) { - const auto* params = attrs.as(); - Layout layout{params->layout}; - return InferCorrectLayoutOutput({layout, layout}, {layout}, attrs); -} - -// Positional relay function to create correlation operator -// used by frontend FFI. -Expr MakeCorrelation(Expr data1, Expr data2, int kernel_size, int max_displacement, int stride1, - int stride2, Array padding, bool is_multiply, String layout) { - auto attrs = make_object(); - attrs->kernel_size = kernel_size; - attrs->max_displacement = max_displacement; - attrs->stride1 = stride1; - attrs->stride2 = stride2; - attrs->padding = std::move(padding); - attrs->is_multiply = is_multiply; - attrs->layout = std::move(layout); - static const Op& op = Op::Get("nn.correlation"); - return Call(op, {data1, data2}, Attrs(attrs), {}); -} - -bool CorrelationRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - ICHECK_EQ(types.size(), 3); - const auto* data1 = types[0].as(); - const auto* data2 = types[1].as(); - if (data1 == nullptr || data2 == nullptr) return false; - - const CorrelationAttrs* param = attrs.as(); - ICHECK(param != nullptr); - ICHECK_EQ(param->layout, "NCHW") << "layout not supported."; - IndexExpr pad_h, pad_w; - GetPaddingHeightWidth(param->padding, &pad_h, &pad_w); - IndexExpr padded_height = data1->shape[2] + pad_h; - IndexExpr padded_width = data2->shape[3] + pad_w; - int kernel_radius = (param->kernel_size - 1) / 2; - int border_size = param->max_displacement + kernel_radius; - int displacement_radius = param->max_displacement / param->stride2; - int displacement_size = 2 * displacement_radius + 1; - int out_channel = displacement_size * displacement_size; - IndexExpr out_height = - indexdiv((padded_height - 2 * border_size + param->stride1 - 1), param->stride1); - IndexExpr out_width = - indexdiv((padded_width - 2 * border_size + param->stride1 - 1), param->stride1); - Array oshape{data1->shape[0], out_channel, out_height, out_width}; - // assign output type - reporter->Assign(types[2], TensorType(oshape, data1->dtype)); - return true; -} - -TVM_REGISTER_GLOBAL("relay.op.nn._make.correlation").set_body_typed(MakeCorrelation); - -RELAY_REGISTER_OP("nn.correlation") - .describe(R"code(Applies correlation to inputs. - -The correlation layer performs multiplicative patch comparisons between two feature maps. -Given two multi-channel feature maps :math:`f_{1}, f_{2}`, with :math:`w`, :math:`h`, and :math:`c` being their width, height, and number of channels, -the correlation layer lets the network compare each patch from :math:`f_{1}` with each patch from :math:`f_{2}`. - -For now we consider only a single comparison of two patches. The 'correlation' of two patches centered at :math:`x_{1}` in the first map and -:math:`x_{2}` in the second map is then defined as: - -.. math:: - c(x_{1}, x_{2}) = \sum_{o \in [-k,k] \times [-k,k]} - -for a square patch of size :math:`K:=2k+1`. - -Note that the equation above is identical to one step of a convolution in neural networks, but instead of convolving data with a filter, it convolves data with other -data. For this reason, it has no training weights. - -Computing :math:`c(x_{1}, x_{2})` involves :math:`c * K^{2}` multiplications. Comparing all patch combinations involves :math:`w^{2}*h^{2}` such computations. - -Given a maximum displacement :math:`d`, for each location :math:`x_{1}` it computes correlations :math:`c(x_{1}, x_{2})` only in a neighborhood of size :math:`D:=2d+1`, -by limiting the range of :math:`x_{2}`. We use strides :math:`s_{1}, s_{2}`, to quantize :math:`x_{1}` globally and to quantize :math:`x_{2}` within the neighborhood -centered around :math:`x_{1}`. - -The final output is defined by the following expression: - -.. math:: - out[n, q, i, j] = c(x_{i, j}, x_{q}) - -where :math:`i` and :math:`j` enumerate spatial locations in :math:`f_{1}`, and :math:`q` denotes the :math:`q^{th}` neighborhood of :math:`x_{i,j}`. -)code" TVM_ADD_FILELINE) - .set_attrs_type() - .set_num_inputs(2) - .add_argument("data1", "Tensor", "Input data1 to the correlation.") - .add_argument("data2", "Tensor", "Input data2 to the correlation.") - .set_support_level(2) - .set_attr("FInferCorrectLayout", CorrelationInferCorrectLayout) - .add_type_rel("Correlation", CorrelationRel) - .set_attr("TOpPattern", kOutEWiseFusable); - -} // namespace relay -} // namespace tvm diff --git a/src/relay/op/nn/nn.cc b/src/relay/op/nn/nn.cc deleted file mode 100644 index ccc973485529..000000000000 --- a/src/relay/op/nn/nn.cc +++ /dev/null @@ -1,1680 +0,0 @@ -/* - * 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 nn.cc - * \brief Property def of nn operators. - */ - -#include "nn.h" - -#include -#include -#include -#include -#include -#include -#include -#include -#include - -#include -#include -#include - -#include "../../transforms/infer_layout_utils.h" -#include "../make_op.h" -#include "../op_common.h" -#include "../type_relations.h" - -namespace tvm { -namespace relay { - -// relay.nn.bias_add -TVM_REGISTER_NODE_TYPE(BiasAddAttrs); - -bool BiasAddRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - ICHECK_EQ(types.size(), 3); - const auto* data = types[0].as(); - if (data == nullptr) return false; - - const BiasAddAttrs* param = attrs.as(); - ICHECK(param != nullptr); - int axis = param->axis; - if (axis < 0) { - axis = data->shape.size() + axis; - } - if (axis >= static_cast(data->shape.size()) || axis < 0) { - reporter->GetDiagCtx().EmitFatal(Diagnostic::Error(reporter->GetSpan()) - << "The axis in bias_add must be in range for the shape; " - << "attempted to access index " << param->axis << " of " - << PrettyPrint(data->shape)); - return false; - } - - // assign output type - reporter->Assign(types[1], TensorType({data->shape[axis]}, data->dtype)); - reporter->Assign(types[2], types[0]); - return true; -} - -// Positional relay function to create dense operator used by frontend FFI. -Expr MakeBiasAdd(Expr data, Expr bias, int axis) { - auto attrs = make_object(); - attrs->axis = axis; - static const Op& op = Op::Get("nn.bias_add"); - return Call(op, {data, bias}, Attrs(attrs), {}); -} - -TVM_REGISTER_GLOBAL("relay.op.nn._make.bias_add").set_body_typed(MakeBiasAdd); - -RELAY_REGISTER_OP("nn.bias_add") - .describe(R"code(Add bias to an axis of the input. - -)code" TVM_ADD_FILELINE) - .set_attrs_type() - .set_num_inputs(2) - .add_argument("data", "nD Tensor", "Input data.") - .add_argument("bias", "1D Tensor", "Bias.") - .set_support_level(1) - .add_type_rel("BiasAdd", BiasAddRel) - .set_attr("TOpPattern", kBroadcast) - .set_attr("FTVMCompute", [](const Attrs& attrs, const Array& inputs, - const Type& out_type) { - const auto* param = attrs.as(); - return tvm::Array{topi::nn::bias_add(inputs[0], inputs[1], param->axis)}; - }); - -// relay.nn.fifo_buffer -TVM_REGISTER_NODE_TYPE(FIFOBufferAttrs); - -Expr MakeFIFOBuffer(Expr input, Expr buffer, int axis) { - auto attrs = make_object(); - attrs->axis = axis; - static const Op& op = Op::Get("nn.fifo_buffer"); - return Call(op, {input, buffer}, Attrs(attrs), {}); -} - -bool FIFOBufferRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - ICHECK_EQ(types.size(), 3); - const auto* input = types[0].as(); - const auto* buffer = types[1].as(); - const FIFOBufferAttrs* param = attrs.as(); - if (input == nullptr || buffer == nullptr) { - return false; - } - ICHECK(param != nullptr); - ICHECK_EQ(input->shape.size(), buffer->shape.size()); - - const size_t buffer_axis = static_cast( - param->axis < 0 ? static_cast(buffer->shape.size()) + param->axis : param->axis); - - reporter->Assert(buffer_axis < buffer->shape.size()); - for (size_t i = 0; i < buffer->shape.size(); ++i) { - if (i != buffer_axis) { - reporter->AssertEQ(input->shape[i], buffer->shape[i]); - } - } - reporter->Assert(input->shape[buffer_axis] < buffer->shape[buffer_axis]); - - Array oshape = buffer->shape; - - reporter->Assign(types[2], TensorType(oshape, buffer->dtype)); - return true; -} - -TVM_REGISTER_GLOBAL("relay.op.nn._make.fifo_buffer").set_body_typed(MakeFIFOBuffer); - -RELAY_REGISTER_OP("nn.fifo_buffer") - .describe(R"code(FIFO buffer -Compute equivalent of - -``` -concat(buffer, data, axis=axis) \ -.slice_axis(axis=axis, begin=data.shape[axis], end=data.shape[axis]+buffer.shape[axis]) -``` - -Useful for -* Encoding explicit re-use of computation in convolution ops operated on a sliding window input -* Implementing a FIFO queue to cache intermediate results, e.g. as in Fast WaveNet. -)code" TVM_ADD_FILELINE) - .set_attrs_type() - .set_num_inputs(2) - .add_argument("data", "Tensor", "Latest input") - .add_argument("buffer", "Tensor", "Buffer storing latest [length_buffer] inputs") - .set_support_level(3) - .add_type_rel("FIFOBuffer", FIFOBufferRel) - .set_attr("TOpPattern", kOpaque); - -// ------------------- relay.nn.matmul -TVM_REGISTER_NODE_TYPE(MatmulAttrs); - -Expr MakeMatmul(Expr tensor_a, Expr tensor_b, IndexExpr units, DataType out_dtype, bool transpose_a, - bool transpose_b) { - auto attrs = make_object(); - attrs->units = units; - attrs->out_dtype = out_dtype; - attrs->transpose_a = transpose_a; - attrs->transpose_b = transpose_b; - static const Op& matmul_op = Op::Get("nn.matmul"); - return Call(matmul_op, {tensor_a, tensor_b}, Attrs(attrs), {}); -} - -TVM_REGISTER_GLOBAL("relay.op.nn._make.matmul").set_body_typed(MakeMatmul); - -RELAY_REGISTER_OP("nn.matmul") - .describe(R"code(Applies a linear transformation: :math:`C = A * B`. A & B can be transposed. - -- **tensor_a**: `(x1, x2, ..., xn, input_dim)` or `(x1, x2, ..., input_dim, xn)` -- **tensor_b**: `(input_dim, units)` or `(units, input_dim)` -- **out**: `(x1, x2, ..., xn, units)`. - -)code" TVM_ADD_FILELINE) - .set_attrs_type() - .set_num_inputs(2) - .add_argument("tensor_a", "nD Tensor", "The first input Tensor.") - .add_argument("tensor_b", "2D Tensor", "The second input Tensor.") - .set_support_level(1) - .set_attr("FInferCorrectLayout", DenseInferCorrectLayout) - .add_type_rel("Matmul", MatmulRel) - .set_attr("TOpPattern", kOutEWiseFusable); - -// ------------------- relay.nn.matmul - -// ------------------- relay.nn.dense -TVM_REGISTER_NODE_TYPE(DenseAttrs); - -// Positional relay function to create dense operator used by frontend FFI. -Expr MakeDense(Expr data, Expr weight, IndexExpr units, DataType out_dtype) { - auto attrs = make_object(); - attrs->units = units; - attrs->out_dtype = out_dtype; - static const Op& op = Op::Get("nn.dense"); - return Call(op, {data, weight}, Attrs(attrs), {}); -} - -InferCorrectLayoutOutput DenseInferCorrectLayout(const Attrs& attrs, - const Array& new_in_layouts, - const Array& old_in_layouts, - const Array& old_in_types) { - return InferCorrectLayoutOutput({"NC", "NC"}, {"NC"}, attrs); -} - -TVM_REGISTER_GLOBAL("relay.op.nn._make.dense").set_body_typed(MakeDense); - -RELAY_REGISTER_OP("nn.dense") - .describe(R"code(Applies a linear transformation: :math:`Y = XW^T`. - -- **data**: `(x1, x2, ..., xn, input_dim)` -- **weight**: `(units, input_dim)` -- **out**: `(x1, x2, ..., xn, units)`. - -)code" TVM_ADD_FILELINE) - .set_attrs_type() - .set_num_inputs(2) - .add_argument("data", "nD Tensor", "Input data.") - .add_argument("weight", "2D Tensor", "Weight matrix.") - .set_support_level(1) - .set_attr("FInferCorrectLayout", DenseInferCorrectLayout) - .add_type_rel("Dense", MatmulRel) - .set_attr("TOpPattern", kOutEWiseFusable); -// ------------------- relay.nn.dense - -// ------------------- relay.nn.contrib_dense_pack -TVM_REGISTER_NODE_TYPE(DensePackAttrs); - -// Positional relay function to create dense_pack operator used by frontend FFI. -Expr MakeDensePack(Expr data, Expr weight, tvm::String weight_layout, IndexExpr units, - DataType out_dtype) { - auto attrs = make_object(); - attrs->units = units; - attrs->out_dtype = out_dtype; - attrs->weight_layout = std::move(weight_layout); - static const Op& op = Op::Get("nn.contrib_dense_pack"); - return Call(op, {data, weight}, Attrs(attrs), {}); -} - -TVM_REGISTER_GLOBAL("relay.op.nn._make.contrib_dense_pack").set_body_typed(MakeDensePack); - -bool DensePackRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - ICHECK_EQ(types.size(), 3); - const auto* data = types[0].as(); - const auto* weight = types[1].as(); - if (data == nullptr || weight == nullptr) return false; - - const DensePackAttrs* param = attrs.as(); - ICHECK(param != nullptr); - - ICHECK_EQ(data->shape.size(), 2) << "Only 2D data is supported"; - ICHECK(weight->shape.size() == 3 || weight->shape.size() == 4) << "Expect weight to be 3D or 4D"; - - Array oshape = data->shape; - oshape.Set(1, weight->shape[0] * weight->shape[2]); - - DataType out_dtype = param->out_dtype; - if (out_dtype.bits() == 0) { - out_dtype = data->dtype; - } - // assign output type - reporter->Assign(types[2], TensorType(oshape, out_dtype)); - return true; -} - -InferCorrectLayoutOutput DensePackInferCorrectLayout(const Attrs& attrs, - const Array& new_in_layouts, - const Array& old_in_layouts, - const Array& old_in_types) { - auto params = attrs.as(); - ICHECK(params); - return InferCorrectLayoutOutput({"NC", params->weight_layout}, {"NC"}, attrs); -} - -RELAY_REGISTER_OP("nn.contrib_dense_pack") - .describe(R"code(Applies a linear transformation: :math:`Y = XW^T`. - -- **data**: `(batch, input_dim)` -- **weight**: `(units // pack_weight_tile, input_dim, pack_weight_tile)` -- **out**: `(batch, units)`. - -)code" TVM_ADD_FILELINE) - .set_attrs_type() - .set_num_inputs(2) - .add_argument("data", "2D Tensor", "Input data.") - .add_argument("weight", "3D Tensor", "Packed weight matrix.") - .set_support_level(10) - .set_attr("FInferCorrectLayout", DensePackInferCorrectLayout) - .add_type_rel("DensePack", DensePackRel) - .set_attr("TOpPattern", kOutEWiseFusable); - -// ------------------- relay.nn.contrib_dense_pack - -// relay.leaky_relu -TVM_REGISTER_NODE_TYPE(LeakyReluAttrs); - -// Positional relay function to create leaky relu operator used by frontend FFI. -Expr MakeLeakyRelu(Expr data, double alpha) { - auto attrs = make_object(); - attrs->alpha = alpha; - static const Op& op = Op::Get("nn.leaky_relu"); - return Call(op, {data}, Attrs(attrs), {}); -} - -TVM_REGISTER_GLOBAL("relay.op.nn._make.leaky_relu").set_body_typed(MakeLeakyRelu); - -RELAY_REGISTER_OP("nn.leaky_relu") - .describe(R"code(Leaky version of a Rectified Linear Unit. - -`y = x > 0 ? x : alpha * x` - -)code" TVM_ADD_FILELINE) - .set_attrs_type() - .set_num_inputs(1) - .add_argument("data", "Tensor", "Input data.") - .set_support_level(3) - .add_type_rel("Identity", IdentityRel) - .set_attr("FInferCorrectLayout", ElemwiseArbitraryLayout) - .set_attr("TOpPattern", kElemWise) - .set_attr("FTVMCompute", [](const Attrs& attrs, const Array& inputs, - const Type& out_type) { - const auto* param = attrs.as(); - return Array{topi::leaky_relu(inputs[0], param->alpha)}; - }); - -// relay.prelu -TVM_REGISTER_NODE_TYPE(PReluAttrs); - -bool PReluRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - ICHECK_EQ(types.size(), 3); - const auto* data = types[0].as(); - if (data == nullptr) return false; - - const PReluAttrs* param = attrs.as(); - ICHECK(param != nullptr); - - ICHECK(param->axis < static_cast(data->shape.size())) - << "Wrong axis (" << param->axis << ")value."; - - // assign alpha type - Array alpha_shape({data->shape[param->axis]}); - reporter->Assign(types[1], TensorType(alpha_shape, data->dtype)); - - // assign output type - reporter->Assign(types[2], TensorType(data->shape, data->dtype)); - return true; -} - -InferCorrectLayoutOutput PReluInferCorrectLayout(const Attrs& attrs, - const Array& new_in_layouts, - const Array& old_in_layouts, - const Array& old_in_types) { - ICHECK_EQ(old_in_layouts.size(), 2U); - ICHECK_EQ(old_in_types.size(), 2U); - Layout data_layout = old_in_layouts[0]; - if (new_in_layouts.defined()) { - ICHECK_EQ(new_in_layouts.size(), 2U); - } - return InferCorrectLayoutOutput({data_layout, Layout("C")}, {data_layout}, attrs); -} - -// Positional relay function to create prelu operator used by frontend FFI. -Expr MakePRelu(Expr data, Expr alpha, int axis) { - auto attrs = make_object(); - attrs->axis = axis; - static const Op& op = Op::Get("nn.prelu"); - return Call(op, {data, alpha}, Attrs(attrs), {}); -} - -TVM_REGISTER_GLOBAL("relay.op.nn._make.prelu").set_body_typed(MakePRelu); - -RELAY_REGISTER_OP("nn.prelu") - .describe(R"code(Parametric version of a Rectified Linear Unit. -It accepts two arguments: an input ``x`` and a channelwise slope ``alpha`` -and computes the output as :math:`PReLU(x) y = x > 0 ? x : alpha * x`, -where :math:`*` is an channelwise multiplication for each sample in the batch. -)code" TVM_ADD_FILELINE) - .set_attrs_type() - .set_num_inputs(2) - .add_argument("data", "Tensor", "Input data.") - .add_argument("alpha", "Tensor", "Input channelwise alpha.") - .set_support_level(3) - .add_type_rel("PRelu", PReluRel) - .set_attr("FInferCorrectLayout", PReluInferCorrectLayout) - .set_attr("TOpPattern", kBroadcast) - .set_attr("FTVMCompute", [](const Attrs& attrs, const Array& inputs, - const Type& out_type) { - const auto* param = attrs.as(); - return Array{topi::prelu(inputs[0], inputs[1], param->axis)}; - }); - -// relay.softmax -TVM_REGISTER_NODE_TYPE(SoftmaxAttrs); - -bool SoftmaxRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - ICHECK_EQ(types.size(), 2); - const auto* data = types[0].as(); - if (data == nullptr) return false; - - const SoftmaxAttrs* param = attrs.as(); - ICHECK(param != nullptr); - int axis = param->axis; - int ndim = static_cast(data->shape.size()); - if (axis >= ndim || axis < -ndim) { - reporter->GetDiagCtx().EmitFatal(Diagnostic::Error(reporter->GetSpan()) - << "Wrong axis (" << axis << ") not in expected range: [" - << -ndim << ", " << ndim << ")"); - return false; - } - - reporter->Assign(types[1], types[0]); - return true; -} - -TVM_REGISTER_GLOBAL("relay.op.nn._make.softmax").set_body_typed([](Expr data, int axis) { - auto attrs = make_object(); - attrs->axis = axis; - static const Op& op = Op::Get("nn.softmax"); - return Call(op, {data}, Attrs(attrs), {}); -}); - -RELAY_REGISTER_OP("nn.softmax") - .describe(R"code(Softmax layer. - -.. math:: \text{softmax}(x)_i = \frac{exp(x_i)}{\sum_j exp(x_j)} - -.. note:: - This operator can be optimized away for inference. - -- **data**: The input data -)code" TVM_ADD_FILELINE) - .set_attrs_type() - .set_num_inputs(1) - .add_argument("data", "Tensor", "The input tensor.") - .set_support_level(1) - .add_type_rel("Softmax", SoftmaxRel) - .set_attr("TOpPattern", kOutEWiseFusable); - -// relay.fast_softmax -TVM_REGISTER_NODE_TYPE(SoftmaxAttrs); - -TVM_REGISTER_GLOBAL("relay.op.nn._make.fast_softmax").set_body_typed([](Expr data, int axis) { - auto attrs = make_object(); - attrs->axis = axis; - static const Op& op = Op::Get("nn.fast_softmax"); - return Call(op, {data}, Attrs(attrs), {}); -}); - -RELAY_REGISTER_OP("nn.fast_softmax") - .describe(R"code(Softmax layer. - Use approximation to compute exponent for faster speed. - -.. math:: \text{softmax}(x)_i = \frac{exp(x_i)}{\sum_j exp(x_j)} - -.. note:: - This operator can be optimized away for inference. - -- **data**: The input data -)code" TVM_ADD_FILELINE) - .set_attrs_type() - .set_num_inputs(1) - .add_argument("data", "Tensor", "The input tensor.") - .set_support_level(1) - .add_type_rel("Softmax", SoftmaxRel) - .set_attr("TOpPattern", kOutEWiseFusable); - -// relay.nn.log_softmax -TVM_REGISTER_GLOBAL("relay.op.nn._make.log_softmax").set_body_typed([](Expr data, int axis) { - auto attrs = make_object(); - attrs->axis = axis; - static const Op& op = Op::Get("nn.log_softmax"); - return Call(op, {data}, Attrs(attrs), {}); -}); - -RELAY_REGISTER_OP("nn.log_softmax") - .describe(R"code(Computes log softmax. - -.. math:: \text{log_softmax}(x)_i = \log \frac{exp(x_i)}{\sum_j exp(x_j)} - -.. note:: - This operator can be optimized away for inference. - -- **data**: The input data -)code" TVM_ADD_FILELINE) - .set_attrs_type() - .set_num_inputs(1) - .add_argument("data", "Tensor", "The input tensor.") - .set_support_level(1) - .add_type_rel("Softmax", SoftmaxRel) - .set_attr("TOpPattern", kOutEWiseFusable) - .set_attr("FTVMCompute", [](const Attrs& attrs, const Array& inputs, - const Type& out_type) { - const auto* param = attrs.as(); - ICHECK(param != nullptr); - ICHECK(param->axis == -1 || param->axis == static_cast(inputs[0].ndim()) - 1) - << "log_softmax currently only works on last dimension"; - return Array{topi::nn::log_softmax(inputs[0])}; - }); - -// relay.nn.batch_flatten -bool BatchFlattenRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - ICHECK_EQ(types.size(), 2); - const auto* data = types[0].as(); - if (data == nullptr) return false; - if (data->shape.size() == 0) return false; - - auto target_dim = tir::make_const(DataType::Int(32), 1); - - for (uint32_t i = 1; i < data->shape.size(); ++i) { - if (!data->shape[i].as()) { - target_dim = target_dim * data->shape[i]; - } else { - target_dim = data->shape[i]; - break; - } - } - - std::vector oshape({data->shape[0], target_dim}); - - // assign output type - reporter->Assign(types[1], TensorType(oshape, data->dtype)); - return true; -} - -Expr MakeBatchFlatten(Expr data) { - static const Op& op = Op::Get("nn.batch_flatten"); - return Call(op, {data}, Attrs(), {}); -} - -TVM_REGISTER_GLOBAL("relay.op.nn._make.batch_flatten").set_body_typed(MakeBatchFlatten); - -RELAY_REGISTER_OP("nn.batch_flatten") - .describe(R"code(Flattens the input into a 2-D array. - -For an input array with shape ``(d1, d2, ..., dk)``, `batch_flatten` operation reshapes -the input array into an output array of shape ``(d1, d2*...*dk)``. - -Example:: - - x = [[ - [1,2,3], - [4,5,6], - [7,8,9] - ], - [ [1,2,3], - [4,5,6], - [7,8,9] - ]], - - batch_flatten(x) = [[ 1., 2., 3., 4., 5., 6., 7., 8., 9.], - [ 1., 2., 3., 4., 5., 6., 7., 8., 9.]] - -)code" TVM_ADD_FILELINE) - .set_num_inputs(1) - .add_argument("data", "Tensor", "The input tensor.") - .set_support_level(2) - .add_type_rel("BatchFlatten", BatchFlattenRel) - .set_attr("TOpPattern", kInjective) - .set_attr("FTVMCompute", - [](const Attrs& attrs, const Array& inputs, - const Type& out_type) { - return Array{topi::nn::flatten(inputs[0])}; - }) - .set_attr("TReshapeOp", true); - -// relu -TVM_REGISTER_GLOBAL("relay.op.nn._make.relu").set_body_typed([](Expr data) { - static const Op& op = Op::Get("nn.relu"); - return Call(op, {data}, Attrs(), {}); -}); - -RELAY_REGISTER_OP("nn.relu") - .describe(R"code(Returns the relu input array, computed element-wise. - -.. math:: - max(x, 0) - -)code" TVM_ADD_FILELINE) - .set_num_inputs(1) - .add_argument("data", "Tensor", "The input tensor.") - .set_support_level(1) - .add_type_rel("Identity", IdentityRel) - .set_attr("FInferCorrectLayout", ElemwiseArbitraryLayout) - .set_attr("TOpPattern", kElemWise) - .set_attr("FTVMCompute", [](const Attrs& attrs, const Array& inputs, - const Type& out_type) { - return Array{topi::relu(inputs[0], 0.0f)}; - }); - -// Positional relay function to create LRN operator used by frontend FFI. -TVM_REGISTER_NODE_TYPE(LRNAttrs); - -Expr MakeLRN(Expr data, int size, int axis, double alpha, double beta, double bias) { - auto attrs = make_object(); - attrs->size = size; - attrs->axis = axis; - attrs->alpha = alpha; - attrs->beta = beta; - attrs->bias = bias; - static const Op& op = Op::Get("nn.lrn"); - return Call(op, {data}, Attrs(attrs), {}); -} - -TVM_REGISTER_GLOBAL("relay.op.nn._make.lrn").set_body_typed(MakeLRN); - -RELAY_REGISTER_OP("nn.lrn") - .describe(R"code(LRN layer. - -Normalize the input in a local region across or within feature maps. -Each input value is divided by (1 + (\alpha/n) \sum_i x_i^2)^\beta, -where n is the size of each local region, and the sum is taken over the region -centered at that value (zero padding is added where necessary). - -.. math:: - - data / (bias + (alpha * sum_data ^2 /size))^beta - -- **data**: The input tensor. -)code" TVM_ADD_FILELINE) - .set_attrs_type() - .set_num_inputs(1) - .add_argument("data", "Tensor", "The input tensor.") - .set_support_level(2) - .add_type_rel("Identity", IdentityRel) - .set_attr("TOpPattern", kOpaque); - -// Positional relay function to create L2Normalize operator used by frontend FFI. -TVM_REGISTER_NODE_TYPE(L2NormalizeAttrs); - -Expr MakeL2Normalize(Expr data, double eps, Array axis) { - auto attrs = make_object(); - attrs->eps = eps; - attrs->axis = std::move(axis); - static const Op& op = Op::Get("nn.l2_normalize"); - return Call(op, {data}, Attrs(attrs), {}); -} - -InferCorrectLayoutOutput L2NormalizeInferCorrectLayout( - const Attrs& attrs, const Array& new_in_layouts, const Array& old_in_layouts, - const Array& old_in_types) { - const auto* attrs_ptr = attrs.as(); - ICHECK(attrs_ptr); - ObjectPtr param = make_object(*attrs_ptr); - - Array> old_in_shapes; - for (auto old_in_t : old_in_types) { - ICHECK(old_in_t.as()); - old_in_shapes.push_back(old_in_t.as()->shape); - } - std::vector axis_list; - for (auto i : param->axis) { - int64_t axis = i->value; - if (axis < 0) { - axis = axis + static_cast(old_in_shapes[0].size()); - } - axis_list.emplace_back(axis); - } - - Layout ret = Layout::Undef(); - if (new_in_layouts.defined() && old_in_layouts.defined()) { - for (size_t i = 0; i < axis_list.size(); ++i) { - const auto& axis_dim = old_in_layouts[0][axis_list[i]]; - auto axis_index = new_in_layouts[0].IndexOf(axis_dim); - param->axis.Set(i, axis_index); - } - ret = new_in_layouts[0]; - } else if (old_in_layouts.defined()) { - ret = old_in_layouts[0]; - } - - return InferCorrectLayoutOutput({ret}, {ret}, Attrs(param)); -} - -TVM_REGISTER_GLOBAL("relay.op.nn._make.l2_normalize").set_body_typed(MakeL2Normalize); - -RELAY_REGISTER_OP("nn.l2_normalize") - .describe(R"code(L2 Normalization layer. - -Normalizes along dimension axis using an L2 norm - -.. math:: - output = x / sqrt(max(sum(x^2), epsilon)) - -- **data**: The input tensor. -)code" TVM_ADD_FILELINE) - .set_attrs_type() - .set_num_inputs(1) - .add_argument("data", "Tensor", "The input tensor.") - .set_support_level(2) - .set_attr("FInferCorrectLayout", L2NormalizeInferCorrectLayout) - .add_type_rel("Identity", IdentityRel); - -// Dropout -TVM_REGISTER_NODE_TYPE(DropoutAttrs); - -bool DropoutRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - ICHECK_EQ(types.size(), 2); - const auto* data = types[0].as(); - if (data == nullptr) return false; - - // dropout returns the original tensor with dropout applied - // and a mask tensor (1.0 where element not dropped, 0.0 where dropped) - auto ret_type = TensorType(data->shape, data->dtype); - reporter->Assign(types[1], TupleType(Array({ret_type, ret_type}))); - return true; -} - -Expr MakeDropout(Expr data, double rate) { - auto attrs = make_object(); - attrs->rate = rate; - static const Op& op = Op::Get("nn.dropout"); - return Call(op, {data}, Attrs(attrs), {}); -} - -TVM_REGISTER_GLOBAL("relay.op.nn._make.dropout").set_body_typed(MakeDropout); - -RELAY_REGISTER_OP("nn.dropout") - .describe(R"code(Applies the dropout operation to the input array. - -During training, each element of the input is set to zero with probability ``p``. -The whole array is rescaled by ``1/(1-p)`` to keep the expected sum of the input unchanged. - -)code" TVM_ADD_FILELINE) - .set_attrs_type() - .set_num_inputs(1) - .add_argument("data", "Tensor", "Input to which dropout will be applied.") - .set_support_level(1) - .set_attr("FInferCorrectLayout", ElemwiseArbitraryLayout) - .set_attr("TOpPattern", kOpaque) - .add_type_rel("Dropout", DropoutRel) - .set_attr("TOpIsStateful", true); - -// batch_norm -TVM_REGISTER_NODE_TYPE(BatchNormAttrs); - -InferCorrectLayoutOutput BatchNormInferCorrectLayout(const Attrs& attrs, - const Array& new_in_layouts, - const Array& old_in_layouts, - const Array& old_in_types) { - const auto* attrs_ptr = attrs.as(); - ICHECK(attrs_ptr); - ObjectPtr param = make_object(*attrs_ptr); - - Array> old_in_shapes; - for (auto old_in_t : old_in_types) { - ICHECK(old_in_t.as()); - old_in_shapes.push_back(old_in_t.as()->shape); - } - - size_t axis = - param->axis < 0 ? param->axis + old_in_shapes[0].size() : static_cast(param->axis); - - Layout ret = Layout::Undef(); - - // If new_in_layouts are defined, this code tries to modify the layout. - if (new_in_layouts.defined() && old_in_layouts.defined()) { - // Get the new C axis. Extract the dim in old layout. Find the index of that dim in next layout. - const auto& bn_dim = old_in_layouts[0][axis]; - auto new_index = new_in_layouts[0].IndexOf(bn_dim); - param->axis = new_index; - ret = new_in_layouts[0]; - } else if (old_in_layouts.defined()) { - ret = old_in_layouts[0]; - } - // BN has 5 inputs, 3 outputs. The last 4 inputs and last 2 outputs have "C" layout. - Layout c_layout = Layout("C"); - return InferCorrectLayoutOutput({ret, c_layout, c_layout, c_layout, c_layout}, - {ret, c_layout, c_layout}, Attrs(param)); -} - -bool BatchNormRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - ICHECK_EQ(types.size(), 6); - const auto* data = types[0].as(); - if (data == nullptr) return false; - - const BatchNormAttrs* param = attrs.as(); - - // axis of -1 means use the last dimension - ICHECK(param->axis >= -1 && param->axis < (int)data->shape.size()); - int axis = (param->axis != -1) ? param->axis : data->shape.size() - 1; - auto axis_size = data->shape[axis]; - - // if we are using beta and gamma, they need to be of shape (dim,) - reporter->Assign(types[1], TensorType({axis_size}, data->dtype)); - reporter->Assign(types[2], TensorType({axis_size}, data->dtype)); - reporter->Assign(types[3], TensorType({axis_size}, data->dtype)); - reporter->Assign(types[4], TensorType({axis_size}, data->dtype)); - - // output is a tuple of the normed data (same shape as input), new running mean, - // new running variance, saved mean and saved variance (the latter are all - // vectors of length dim) - std::vector fields; - auto vec_ty = TensorType(Array({data->shape[axis]}), data->dtype); - fields.push_back(TensorType(data->shape, data->dtype)); - fields.push_back(vec_ty); - fields.push_back(vec_ty); - reporter->Assign(types[5], TupleType(Array(fields))); - return true; -} - -Expr MakeBatchNorm(Expr data, Expr gamma, Expr beta, Expr moving_mean, Expr moving_var, int axis, - double epsilon, bool center, bool scale) { - auto attrs = make_object(); - attrs->axis = axis; - attrs->epsilon = epsilon; - attrs->center = center; - attrs->scale = scale; - static const Op& op = Op::Get("nn.batch_norm"); - return Call(op, {data, gamma, beta, moving_mean, moving_var}, Attrs(attrs), {}); -} - -TVM_REGISTER_GLOBAL("relay.op.nn._make.batch_norm").set_body_typed(MakeBatchNorm); - -RELAY_REGISTER_OP("nn.batch_norm") - .describe(R"code(Batch normalization layer (Ioffe and Szegedy, 2014). -Normalizes the input at each batch, i.e. applies a transformation -that maintains the mean activation close to 0 and the activation -standard deviation close to 1. - -.. math:: - - data\_mean[i] = mean(data[:,i,:,...]) \\ - data\_var[i] = var(data[:,i,:,...]) - -Then compute the normalized output, which has the same shape as input, as following: - -.. math:: - - out[:,i,:,...] = \frac{data[:,i,:,...] - data\_mean[i]}{\sqrt{data\_var[i]+\epsilon}} \ -* gamma[i] + beta[i] - -Both *mean* and *var* returns a scalar by treating the input as a vector. - -Assume the input has size *k* on axis 1, then both ``gamma`` and ``beta`` have shape *(k,)*. - -Besides the inputs and the outputs, this operator accepts two auxiliary -states, ``moving_mean`` and ``moving_var``, which are *k*-length -vectors. They are global statistics for the whole dataset, which are updated -by:: - - moving_mean = moving_mean * momentum + data_mean * (1 - momentum) - moving_var = moving_var * momentum + data_var * (1 - momentum) - -The parameter ``axis`` specifies which axis of the input shape denotes -the 'channel' (separately normalized groups). The default is 1. Specifying -1 sets the channel -axis to be the last item in the input shape. - -.. note:: - This operator can be optimized away for inference. -)code" TVM_ADD_FILELINE) - .set_attrs_type() - .set_num_inputs(5) - .add_argument("data", "Tensor", "Input to which batch_norm will be applied.") - .add_argument("gamma", "Tensor", "The gamma scale factor.") - .add_argument("beta", "Tensor", "The beta offset factor.") - .add_argument("moving_mean", "Tensor", "Running mean of input.") - .add_argument("moving_var", "Tensor", "Running variance of input.") - .set_attr("FInferCorrectLayout", BatchNormInferCorrectLayout) - .set_support_level(1) - .add_type_rel("BatchNorm", BatchNormRel) - .set_attr("TOpPattern", kOutEWiseFusable); - -// instance_norm -TVM_REGISTER_NODE_TYPE(InstanceNormAttrs); - -template -InferCorrectLayoutOutput NormalizationInferCorrectLayout( - const Attrs& attrs, const Array& new_in_layouts, const Array& old_in_layouts, - const Array& old_in_types) { - const auto* attrs_ptr = attrs.as(); - ICHECK(attrs_ptr); - ObjectPtr param = make_object(*attrs_ptr); - - Array> old_in_shapes; - for (auto old_in_t : old_in_types) { - ICHECK(old_in_t.as()); - old_in_shapes.push_back(old_in_t.as()->shape); - } - - size_t axis = - param->axis < 0 ? param->axis + old_in_shapes[0].size() : static_cast(param->axis); - - Layout ret = Layout::Undef(); - - // If new_in_layouts are defined, this code tries to modify the layout. - if (new_in_layouts.defined() && old_in_layouts.defined()) { - // Get the new C axis. Extract the dim in old layout. Find the index of that dim in next layout. - const auto& ln_dim = old_in_layouts[0][axis]; - auto new_index = new_in_layouts[0].IndexOf(ln_dim); - param->axis = new_index; - ret = new_in_layouts[0]; - } else if (old_in_layouts.defined()) { - ret = old_in_layouts[0]; - } - - // For normalization has 3 inputs, 1 outputs. The last 2 inputs have "C" layout. - Layout c_layout = Layout("C"); - return InferCorrectLayoutOutput({ret, c_layout, c_layout}, {ret}, Attrs(param)); -} - -bool InstanceNormRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - ICHECK_EQ(types.size(), 4); - const auto* data = types[0].as(); - if (data == nullptr) return false; - ICHECK_GT(data->shape.size(), 2); - const InstanceNormAttrs* param = attrs.as(); - int axis = param->axis >= 0 ? param->axis : param->axis + data->shape.size(); - ICHECK(axis >= 0 && axis < (int)data->shape.size()); - reporter->Assign(types[1], TensorType({data->shape[axis]}, data->dtype)); - reporter->Assign(types[2], TensorType({data->shape[axis]}, data->dtype)); - reporter->Assign(types[3], TensorType(data->shape, data->dtype)); - - return true; -} - -Expr MakeInstanceNorm(Expr data, Expr gamma, Expr beta, int axis, double epsilon, bool center, - bool scale) { - auto attrs = make_object(); - attrs->axis = axis; - attrs->epsilon = epsilon; - attrs->center = center; - attrs->scale = scale; - static const Op& op = Op::Get("nn.instance_norm"); - return Call(op, {data, gamma, beta}, Attrs(attrs), {}); -} - -TVM_REGISTER_GLOBAL("relay.op.nn._make.instance_norm").set_body_typed(MakeInstanceNorm); - -RELAY_REGISTER_OP("nn.instance_norm") - .describe(R"code(Instance Normalization (Ulyanov and et al., 2016) -Applies instance normalization to the n-dimensional input array. - -.. math:: - - out = \frac{data - mean(data)}{\sqrt{var(data)+\epsilon}} - * gamma + beta - -The instance normalization is similar to batch normalization, but unlike -batch normalization, the mean and var are calculated per-dimension -separately for each object(instance) in a mini-batch, not over a batch. -And the same normalization is applied both at test and train time. - -Assume the input has size *k* on axis 1, then both ``gamma`` and ``beta`` -have shape *(k,)*. - -The parameter ``axis`` specifies which axis of the input shape denotes -the 'channel'. The default is 1. Specifying -1 sets the channel axis -to be the last item in the input shape. - -.. note:: - - This operator can be optimized away for inference. -)code" TVM_ADD_FILELINE) - .set_attrs_type() - .set_num_inputs(3) - .add_argument("data", "Tensor", "Input to which instance_norm will be applied.") - .add_argument("gamma", "Tensor", "The gamma scale factor.") - .add_argument("beta", "Tensor", "The beta offset factor.") - .set_attr("FInferCorrectLayout", - NormalizationInferCorrectLayout) - .set_support_level(1) - .add_type_rel("InstanceNorm", InstanceNormRel); - -// layer_norm -TVM_REGISTER_NODE_TYPE(LayerNormAttrs); - -bool LayerNormRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - ICHECK_EQ(types.size(), 4); - const auto* data = types[0].as(); - if (data == nullptr) return false; - const LayerNormAttrs* param = attrs.as(); - int axis = param->axis >= 0 ? param->axis : param->axis + data->shape.size(); - ICHECK(axis >= 0 && axis < (int)data->shape.size()); - reporter->Assign(types[1], TensorType({data->shape[axis]}, data->dtype)); - reporter->Assign(types[2], TensorType({data->shape[axis]}, data->dtype)); - reporter->Assign(types[3], TensorType(data->shape, data->dtype)); - - return true; -} - -Expr MakeLayerNorm(Expr data, Expr gamma, Expr beta, int axis, double epsilon, bool center, - bool scale) { - auto attrs = make_object(); - attrs->axis = axis; - attrs->epsilon = epsilon; - attrs->center = center; - attrs->scale = scale; - static const Op& op = Op::Get("nn.layer_norm"); - return Call(op, {data, gamma, beta}, Attrs(attrs), {}); -} - -TVM_REGISTER_GLOBAL("relay.op.nn._make.layer_norm").set_body_typed(MakeLayerNorm); - -RELAY_REGISTER_OP("nn.layer_norm") - .describe(R"code( -)code" TVM_ADD_FILELINE) - .set_attrs_type() - .set_num_inputs(3) - .add_argument("data", "Tensor", "Input to which layer_norm will be applied.") - .add_argument("gamma", "Tensor", "The gamma scale factor.") - .add_argument("beta", "Tensor", "The beta offset factor.") - .set_attr("FInferCorrectLayout", - NormalizationInferCorrectLayout) - .set_support_level(1) - .add_type_rel("LayerNorm", LayerNormRel); - -// group_norm -TVM_REGISTER_NODE_TYPE(GroupNormAttrs); - -bool GroupNormRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - ICHECK_EQ(types.size(), 4); - const auto* data = types[0].as(); - if (data == nullptr) return false; - const GroupNormAttrs* param = attrs.as(); - int axis = param->axis >= 0 ? param->axis : param->axis + data->shape.size(); - ICHECK(axis >= 0 && axis < (int)data->shape.size()); - reporter->Assign(types[1], TensorType({data->shape[axis]}, data->dtype)); - reporter->Assign(types[2], TensorType({data->shape[axis]}, data->dtype)); - reporter->Assign(types[3], TensorType(data->shape, data->dtype)); - - return true; -} - -Expr MakeGroupNorm(Expr data, Expr gamma, Expr beta, int num_groups, int axis, double epsilon, - bool center, bool scale) { - auto attrs = make_object(); - attrs->num_groups = num_groups; - attrs->axis = axis; - attrs->epsilon = epsilon; - attrs->center = center; - attrs->scale = scale; - static const Op& op = Op::Get("nn.group_norm"); - return Call(op, {data, gamma, beta}, Attrs(attrs), {}); -} - -TVM_REGISTER_GLOBAL("relay.op.nn._make.group_norm").set_body_typed(MakeGroupNorm); - -RELAY_REGISTER_OP("nn.group_norm") - .describe(R"code( -Group normalization normalizes over group of channels for each training examples. -We can say that, Group Norm is in between Instance Norm and Layer Norm. When we put -all the channels into a single group, group normalization becomes Layer normalization. -And, when we put each channel into different groups it becomes Instance normalization - -https://arxiv.org/pdf/1803.08494.pdf - -Applies group normalization to the n-dimensional input array by seperating the input channels -into 'num_groups' groups, each containing 'num_channels / num_groups' channels. -The mean and standard-deviation are calculated separately over the each group. gamma and -beta are learnable per-channel affine transform parameter vectors of size num_channels. - -.. math:: - - out = \frac{data - mean(data, axis)}{\sqrt{var(data, axis)+\epsilon}} - * gamma + beta - -Unlike batch normalization, the mean and var are computed along a group of channels. - -If the input has size k on axis 1, then both gamma and beta have shape (k,). - -.. note:: - - This operator can be optimized away for inference. - -)code" TVM_ADD_FILELINE) - .set_attrs_type() - .set_num_inputs(3) - .add_argument("data", "Tensor", "Input to which group_norm will be applied.") - .add_argument("gamma", "Tensor", "The gamma scale factor.") - .add_argument("beta", "Tensor", "The beta offset factor.") - .set_support_level(1) - .add_type_rel("GroupNorm", GroupNormRel); - -// ------------------- relay.nn.batch_matmul -TVM_REGISTER_NODE_TYPE(BatchMatmulAttrs); - -// Positional relay function to create batch_matmul operator used by frontend FFI. -Expr MakeBatchMatmul(Expr tensor_a, Expr tensor_b, DataType out_dtype, bool transpose_a, - bool transpose_b) { - auto attrs = make_object(); - attrs->out_dtype = out_dtype; - attrs->transpose_a = transpose_a; - attrs->transpose_b = transpose_b; - static const Op& op = Op::Get("nn.batch_matmul"); - return Call(op, {tensor_a, tensor_b}, Attrs(attrs), {}); -} - -TVM_REGISTER_GLOBAL("relay.op.nn._make.batch_matmul").set_body_typed(MakeBatchMatmul); - -RELAY_REGISTER_OP("nn.batch_matmul") - .describe(R"code(Compute batch matrix multiplication of `tensor_a` and `tensor_b`. - -Both `tensor_a` and `tensor_b` can be transposed. For legacy reason, we use NT format -(transpose_a=False, transpose_b=True) by default. - -.. math:: - - batch\_matmul(A, B)[i, :, :] = matmul(A[i, :, :], B[i, :, :]^T) - -- **tensor_a**: `(b, m, k)` or `(b, k, m)` -- **tensor_b**: `(b, k, n)` or `(b, n, k)` -- **out**: `(b, m, n)`. - -)code" TVM_ADD_FILELINE) - .set_attrs_type() - .set_num_inputs(2) - .add_argument("tensor_a", "3D Tensor", "The first input.") - .add_argument("tensor_b", "3D Tensor", "The second input.") - .set_support_level(10) - .add_type_rel("BatchMatmul", BatchMatmulRel) - .set_attr("TOpPattern", kOutEWiseFusable); - -// ------------------- relay.nn.batch_matmul - -// relay.nn.cross_entropy -bool CrossEntropyRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - ICHECK_EQ(types.size(), 3); - const auto* x = types[0].as(); - const auto* y = types[1].as(); - if (x == nullptr || y == nullptr) return false; - ICHECK(x->shape.size() == 2 && y->shape.size() == 2) - << "CrossEntropy: shapes of x and y is inconsistent, " - << "x shape = " << x->shape << ", " - << "y shape = " << y->shape; - ICHECK(reporter->AssertEQ(x->shape[0], y->shape[0])) - << "CrossEntropy: shapes of x and y is inconsistent, " - << "x shape = " << x->shape << ", " - << "y shape = " << y->shape; - ICHECK(reporter->AssertEQ(x->shape[1], y->shape[1])) - << "CrossEntropy: shapes of x and y is inconsistent, " - << "x shape = " << x->shape << ", " - << "y shape = " << y->shape; - // assign output type - reporter->Assign(types[2], TensorType({}, x->dtype)); - return true; -} - -// Positional relay function to create cross_entropy operator used by frontend FFI. -Expr MakeCrossEntropy(Expr predictions, Expr targets) { - static const Op& op = Op::Get("nn.cross_entropy"); - return Call(op, {predictions, targets}, Attrs(), {}); -} - -TVM_REGISTER_GLOBAL("relay.op.nn._make.cross_entropy").set_body_typed(MakeCrossEntropy); - -RELAY_REGISTER_OP("nn.cross_entropy") - .describe(R"code( -Computes cross entropy given predictions and targets. -Do log on the data - do not accept logits. -)code" TVM_ADD_FILELINE) - .set_num_inputs(2) - .add_argument("x", "1D Tensor", "Predictions.") - .add_argument("y", "1D Tensor", "Targets.") - .set_support_level(10) - .add_type_rel("CrossEntropy", CrossEntropyRel) - .set_attr("TOpPattern", kOpaque); - -// relay.nn.dilate -TVM_REGISTER_NODE_TYPE(DilateAttrs); - -bool DilateRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - ICHECK_EQ(types.size(), 2); - const auto* x = types[0].as(); - const DilateAttrs* param = attrs.as(); - if (x == nullptr) return false; - ICHECK_EQ(x->shape.size(), param->strides.size()); - - std::vector oshape; - for (size_t i = 0; i < param->strides.size(); ++i) { - if (!x->shape[i].as()) { - oshape.push_back((x->shape[i] - 1) * param->strides[i] + 1); - } else { - oshape.push_back(x->shape[i]); - } - } - - reporter->Assign(types[1], TensorType(Array(oshape), x->dtype)); - return true; -} - -// Positional relay function to create dilate operator used by frontend FFI. -Expr MakeDilate(Expr data, Array strides, double dilation_value = 0.0) { - auto attrs = make_object(); - attrs->strides = std::move(strides); - attrs->dilation_value = std::move(dilation_value); - static const Op& op = Op::Get("nn.dilate"); - return Call(op, {data}, Attrs(attrs), {}); -} - -TVM_REGISTER_GLOBAL("relay.op.nn._make.dilate").set_body_typed(MakeDilate); - -RELAY_REGISTER_OP("nn.dilate") - .describe(R"code( -Dilate data with given dilation value (0 by default). -)code" TVM_ADD_FILELINE) - .set_num_inputs(1) - .add_argument("x", "1D Tensor", "Data to dilate.") - .set_support_level(10) - .add_type_rel("Dilate", DilateRel) - .set_attr("TOpPattern", kInjective); - -// relay.nn.cross_entropy_with_logits -// Positional relay function to create cross_entropy_with_logits operator used by frontend FFI. -Expr MakeCrossEntropyWithLogits(Expr predictions, Expr targets) { - static const Op& op = Op::Get("nn.cross_entropy_with_logits"); - return Call(op, {predictions, targets}, Attrs(), {}); -} - -TVM_REGISTER_GLOBAL("relay.op.nn._make.cross_entropy_with_logits") - .set_body_typed(MakeCrossEntropyWithLogits); - -RELAY_REGISTER_OP("nn.cross_entropy_with_logits") - .describe(R"code( -Computes cross entropy given predictions and targets. -Accept logits. -)code" TVM_ADD_FILELINE) - .set_num_inputs(2) - .add_argument("x", "1D Tensor", "Predictions.") - .add_argument("y", "1D Tensor", "Targets.") - .set_support_level(10) - .add_type_rel("CrossEntropy", CrossEntropyRel) - .set_attr("TOpPattern", kOpaque); - -// Depth to space and space to depth -TVM_REGISTER_NODE_TYPE(SubPixelAttrs); - -// relay.nn.nll_loss -TVM_REGISTER_NODE_TYPE(NLLLossAttrs); - -bool NLLLossRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - ICHECK_EQ(types.size(), 4) << "NLLLossRel expects 4 types, but " << types.size() - << " were provided."; - const auto* predictions = types[0].as(); - const auto* targets = types[1].as(); - const auto* weights = types[2].as(); - const NLLLossAttrs* param = attrs.as(); - if (predictions == nullptr || targets == nullptr || weights == nullptr) return false; - if (!(predictions->shape.size() - targets->shape.size() == 1)) { - reporter->GetDiagCtx().EmitFatal(Diagnostic::Error(reporter->GetSpan()) - << "NLLLossRel: predictions should be one" - << " dimension larger than targets," - << "predictions shape = " << predictions->shape - << ", targets shape = " << targets->shape); - return false; - } - if (!(weights->shape.size() == 1)) { - reporter->GetDiagCtx().EmitFatal(Diagnostic::Error(reporter->GetSpan()) - << "NLLLossRel: weights should be a one dimension" - << " Tensor with its length the number of classes," - << " but Tensor of dimension " << weights->shape.size() - << " were provided."); - return false; - } - if (!reporter->AssertEQ(predictions->shape[1], weights->shape[0])) { - reporter->GetDiagCtx().EmitFatal(Diagnostic::Error(reporter->GetSpan()) - << "NLLLossRel: the second dimension of predictions" - << " should be the number of classes, " - << "which is the length of weights, " - << "predictions shape = " << predictions->shape - << ", weights shape = " << weights->shape); - return false; - } - if (!(predictions->dtype == weights->dtype && - (predictions->dtype.is_float() || predictions->dtype.is_bfloat16()))) { - reporter->GetDiagCtx().EmitFatal(Diagnostic::Error(reporter->GetSpan()) - << "NLLLossRel: predictions and weights should" - << " be of the same floating type."); - return false; - } - if (!targets->dtype.is_int()) { - reporter->GetDiagCtx().EmitFatal(Diagnostic::Error(reporter->GetSpan()) - << "NLLLossRel: targets should be of int type."); - return false; - } - // assign output type - if (param->reduction == "none") { - reporter->Assign(types[3], TensorType(targets->shape, predictions->dtype)); - } else { - reporter->Assign(types[3], TensorType({}, predictions->dtype)); - } - return true; -} - -// Handler to create a call to the padding op used by front-end FFI -Expr MakeNLLLoss(Expr predictions, Expr targets, Expr weights, String reduction, int ignore_index) { - auto attrs = make_object(); - attrs->reduction = reduction; - attrs->ignore_index = ignore_index; - static const Op& op = Op::Get("nn.nll_loss"); - return Call(op, {predictions, targets, weights}, Attrs(attrs), {}); -} - -TVM_REGISTER_GLOBAL("relay.op.nn._make.nll_loss").set_body_typed(MakeNLLLoss); - -RELAY_REGISTER_OP("nn.nll_loss") - .describe(R"code( -Negative log likelihood loss for given prediction and target. -)code" TVM_ADD_FILELINE) - .set_attrs_type() - .set_num_inputs(3) - .add_argument("predictions", "Tensor", "The prediction tensor.") - .add_argument("targets", "Tensor", "The target tensor.") - .add_argument("weights", "Tensor", "The weight of each target values.") - .add_type_rel("NLLLoss", NLLLossRel) - .set_attr("TOpPattern", kOutEWiseFusable); - -bool DepthToSpaceRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - ICHECK_EQ(types.size(), 2); - const auto* data = types[0].as(); - if (data == nullptr) return false; - - static const Layout kNCHW("NCHW"); - - const SubPixelAttrs* param = attrs.as(); - ICHECK(param != nullptr); - const int block_size = param->block_size; - const Layout in_layout(param->layout); - auto layout_converter = tir::BijectiveLayout(in_layout, kNCHW); - ICHECK(layout_converter.defined()) - << "DepthToSpace only support input layouts that are convertible from NCHW." - << " But got " << in_layout; - - auto oshape = layout_converter.ForwardShape(data->shape); - if (!oshape[1].as()) { - oshape.Set(1, indexdiv(oshape[1], (block_size * block_size))); - } - if (!oshape[2].as()) { - oshape.Set(2, oshape[2] * block_size); - } - if (!oshape[3].as()) { - oshape.Set(3, oshape[3] * block_size); - } - - // Assign output type - reporter->Assign(types[1], TensorType(layout_converter.BackwardShape(oshape), data->dtype)); - - return true; -} - -// Positional relay function to create DepthToSpace operator -// used by frontend FFI -Expr MakeDepthToSpace(Expr data, int block_size, String layout, String mode) { - auto attrs = make_object(); - attrs->block_size = block_size; - attrs->layout = std::move(layout); - attrs->mode = std::move(mode); - static const Op& op = Op::Get("nn.depth_to_space"); - return Call(op, {data}, Attrs(attrs), {}); -} - -TVM_REGISTER_GLOBAL("relay.op.nn._make.depth_to_space").set_body_typed(MakeDepthToSpace); - -RELAY_REGISTER_OP("nn.depth_to_space") - .describe(R"code(Rearrange input channels into spatial pixels. - -- **data**: data is a 4D array of shape - (batch, in_channels, in_height, in_width) for NCHW - -- **out**: Output is a 4D array of shape - (batch, in_channels / block_size * block_size, in_height * block_size, in_width * block_size) for NCHW. - -)code" TVM_ADD_FILELINE) - .set_attrs_type() - .set_num_inputs(1) - .add_argument("data", "Tensor", "The input tensor") - .set_support_level(5) - .add_type_rel("DepthToSpace", DepthToSpaceRel) - .set_attr("TOpPattern", kInjective); - -bool SpaceToDepthRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - ICHECK_EQ(types.size(), 2); - const auto* data = types[0].as(); - if (data == nullptr) return false; - - static const Layout kNCHW("NCHW"); - - const SubPixelAttrs* param = attrs.as(); - ICHECK(param != nullptr); - const int block_size = param->block_size; - const Layout in_layout(param->layout); - auto layout_converter = tir::BijectiveLayout(in_layout, kNCHW); - ICHECK(layout_converter.defined()) - << "SpaceToDepth only support input layouts that are convertible from NCHW." - << " But got " << in_layout; - - auto oshape = layout_converter.ForwardShape(data->shape); - if (!oshape[1].as()) { - oshape.Set(1, oshape[1] * (block_size * block_size)); - } - if (!oshape[2].as()) { - oshape.Set(2, indexdiv(oshape[2], block_size)); - } - if (!oshape[3].as()) { - oshape.Set(3, indexdiv(oshape[3], block_size)); - } - - // Assign output type - reporter->Assign(types[1], TensorType(layout_converter.BackwardShape(oshape), data->dtype)); - - return true; -} - -// Positional relay function to create SpaceToDepth operator -// used by frontend FFI -Expr MakeSpaceToDepth(Expr data, int block_size, String layout) { - auto attrs = make_object(); - attrs->block_size = block_size; - attrs->layout = std::move(layout); - static const Op& op = Op::Get("nn.space_to_depth"); - return Call(op, {data}, Attrs(attrs), {}); -} - -TVM_REGISTER_GLOBAL("relay.op.nn._make.space_to_depth").set_body_typed(MakeSpaceToDepth); - -RELAY_REGISTER_OP("nn.space_to_depth") - .describe(R"code(Rearrange spatial pixels into new output channels. - -- **data**: data is a 4D array of shape - (batch, in_channels, in_height, in_width) for NCHW - -- **out**: Output is a 4D array of shape - (batch, in_channels * block_size * block_size, in_height / block_size, in_width / block_size) for NCHW. - -)code" TVM_ADD_FILELINE) - .set_attrs_type() - .set_num_inputs(1) - .add_argument("data", "Tensor", "The input tensor") - .set_support_level(5) - .add_type_rel("SpaceToDepth", SpaceToDepthRel) - .set_attr("TOpPattern", kInjective); - -// Positional relay function to create SpaceToBatchND operator -// used by frontend FFI -TVM_REGISTER_NODE_TYPE(SpaceToBatchNDAttrs); - -Expr MakeSpaceToBatchND(Expr data, Array block_shape, Array> paddings, - double pad_value) { - auto attrs = make_object(); - attrs->block_shape = std::move(block_shape); - attrs->paddings = std::move(paddings); - attrs->pad_value = pad_value; - static const Op& op = Op::Get("nn.space_to_batch_nd"); - return Call(op, {data}, Attrs(attrs), {}); -} - -bool SpaceToBatchNDRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - CHECK_EQ(types.size(), 2); - - auto* input = types[0].as(); - // Input must be a TensorType - if (input == nullptr) { - CHECK(types[0].as()) - << "SpaceToBatchND: expect input type to be TensorType but got " << types[0]; - return false; - } - - if (input->shape.size() <= 1) return false; - - const auto* param = attrs.as(); - CHECK(param != nullptr); - - auto block_shape = param->block_shape; - auto paddings = param->paddings; - const int bdims = static_cast(block_shape.size()); - const int pdims = static_cast(paddings.size()); - // Paddings must be provided for each spatial dim. - CHECK(pdims == bdims) << "SpaceToBatchND: Paddings must be provided for each spatial dim"; - - // Apply paddings to input - auto in_shape = input->shape; - std::vector padded_shape(input->shape.begin(), input->shape.end()); - for (size_t i = 0; i < paddings.size(); i++) { - CHECK_EQ(paddings[i].size(), 2U); - auto pad_before = tir::as_const_int(param->paddings[i][0]); - auto pad_after = tir::as_const_int(param->paddings[i][1]); - auto padding = tir::make_const(input->shape[i].dtype(), *pad_before + *pad_after); - padded_shape[i + 1] = in_shape[i + 1] + padding; - } - - auto block_shape_numele = tir::make_const(DataType::Int(32), 1); - for (size_t i = 0; i < block_shape.size(); i++) { - block_shape_numele *= block_shape[i]; - } - - // Construct output shape - std::vector out_shape(padded_shape); - out_shape[0] = in_shape[0] * block_shape_numele; - for (size_t i = 1; i <= block_shape.size(); i++) { - out_shape[i] = div(padded_shape[i], block_shape[i - 1]); - } - - // Assign output shape - reporter->Assign(types[1], TensorType(Array(out_shape), input->dtype)); - return true; -} - -Array SpaceToBatchNDCompute(const Attrs& attrs, const Array& inputs, - const Type& out_type) { - const auto* param = attrs.as(); - CHECK(param != nullptr); - - auto b_shape = param->block_shape; - auto paddings = param->paddings; - Array pad_before; - Array pad_after; - - for (size_t i = 0; i < paddings.size(); ++i) { - pad_before.push_back(paddings[i][0]); - } - for (size_t i = 0; i < paddings.size(); ++i) { - pad_after.push_back(paddings[i][1]); - } - const auto* out_ttype = out_type.as(); - return Array{ - topi::space_to_batch_nd(inputs[0], b_shape, pad_before, pad_after, - tvm::tir::make_const(out_ttype->dtype, param->pad_value))}; -} - -TVM_REGISTER_GLOBAL("relay.op.nn._make.space_to_batch_nd").set_body_typed(MakeSpaceToBatchND); - -RELAY_REGISTER_OP("nn.space_to_batch_nd") - .describe(R"code(Divide spatial dimensions of the input into a grid of blocks -and interleave them into batch dim. - -- **data**: data is a ND array of shape - (batch, spatial_shapes, remaining_shapes) for NHWC - -- **out**: Output is a ND array of shape - (batch * prod(block_shape), padded_data[1] / block_shape[0], ..., padded_data[M] / block_shape[M-1], - remaining_shape) for NHWC, where M is the number of spatial dimensions. - -Example:: - - x = [[[[1], [2]], [[3], [4]]]] - - space_to_batch_nd(x, block_shape = [2, 2]) = - [[[[1]]], [[[2]]], [[[3]]], [[[4]]]] - -)code" TVM_ADD_FILELINE) - .set_num_inputs(1) - .add_argument("data", "Tensor", "The input tensor.") - .set_attrs_type() - .set_support_level(5) - .add_type_rel("SpaceToBatchND", SpaceToBatchNDRel) - .set_attr("FTVMCompute", SpaceToBatchNDCompute) - .set_attr("TOpPattern", kInjective); - -/*****************************************************************/ - -// Positional relay function to create BatchToSpaceND operator -// used by frontend FFI -TVM_REGISTER_NODE_TYPE(BatchToSpaceNDAttrs); - -Expr MakeBatchToSpaceND(Expr data, Array block_shape, Array> crops) { - auto attrs = make_object(); - attrs->block_shape = std::move(block_shape); - attrs->crops = std::move(crops); - static const Op& op = Op::Get("nn.batch_to_space_nd"); - return Call(op, {data}, Attrs(attrs), {}); -} - -bool BatchToSpaceNDRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - CHECK_EQ(types.size(), 2); - - auto* input = types[0].as(); - // Input must be a TensorType - if (input == nullptr) { - CHECK(types[0].as()) - << "BatchToSpaceND: expect input type to be TensorType but got " << types[0]; - return false; - } - - if (input->shape.size() <= 1) return false; - - const auto* param = attrs.as(); - CHECK(param != nullptr); - - auto block_shape = param->block_shape; - auto crops = param->crops; - const int bdims = static_cast(block_shape.size()); - const int cdims = static_cast(crops.size()); - const int indims = static_cast(input->shape.size()); - // crops must be provided for each spatial dim. - CHECK(cdims == bdims) << "BatchToSpaceND: crops must be provided for each spatial dim"; - CHECK(bdims < indims) << "BatchToSpaceND: block_shape must be less than input shape"; - - auto block_shape_numele = tir::make_const(DataType::Int(32), 1); - for (size_t i = 0; i < block_shape.size(); i++) { - block_shape_numele *= block_shape[i]; - } - - auto in_shape = input->shape; - - // Construct output shape - // Start with input shape, only batch and spatial dims shapes are modified. - std::vector out_shape(input->shape.begin(), input->shape.end()); - out_shape[0] = in_shape[0] / block_shape_numele; - for (size_t i = 1; i <= block_shape.size(); i++) { - out_shape[i] = (in_shape[i] * block_shape[i - 1]) - crops[i - 1][0] - crops[i - 1][1]; - } - for (int i = bdims + 1; i < indims; i++) { - out_shape[i] = in_shape[i]; - } - - // Assign output shape - reporter->Assign(types[1], TensorType(Array(out_shape), input->dtype)); - return true; -} - -Array BatchToSpaceNDCompute(const Attrs& attrs, const Array& inputs, - const Type& out_type) { - const auto* param = attrs.as(); - CHECK(param != nullptr); - - auto b_shape = param->block_shape; - auto crops = param->crops; - Array crop_begin_list, crop_end_list; - for (size_t i = 0; i < crops.size(); ++i) { - crop_begin_list.push_back(crops[i][0]); - crop_end_list.push_back(crops[i][1]); - } - - return Array{ - topi::batch_to_space_nd(inputs[0], b_shape, crop_begin_list, crop_end_list)}; -} - -TVM_REGISTER_GLOBAL("relay.op.nn._make.batch_to_space_nd").set_body_typed(MakeBatchToSpaceND); - -RELAY_REGISTER_OP("nn.batch_to_space_nd") - .describe(R"code(Reshape the batch dimension into spatial dimensions. - -Example:: - - x = [[[[1]]], [[[2]]], [[[3]]], [[[4]]]] - - batch_to_space_nd(x, block_shape = [2, 2]) = - [[[[1], [2]], [[3], [4]]]] - -)code" TVM_ADD_FILELINE) - .set_num_inputs(1) - .add_argument("data", "Tensor", "The input tensor.") - .set_attrs_type() - .set_support_level(5) - .add_type_rel("BatchToSpaceND", BatchToSpaceNDRel) - .set_attr("FTVMCompute", BatchToSpaceNDCompute) - .set_attr("TOpPattern", kInjective); - -} // namespace relay -} // namespace tvm diff --git a/src/relay/op/nn/nn.h b/src/relay/op/nn/nn.h deleted file mode 100644 index 3ebef29776cd..000000000000 --- a/src/relay/op/nn/nn.h +++ /dev/null @@ -1,230 +0,0 @@ -/* - * 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 src/relay/op/nn/nn.h - * \brief Properties def of nn operators for sharing. - */ -#ifndef TVM_RELAY_OP_NN_NN_H_ -#define TVM_RELAY_OP_NN_NN_H_ - -#include -#include -#include -#include - -#include -#include -#include - -#include "../op_common.h" - -namespace tvm { -namespace relay { - -template -bool MatmulRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - ICHECK_EQ(types.size(), 3); - const auto* tensor_a = types[0].as(); - const auto* tensor_b = types[1].as(); - if (tensor_a == nullptr) return false; - ICHECK(static_cast(tensor_a->shape.size()) != 0); - - const AttrType* param = attrs.as(); - ICHECK(param != nullptr); - TensorType meta_schedule_tensor_b{nullptr}; - if (param->meta_schedule_original_shape.size() > 0) { - meta_schedule_tensor_b = TensorType(param->meta_schedule_original_shape, - tensor_b == nullptr ? tensor_a->dtype : tensor_b->dtype); - tensor_b = meta_schedule_tensor_b.get(); - } - // Default set to dense layout - bool transpose_a = false; - bool transpose_b = true; - const auto& mattrs = attrs.as(); - if (mattrs != nullptr) { - transpose_a = mattrs->transpose_a; - transpose_b = mattrs->transpose_b; - } - - const Array& dshape = tensor_a->shape; - Array oshape = dshape; - tvm::PrimExpr reduce = dshape[dshape.size() - 1]; - if (transpose_a) { - reduce = dshape[dshape.size() - 2]; - oshape.Set((oshape.size() - 2), dshape[oshape.size() - 1]); - } - auto tensor_b_dtype = (tensor_b == nullptr ? tensor_a->dtype : tensor_b->dtype); - if (param->units.defined()) { - // validate the tensor_b shape is proper if defined - // Assign tensor_b type - const Array& wshape = transpose_b ? Array({param->units, reduce}) - : Array({reduce, param->units}); - // It is possible for tensor_b to be nullptr in which case we will use - // data dtype as the tensor_b dtype. However if tensor_b dtype is explicitly - // present we will use that. - if (param->auto_scheduler_rewritten_layout.size() != 0) { - // If the layout is rewritten by auto-scheduler or meta-schedule, - // we just forcefully apply the layout provided by auto-scheduler and - // skip the normal inference logic. - {} // do nothing - } else if (param->meta_schedule_original_shape.size() == 0) { - // Normal case: assign result to reporter - reporter->Assign(types[1], TensorType(wshape, tensor_b_dtype)); - } - oshape.Set((oshape.size() - 1), param->units); - } else { - if (tensor_b == nullptr) return false; - const Array& wshape = tensor_b->shape; - // When tensor_b's layout has been rewritten, figure it out based on the - // total number of elements and input dimensions. - if (param->auto_scheduler_rewritten_layout.size() != 0) { - PrimExpr tensor_b_elements = 1; - for (size_t i = 0; i < wshape.size(); i++) { - tensor_b_elements = tensor_b_elements * wshape[i]; - } - oshape.Set(oshape.size() - 1, tensor_b_elements / dshape[dshape.size() - 1]); - // Otherwise just pull it out of the tensor_b shape directly. - } else { - ICHECK(static_cast(tensor_b->shape.size()) == 2); - if (param->auto_scheduler_rewritten_layout.size() == 0 && - param->meta_schedule_original_shape.size() == 0) { - // ensure inner dimension matches between data and weight. If one inner - // dimension is dynamic then it is inferred to match the other inner - // dimension. - std::vector A_shape(tensor_a->shape.begin(), tensor_a->shape.end()); - std::vector B_shape(tensor_b->shape.begin(), tensor_b->shape.end()); - auto sa = A_shape.size(); - auto sb = B_shape.size(); - size_t index_swap_A; - size_t index_swap_B; - if (transpose_a && transpose_b) { - index_swap_A = sa - 2; - index_swap_B = sb - 1; - } else if (transpose_a) { - index_swap_A = sa - 2; - index_swap_B = sb - 2; - } else if (transpose_b) { - index_swap_A = sa - 1; - index_swap_B = sb - 1; - } else { - index_swap_A = sa - 1; - index_swap_B = sb - 2; - } - - // Rewrite dynamic axes to static where constraints allow. - auto tmp = A_shape[index_swap_A]; - if (A_shape[index_swap_A].as()) { - A_shape[index_swap_A] = B_shape[index_swap_B]; - } - if (B_shape[index_swap_B].as()) { - B_shape[index_swap_B] = tmp; - } - - // Update input types with new constrained shapes. - reporter->Assign(types[0], TensorType(A_shape, tensor_a->dtype)); - reporter->Assign(types[1], TensorType(B_shape, tensor_b_dtype)); - } - oshape.Set(oshape.size() - 1, transpose_b ? wshape[0] : wshape[1]); - } - } - - DataType out_dtype = param->out_dtype; - if (out_dtype.bits() == 0) { - out_dtype = tensor_a->dtype; - } - // assign output type - reporter->Assign(types[2], TensorType(oshape, out_dtype)); - return true; -} - -template -bool BatchMatmulRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - ICHECK_EQ(types.size(), 3); - const auto* x = types[0].as(); - const auto* y = types[1].as(); - if (x == nullptr || y == nullptr) return false; - - const AttrType* param = attrs.as(); - DataType out_dtype = param->out_dtype; - if (out_dtype.bits() == 0) { - out_dtype = x->dtype; - if (x->dtype.bits() == 0) { - out_dtype = y->dtype; - } - } - TensorType meta_schedule_y{nullptr}; - if (param->meta_schedule_original_shape.size() != 0) { - meta_schedule_y = TensorType(param->meta_schedule_original_shape, out_dtype); - y = meta_schedule_y.get(); - } - ICHECK(param != nullptr); - bool transpose_a = param->transpose_a; - bool transpose_b = param->transpose_b; - Array y_shape{nullptr}; - if (param->auto_scheduler_rewritten_layout.size() != 0) { - y_shape = auto_scheduler::GetShapeFromRewrittenLayout( - param->auto_scheduler_rewritten_layout, - transpose_b ? tvm::runtime::Array({"b", "j", "k"}) - : tvm::runtime::Array({"b", "k", "j"})); - } else if (param->meta_schedule_original_shape.size() != 0) { - y_shape = param->meta_schedule_original_shape; - } else { - y_shape = y->shape; - } - ICHECK(x->shape.size() == 3 && y_shape.size() == 3); - const PrimExpr& xb = x->shape[0]; - const PrimExpr& xi = x->shape[transpose_a ? 2 : 1]; - const PrimExpr& xk = x->shape[transpose_a ? 1 : 2]; - const PrimExpr& yb = y_shape[0]; - const PrimExpr& yk = y_shape[transpose_b ? 2 : 1]; - const PrimExpr& yj = y_shape[transpose_b ? 1 : 2]; - - bool is_dyn = false; - for (size_t i = 0; i < 3; ++i) { - if (x->shape[i].as() != nullptr || y_shape[i].as() != nullptr) { - is_dyn = true; - break; - } - } - if (!is_dyn) { - ICHECK(reporter->AssertEQ(xb, yb) || reporter->AssertEQ(xb, 1) || reporter->AssertEQ(yb, 1)) - << "BatchDot: batch dimensions don't match, " - << " x shape=" << x->shape << ", y shape=" << y_shape; - ICHECK(reporter->AssertEQ(xk, yk)) << "BatchDot: shapes of x and y is inconsistent, " - << " x shape=" << x->shape << ", y shape=" << y_shape; - } - - // assign output type - const auto& out_b = - xb->IsInstance() || yb->IsInstance() ? tir::Any() : max(xb, yb); - reporter->Assign(types[2], TensorType(Array({out_b, xi, yj}), out_dtype)); - return true; -} - -InferCorrectLayoutOutput DenseInferCorrectLayout(const Attrs& attrs, - const Array& new_in_layouts, - const Array& old_in_layouts, - const Array& old_in_types); - -} // namespace relay -} // namespace tvm -#endif // TVM_RELAY_OP_NN_NN_H_ diff --git a/src/relay/op/nn/pad.cc b/src/relay/op/nn/pad.cc deleted file mode 100644 index 8cfb369901ad..000000000000 --- a/src/relay/op/nn/pad.cc +++ /dev/null @@ -1,276 +0,0 @@ -/* - * 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 pad.cc - * \brief Implementation of operator pad - */ -#include -#include -#include -#include -#include -#include - -#include - -#include "../make_op.h" -#include "../op_common.h" - -namespace tvm { -namespace relay { - -// relay.nn.pad -TVM_REGISTER_NODE_TYPE(PadAttrs); - -InferCorrectLayoutOutput PadInferCorrectLayout(const Attrs& attrs, - const Array& new_in_layouts, - const Array& old_in_layouts, - const Array& old_in_types) { - const auto* attrs_ptr = attrs.as(); - CHECK(attrs_ptr); - ObjectPtr params = make_object(*attrs_ptr); - - Layout ret_data; - // If new_in_layouts are defined, this code tries to modify the layout. - bool is_layout_modified = new_in_layouts.defined(); - if (new_in_layouts.defined()) { - // Create a map of axis to param_width. For the new layout, a new param_width is generated using - // the map. The new layout is rejected, if the padding is happening along the axis which was - // split. - - // 1) Create a map from axis to param_width using old layout. - std::map> axis_pad_width; - int index_counter = 0; - ICHECK_EQ(new_in_layouts.size(), 2); - ICHECK_EQ(old_in_layouts.size(), 2); - for (auto iter_var : old_in_layouts[0]->axes) { - const auto& old_layout_axis = LayoutAxis::Get(iter_var); - axis_pad_width.emplace(old_layout_axis.name(), params->pad_width[index_counter]); - index_counter++; - } - - // 2) Create new pad width by walking over the new layout and using the map. - tvm::Array> new_pad_width; - for (auto iter_var : new_in_layouts[0]->axes) { - const auto& new_layout_axis = LayoutAxis::Get(iter_var); - auto axis_name = new_layout_axis.name(); - if (axis_pad_width.count(axis_name) != 0 && new_layout_axis.IsPrimal()) { - // This is primal axis. So, directly use the original pad_width. - new_pad_width.push_back(axis_pad_width.at(axis_name)); - } else { - // This is the axis that got split. So, check that pad_width was [0, 0] originally. - const auto& dual_axis = new_layout_axis.ToPrimal(); - auto dual_axis_name = dual_axis.name(); - ICHECK(axis_pad_width.count(dual_axis_name)) - << "Missing axis " << dual_axis << " in " << old_in_layouts[0].name(); - new_pad_width.push_back(axis_pad_width.at(dual_axis_name)); - - // If any pad_width element is not zero, do not change the layout. - for (auto width : axis_pad_width.at(dual_axis_name)) { - if (auto* width_imm = width.as()) { - if (width_imm->value != 0) { - is_layout_modified = false; - } - } else { - is_layout_modified = false; - } - } - } - } - - // If the above conditions satisfied, we can set the newly created pad_width and use the new - // layout. - if (is_layout_modified) { - ret_data = new_in_layouts[0]; - params->pad_width = new_pad_width; - } - } - - if (!is_layout_modified) { - if (old_in_layouts.defined()) { - ICHECK_EQ(old_in_layouts.size(), 2); - ret_data = old_in_layouts[0]; - } else { - ret_data = Layout::Undef(); - } - } - - // The pad value is always a scalar - Layout ret_pad_value = Layout("1"); - return InferCorrectLayoutOutput({ret_data, ret_pad_value}, {ret_data}, Attrs(params)); -} - -bool PadRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - // types = [pad_data_type, pad_value_type, ret_type] - ICHECK_EQ(types.size(), 3); - const auto* data = types[0].as(); - if (data == nullptr) return false; - - const PadAttrs* param = attrs.as(); - ICHECK(param != nullptr); - - // check that pad widths match lengths - ICHECK(data->shape.size() == param->pad_width.size()) - << "There should be as many pad width pairs as shape dimensions " - << "but the shape has " << data->shape.size() << " dimensions " - << "and there are " << param->pad_width.size() << " pad width pairs."; - - // each pad width element should be a pair of positive integers - std::vector oshape; - for (size_t i = 0; i < param->pad_width.size(); i++) { - ICHECK(param->pad_width[i].size() == 2) - << "Each pad width element should be a pair but at index " << i << " there are " - << param->pad_width[i].size() << " elements."; - - auto width1 = tir::as_const_int(param->pad_width[i][0]); - auto width2 = tir::as_const_int(param->pad_width[i][1]); - ICHECK(width1 != nullptr); - ICHECK(width2 != nullptr); - - if (!data->shape[i].as()) { - auto padding = tir::make_const(data->shape[i].dtype(), *width1 + *width2); - oshape.push_back(data->shape[i] + padding); - if (tir::as_const_int(data->shape[i])) { - ICHECK(topi::detail::GetConstInt(data->shape[i] + padding) >= 0) - << "Output shape post padding should be positive but got " << data->shape[i] + padding; - } - } else { - oshape.push_back(data->shape[i]); - } - } - - reporter->Assign(types[2], TensorType(Array(oshape), data->dtype)); - return true; -} - -Array PadCompute(const Attrs& attrs, const Array& inputs, - const Type& out_type) { - const auto* param = attrs.as(); - ICHECK(param != nullptr); - - auto pad_width = param->pad_width; - ICHECK(pad_width.size() == inputs[0].ndim() && pad_width[0].size() == 2) << "Illegal pad_width"; - Array pad_before; - for (size_t i = 0; i < pad_width.size(); ++i) { - pad_before.push_back(pad_width[i][0]); - } - Array pad_after; - for (size_t i = 0; i < pad_width.size(); ++i) { - pad_after.push_back(pad_width[i][1]); - } - te::Tensor cast_pad_value = topi::cast(inputs[1], inputs[0]->dtype); - const PrimExpr& pad_value = cast_pad_value(Array(inputs[1]->shape.size(), 0)); - return Array{topi::pad(inputs[0], pad_before, pad_after, pad_value, "T_pad", - topi::kElementWise, param->pad_mode)}; -} - -// Handler to create a call to the padding op used by front-end FFI -Expr MakePad(Expr data, Array> pad_width, Expr pad_value, String pad_mode) { - auto attrs = make_object(); - attrs->pad_width = std::move(pad_width); - attrs->pad_mode = std::move(pad_mode); - static const Op& op = Op::Get("nn.pad"); - return Call(op, {data, pad_value}, Attrs(attrs), {}); -} - -TVM_REGISTER_GLOBAL("relay.op.nn._make.pad").set_body_typed(MakePad); - -RELAY_REGISTER_OP("nn.pad") - .describe(R"code(Pad for n-D tensor. - -)code" TVM_ADD_FILELINE) - .set_attrs_type() - .set_num_inputs(2) - .add_argument("data", "Tensor", "The input tensor.") - .add_argument("pad_val", "Tensor", "The value to fill the padded area with") - .set_support_level(2) - .add_type_rel("Pad", PadRel) - .set_attr("FInferCorrectLayout", PadInferCorrectLayout) - .set_attr("TOpPattern", kInjective) - .set_attr("FTVMCompute", PadCompute); - -// relay.nn.mirror_pad -TVM_REGISTER_NODE_TYPE(MirrorPadAttrs); - -bool MirrorPadRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - ICHECK_EQ(types.size(), 2); - const auto* data = types[0].as(); - if (data == nullptr) return false; - - const MirrorPadAttrs* param = attrs.as(); - ICHECK(param != nullptr); - - // check that pad widths match lengths - ICHECK(data->shape.size() == param->pad_width.size()) - << "There should be as many pad width pairs as shape dimensions " - << "but the shape has " << data->shape.size() << " dimensions " - << "and there are " << param->pad_width.size() << " pad width pairs."; - - // each pad width element should be a pair of positive integers - std::vector oshape; - for (size_t i = 0; i < param->pad_width.size(); i++) { - ICHECK(param->pad_width[i].size() == 2) - << "Each pad width element should be a pair but at index " << i << " there are " - << param->pad_width[i].size() << " elements."; - - auto width1 = tir::as_const_int(param->pad_width[i][0]); - auto width2 = tir::as_const_int(param->pad_width[i][1]); - ICHECK(width1 != nullptr); - ICHECK(width2 != nullptr); - - ICHECK(*width1 >= 0) << "Param width elements should be positive but first pad width at " - << "index " << i << " is " << *width1 << "."; - ICHECK(*width2 >= 0) << "Param width elements should be positive but first pad width at " - << "index " << i << " is " << *width2 << "."; - - auto padding = tir::make_const(data->shape[i].dtype(), *width1 + *width2); - oshape.push_back(data->shape[i] + padding); - } - - reporter->Assign(types[1], TensorType(Array(oshape), data->dtype)); - return true; -} - -// Handler to create a call to the padding op used by front-end FFI -Expr MakeMirrorPad(Expr data, Array> pad_width, String mode) { - auto attrs = make_object(); - attrs->mode = mode; - attrs->pad_width = std::move(pad_width); - static const Op& op = Op::Get("nn.mirror_pad"); - return Call(op, {data}, Attrs(attrs), {}); -} - -TVM_REGISTER_GLOBAL("relay.op.nn._make.mirror_pad").set_body_typed(MakeMirrorPad); - -RELAY_REGISTER_OP("nn.mirror_pad") - .describe(R"code(MirrorPad for n-D tensor. - -)code" TVM_ADD_FILELINE) - .set_attrs_type() - .set_num_inputs(1) - .add_argument("data", "Tensor", "The input tensor.") - .set_support_level(2) - .add_type_rel("MirrorPad", MirrorPadRel) - .set_attr("TOpPattern", kInjective); - -} // namespace relay -} // namespace tvm diff --git a/src/relay/op/nn/pooling.cc b/src/relay/op/nn/pooling.cc deleted file mode 100644 index 1cfbab6e661e..000000000000 --- a/src/relay/op/nn/pooling.cc +++ /dev/null @@ -1,1317 +0,0 @@ -/* - * 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 pooling.cc - * \brief Pooling operators - */ -#include "pooling.h" - -#include -#include -#include -#include -#include - -#include - -#include "../../transforms/infer_layout_utils.h" -#include "pooling_common.h" - -namespace tvm { -namespace relay { - -// relay.nn.max_pool2d & relay.nn.avg_pool2d -TVM_REGISTER_NODE_TYPE(MaxPool2DAttrs); -TVM_REGISTER_NODE_TYPE(AvgPool2DAttrs); - -template -bool Pool2DRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - ICHECK_EQ(types.size(), 2); - const auto* data = types[0].as(); - - if (data == nullptr) return false; - - const auto dshape = data->shape; - ICHECK_GE(dshape.size(), 2U) - << "Pool2D only support input >= 2-D: input must have height and width"; - const auto param = attrs.as(); - ICHECK(param != nullptr); - - Layout layout(param->layout); - ICHECK(layout.Contains(LayoutAxis::Get('H')) && layout.Contains(LayoutAxis::Get('W')) && - !layout.Contains(LayoutAxis::Get('h')) && !layout.Contains(LayoutAxis::Get('w'))) - << "Invalid layout " << layout << ". Pool2D layout must have H and W, which cannot be split"; - - const auto hidx = layout.IndexOf(LayoutAxis::Get('H')); - const auto widx = layout.IndexOf(LayoutAxis::Get('W')); - - IndexExpr pad_h, pad_w; - if (param->padding.size() == 1) { - pad_h = param->padding[0] * 2; - pad_w = param->padding[0] * 2; - } else if (param->padding.size() == 2) { - // (top, left) - pad_h = param->padding[0] * 2; - pad_w = param->padding[1] * 2; - } else if (param->padding.size() == 4) { - // (top, left, bottom, right) - pad_h = param->padding[0] + param->padding[2]; - pad_w = param->padding[1] + param->padding[3]; - } else { - return false; - } - - std::vector oshape(dshape.begin(), dshape.end()); - - if (dshape[hidx].as()) { - oshape[hidx] = dshape[hidx]; - } else { - oshape[hidx] = - calculate_pool_dimension(dshape[hidx], pad_h, param->pool_size[0], param->dilation[0], - param->strides[0], param->ceil_mode); - } - if (dshape[widx].as()) { - oshape[widx] = dshape[widx]; - } else { - oshape[widx] = - calculate_pool_dimension(dshape[widx], pad_w, param->pool_size[1], param->dilation[1], - param->strides[1], param->ceil_mode); - } - - // assign output type - reporter->Assign(types[1], TensorType(oshape, data->dtype)); - return true; -} - -template -Array Pool2DCompute(const Attrs& attrs, const Array& inputs, - const Type& out_type) { - static const Layout kNCHW("NCHW"); - const auto* param = attrs.as(); - ICHECK(param != nullptr); - auto pool_size = param->pool_size; - auto strides = param->strides; - auto dilation = param->dilation; - auto padding = param->padding; - auto ceil_mode = param->ceil_mode; - Layout layout(param->layout); - Layout out_layout(param->out_layout); - - ICHECK(tir::BijectiveLayout(layout, kNCHW).defined()) - << "max_pool2d currently only supports layouts that are convertible from NCHW"; - ICHECK_EQ(layout.IndexOf(LayoutAxis::Get('h')), -1) - << "max_pool2d does not support input split on height"; - ICHECK_EQ(layout.IndexOf(LayoutAxis::Get('w')), -1) - << "max_pool2d does not support input split on width"; - - ICHECK(inputs[0].ndim() == 4U || inputs[0].ndim() == 5U || inputs[0].ndim() == 6U) - << "Pool2D only support 4-D input (e.g., NCHW)" - << " or 5-D input (e.g. NCHWc on for vector instructions)" - << " or 6-D input (e.g. NCHWnc for tensor accelerators)"; - - if (param->padding.size() == 1) { - padding.push_back(padding[0]); - padding.push_back(padding[0]); - padding.push_back(padding[0]); - } else if (param->padding.size() == 2) { - padding.push_back(padding[0]); - padding.push_back(padding[1]); - } - if (mode == topi::nn::kAvgPool) { - bool count_include_pad = reinterpret_cast(param)->count_include_pad; - return Array{topi::nn::pool2d(inputs[0], pool_size, strides, dilation, padding, - mode, ceil_mode, layout.name(), count_include_pad)}; - } else { - return Array{topi::nn::pool2d(inputs[0], pool_size, strides, dilation, padding, - mode, ceil_mode, layout.name())}; - } -} - -TVM_REGISTER_GLOBAL("relay.op.nn._make.max_pool2d") - .set_body_typed([](Expr data, Array pool_size, Array strides, - Array dilation, Array padding, String layout, - String out_layout, bool ceil_mode) { - return MakeMaxPool(data, pool_size, strides, dilation, padding, layout, - out_layout, ceil_mode, "nn.max_pool2d"); - }); - -RELAY_REGISTER_OP("nn.max_pool2d") - .describe(R"code(Max pooling operation for two dimensional data. - -- **data**: This depends on the `layout` parameter. Input is 4D array of shape - (batch_size, channels, height, width) if `layout` is `NCHW`. -- **out**: This depends on the `layout` parameter. Output is 4D array of shape - (batch_size, channels, out_height, out_width) if `layout` is `NCHW`. - out_height and out_width are calculated as:: - - out_height = floor((height+padding[0]+padding[2]-pool_size[0])/strides[0])+1 - out_width = floor((width+padding[1]+padding[3]-pool_size[1])/strides[1])+1 - - where padding will be an expanded array based on number of values passed as:: - one int : all sides same padding used. - two int : bottom, right use same as top and left. - four int: padding width in the order of (top, left, bottom, right). - - When `ceil_mode` is `True`, ceil will be used instead of floor in this - equation. - -)code" TVM_ADD_FILELINE) - .set_attrs_type() - .set_num_inputs(1) - .add_argument("data", "Tensor", "The input tensor.") - .set_support_level(2) - .add_type_rel("MaxPool2D", Pool2DRel) - .set_attr("FInferCorrectLayout", PoolInferCorrectLayout) - .set_attr("TOpPattern", kOutEWiseFusable) - .set_attr("FTVMCompute", Pool2DCompute); - -// AvgPool2D -TVM_REGISTER_GLOBAL("relay.op.nn._make.avg_pool2d") - .set_body_typed([](Expr data, Array pool_size, Array strides, - Array dilation, Array padding, String layout, - String out_layout, bool ceil_mode, bool count_include_pad) { - return MakeAvgPool(data, pool_size, strides, dilation, padding, layout, - out_layout, ceil_mode, count_include_pad, "nn.avg_pool2d"); - }); - -RELAY_REGISTER_OP("nn.avg_pool2d") - .describe(R"code( -Average pooling operation for one dimensional data. - -- **data**: This depends on the `layout` parameter. Input is 4D array of shape - (batch_size, channels, height, width) if `layout` is `NCHW`. -- **out**: This depends on the `layout` parameter. Output is 4D array of shape - (batch_size, channels, out_height, out_width) if `layout` is `NCHW`. - out_height and out_width are calculated as:: - - out_height = floor((height+padding[0]+padding[2]-pool_size[0])/strides[0])+1 - out_width = floor((width+padding[1]+padding[3]-pool_size[1])/strides[1])+1 - - where padding will be an expanded array based on number of values passed as:: - one int : all sides same padding used. - two int : bottom, right use same as top and left. - four int: padding width in the order of (top, left, bottom, right). - - When `ceil_mode` is `True`, ceil will be used instead of floor in this - equation. - -)code" TVM_ADD_FILELINE) - .set_attrs_type() - .set_num_inputs(1) - .add_argument("data", "Tensor", "The input tensor.") - .set_support_level(2) - .add_type_rel("AvgPool2D", Pool2DRel) - .set_attr("FInferCorrectLayout", PoolInferCorrectLayout) - .set_attr("TOpPattern", kOutEWiseFusable) - .set_attr("FTVMCompute", Pool2DCompute); - -// relay.nn.global_pool_2d & relay.nn.max_pool_2d -TVM_REGISTER_NODE_TYPE(GlobalPool2DAttrs); - -bool GlobalPool2DRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - ICHECK_EQ(types.size(), 2); - const auto* data = types[0].as(); - if (data == nullptr) { - return false; - } - const auto dshape = data->shape; - ICHECK_GE(dshape.size(), 2U) - << "Pool2D only support input >= 2-D: input must have height and width"; - const auto param = attrs.as(); - ICHECK(param != nullptr); - - Layout layout(param->layout); - ICHECK(layout.Contains(LayoutAxis::Get('H')) && layout.Contains(LayoutAxis::Get('W')) && - !layout.Contains(LayoutAxis::Get('h')) && !layout.Contains(LayoutAxis::Get('w'))) - << "Invalid layout " << layout << ". Pool2D layout must have H and W, which cannot be split"; - - const auto hidx = layout.IndexOf(LayoutAxis::Get('H')); - const auto widx = layout.IndexOf(LayoutAxis::Get('W')); - Array oshape(dshape); - oshape.Set(hidx, 1); - oshape.Set(widx, 1); - - // assign output type - reporter->Assign(types[1], TensorType(oshape, data->dtype)); - return true; -} - -template -Array GlobalPool2DCompute(const Attrs& attrs, const Array& inputs, - const Type& out_type) { - static const Layout kNCHW("NCHW"); - const auto* param = attrs.as(); - ICHECK(param != nullptr); - Layout layout(param->layout); - ICHECK(tir::BijectiveLayout(layout, kNCHW).defined()) - << "global_avg_pool2d currently only supports layouts that are convertible from NCHW"; - ICHECK_EQ(layout.IndexOf(LayoutAxis::Get('h')), -1) - << "global_avg_pool2d does not support input split on height"; - ICHECK_EQ(layout.IndexOf(LayoutAxis::Get('w')), -1) - << "global_avg_pool2d does not support input split on width"; - - ICHECK(inputs[0].ndim() == 4U || inputs[0].ndim() == 5U) - << "Pool2D only support 4-D input (e.g., NCHW)" - << " or 5-D input (last dimension is a split of channel)"; - return Array{topi::nn::global_pool(inputs[0], mode, layout.name())}; -} - -Expr MakeGlobalAvgPool2D(Expr data, String layout, String out_layout) { - auto attrs = make_object(); - attrs->layout = std::move(layout); - attrs->out_layout = std::move(out_layout); - static const Op& op = Op::Get("nn.global_avg_pool2d"); - return Call(op, {data}, Attrs(attrs), {}); -} - -TVM_REGISTER_GLOBAL("relay.op.nn._make.global_avg_pool2d").set_body_typed(MakeGlobalAvgPool2D); - -// GlobalAvgPool -RELAY_REGISTER_OP("nn.global_avg_pool2d") - .describe(R"code(Global average pooling operation for 2D data. - -- **data**: This depends on the `layout` parameter. Input is 4D array of shape - (batch_size, channels, height, width) if `layout` is `NCHW`. -- **out**: This depends on the `layout` parameter. Output is 4D array of shape - (batch_size, channels, 1, 1) if `layout` is `NCHW`. - -)code" TVM_ADD_FILELINE) - .set_attrs_type() - .set_num_inputs(1) - .add_argument("data", "Tensor", "The input tensor.") - .set_support_level(2) - .add_type_rel("GlobalAvgPool2D", GlobalPool2DRel) - .set_attr("FInferCorrectLayout", PoolInferCorrectLayout) - .set_attr("TOpPattern", kOutEWiseFusable) - .set_attr("FTVMCompute", GlobalPool2DCompute); - -// GlobalMaxPool -Expr MakeGlobalMaxPool2D(Expr data, String layout, String out_layout) { - auto attrs = make_object(); - attrs->layout = std::move(layout); - attrs->out_layout = std::move(out_layout); - static const Op& op = Op::Get("nn.global_max_pool2d"); - return Call(op, {data}, Attrs(attrs), {}); -} - -TVM_REGISTER_GLOBAL("relay.op.nn._make.global_max_pool2d").set_body_typed(MakeGlobalMaxPool2D); - -RELAY_REGISTER_OP("nn.global_max_pool2d") - .describe(R"code(Global max pooling operation for 2D data. - -- **data**: This depends on the `layout` parameter. Input is 4D array of shape - (batch_size, channels, height, width) if `layout` is `NCHW`. -- **out**: This depends on the `layout` parameter. Output is 4D array of shape - (batch_size, channels, 1, 1) if `layout` is `NCHW`. - -)code" TVM_ADD_FILELINE) - .set_attrs_type() - .set_num_inputs(1) - .add_argument("data", "Tensor", "The input tensor.") - .set_support_level(2) - .add_type_rel("GlobalMaxPool2D", GlobalPool2DRel) - .set_attr("FInferCorrectLayout", PoolInferCorrectLayout) - .set_attr("TOpPattern", kOutEWiseFusable) - .set_attr("FTVMCompute", GlobalPool2DCompute); - -// relay.nn.adaptive_pool_1d -TVM_REGISTER_NODE_TYPE(AdaptivePool1DAttrs); - -bool AdaptivePool1DRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - ICHECK_EQ(types.size(), 2); - const auto* data = types[0].as(); - if (data == nullptr) { - return false; - } - const auto dshape = data->shape; - ICHECK_GE(dshape.size(), 1U) << "Pool2D only support input >= 1-D: input must have width"; - const auto* param = attrs.as(); - ICHECK(param != nullptr); - - Layout layout(param->layout); - ICHECK(layout.Contains(LayoutAxis::Get('W')) && !layout.Contains(LayoutAxis::Get('w'))) - << "Invalid layout " << layout << ". Pool1D layout must have W, which cannot be split"; - - const auto widx = layout.IndexOf(LayoutAxis::Get('W')); - Array oshape(dshape); - auto output_size = param->output_size; - ICHECK_LE(output_size.size(), 1U) << "output_size must have 1 element."; - IndexExpr output_width; - if (output_size.empty()) { - output_width = dshape[widx]; - } else { - output_width = output_size[0]; - } - - oshape.Set(widx, output_width); - - // assign output type - reporter->Assign(types[1], TensorType(oshape, data->dtype)); - return true; -} - -template -Array AdaptivePool1DCompute(const Attrs& attrs, const Array& inputs, - const Type& out_type) { - static const Layout kNCW("NCW"); - const auto* param = attrs.as(); - ICHECK(param != nullptr); - Layout layout(param->layout); - ICHECK(tir::BijectiveLayout(layout, kNCW).defined()) - << "Adaptive pool1d currently only supports layouts that are convertible from NCW"; - ICHECK_EQ(layout.IndexOf(LayoutAxis::Get('w')), -1) - << "Adaptive pool2d does not support input split on width"; - - ICHECK(inputs[0].ndim() == 3U || inputs[0].ndim() == 4U) - << "Pool1D only support 3-D input (e.g., NCW)" - << " or 4-D input (last dimension is a split of channel)"; - - auto output_size = param->output_size; - const auto widx = layout.IndexOf(LayoutAxis::Get('W')); - IndexExpr output_width; - if (output_size.empty()) { - output_width = inputs[0]->shape[widx]; - } else { - output_width = output_size[0]; - } - return Array{ - topi::nn::adaptive_pool1d(inputs[0], Array{output_width}, mode, layout.name())}; -} - -// relay.nn.adaptive_avg_pool1d -Expr MakeAdaptiveAvgPool1D(Expr data, Array output_size, String layout, - String out_layout) { - auto attrs = make_object(); - attrs->output_size = std::move(output_size); - attrs->layout = std::move(layout); - attrs->out_layout = std::move(out_layout); - static const Op& op = Op::Get("nn.adaptive_avg_pool1d"); - return Call(op, {data}, Attrs(attrs), {}); -} - -TVM_REGISTER_GLOBAL("relay.op.nn._make.adaptive_avg_pool1d").set_body_typed(MakeAdaptiveAvgPool1D); - -RELAY_REGISTER_OP("nn.adaptive_avg_pool1d") - .describe(R"code(Adaptive average pooling operation for 1D data. - -- **data**: This depends on the `layout` parameter. Input is 3D array of shape - (batch_size, channels, width) if `layout` is `NCW`. -- **output_size**: If this argument is not provided, input width will be used - as output width. - If an integer is provided for output_size, the output size is - (N x C x output_size) for any input (NCW). -- **out**: This depends on the `layout` parameter. Output is 3D array of shape - (batch_size, channels, output_width) if `layout` is `NCW`. - -)code" TVM_ADD_FILELINE) - .set_attrs_type() - .set_num_inputs(1) - .add_argument("data", "Tensor", "The input tensor.") - .set_support_level(10) - .add_type_rel("AdaptiveAvgPool1D", AdaptivePool1DRel) - .set_attr("FInferCorrectLayout", - PoolInferCorrectLayout) - .set_attr("TOpPattern", kOutEWiseFusable) - .set_attr("FTVMCompute", AdaptivePool1DCompute); - -// relay.nn.adaptive_max_pool1d -Expr MakeAdaptiveMaxPool1D(Expr data, Array output_size, String layout, - String out_layout) { - auto attrs = make_object(); - attrs->output_size = std::move(output_size); - attrs->layout = std::move(layout); - attrs->out_layout = std::move(out_layout); - static const Op& op = Op::Get("nn.adaptive_max_pool1d"); - return Call(op, {data}, Attrs(attrs), {}); -} - -TVM_REGISTER_GLOBAL("relay.op.nn._make.adaptive_max_pool1d").set_body_typed(MakeAdaptiveMaxPool1D); - -RELAY_REGISTER_OP("nn.adaptive_max_pool1d") - .describe(R"code(Adaptive max pooling operation for 1D data. - -- **data**: This depends on the `layout` parameter. Input is 3D array of shape - (batch_size, channels, width) if `layout` is `NCW`. -- **output_size**: If this argument is not provided, input width will be used - as output width. - If an integer is provided for output_size, the output size is - (N x C x output_size) for any input (NCW). -- **out**: This depends on the `layout` parameter. Output is 3D array of shape - (batch_size, channels, output_width) if `layout` is `NCW`. - -)code" TVM_ADD_FILELINE) - .set_attrs_type() - .set_num_inputs(1) - .add_argument("data", "Tensor", "The input tensor.") - .set_support_level(10) - .add_type_rel("AdaptiveMaxPool1D", AdaptivePool1DRel) - .set_attr("FInferCorrectLayout", - PoolInferCorrectLayout) - .set_attr("TOpPattern", kOutEWiseFusable) - .set_attr("FTVMCompute", AdaptivePool1DCompute); - -// relay.nn.adaptive_pool_2d -TVM_REGISTER_NODE_TYPE(AdaptivePool2DAttrs); - -bool AdaptivePool2DRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - ICHECK_EQ(types.size(), 2); - const auto* data = types[0].as(); - if (data == nullptr) { - return false; - } - const auto dshape = data->shape; - ICHECK_GE(dshape.size(), 2U) - << "Pool2D only support input >= 2-D: input must have height and width"; - const auto* param = attrs.as(); - ICHECK(param != nullptr); - - Layout layout(param->layout); - ICHECK(layout.Contains(LayoutAxis::Get('H')) && layout.Contains(LayoutAxis::Get('W')) && - !layout.Contains(LayoutAxis::Get('h')) && !layout.Contains(LayoutAxis::Get('w'))) - << "Invalid layout " << layout << ". Pool2D layout must have H and W, which cannot be split"; - - const auto hidx = layout.IndexOf(LayoutAxis::Get('H')); - const auto widx = layout.IndexOf(LayoutAxis::Get('W')); - Array oshape(dshape); - auto output_size = param->output_size; - ICHECK_LE(output_size.size(), 2U) << "output_size can have up to 2 elements."; - IndexExpr output_height, output_width; - if (output_size.empty()) { - output_height = dshape[hidx]; - output_width = dshape[widx]; - } else if (output_size.size() == 1) { - output_height = output_size[0]; - output_width = output_size[0]; - } else { - output_height = output_size[0]; - output_width = output_size[1]; - } - - oshape.Set(hidx, output_height); - oshape.Set(widx, output_width); - - // assign output type - reporter->Assign(types[1], TensorType(oshape, data->dtype)); - return true; -} - -template -Array AdaptivePool2DCompute(const Attrs& attrs, const Array& inputs, - const Type& out_type) { - static const Layout kNCHW("NCHW"); - const auto* param = attrs.as(); - ICHECK(param != nullptr); - Layout layout(param->layout); - ICHECK(tir::BijectiveLayout(layout, kNCHW).defined()) - << "Adaptive pool2d currently only supports layouts that are convertible from NCHW"; - ICHECK_EQ(layout.IndexOf(LayoutAxis::Get('h')), -1) - << "Adaptive pool2d does not support input split on height"; - ICHECK_EQ(layout.IndexOf(LayoutAxis::Get('w')), -1) - << "Adaptive pool2d does not support input split on width"; - - ICHECK(inputs[0].ndim() == 4U || inputs[0].ndim() == 5U) - << "Pool2D only support 4-D input (e.g., NCHW)" - << " or 5-D input (last dimension is a split of channel)"; - - auto output_size = param->output_size; - const auto hidx = layout.IndexOf(LayoutAxis::Get('H')); - const auto widx = layout.IndexOf(LayoutAxis::Get('W')); - IndexExpr output_height, output_width; - if (output_size.empty()) { - output_height = inputs[0]->shape[hidx]; - output_width = inputs[0]->shape[widx]; - } else if (output_size.size() == 1) { - output_height = output_size[0]; - output_width = output_size[0]; - } else { - output_height = output_size[0]; - output_width = output_size[1]; - } - return Array{topi::nn::adaptive_pool( - inputs[0], Array{output_height, output_width}, mode, layout.name())}; -} - -// relay.nn.adaptive_avg_pool2d -Expr MakeAdaptiveAvgPool2D(Expr data, Array output_size, String layout, - String out_layout) { - auto attrs = make_object(); - attrs->output_size = std::move(output_size); - attrs->layout = std::move(layout); - attrs->out_layout = std::move(out_layout); - static const Op& op = Op::Get("nn.adaptive_avg_pool2d"); - return Call(op, {data}, Attrs(attrs), {}); -} - -TVM_REGISTER_GLOBAL("relay.op.nn._make.adaptive_avg_pool2d").set_body_typed(MakeAdaptiveAvgPool2D); - -RELAY_REGISTER_OP("nn.adaptive_avg_pool2d") - .describe(R"code(Adaptive average pooling operation for 2D data. - -- **data**: This depends on the `layout` parameter. Input is 4D array of shape - (batch_size, channels, height, width) if `layout` is `NCHW`. -- **output_size**: If this argument is not provided, input height and width will be used - as output height and width. - If a single integer is provided for output_size, the output size is - (N x C x output_size x output_size) for any input (NCHW). - If a tuple of integers (height, width) are provided for output_size, - the output size is (N x C x height x width) for any input (NCHW). -- **out**: This depends on the `layout` parameter. Output is 4D array of shape - (batch_size, channels, output_height, output_width) if `layout` is `NCHW`. - -)code" TVM_ADD_FILELINE) - .set_attrs_type() - .set_num_inputs(1) - .add_argument("data", "Tensor", "The input tensor.") - .set_support_level(10) - .add_type_rel("AdaptiveAvgPool2D", AdaptivePool2DRel) - .set_attr("FInferCorrectLayout", - PoolInferCorrectLayout) - .set_attr("TOpPattern", kOutEWiseFusable) - .set_attr("FTVMCompute", AdaptivePool2DCompute); - -// relay.nn.adaptive_max_pool2d -Expr MakeAdaptiveMaxPool2D(Expr data, Array output_size, String layout, - String out_layout) { - auto attrs = make_object(); - attrs->output_size = std::move(output_size); - attrs->layout = std::move(layout); - attrs->out_layout = std::move(out_layout); - static const Op& op = Op::Get("nn.adaptive_max_pool2d"); - return Call(op, {data}, Attrs(attrs), {}); -} - -TVM_REGISTER_GLOBAL("relay.op.nn._make.adaptive_max_pool2d").set_body_typed(MakeAdaptiveMaxPool2D); - -RELAY_REGISTER_OP("nn.adaptive_max_pool2d") - .describe(R"code(Adaptive max pooling operation for 2D data. - -- **data**: This depends on the `layout` parameter. Input is 4D array of shape - (batch_size, channels, height, width) if `layout` is `NCHW`. -- **output_size**: If this argument is not provided, input height and width will be used - as output height and width. - If a single integer is provided for output_size, the output size is - (N x C x output_size x output_size) for any input (NCHW). - If a tuple of integers (height, width) are provided for output_size, - the output size is (N x C x height x width) for any input (NCHW). -- **out**: This depends on the `layout` parameter. Output is 4D array of shape - (batch_size, channels, output_height, output_width) if `layout` is `NCHW`. - -)code" TVM_ADD_FILELINE) - .set_attrs_type() - .set_num_inputs(1) - .add_argument("data", "Tensor", "The input tensor.") - .set_support_level(10) - .add_type_rel("AdaptiveMaxPool2D", AdaptivePool2DRel) - .set_attr("FInferCorrectLayout", - PoolInferCorrectLayout) - .set_attr("TOpPattern", kOutEWiseFusable) - .set_attr("FTVMCompute", AdaptivePool2DCompute); - -// relay.nn.adaptive_pool3d -TVM_REGISTER_NODE_TYPE(AdaptivePool3DAttrs); - -bool AdaptivePool3DRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - ICHECK_EQ(types.size(), 2); - const auto* data = types[0].as(); - if (data == nullptr) { - return false; - } - const auto dshape = data->shape; - ICHECK_GE(dshape.size(), 3U) - << "Pool3D only support input >= 3-D: input must have depth, height and width"; - const auto* param = attrs.as(); - ICHECK(param != nullptr); - - Layout layout(param->layout); - ICHECK(layout.Contains(LayoutAxis::Get('D')) && layout.Contains(LayoutAxis::Get('H')) && - layout.Contains(LayoutAxis::Get('W')) && !layout.Contains(LayoutAxis::Get('d')) && - !layout.Contains(LayoutAxis::Get('h')) && !layout.Contains(LayoutAxis::Get('w'))) - << "Invalid layout " << layout - << ". Pool3D layout must have D, H and W, which cannot be split"; - - const auto didx = layout.IndexOf(LayoutAxis::Get('D')); - const auto hidx = layout.IndexOf(LayoutAxis::Get('H')); - const auto widx = layout.IndexOf(LayoutAxis::Get('W')); - Array oshape(dshape); - auto output_size = param->output_size; - ICHECK_LE(output_size.size(), 3U) << "output_size can have up to 3 elements."; - IndexExpr output_depth, output_height, output_width; - if (output_size.empty()) { - output_depth = dshape[didx]; - output_height = dshape[hidx]; - output_width = dshape[widx]; - } else if (output_size.size() == 1) { - output_depth = output_size[0]; - output_height = output_size[0]; - output_width = output_size[0]; - } else { - output_depth = output_size[0]; - output_height = output_size[1]; - output_width = output_size[2]; - } - - oshape.Set(didx, output_depth); - oshape.Set(hidx, output_height); - oshape.Set(widx, output_width); - - // assign output type - reporter->Assign(types[1], TensorType(oshape, data->dtype)); - return true; -} - -template -Array AdaptivePool3DCompute(const Attrs& attrs, const Array& inputs, - const Type& out_type) { - static const Layout kNCDHW("NCDHW"); - const auto* param = attrs.as(); - ICHECK(param != nullptr); - Layout layout(param->layout); - Layout out_layout(param->out_layout); - ICHECK(tir::BijectiveLayout(layout, kNCDHW).defined()) - << "Adaptive pool3d currently only supports layouts that are convertible from NCDHW"; - ICHECK_EQ(layout.IndexOf(LayoutAxis::Get('d')), -1) - << "Adaptive pool3d does not support input split on depth"; - ICHECK_EQ(layout.IndexOf(LayoutAxis::Get('h')), -1) - << "Adaptive pool3d does not support input split on height"; - ICHECK_EQ(layout.IndexOf(LayoutAxis::Get('w')), -1) - << "Adaptive pool3d does not support input split on width"; - - ICHECK(inputs[0].ndim() == 5U || inputs[0].ndim() == 6U) - << "Pool3D only support 5-D input (e.g., NCDHW)" - << " or 6-D input (last dimension is a split of channel)"; - - auto output_size = param->output_size; - const auto didx = layout.IndexOf(LayoutAxis::Get('D')); - const auto hidx = layout.IndexOf(LayoutAxis::Get('H')); - const auto widx = layout.IndexOf(LayoutAxis::Get('W')); - IndexExpr output_depth, output_height, output_width; - if (output_size.empty()) { - output_depth = inputs[0]->shape[didx]; - output_height = inputs[0]->shape[hidx]; - output_width = inputs[0]->shape[widx]; - } else if (output_size.size() == 1) { - output_depth = output_size[0]; - output_height = output_size[0]; - output_width = output_size[0]; - } else { - output_depth = output_size[0]; - output_height = output_size[1]; - output_width = output_size[2]; - } - - auto osize = Array{output_depth, output_height, output_width}; - return Array{topi::nn::adaptive_pool3d(inputs[0], osize, mode, layout.name())}; -} - -// relay.nn.adaptive_max_pool3d -Expr MakeAdaptiveMaxPool3D(Expr data, Array output_size, String layout, - String out_layout) { - auto attrs = make_object(); - attrs->output_size = std::move(output_size); - attrs->layout = std::move(layout); - attrs->out_layout = std::move(out_layout); - static const Op& op = Op::Get("nn.adaptive_max_pool3d"); - return Call(op, {data}, Attrs(attrs), {}); -} - -TVM_REGISTER_GLOBAL("relay.op.nn._make.adaptive_max_pool3d").set_body_typed(MakeAdaptiveMaxPool3D); - -RELAY_REGISTER_OP("nn.adaptive_max_pool3d") - .describe(R"code(Adaptive max pooling operation for 3D data. - -- **data**: This depends on the `layout` parameter. Input is 5D array of shape - (batch_size, channels, depth, height, width) if `layout` is `NCDHW`. -- **output_size**: If this argument is not provided, input depth, height and width will be used - as output depth, height and width. - If a single integer is provided for output_size, the output size is - (N x C x output_size x output_size x output_size) for any input (NCDHW). - If a tuple of integers (depth, height, width) are provided for output_size, - the output size is (N x C x depth x height x width) for any input (NCDHW). -- **out**: This depends on the `layout` parameter. Output is 5D array of shape - (batch_size, channels, output_depth, output_height, output_width) if `layout` is `NCDHW`. - -)code" TVM_ADD_FILELINE) - .set_attrs_type() - .set_num_inputs(1) - .add_argument("data", "Tensor", "The input tensor.") - .set_support_level(10) - .add_type_rel("AdaptiveMaxPool3D", AdaptivePool3DRel) - .set_attr("FInferCorrectLayout", - PoolInferCorrectLayout) - .set_attr("TOpPattern", kOutEWiseFusable) - .set_attr("FTVMCompute", AdaptivePool3DCompute); - -// relay.nn.adaptive_max_pool3d -Expr MakeAdaptiveAvgPool3D(Expr data, Array output_size, String layout, - String out_layout) { - auto attrs = make_object(); - attrs->output_size = std::move(output_size); - attrs->layout = std::move(layout); - attrs->out_layout = std::move(out_layout); - static const Op& op = Op::Get("nn.adaptive_avg_pool3d"); - return Call(op, {data}, Attrs(attrs), {}); -} - -TVM_REGISTER_GLOBAL("relay.op.nn._make.adaptive_avg_pool3d").set_body_typed(MakeAdaptiveAvgPool3D); - -RELAY_REGISTER_OP("nn.adaptive_avg_pool3d") - .describe(R"code(Adaptive avg pooling operation for 3D data. -- **data**: This depends on the `layout` parameter. Input is 5D array of shape - (batch_size, channels, depth, height, width) if `layout` is `NCDHW`. -- **output_size**: If this argument is not provided, input depth, height and width will be used - as output depth, height and width. - If a single integer is provided for output_size, the output size is - (N x C x output_size x output_size x output_size) for any input (NCDHW). - If a tuple of integers (depth, height, width) are provided for output_size, - the output size is (N x C x depth x height x width) for any input (NCDHW). -- **out**: This depends on the `layout` parameter. Output is 5D array of shape - (batch_size, channels, output_depth, output_height, output_width) if `layout` is `NCDHW`. -)code" TVM_ADD_FILELINE) - .set_attrs_type() - .set_num_inputs(1) - .add_argument("data", "Tensor", "The input tensor.") - .set_support_level(10) - .add_type_rel("AdaptiveAvgPool3D", AdaptivePool3DRel) - .set_attr("FInferCorrectLayout", - PoolInferCorrectLayout) - .set_attr("TOpPattern", kOutEWiseFusable) - .set_attr("FTVMCompute", AdaptivePool3DCompute); - -bool Pool2DGradRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - ICHECK_EQ(types.size(), 3); - const auto* data = types[1].as(); - - if (data == nullptr) return false; - - // assign output type - reporter->Assign(types[2], types[1]); - return true; -} - -template -Array Pool2DGradCompute(const Attrs& attrs, const Array& inputs, - const Type& out_type) { - static const Layout kNCHW("NCHW"); - const auto* param = attrs.as(); - ICHECK(param != nullptr); - ICHECK_EQ(inputs.size(), 2); - auto pool_size = param->pool_size; - auto strides = param->strides; - auto padding = param->padding; - auto ceil_mode = param->ceil_mode; - Layout layout(param->layout); - - ICHECK(tir::BijectiveLayout(layout, kNCHW).defined()) - << "pool2d_grad currently only supports layouts that are convertible from NCHW"; - ICHECK_EQ(layout.IndexOf(LayoutAxis::Get('h')), -1) - << "pool2d_grad does not support input split on height"; - ICHECK_EQ(layout.IndexOf(LayoutAxis::Get('w')), -1) - << "pool2d_grad does not support input split on width"; - - ICHECK(inputs[0].ndim() == 4U || inputs[0].ndim() == 5U) - << "Pool2DGrad only support 4-D output gradient (e.g., NCHW)" - << " or 5-D output gradient (last dimension is a split of channel)"; - - ICHECK(inputs[1].ndim() == 4U || inputs[1].ndim() == 5U) - << "Pool2DGrad only support 4-D input (e.g., NCHW)" - << " or 5-D input (last dimension is a split of channel)"; - - if (param->padding.size() == 1) { - padding.push_back(padding[0]); - padding.push_back(padding[0]); - padding.push_back(padding[0]); - } else if (param->padding.size() == 2) { - padding.push_back(padding[0]); - padding.push_back(padding[1]); - } - if (mode == topi::nn::kAvgPool) { - bool count_include_pad = reinterpret_cast(param)->count_include_pad; - return Array{topi::nn::pool_grad(inputs[0], inputs[1], pool_size, strides, padding, - mode, ceil_mode, layout.name(), - count_include_pad)}; - } else { - return Array{topi::nn::pool_grad(inputs[0], inputs[1], pool_size, strides, padding, - mode, ceil_mode, layout.name())}; - } -} - -// MaxPool2DGrad -Expr MakeMaxPool2DGrad(Expr out_grad, Expr data, Array pool_size, - Array strides, Array padding, String layout, - String out_layout, bool ceil_mode) { - auto attrs = make_object(); - attrs->pool_size = std::move(pool_size); - attrs->strides = std::move(strides); - attrs->padding = std::move(padding); - attrs->layout = std::move(layout); - attrs->out_layout = std::move(out_layout); - attrs->ceil_mode = ceil_mode; - static const Op& op = Op::Get("nn.max_pool2d_grad"); - return Call(op, {out_grad, data}, Attrs(attrs), {}); -} - -TVM_REGISTER_GLOBAL("relay.op.nn._make.max_pool2d_grad").set_body_typed(MakeMaxPool2DGrad); - -RELAY_REGISTER_OP("nn.max_pool2d_grad") - .describe(R"code(Gradient of max pooling operation for two dimensional data. - -- **out_grad**: This depends on the `layout` parameter. Output gradient is 4D array of - shape (batch_size, channels, out_height, out_width) if `layout` is `NCHW`. - out_height and out_width are the output size of the pooling operation, - which are calculated as:: - out_height = floor((height+padding[0]+padding[2]-pool_size[0])/strides[0])+1 - out_width = floor((width+padding[1]+padding[3]-pool_size[1])/strides[1])+1 - - where padding will be an expanded array based on number of values passed as:: - one int : all sides same padding used. - two int : bottom, right use same as top and left. - four int: padding width in the order of (top, left, bottom, right). - - When `ceil_mode` is `True`, ceil will be used instead of floor in this - equation. -- **data**: This depends on the `layout` parameter. Input is 4D array of shape - (batch_size, channels, height, width) if `layout` is `NCHW`. -- **grad**: This depends on the `layout` parameter. Grad is 4D array of shape - (batch_size, channels, height, width) if `layout` is `NCHW`. - -)code" TVM_ADD_FILELINE) - .set_attrs_type() - .set_num_inputs(2) - .add_argument("data", "Tensor", "The input tensor.") - .add_argument("grad", "Tensor", "The grad tensor.") - .set_support_level(2) - .add_type_rel("MaxPool2DGrad", Pool2DGradRel) - .set_attr("TOpPattern", kOutEWiseFusable) - .set_attr("FTVMCompute", Pool2DGradCompute); - -// AvgPool2DGrad -Expr MakeAvgPool2DGrad(Expr out_grad, Expr data, Array pool_size, - Array strides, Array padding, String layout, - String out_layout, bool ceil_mode, bool count_include_pad) { - auto attrs = make_object(); - attrs->pool_size = std::move(pool_size); - attrs->strides = std::move(strides); - attrs->padding = std::move(padding); - attrs->layout = std::move(layout); - attrs->out_layout = std::move(out_layout); - attrs->ceil_mode = ceil_mode; - attrs->count_include_pad = count_include_pad; - static const Op& op = Op::Get("nn.avg_pool2d_grad"); - return Call(op, {out_grad, data}, Attrs(attrs), {}); -} - -TVM_REGISTER_GLOBAL("relay.op.nn._make.avg_pool2d_grad").set_body_typed(MakeAvgPool2DGrad); - -RELAY_REGISTER_OP("nn.avg_pool2d_grad") - .describe(R"code(Gradient of average pooling operation for two dimensional data. - -- **out_grad**: This depends on the `layout` parameter. Output gradient is 4D array of - shape (batch_size, channels, out_height, out_width) if `layout` is `NCHW`. - out_height and out_width are the output size of the pooling operation, - which are calculated as:: - out_height = floor((height+padding[0]+padding[2]-pool_size[0])/strides[0])+1 - out_width = floor((width+padding[1]+padding[3]-pool_size[1])/strides[1])+1 - - where padding will be an expanded array based on number of values passed as:: - one int : all sides same padding used. - two int : bottom, right use same as top and left. - four int: padding width in the order of (top, left, bottom, right). - - When `ceil_mode` is `True`, ceil will be used instead of floor in this - equation. -- **data**: This depends on the `layout` parameter. Input is 4D array of shape - (batch_size, channels, height, width) if `layout` is `NCHW`. -- **grad**: This depends on the `layout` parameter. Grad is 4D array of shape - (batch_size, channels, height, width) if `layout` is `NCHW`. - -)code" TVM_ADD_FILELINE) - .set_attrs_type() - .set_num_inputs(2) - .add_argument("data", "Tensor", "The input tensor.") - .add_argument("grad", "Tensor", "The grad tensor.") - .set_support_level(2) - .add_type_rel("MaxPool2DGrad", Pool2DGradRel) - .set_attr("TOpPattern", kOutEWiseFusable) - .set_attr("FTVMCompute", Pool2DGradCompute); - -// relay.nn.max_pool1d & relay.nn.avg_pool1d -TVM_REGISTER_NODE_TYPE(MaxPool1DAttrs); -TVM_REGISTER_NODE_TYPE(AvgPool1DAttrs); - -template -bool Pool1DRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - ICHECK_EQ(types.size(), 2); - const auto* data = types[0].as(); - - if (data == nullptr) return false; - - const auto dshape = data->shape; - ICHECK_GE(dshape.size(), 1U) << "Pool1D only support input >= 1-D: input must have width"; - const auto param = attrs.as(); - ICHECK(param != nullptr); - - Layout layout(param->layout); - Layout out_layout(param->out_layout); - ICHECK(layout.Contains(LayoutAxis::Get('W')) && !layout.Contains(LayoutAxis::Get('w'))) - << "Invalid layout " << layout << ". Pool1D layout must have W, which cannot be split"; - - const auto widx = layout.IndexOf(LayoutAxis::Get('W')); - - IndexExpr pad_w; - if (param->padding.size() == 1) { - pad_w = param->padding[0] * 2; - } else if (param->padding.size() == 2) { - // (left, right) - pad_w = param->padding[0] + param->padding[1]; - } else { - return false; - } - - std::vector oshape(dshape.begin(), dshape.end()); - - if (dshape[widx].as()) { - oshape[widx] = dshape[widx]; - } else { - oshape[widx] = - calculate_pool_dimension(dshape[widx], pad_w, param->pool_size[0], param->dilation[0], - param->strides[0], param->ceil_mode); - } - - // assign output type - reporter->Assign(types[1], TensorType(oshape, data->dtype)); - return true; -} - -template -Array Pool1DCompute(const Attrs& attrs, const Array& inputs, - const Type& out_type) { - static const Layout kNCW("NCW"); - const auto* param = attrs.as(); - ICHECK(param != nullptr); - auto pool_size = param->pool_size; - auto strides = param->strides; - auto dilation = param->dilation; - auto padding = param->padding; - auto ceil_mode = param->ceil_mode; - Layout layout(param->layout); - Layout out_layout(param->out_layout); - - ICHECK(tir::BijectiveLayout(layout, kNCW).defined()) - << "max_pool1d currently only supports layouts that are convertible from NCW"; - ICHECK_EQ(layout.IndexOf(LayoutAxis::Get('w')), -1) - << "max_pool1d does not support input split on width"; - - ICHECK(inputs[0].ndim() == 3U || inputs[0].ndim() == 4U || inputs[0].ndim() == 5U) - << "Pool1D only support 3-D input (e.g., NCW)" - << " or 4-D input (e.g. NCWc on for vector instructions)" - << " or 5-D input (e.g. NCWnc for tensor accelerators)"; - - if (param->padding.size() == 1) { - padding.push_back(padding[0]); - } - - if (mode == topi::nn::kAvgPool) { - bool count_include_pad = reinterpret_cast(param)->count_include_pad; - return Array{topi::nn::pool1d(inputs[0], pool_size, strides, dilation, padding, - mode, ceil_mode, layout.name(), count_include_pad)}; - } else { - return Array{topi::nn::pool1d(inputs[0], pool_size, strides, dilation, padding, - mode, ceil_mode, layout.name())}; - } -} - -TVM_REGISTER_GLOBAL("relay.op.nn._make.max_pool1d") - .set_body_typed([](Expr data, Array pool_size, Array strides, - Array dilation, Array padding, String layout, - String out_layout, bool ceil_mode) { - return MakeMaxPool(data, pool_size, strides, dilation, padding, layout, - out_layout, ceil_mode, "nn.max_pool1d"); - }); - -RELAY_REGISTER_OP("nn.max_pool1d") - .describe(R"code(Max pooling operation for one dimensional data. - -- **data**: This depends on the `layout` parameter. Input is 3D array of shape - (batch_size, channels, width) if `layout` is `NCW`. -- **out**: This depends on the `layout` parameter. Output is 3D array of shape - (batch_size, channels, , out_width) if `layout` is `NCW`. - out_width is calculated as:: - - out_width = floor((width+padding[0]+padding[1]-pool_size[0])/strides[0])+1 - - where padding will be an expanded array based on number of values passed as:: - one int : all sides same padding used. - two int: padding width in the order of (left, right). - - When `ceil_mode` is `True`, ceil will be used instead of floor in this - equation. - -)code" TVM_ADD_FILELINE) - .set_attrs_type() - .set_num_inputs(1) - .add_argument("data", "Tensor", "The input tensor.") - .set_support_level(2) - .add_type_rel("MaxPool1D", Pool1DRel) - .set_attr("FInferCorrectLayout", PoolInferCorrectLayout) - .set_attr("TOpPattern", kOutEWiseFusable) - .set_attr("FTVMCompute", Pool1DCompute); - -// AvgPool1D -TVM_REGISTER_GLOBAL("relay.op.nn._make.avg_pool1d") - .set_body_typed([](Expr data, Array pool_size, Array strides, - Array dilation, Array padding, String layout, - String out_layout, bool ceil_mode, bool count_include_pad) { - return MakeAvgPool(data, pool_size, strides, dilation, padding, layout, - out_layout, ceil_mode, count_include_pad, "nn.avg_pool1d"); - }); - -RELAY_REGISTER_OP("nn.avg_pool1d") - .describe(R"code( -Average pooling operation for one dimensional data. - -- **data**: This depends on the `layout` parameter. Input is 3D array of shape - (batch_size, channels, width) if `layout` is `NCW`. -- **out**: This depends on the `layout` parameter. Output is 3D array of shape - (batch_size, channels, out_width) if `layout` is `NCW`. - out_width is calculated as:: - - out_width = floor((width+padding[0]+padding[1]-pool_size[0])/strides[0])+1 - - where padding will be an expanded array based on number of values passed as:: - one int : all sides same padding used. - two int: padding width in the order of (left, right). - - When `ceil_mode` is `True`, ceil will be used instead of floor in this - equation. - -)code" TVM_ADD_FILELINE) - .set_attrs_type() - .set_num_inputs(1) - .add_argument("data", "Tensor", "The input tensor.") - .set_support_level(2) - .add_type_rel("AvgPool1D", Pool1DRel) - .set_attr("FInferCorrectLayout", PoolInferCorrectLayout) - .set_attr("TOpPattern", kOutEWiseFusable) - .set_attr("FTVMCompute", Pool1DCompute); - -// relay.nn.max_pool3d & relay.nn.avg_pool3d -TVM_REGISTER_NODE_TYPE(MaxPool3DAttrs); -TVM_REGISTER_NODE_TYPE(AvgPool3DAttrs); - -template -bool Pool3DRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - ICHECK_EQ(types.size(), 2); - const auto* data = types[0].as(); - - if (data == nullptr) return false; - - const auto dshape = data->shape; - ICHECK_GE(dshape.size(), 3U) - << "Pool3D only support input >= 3-D: input must have depth, height and width"; - const auto param = attrs.as(); - ICHECK(param != nullptr); - - Layout layout(param->layout); - Layout out_layout(param->out_layout); - ICHECK(layout.Contains(LayoutAxis::Get('D')) && layout.Contains(LayoutAxis::Get('H')) && - layout.Contains(LayoutAxis::Get('W')) && !layout.Contains(LayoutAxis::Get('d')) && - !layout.Contains(LayoutAxis::Get('h')) && !layout.Contains(LayoutAxis::Get('w'))) - << "Invalid layout " << layout - << ". Pool3D layout must have D, H and W, which cannot be split"; - - const auto didx = layout.IndexOf(LayoutAxis::Get('D')); - const auto hidx = layout.IndexOf(LayoutAxis::Get('H')); - const auto widx = layout.IndexOf(LayoutAxis::Get('W')); - - IndexExpr pad[3]; - if (param->padding.size() == 1) { - pad[0] = param->padding[0] * 2; - pad[1] = param->padding[0] * 2; - pad[2] = param->padding[0] * 2; - } else if (param->padding.size() == 3) { - // (front, top, left) - pad[0] = param->padding[0] * 2; - pad[1] = param->padding[1] * 2; - pad[2] = param->padding[2] * 2; - } else if (param->padding.size() == 6) { - // (front, top, left, back, bottom, right) - pad[0] = param->padding[0] + param->padding[3]; - pad[1] = param->padding[1] + param->padding[4]; - pad[2] = param->padding[2] + param->padding[5]; - } else { - return false; - } - - std::vector oshape(dshape.begin(), dshape.end()); - - int idxes[3] = {didx, hidx, widx}; - for (int i = 0; i < 3; i++) { - int ii = idxes[i]; - if (dshape[ii].as()) { - oshape[ii] = dshape[ii]; - } else { - oshape[ii] = - calculate_pool_dimension(dshape[ii], pad[i], param->pool_size[i], param->dilation[i], - param->strides[i], param->ceil_mode); - } - } - - // assign output type - reporter->Assign(types[1], TensorType(oshape, data->dtype)); - return true; -} - -template -Array Pool3DCompute(const Attrs& attrs, const Array& inputs, - const Type& out_type) { - static const Layout kNCDHW("NCDHW"); - const auto* param = attrs.as(); - ICHECK(param != nullptr); - auto pool_size = param->pool_size; - auto strides = param->strides; - auto dilation = param->dilation; - auto padding = param->padding; - auto ceil_mode = param->ceil_mode; - Layout layout(param->layout); - Layout out_layout(param->out_layout); - - ICHECK(tir::BijectiveLayout(layout, kNCDHW).defined()) - << "max_pool3d currently only supports layouts that are convertible from NCDHW"; - ICHECK_EQ(layout.IndexOf(LayoutAxis::Get('d')), -1) - << "max_pool3d does not support input split on depth"; - ICHECK_EQ(layout.IndexOf(LayoutAxis::Get('h')), -1) - << "max_pool3d does not support input split on height"; - ICHECK_EQ(layout.IndexOf(LayoutAxis::Get('w')), -1) - << "max_pool3d does not support input split on width"; - - ICHECK(inputs[0].ndim() == 4U || inputs[0].ndim() == 5U || inputs[0].ndim() == 6U) - << "Pool3D only support 5-D input (e.g., NCDHW)" - << " or 6-D input (e.g. NCDHWc on for vector instructions)" - << " or 7-D input (e.g. NCDHWnc for tensor accelerators)"; - - if (param->padding.size() == 1) { - padding.push_back(padding[0]); - padding.push_back(padding[0]); - padding.push_back(padding[0]); - } else if (param->padding.size() == 3) { - padding.push_back(padding[0]); - padding.push_back(padding[1]); - padding.push_back(padding[2]); - } - if (mode == topi::nn::kAvgPool) { - bool count_include_pad = reinterpret_cast(param)->count_include_pad; - return Array{topi::nn::pool3d(inputs[0], pool_size, strides, dilation, padding, - mode, ceil_mode, layout.name(), count_include_pad)}; - } else { - return Array{topi::nn::pool3d(inputs[0], pool_size, strides, dilation, padding, - mode, ceil_mode, layout.name())}; - } -} - -TVM_REGISTER_GLOBAL("relay.op.nn._make.max_pool3d") - .set_body_typed([](Expr data, Array pool_size, Array strides, - Array dilation, Array padding, String layout, - String out_layout, bool ceil_mode) { - return MakeMaxPool(data, pool_size, strides, dilation, padding, layout, - out_layout, ceil_mode, "nn.max_pool3d"); - }); - -RELAY_REGISTER_OP("nn.max_pool3d") - .describe(R"code(Max pooling operation for three dimensional data. - -- **data**: This depends on the `layout` parameter. Input is 5D array of shape - (batch_size, channels, depth, height, width) if `layout` is `NCDHW`. -- **out**: This depends on the `layout` parameter. Output is 5D array of shape - (batch_size, channels, out_depth, out_height, out_width) if `layout` is `NCDHW`. - out_depth, out_height and out_width are calculated as:: - - out_depth = floor((depth+padding[0]+padding[3]-pool_size[0])/strides[0])+1 - out_height = floor((height+padding[1]+padding[4]-pool_size[1])/strides[1])+1 - out_width = floor((width+padding[2]+padding[5]-pool_size[2])/strides[2])+1 - - where padding will be an expanded array based on number of values passed as:: - one int : all sides same padding used. - three int : front, bottom, right use same as back, top and left. - six int: padding width in the order of (front, top, left, back, bottom, right). - - When `ceil_mode` is `True`, ceil will be used instead of floor in this - equation. - -)code" TVM_ADD_FILELINE) - .set_attrs_type() - .set_num_inputs(1) - .add_argument("data", "Tensor", "The input tensor.") - .set_support_level(2) - .add_type_rel("MaxPool3D", Pool3DRel) - .set_attr("FInferCorrectLayout", PoolInferCorrectLayout) - .set_attr("TOpPattern", kOutEWiseFusable) - .set_attr("FTVMCompute", Pool3DCompute); - -// AvgPool3D -TVM_REGISTER_GLOBAL("relay.op.nn._make.avg_pool3d") - .set_body_typed([](Expr data, Array pool_size, Array strides, - Array dilation, Array padding, String layout, - String out_layout, bool ceil_mode, bool count_include_pad) { - return MakeAvgPool(data, pool_size, strides, dilation, padding, layout, - out_layout, ceil_mode, count_include_pad, "nn.avg_pool3d"); - }); - -RELAY_REGISTER_OP("nn.avg_pool3d") - .describe(R"code( -Average pooling operation for three dimensional data. - -- **data**: This depends on the `layout` parameter. Input is 5D array of shape - (batch_size, channels, depth, height, width) if `layout` is `NCDHW`. -- **out**: This depends on the `layout` parameter. Output is 5D array of shape - (batch_size, channels, out_depth, out_height, out_width) if `layout` is `NCDHW`. - out_depth, out_height and out_width are calculated as:: - - out_depth = floor((depth+padding[0]+padding[3]-pool_size[0])/strides[0])+1 - out_height = floor((height+padding[1]+padding[4]-pool_size[1])/strides[1])+1 - out_width = floor((width+padding[2]+padding[5]-pool_size[2])/strides[2])+1 - - where padding will be an expanded array based on number of values passed as:: - one int : all sides same padding used. - three int : front, bottom, right use same as back, top and left. - six int: padding width in the order of (front, top, left, back, bottom, right). - - When `ceil_mode` is `True`, ceil will be used instead of floor in this - equation. - -)code" TVM_ADD_FILELINE) - .set_attrs_type() - .set_num_inputs(1) - .add_argument("data", "Tensor", "The input tensor.") - .set_support_level(2) - .add_type_rel("AvgPool3D", Pool3DRel) - .set_attr("FInferCorrectLayout", PoolInferCorrectLayout) - .set_attr("TOpPattern", kOutEWiseFusable) - .set_attr("FTVMCompute", Pool3DCompute); - -} // namespace relay -} // namespace tvm diff --git a/src/relay/op/nn/pooling.h b/src/relay/op/nn/pooling.h deleted file mode 100644 index 123cfcd07570..000000000000 --- a/src/relay/op/nn/pooling.h +++ /dev/null @@ -1,70 +0,0 @@ -/* - * 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 src/relay/op/nn/pooling.h - * \brief utilities for creating pool ops - */ -#ifndef TVM_RELAY_OP_NN_POOLING_H_ -#define TVM_RELAY_OP_NN_POOLING_H_ - -#include -#include - -#include - -namespace tvm { -namespace relay { - -template -inline Expr MakeMaxPool(Expr data, Array pool_size, Array strides, - Array dilation, Array padding, String layout, - String out_layout, bool ceil_mode, String op_name) { - auto attrs = make_object(); - attrs->pool_size = std::move(pool_size); - attrs->strides = std::move(strides); - attrs->dilation = std::move(dilation); - attrs->padding = std::move(padding); - attrs->layout = std::move(layout); - attrs->out_layout = std::move(out_layout); - attrs->ceil_mode = ceil_mode; - static const Op& op = Op::Get(op_name); - return Call(op, {data}, Attrs(attrs), {}); -} - -template -inline Expr MakeAvgPool(Expr data, Array pool_size, Array strides, - Array dilation, Array padding, String layout, - String out_layout, bool ceil_mode, bool count_include_pad, String op_name) { - auto attrs = make_object(); - attrs->pool_size = std::move(pool_size); - attrs->strides = std::move(strides); - attrs->dilation = std::move(dilation); - attrs->padding = std::move(padding); - attrs->layout = std::move(layout); - attrs->out_layout = std::move(out_layout); - attrs->ceil_mode = ceil_mode; - attrs->count_include_pad = count_include_pad; - static const Op& op = Op::Get(op_name); - return Call(op, {data}, Attrs(attrs), {}); -} - -} // namespace relay -} // namespace tvm -#endif // TVM_RELAY_OP_NN_POOLING_H_ diff --git a/src/relay/op/nn/pooling_common.h b/src/relay/op/nn/pooling_common.h deleted file mode 100644 index 1193d36ebe88..000000000000 --- a/src/relay/op/nn/pooling_common.h +++ /dev/null @@ -1,78 +0,0 @@ -/* - * 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 src/relay/op/nn/pooling_common.h - * \brief Common functions for pooling operator definition. - */ -#ifndef TVM_RELAY_OP_NN_POOLING_COMMON_H_ -#define TVM_RELAY_OP_NN_POOLING_COMMON_H_ - -#include -#include -#include - -#include -#include -#include - -#include "../op_common.h" - -namespace tvm { -namespace relay { - -inline IndexExpr calculate_pool_dimension(IndexExpr in_dimension, IndexExpr pad_amount, - IndexExpr pool_size, IndexExpr dilation, - IndexExpr stride_size, bool ceil_mode) { - IndexExpr numerator = in_dimension + pad_amount - ((pool_size - 1) * dilation + 1); - IndexExpr denominator = stride_size; - - // Emulate the behavior of running ceil on numerator / denominator rather than floor - if (ceil_mode) { - numerator += denominator - 1; - } - - return numerator / denominator + 1; -} - -template -InferCorrectLayoutOutput PoolInferCorrectLayout(const Attrs& attrs, - const Array& new_in_layouts, - const Array& old_in_layouts, - const Array& old_in_types) { - const auto* attrs_ptr = attrs.as(); - ICHECK(attrs_ptr); - ObjectPtr params = make_object(*attrs_ptr); - - if (params->out_layout != "") { - // when users specify the out_layout of pooling, follow user's preference - ICHECK_EQ(params->layout, params->out_layout) - << "Pooling input/output layouts mismatch: " << params->layout << " vs. " - << params->out_layout; - } else if (new_in_layouts.defined()) { - // the pooling is using an inferred layout (i.e., new_in_layouts[0]) given by relay caller - // ICHECK_EQ(new_in_layouts.size(), 1); - params->layout = new_in_layouts[0].name(); - } - - return InferCorrectLayoutOutput({params->layout}, {params->layout}, Attrs(params)); -} -} // namespace relay -} // namespace tvm -#endif // TVM_RELAY_OP_NN_POOLING_COMMON_H_ diff --git a/src/relay/op/nn/sparse.cc b/src/relay/op/nn/sparse.cc deleted file mode 100644 index 60c03895da46..000000000000 --- a/src/relay/op/nn/sparse.cc +++ /dev/null @@ -1,309 +0,0 @@ -/* - * 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 sparse.cc - * \brief Property def of nn.sparse_dense operator. - */ - -#include -#include -#include - -#include -#include - -#include "../../transforms/infer_layout_utils.h" - -namespace tvm { -namespace relay { - -// relay.nn.sparse_dense -TVM_REGISTER_NODE_TYPE(SparseDenseAttrs); - -bool SparseDenseRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - ICHECK_EQ(types.size(), 5); - const auto* param = attrs.as(); - ICHECK(param != nullptr); - - if (param->sparse_lhs) { - const auto* weight = types[0].as(); - const auto* data_data = types[1].as(); - ICHECK(data_data->shape.size() == 1 || data_data->shape.size() == 3); - const auto* data_indptr = types[3].as(); - if (weight == nullptr) return false; - - if (data_data->shape.size() == 1) { - // CSR case. - Array oshape({data_indptr->shape[0] - 1, weight->shape[0]}); - reporter->Assign(types[4], TensorType(oshape, weight->dtype)); - return true; - } - - if (data_data->shape.size() == 3) { - // BSR case. - Array oshape( - {(data_indptr->shape[0] - 1) * data_data->shape[1], weight->shape[0]}); - reporter->Assign(types[4], TensorType(oshape, weight->dtype)); - return true; - } - LOG(FATAL) << "Unknown data ndim for nn.sparse_dense, should be 1 (CSR) or 3 (BSR)"; - - } else { - const auto* data = types[0].as(); - const auto* weight_data = types[1].as(); - ICHECK(weight_data->shape.size() == 1 || weight_data->shape.size() == 3); - const auto* weight_indptr = types[3].as(); - if (data == nullptr) return false; - - if (weight_data->shape.size() == 1) { - // CSR case. - Array oshape({data->shape[0], weight_indptr->shape[0] - 1}); - reporter->Assign(types[4], TensorType(oshape, data->dtype)); - return true; - } - - if (weight_data->shape.size() == 3) { - // BSR case. - Array oshape( - {data->shape[0], (weight_indptr->shape[0] - 1) * weight_data->shape[1]}); - reporter->Assign(types[4], TensorType(oshape, data->dtype)); - return true; - } - LOG(FATAL) << "Unknown weight ndim for nn.sparse_dense, should be 1 (CSR) or 3 (BSR)"; - } -} - -// Positional relay function to create dense operator used by frontend FFI. -Expr MakeSparseDense(Expr data, Expr weight_data, Expr weight_indices, Expr weight_indptr, - bool sparse_lhs) { - auto attrs = make_object(); - attrs->sparse_lhs = std::move(sparse_lhs); - static const Op& op = Op::Get("nn.sparse_dense"); - return Call(op, {data, weight_data, weight_indices, weight_indptr}, Attrs(attrs), {}); -} - -TVM_REGISTER_GLOBAL("relay.op.nn._make.sparse_dense").set_body_typed(MakeSparseDense); - -RELAY_REGISTER_OP("nn.sparse_dense") - .describe( - R"code(Applies a sparse linear transformation: :math:`Y = XW^T` with either X or W sparse. - -- **data**: `(x1, x2, ..., xn, input_dim)` -- **weight**: `(units, input_dim)` -- **out**: `(x1, x2, ..., xn, units)`. - -)code" TVM_ADD_FILELINE) - .set_attrs_type() - .set_num_inputs(4) - .add_argument("dense_data", "nD Tensor", "Input dense data.") - .add_argument("sparse_data", "1D or 3D Tensor", "Sparse data matrix.") - .add_argument("sparse_indices", "1D Tensor", "Sparse indices matrix.") - .add_argument("sparse_indptr", "1D Tensor", "Sparse indptr matrix.") - .set_support_level(1) - .add_type_rel("SparseDense", SparseDenseRel) - .set_attr("TOpPattern", kOutEWiseFusable); - -Expr MakeSparseDensePadded(Expr data, Expr weight_data, Expr weight_indices, Expr weight_indptr) { - auto attrs = make_object(); - static const Op& op = Op::Get("nn.internal.sparse_dense_padded"); - return Call(op, {data, weight_data, weight_indices, weight_indptr}, Attrs(attrs), {}); -} - -TVM_REGISTER_GLOBAL("relay.op.nn._make.sparse_dense_padded").set_body_typed(MakeSparseDensePadded); - -RELAY_REGISTER_OP("nn.internal.sparse_dense_padded") - .describe( - R"code(Applies a sparse linear transformation: :math:`Y = XW^T` with W -sparse. This variation uses a matrix with row lengths padded to a -multiple of 32 for better GPU performance. - -This op should not be directly used by a user. Instead, use `sparse_dense` -which will be converted to this op when running on the GPU. - -- **data**: `(x1, x2, ..., xn, input_dim)` -- **weight**: `(units, input_dim)` -- **out**: `(x1, x2, ..., xn, units)`. - -)code" TVM_ADD_FILELINE) - .set_attrs_type() - .set_num_inputs(4) - .add_argument("data", "nD Tensor", "Input data.") - .add_argument("weight_data", "1D Tensor", "Weight data matrix.") - .add_argument("weight_indices", "1D Tensor", "Weight indices matrix.") - .add_argument("weight_indptr", "1D Tensor", "Weight indptr matrix.") - .set_support_level(1) - .add_type_rel("SparseDense", SparseDenseRel) - .set_attr("TOpPattern", kOutEWiseFusable); - -// relay.nn.sparse_transpose -TVM_REGISTER_NODE_TYPE(SparseTransposeAttrs); - -bool SparseTransposeRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - ICHECK_EQ(types.size(), 4); - const auto* sparse_data = types[0].as(); - ICHECK_EQ(sparse_data->shape.size(), 1); - const auto* sparse_indices = types[1].as(); - ICHECK_EQ(sparse_indices->shape.size(), 1); - const auto* sparse_indptr = types[2].as(); - - std::vector output_types; - output_types.push_back(TensorType(sparse_data->shape, sparse_data->dtype)); - output_types.push_back(TensorType(sparse_indices->shape, sparse_indices->dtype)); - output_types.push_back(TensorType(sparse_indptr->shape, sparse_indptr->dtype)); - - reporter->Assign(types[3], TupleType(Array(output_types))); - return true; -} - -Expr MakeSparseTranspose(Expr sparse_data, Expr sparse_indices, Expr sparse_indptr) { - auto attrs = make_object(); - static const Op& op = Op::Get("nn.sparse_transpose"); - return Call(op, {sparse_data, sparse_indices, sparse_indptr}, Attrs(attrs), {}); -} - -TVM_REGISTER_GLOBAL("relay.op.nn._make.sparse_transpose").set_body_typed(MakeSparseTranspose); - -RELAY_REGISTER_OP("nn.sparse_transpose") - .describe(R"code(Transpose a sparse matrix X. Only support square sparse matrix - -- **input**: `(N, N)` -- **out**: `(N, N)`. - -)code" TVM_ADD_FILELINE) - .set_attrs_type() - .set_num_inputs(3) - .add_argument("sparse_data", "1D Tensor", "Sparse data matrix.") - .add_argument("sparse_indices", "1D Tensor", "Sparse indices matrix.") - .add_argument("sparse_indptr", "1D Tensor", "Sparse index pointer matrix.") - .set_support_level(1) - .add_type_rel("SparseTranspose", SparseTransposeRel) - .set_attr("TOpPattern", kOutEWiseFusable); - -// relay.nn.sparse_add -bool SparseAddRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - ICHECK_EQ(types.size(), 5) << "expecting 4 inputs and 1 output."; - const auto* dense_data = types[0].as(); - const auto* sparse_data = types[1].as(); - ICHECK(reporter->Assert(sparse_data->dtype == dense_data->dtype)) - << "sparse tensor and dense tensor datatype should match."; - ICHECK(reporter->Assert(sparse_data->shape.size() == 1)) << "sparse data tensor should be 1D."; - const auto* sparse_indices = types[2].as(); - ICHECK(reporter->Assert(sparse_indices->shape.size() == 1)) - << "sparse indices tensor should be 1D."; - - reporter->Assign(types[4], TensorType(dense_data->shape, dense_data->dtype)); - return true; -} - -Expr MakeSparseAdd(Expr dense_data, Expr sparse_data, Expr sparse_indices, Expr sparse_indptr) { - static const Op& op = Op::Get("nn.sparse_add"); - return Call(op, {dense_data, sparse_data, sparse_indices, sparse_indptr}, Attrs(), {}); -} - -TVM_REGISTER_GLOBAL("relay.op.nn._make.sparse_add").set_body_typed(MakeSparseAdd); - -RELAY_REGISTER_OP("nn.sparse_add") - .describe(R"code(Add a dense matrix X with sparse matrix Y. - -- **dense**: `(M, N)` -- **sparse**: `(M, N)` - -- **out**: `(M, N)`. - -)code" TVM_ADD_FILELINE) - .set_num_inputs(4) - .add_argument("dense_data", "2D Tensor", "Dense data matrix.") - .add_argument("sparse_data", "1D Tensor", "Sparse data vector.") - .add_argument("sparse_indices", "1D Tensor", "Sparse indices vector.") - .add_argument("sparse_indptr", "1D Tensor", "Sparse index pointer vector.") - .set_support_level(1) - .add_type_rel("SparseAdd", SparseAddRel) - .set_attr("TOpPattern", kOpaque); - -TVM_REGISTER_NODE_TYPE(SparseConv2DAttrs); - -bool SparseConv2dRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - ICHECK_EQ(types.size(), 5); - const auto* param = attrs.as(); - ICHECK(param != nullptr); - - const auto* data = types[0].as(); - const auto* weight_data = types[1].as(); - ICHECK(weight_data->shape.size() == 1 || weight_data->shape.size() == 2 || - weight_data->shape.size() == 3); - const auto* weight_indptr = types[3].as(); - if (data == nullptr) return false; - - if (weight_data->shape.size() == 2 || weight_data->shape.size() == 3) { - // BSR case. - if (param->layout == "NHWC") { - Array oshape({data->shape[0], data->shape[1], data->shape[2], - (weight_indptr->shape[0] - 1) * weight_data->shape[1]}); - reporter->Assign(types[4], TensorType(oshape, data->dtype)); - return true; - } else if (param->layout == "NCHW") { - Array oshape({data->shape[0], - (weight_indptr->shape[0] - 1) * weight_data->shape[1], - data->shape[2], data->shape[3]}); - reporter->Assign(types[4], TensorType(oshape, data->dtype)); - return true; - } - } - LOG(FATAL) << "Unknown weight ndim " << weight_data->shape.size() - << " for nn.sparse_conv2d, should be 2 or 3 (BSR)"; - return false; -} - -Expr MakeSparseConv2d(Expr data, Expr weight_data, Expr weight_indices, Expr weight_indptr, - std::string layout, Array kernel_size) { - static const Op& op = Op::Get("nn.sparse_conv2d"); - auto attrs = make_object(); - attrs->layout = std::move(layout); - attrs->kernel_size = std::move(kernel_size); - return Call(op, {data, weight_data, weight_indices, weight_indptr}, Attrs(attrs), {}); -} - -TVM_REGISTER_GLOBAL("relay.op.nn._make.sparse_conv2d").set_body_typed(MakeSparseConv2d); - -RELAY_REGISTER_OP("nn.sparse_conv2d") - .describe( - R"code(Applies a sparse convolution :math:`Y = X*W^T` with W sparse. - -- **data**: `(x1, x2, ..., xn, input_dim)` -- **weight**: `(units, input_dim)` -- **out**: `(x1, x2, ..., xn, units)`. - -)code" TVM_ADD_FILELINE) - .set_attrs_type() - .set_num_inputs(4) - .add_argument("dense_data", "nD Tensor", "Input dense data.") - .add_argument("sparse_data", "1D or 3D Tensor", "Sparse data matrix.") - .add_argument("sparse_indices", "1D Tensor", "Sparse indices matrix.") - .add_argument("sparse_indptr", "1D Tensor", "Sparse indptr matrix.") - .set_support_level(1) - .add_type_rel("SparseConv2d", SparseConv2dRel) - .set_attr("TOpPattern", kOutEWiseFusable); - -} // namespace relay -} // namespace tvm diff --git a/src/relay/op/nn/upsampling.cc b/src/relay/op/nn/upsampling.cc deleted file mode 100644 index 3b0139b16b1b..000000000000 --- a/src/relay/op/nn/upsampling.cc +++ /dev/null @@ -1,181 +0,0 @@ -/* - * 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 upsampling.cc - * \brief upsampling operator - */ - -#include "upsampling.h" - -#include -#include -#include -#include - -#include -#include - -#include "../op_common.h" - -namespace tvm { -namespace relay { - -TVM_REGISTER_NODE_TYPE(UpSamplingAttrs); -TVM_REGISTER_NODE_TYPE(UpSampling3DAttrs); - -bool UpSamplingRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - ICHECK_EQ(types.size(), 2); - const auto* data = types[0].as(); - if (data == nullptr) return false; - - static const Layout kNCHW("NCHW"); - - const UpSamplingAttrs* param = attrs.as(); - ICHECK(param != nullptr); - const Layout in_layout(param->layout); - - auto layout_converter = tir::BijectiveLayout(in_layout, kNCHW); - ICHECK(layout_converter.defined()) - << "UpSampling only support input layouts that are convertible from NCHW." - << " But got " << in_layout; - - auto oshape = layout_converter.ForwardShape(data->shape); - oshape.Set(2, tir::Cast(oshape[2].dtype(), tvm::round(oshape[2] * param->scale_h))); - oshape.Set(3, tir::Cast(oshape[3].dtype(), tvm::round(oshape[3] * param->scale_w))); - - // assign output type - reporter->Assign(types[1], TensorType(layout_converter.BackwardShape(oshape), data->dtype)); - return true; -} - -// Positional relay function to create upsampling operator -// used by frontend FFI. -Expr MakeUpSampling(Expr data, double scale_h, double scale_w, String layout, String method, - bool align_corners) { - auto attrs = make_object(); - attrs->layout = std::move(layout); - attrs->method = std::move(method); - attrs->scale_h = scale_h; - attrs->scale_w = scale_w; - attrs->align_corners = align_corners; - static const Op& op = Op::Get("nn.upsampling"); - return Call(op, {data}, Attrs(attrs), {}); -} - -TVM_REGISTER_GLOBAL("relay.op.nn._make.upsampling").set_body_typed(MakeUpSampling); - -RELAY_REGISTER_OP("nn.upsampling") - .describe( - R"code(Perform upsampling on input array with nearest neighbour or bilinear interpolation. - -- **data**: data is 4D array of shape - (batch_size, channels, in_height, in_width) for NCHW - (batch_size, in_height, in_width, channels) for NHWC - -- **out**: Output is 4D array of shape - for layout NCHW - (batch_size, channels, in_height*scale, in_width*scale) - - for layout NHWC - (batch_size, in_height*scale, in_width*scale, channels) - -)code" TVM_ADD_FILELINE) - .set_attrs_type() - .set_num_inputs(1) - .add_argument("data", "Tensor", "The input tensor.") - .set_support_level(2) - .add_type_rel("UpSampling", UpSamplingRel) - .set_attr("FInferCorrectLayout", - UpsamplingInferCorrectLayout) - .set_attr("TOpPattern", kInjective); - -// UpSampling3D -bool UpSampling3DRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - ICHECK_EQ(types.size(), 2); - const auto* data = types[0].as(); - if (data == nullptr) return false; - - static const Layout kNCDHW("NCDHW"); - - const UpSampling3DAttrs* param = attrs.as(); - ICHECK(param != nullptr); - const Layout in_layout(param->layout); - - auto layout_converter = tir::BijectiveLayout(in_layout, kNCDHW); - ICHECK(layout_converter.defined()) - << "UpSampling3D only support input layouts that are convertible from NCDHW." - << " But got " << in_layout; - - auto oshape = layout_converter.ForwardShape(data->shape); - oshape.Set(2, tir::Cast(oshape[2].dtype(), tvm::round(oshape[2] * param->scale_d))); - oshape.Set(3, tir::Cast(oshape[3].dtype(), tvm::round(oshape[3] * param->scale_h))); - oshape.Set(4, tir::Cast(oshape[4].dtype(), tvm::round(oshape[4] * param->scale_w))); - - // assign output type - reporter->Assign(types[1], TensorType(layout_converter.BackwardShape(oshape), data->dtype)); - return true; -} - -// Positional relay function to create upsampling3d operator -// used by frontend FFI. -Expr MakeUpSampling3D(Expr data, double scale_d, double scale_h, double scale_w, String layout, - String method, String coordinate_transformation_mode) { - auto attrs = make_object(); - attrs->layout = std::move(layout); - attrs->method = std::move(method); - attrs->scale_d = scale_d; - attrs->scale_h = scale_h; - attrs->scale_w = scale_w; - attrs->coordinate_transformation_mode = coordinate_transformation_mode; - static const Op& op = Op::Get("nn.upsampling3d"); - return Call(op, {data}, Attrs(attrs), {}); -} - -TVM_REGISTER_GLOBAL("relay.op.nn._make.upsampling3d").set_body_typed(MakeUpSampling3D); - -RELAY_REGISTER_OP("nn.upsampling3d") - .describe(R"code(Perform upsampling on input array with nearest neighbour or -bilinear interpolation. - -- **data**: data is 5D array of shape - (batch_size, channels, in_depth, in_height, in_width) for NCDHW - (batch_size, in_depth, in_height, in_width, channels) for NDHWC - -- **out**: Output is 5D array of shape - for layout NCDHW - (batch_size, channels, in_depth*scale, in_height*scale, in_width*scale) - - for layout NDHWC - (batch_size, in_depth*scale, in_height*scale, in_width*scale, channels) - -)code" TVM_ADD_FILELINE) - .set_attrs_type() - .set_num_inputs(1) - .add_argument("data", "Tensor", "The input tensor.") - .set_support_level(2) - .add_type_rel("UpSampling3D", UpSampling3DRel) - .set_attr("FInferCorrectLayout", - UpsamplingInferCorrectLayout) - .set_attr("TOpPattern", kInjective); - -} // namespace relay -} // namespace tvm diff --git a/src/relay/op/nn/upsampling.h b/src/relay/op/nn/upsampling.h deleted file mode 100644 index f756eddb287d..000000000000 --- a/src/relay/op/nn/upsampling.h +++ /dev/null @@ -1,67 +0,0 @@ -/* - * 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 src/relay/op/nn/upsampling.h - * \brief implementation of the InferCorrectLayout pass for upsampling - */ - -#ifndef TVM_RELAY_OP_NN_UPSAMPLING_H_ -#define TVM_RELAY_OP_NN_UPSAMPLING_H_ - -#include -#include - -#include "../op_common.h" - -namespace tvm { -namespace relay { - -template -InferCorrectLayoutOutput UpsamplingInferCorrectLayout(const Attrs& attrs, - const Array& new_in_layouts, - const Array& old_in_layouts, - const Array& old_in_types) { - const auto* attrs_ptr = attrs.as(); - ICHECK(attrs_ptr); - ObjectPtr params = make_object(*attrs_ptr); - - if (new_in_layouts.defined()) { - ICHECK_EQ(new_in_layouts.size(), 1); - - Layout raw_layout(params->layout); - Layout input = new_in_layouts[0]; - if (input.IndexOf(LayoutAxis::Get('W')) == raw_layout.IndexOf(LayoutAxis::Get('W')) && - input.IndexOf(LayoutAxis::Get('H')) == raw_layout.IndexOf(LayoutAxis::Get('H')) && - !input.Contains(LayoutAxis::Get('w')) && !input.Contains(LayoutAxis::Get('h')) && - (input.IndexOf(LayoutAxis::Get('D')) == -1 || - (input.IndexOf(LayoutAxis::Get('D')) == raw_layout.IndexOf(LayoutAxis::Get('D')) && - !input.Contains(LayoutAxis::Get('d'))))) { - params->layout = input.name(); // modify self to follow the input layout - } - } - - return InferCorrectLayoutOutput({params->layout}, {params->layout}, Attrs(params)); -} - -} // namespace relay -} // namespace tvm - -#endif // TVM_RELAY_OP_NN_UPSAMPLING_H_ diff --git a/src/relay/op/op_common.h b/src/relay/op/op_common.h deleted file mode 100644 index 6c2c6b2cce69..000000000000 --- a/src/relay/op/op_common.h +++ /dev/null @@ -1,198 +0,0 @@ -/* - * 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 op_common.h - * \brief A set of utilities and common functionality - * for relay ops. - */ -#ifndef TVM_RELAY_OP_OP_COMMON_H_ -#define TVM_RELAY_OP_OP_COMMON_H_ - -#include -#include -#include - -#include -#include -#include - -#include "../transforms/infer_layout_utils.h" -#include "type_relations.h" - -namespace tvm { -namespace relay { - -/*! Quick helper macro - * - Expose a positional make function to construct the node. - * - Register op to the registry. - * - * We make the decision to always only expose positional argument. - * We will do rewrapping in the frontend to support language - * sugars such as keyword arguments and default value. - - * \param OpName the name of registry. - */ -#define RELAY_REGISTER_UNARY_OP(OpName) \ - TVM_REGISTER_GLOBAL("relay.op._make." OpName).set_body_typed([](Expr data) { \ - static const Op& op = Op::Get(OpName); \ - return Call(op, {data}, Attrs(), {}); \ - }); \ - RELAY_REGISTER_OP(OpName) \ - .set_num_inputs(1) \ - .add_argument("data", "Tensor", "The input tensor.") \ - .add_type_rel("Identity", IdentityRel) \ - .set_attr("TOpPattern", kElemWise) \ - .set_attr("TOpIsStateful", false) \ - .set_attr("FInferCorrectLayout", ElemwiseArbitraryLayout) - -/*! Quick helper macro - * - Expose a positional make function to construct the node. - * - Register op to the registry. - * - * We make the decision to always only expose positional argument. - * We will do rewrapping in the frontend to support language - * sugars such as keyword arguments and default value. - * - * \param OpName the name of registry. - */ -#define RELAY_REGISTER_BINARY_OP(OpName) \ - TVM_REGISTER_GLOBAL("relay.op._make." OpName).set_body_typed([](Expr lhs, Expr rhs) { \ - static const Op& op = Op::Get(OpName); \ - return Call(op, {lhs, rhs}, Attrs(), {}); \ - }); \ - RELAY_REGISTER_OP(OpName) \ - .set_num_inputs(2) \ - .add_argument("lhs", "Tensor", "The left hand side tensor.") \ - .add_argument("rhs", "Tensor", "The right hand side tensor.") \ - .add_type_rel("Broadcast", BroadcastRel) \ - .set_attr("TOpPattern", kBroadcast) \ - .set_attr("TOpIsStateful", false) \ - .set_attr("FInferCorrectLayout", BinaryBroadcastLayout) - -// Comparisons -#define RELAY_REGISTER_CMP_OP(OpName) \ - TVM_REGISTER_GLOBAL("relay.op._make." OpName).set_body_typed([](Expr lhs, Expr rhs) { \ - static const Op& op = Op::Get(OpName); \ - return Call(op, {lhs, rhs}, Attrs(), {}); \ - }); \ - RELAY_REGISTER_OP(OpName) \ - .set_num_inputs(2) \ - .add_argument("lhs", "Tensor", "The left hand side tensor.") \ - .add_argument("rhs", "Tensor", "The right hand side tensor.") \ - .add_type_rel("BroadcastComp", BroadcastCompRel) \ - .set_attr("TOpPattern", kBroadcast) \ - .set_attr("TOpIsStateful", false) \ - .set_attr("FInferCorrectLayout", BinaryBroadcastLayout) - -/*! \brief A helper class for matching and rewriting operators. */ -template -class OpMatch { - public: - using MatchFunc = - std::function& args, const Attrs& attrs, const Array& type_args)>; - - /*! \brief Match an operator with the given name. - * \param op_name The name of the operator to match. - * \param func The function to execute when it matches. - * \return A self-reference for builder style API. - */ - inline OpMatch& Match(const std::string& op_name, MatchFunc func) { - auto op = Op::Get(op_name); - match_map_.insert({op, func}); - return *this; - } - - /*! \brief Rewrite a call operation based on the operator and the registered - * match functions. - * \param call The call to rewrite. - * \return The result of rewriting. - */ - inline R operator()(const Call& call) { - auto it = match_map_.find(Downcast(call->op)); - if (it != match_map_.end()) { - return it->second(call->args, call->attrs, call->type_args); - } else { - if (default_ != nullptr) { - return default_(call->args, call->attrs, call->type_args); - } else { - LOG(FATAL) << "unexpected operation " << call->op; - } - } - } - - private: - /*! \brief The match function map. */ - std::unordered_map match_map_; - /*! \brief An optional default case. */ - MatchFunc default_; -}; - -/*! \brief A utility function to get padding width from a 1 or 2 ints tuple. */ -inline void GetPaddingWidth(const Array& padding, IndexExpr* pad_w) { - if (padding.size() == 1) { - *pad_w = padding[0] * 2; - } else if (padding.size() == 2) { - *pad_w = padding[0] + padding[1]; - } else { - ICHECK_EQ(padding.size(), 4) << " Expected padding size of 1 or 2, found " << padding.size(); - } -} - -/*! \brief A utility function to get padding height and width from a 1, 2, 4 ints tuple. */ -inline void GetPaddingHeightWidth(const Array& padding, IndexExpr* pad_h, - IndexExpr* pad_w) { - if (padding.size() == 1) { - *pad_h = padding[0] * 2; - *pad_w = padding[0] * 2; - } else if (padding.size() == 2) { - *pad_h = padding[0] * 2; - *pad_w = padding[1] * 2; - } else if (padding.size() == 4) { - *pad_h = padding[0] + padding[2]; - *pad_w = padding[1] + padding[3]; - } else { - ICHECK_EQ(padding.size(), 4) << " Padding size should be 1, 2 or 4, but got " << padding.size(); - } -} - -/*! \brief A utility function to get padding depth, height and width from a 1, 3, 6 ints tuple. */ -inline void GetPaddingDepthHeightWidth(const Array& padding, IndexExpr* pad_d, - IndexExpr* pad_h, IndexExpr* pad_w) { - if (padding.size() == 1) { - *pad_d = padding[0] * 2; - *pad_h = padding[0] * 2; - *pad_w = padding[0] * 2; - } else if (padding.size() == 3) { - *pad_d = padding[0] * 2; - *pad_h = padding[1] * 2; - *pad_w = padding[2] * 2; - } else if (padding.size() == 6) { - *pad_d = padding[0] + padding[3]; - *pad_h = padding[1] + padding[4]; - *pad_w = padding[2] + padding[5]; - } else { - ICHECK_EQ(padding.size(), 6) << " Padding size should be 1, 3 or 6, but got " << padding.size(); - } -} - -} // namespace relay -} // namespace tvm - -#endif // TVM_RELAY_OP_OP_COMMON_H_ diff --git a/src/relay/op/random/kernel.cc b/src/relay/op/random/kernel.cc deleted file mode 100644 index 69a47bf1388a..000000000000 --- a/src/relay/op/random/kernel.cc +++ /dev/null @@ -1,229 +0,0 @@ -/* - * 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 - -namespace tvm { -namespace relay { - -TVM_REGISTER_NODE_TYPE(ThreefryGenerateAttrs); - -static TensorType ThreefryKeyType() { return TensorType({10}, tvm::DataType::UInt(64)); } - -bool ThreefryGenerateRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - const ThreefryGenerateAttrs* param = attrs.as(); - ICHECK_EQ(types.size(), 2) << "ThreefryGenerate should have one input and one output"; - - reporter->Assign(types[0], ThreefryKeyType()); - - std::vector oshape; - for (auto& x : param->out_shape) { - oshape.push_back(x); - } - // generate returns the next key and an array of random values - // TODO(@tkonolige, @altanh): support other output dtypes? - reporter->Assign(types[1], - TupleType({ThreefryKeyType(), TensorType(oshape, tvm::DataType::UInt(64))})); - return true; -} - -Expr MakeThreefryGenerate(Expr key, Array out_shape) { - auto attrs = make_object(); - attrs->out_shape = out_shape; - static const Op& op = Op::Get("random.threefry_generate"); - return Call(op, {key}, Attrs(attrs), {}); -} - -TVM_REGISTER_GLOBAL("relay.op.random._make.threefry_generate").set_body_typed(MakeThreefryGenerate); - -RELAY_REGISTER_OP("random.threefry_generate") - .describe( - R"doc(Generate an array of random numbers using the Threefry algorithm.)doc" TVM_ADD_FILELINE) - .set_num_inputs(1) - .set_attrs_type() - .add_argument("key", "Tensor", "Input Threefry key") - .add_type_rel("ThreefryGenerate", ThreefryGenerateRel); - -bool ThreefrySplitRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - ICHECK_EQ(types.size(), 2) << "ThreefrySplit should have one input and one output"; - - reporter->Assign(types[0], ThreefryKeyType()); - reporter->Assign(types[1], TupleType({ThreefryKeyType(), ThreefryKeyType()})); - - return true; -} - -Expr MakeThreefrySplit(Expr key) { - static const Op& op = Op::Get("random.threefry_split"); - return Call(op, {key}, Attrs(), {}); -} - -TVM_REGISTER_GLOBAL("relay.op.random._make.threefry_split").set_body_typed(MakeThreefrySplit); - -RELAY_REGISTER_OP("random.threefry_split") - .describe(R"doc(Split the input Threefry key into two new ones.)doc" TVM_ADD_FILELINE) - .set_num_inputs(1) - .add_argument("key", "Tensor", "Input Threefry key") - .add_type_rel("ThreefrySplit", ThreefrySplitRel); - -TVM_REGISTER_NODE_TYPE(UniformAttrs); - -bool UniformRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - const UniformAttrs* param = attrs.as(); - ICHECK_EQ(types.size(), 4) << "Uniform should have three inputs and one output"; - - std::vector oshape; - for (auto& x : param->out_shape) { - oshape.push_back(x); - } - DataType out_dtype = param->out_dtype; - // we are supporting float32 and float64 at the moment. - if (!(out_dtype.is_float() && (out_dtype.bits() == 32 || out_dtype.bits() == 64))) { - reporter->GetDiagCtx().EmitFatal(Diagnostic::Error(reporter->GetSpan()) - << "We only support generating uniform random value of " - << "type float32 or float64, got " << out_dtype << "."); - return false; - } - reporter->Assign(types[0], ThreefryKeyType()); - reporter->Assign(types[1], TensorType({}, out_dtype)); - reporter->Assign(types[2], TensorType({}, out_dtype)); - // generate returns the next key and an array of random values - reporter->Assign(types[3], TupleType({ThreefryKeyType(), TensorType(oshape, out_dtype)})); - return true; -} - -Expr MakeUniform(Expr key, Expr low, Expr high, Array out_shape, DataType out_dtype) { - auto attrs = make_object(); - attrs->out_shape = out_shape; - attrs->out_dtype = out_dtype; - static const Op& op = Op::Get("random.uniform"); - return Call(op, {key, low, high}, Attrs(attrs), {}); -} - -TVM_REGISTER_GLOBAL("relay.op.random._make.uniform").set_body_typed(MakeUniform); - -RELAY_REGISTER_OP("random.uniform") - .describe( - R"doc(Generate an array of random numbers under uniform distribution.)doc" TVM_ADD_FILELINE) - .set_num_inputs(3) - .set_attrs_type() - .add_argument("key", "Tensor", "Input Threefry key") - .add_argument("low", "Tensor", "Lower bound of the distribution") - .add_argument("high", "Tensor", "Higher bound of the distribution") - .add_type_rel("Uniform", UniformRel); - -TVM_REGISTER_NODE_TYPE(NormalAttrs); - -bool NormalRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - const NormalAttrs* param = attrs.as(); - ICHECK_EQ(types.size(), 4) << "Normal should have three inputs and one output"; - - std::vector oshape; - for (auto& x : param->out_shape) { - oshape.push_back(x); - } - DataType out_dtype = param->out_dtype; - // we are supporting float32 and float64 at the moment. - if (!(out_dtype.is_float() && (out_dtype.bits() == 32 || out_dtype.bits() == 64))) { - reporter->GetDiagCtx().EmitFatal(Diagnostic::Error(reporter->GetSpan()) - << "We only support generating Normal random value of " - << "type float32 or float64, got " << out_dtype << "."); - return false; - } - reporter->Assign(types[0], ThreefryKeyType()); - reporter->Assign(types[1], TensorType({}, out_dtype)); - reporter->Assign(types[2], TensorType({}, out_dtype)); - // generate returns the next key and an array of random values - reporter->Assign(types[3], TupleType({ThreefryKeyType(), TensorType(oshape, out_dtype)})); - return true; -} - -Expr MakeNormal(Expr key, Expr mean, Expr scale, Array out_shape, DataType out_dtype) { - auto attrs = make_object(); - attrs->out_shape = out_shape; - attrs->out_dtype = out_dtype; - static const Op& op = Op::Get("random.normal"); - return Call(op, {key, mean, scale}, Attrs(attrs), {}); -} - -TVM_REGISTER_GLOBAL("relay.op.random._make.normal").set_body_typed(MakeNormal); - -RELAY_REGISTER_OP("random.normal") - .describe( - R"doc(Generate an array of random numbers under normal distribution.)doc" TVM_ADD_FILELINE) - .set_num_inputs(3) - .set_attrs_type() - .add_argument("key", "Tensor", "Input Threefry key") - .add_argument("mean", "Tensor", "Mean of the distribution") - .add_argument("scale", "Tensor", "Standard deviation of the distribution") - .add_type_rel("Normal", NormalRel); - -TVM_REGISTER_NODE_TYPE(MultinomialAttrs); - -bool MultinomialRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - const MultinomialAttrs* param = attrs.as(); - ICHECK_EQ(types.size(), 3) << "Normal should have two inputs and one output"; - - const auto* data = types[1].as(); - if (data == nullptr) { - ICHECK(types[1].as()) - << "multinomial: expect input type to be TensorType but get " << types[0]; - return false; - } - - std::vector oshape; - for (size_t i = 0; i < data->shape.size() - 1; i++) { - oshape.push_back(data->shape[i]); - } - oshape.push_back(param->num_samples); - - DataType out_dtype = tvm::DataType::Int(32); - - reporter->Assign(types[0], ThreefryKeyType()); - // generate returns the next key and an array of random values - reporter->Assign(types[2], TupleType({ThreefryKeyType(), TensorType(oshape, out_dtype)})); - return true; -} - -Expr MakeMultinomial(Expr key, Expr probs, Integer num_samples) { - auto attrs = make_object(); - attrs->num_samples = num_samples; - static const Op& op = Op::Get("random.multinomial"); - return Call(op, {key, probs}, Attrs(attrs), {}); -} - -TVM_REGISTER_GLOBAL("relay.op.random._make.multinomial").set_body_typed(MakeMultinomial); - -RELAY_REGISTER_OP("random.multinomial") - .describe( - R"doc(Generate an array of random numbers under normal distribution.)doc" TVM_ADD_FILELINE) - .set_num_inputs(2) - .set_attrs_type() - .add_argument("key", "Tensor", "Input Threefry key") - .add_argument("probs", "Tensor", "Array of probabilities for each corresponding index.") - .add_type_rel("Multinomial", MultinomialRel); - -} // namespace relay -} // namespace tvm diff --git a/src/relay/op/tensor/binary.cc b/src/relay/op/tensor/binary.cc deleted file mode 100644 index 81746b8d8719..000000000000 --- a/src/relay/op/tensor/binary.cc +++ /dev/null @@ -1,175 +0,0 @@ -/* - * 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 binary.cc - * \brief binary broadcast operators. - */ -#include -#include -#include - -#include "../op_common.h" -#include "../type_relations.h" - -namespace tvm { -namespace relay { - -#define RELAY_BINARY_COMPUTE(FTOPI) \ - [](const Attrs& attrs, const Array& inputs, \ - const Type& out_type) -> Array { \ - ICHECK_EQ(inputs.size(), 2U); \ - return {FTOPI(inputs[0], inputs[1])}; \ - } - -// Addition -RELAY_REGISTER_BINARY_OP("add") - .describe("Elementwise add with broadcasting") - .set_support_level(1) - .set_attr("FTVMCompute", RELAY_BINARY_COMPUTE(topi::add)); - -// Subtraction -RELAY_REGISTER_BINARY_OP("subtract") - .describe("Elementwise substract with broadcasting") - .set_support_level(1) - .set_attr("FTVMCompute", RELAY_BINARY_COMPUTE(topi::subtract)); - -// Right shift -RELAY_REGISTER_BINARY_OP("right_shift") - .describe("Elementwise right shift with broadcasting") - .set_support_level(4) - .set_attr("FTVMCompute", RELAY_BINARY_COMPUTE(topi::right_shift)); - -RELAY_REGISTER_BINARY_OP("left_shift") - .describe("Elementwise left shift with broadcasting") - .set_support_level(4) - .set_attr("FTVMCompute", RELAY_BINARY_COMPUTE(topi::left_shift)); - -RELAY_REGISTER_BINARY_OP("maximum") - .describe("Elementwise maximum of two tensors with broadcasting") - .set_support_level(4) - .set_attr("FTVMCompute", RELAY_BINARY_COMPUTE(topi::maximum)); - -RELAY_REGISTER_BINARY_OP("minimum") - .describe("Elementwise minimum of two tensors with broadcasting") - .set_support_level(4) - .set_attr("FTVMCompute", RELAY_BINARY_COMPUTE(topi::minimum)); - -RELAY_REGISTER_BINARY_OP("divide") - .describe("Elementwise divide with broadcasting") - .set_support_level(1) - .set_attr("FTVMCompute", RELAY_BINARY_COMPUTE(topi::divide)); - -RELAY_REGISTER_BINARY_OP("trunc_divide") - .describe("Elementwise trunc divide with broadcasting") - .set_support_level(1) - .set_attr("FTVMCompute", RELAY_BINARY_COMPUTE(topi::trunc_divide)); - -RELAY_REGISTER_BINARY_OP("floor_divide") - .describe("Elementwise floor divide with broadcasting") - .set_support_level(1) - .set_attr("FTVMCompute", RELAY_BINARY_COMPUTE(topi::floor_divide)); - -RELAY_REGISTER_BINARY_OP("multiply") - .describe("Elementwise multiply with broadcasting") - .set_support_level(1) - .set_attr("FTVMCompute", RELAY_BINARY_COMPUTE(topi::multiply)); - -RELAY_REGISTER_BINARY_OP("power") - .describe("Elementwise power with broadcasting") - .set_support_level(4) - .set_attr("FTVMCompute", RELAY_BINARY_COMPUTE(topi::power)); - -RELAY_REGISTER_BINARY_OP("mod") - .describe("Elementwise mod with broadcasting") - .set_support_level(1) - .set_attr("FTVMCompute", RELAY_BINARY_COMPUTE(topi::mod)); - -RELAY_REGISTER_BINARY_OP("floor_mod") - .describe("Elementwise floor mod with broadcasting") - .set_support_level(1) - .set_attr("FTVMCompute", RELAY_BINARY_COMPUTE(topi::floor_mod)); - -RELAY_REGISTER_BINARY_OP("trunc_mod") - .describe("Elementwise trunc mod with broadcasting") - .set_support_level(1) - .set_attr("FTVMCompute", RELAY_BINARY_COMPUTE(topi::trunc_mod)); - -RELAY_REGISTER_BINARY_OP("logical_and") - .describe("Elementwise logical AND with broadcasting") - .set_support_level(4) - .set_attr("FTVMCompute", RELAY_BINARY_COMPUTE(topi::logical_and)); - -RELAY_REGISTER_BINARY_OP("logical_or") - .describe("Elementwise logical OR with broadcasting") - .set_support_level(4) - .set_attr("FTVMCompute", RELAY_BINARY_COMPUTE(topi::logical_or)); - -RELAY_REGISTER_BINARY_OP("logical_xor") - .describe("Elementwise logical XOR with broadcasting") - .set_support_level(4) - .set_attr("FTVMCompute", RELAY_BINARY_COMPUTE(topi::logical_xor)); - -RELAY_REGISTER_BINARY_OP("bitwise_and") - .describe("Elementwise bitwise AND with broadcasting") - .set_support_level(4) - .set_attr("FTVMCompute", RELAY_BINARY_COMPUTE(topi::bitwise_and)); - -RELAY_REGISTER_BINARY_OP("bitwise_or") - .describe("Elementwise bitwise OR with broadcasting") - .set_support_level(4) - .set_attr("FTVMCompute", RELAY_BINARY_COMPUTE(topi::bitwise_or)); - -RELAY_REGISTER_BINARY_OP("bitwise_xor") - .describe("Elementwise bitwise XOR with broadcasting") - .set_support_level(4) - .set_attr("FTVMCompute", RELAY_BINARY_COMPUTE(topi::bitwise_xor)); - -RELAY_REGISTER_CMP_OP("equal") - .describe("Elementwise equal compare with broadcasting") - .set_support_level(4) - .set_attr("FTVMCompute", RELAY_BINARY_COMPUTE(topi::equal)); - -RELAY_REGISTER_CMP_OP("not_equal") - .describe("Elementwise not equal with broadcasting") - .set_support_level(4) - .set_attr("FTVMCompute", RELAY_BINARY_COMPUTE(topi::not_equal)); - -RELAY_REGISTER_CMP_OP("less") - .describe("Elementwise less than with broadcasting") - .set_support_level(4) - .set_attr("FTVMCompute", RELAY_BINARY_COMPUTE(topi::less)); - -RELAY_REGISTER_CMP_OP("less_equal") - .describe("Elementwise less than or equal compare with broadcasting") - .set_support_level(4) - .set_attr("FTVMCompute", RELAY_BINARY_COMPUTE(topi::less_equal)); - -RELAY_REGISTER_CMP_OP("greater") - .describe("Elementwise greater than compare with broadcasting") - .set_support_level(4) - .set_attr("FTVMCompute", RELAY_BINARY_COMPUTE(topi::greater)); - -RELAY_REGISTER_CMP_OP("greater_equal") - .describe("Elementwise greater than or equal compare with broadcasting") - .set_support_level(4) - .set_attr("FTVMCompute", RELAY_BINARY_COMPUTE(topi::greater_equal)); - -} // namespace relay -} // namespace tvm diff --git a/src/relay/op/tensor/math.cc b/src/relay/op/tensor/math.cc deleted file mode 100644 index ef3ac8accbf2..000000000000 --- a/src/relay/op/tensor/math.cc +++ /dev/null @@ -1,119 +0,0 @@ -/* - * 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 math.cc - * \brief Math operators. - */ -#include -#include -#include - -#include "../make_op.h" -#include "../op_common.h" -#include "../type_relations.h" - -namespace tvm { -namespace relay { - -// relay.einsum -TVM_REGISTER_NODE_TYPE(EinsumAttrs); - -bool EinsumRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - // Check attrs - const EinsumAttrs* param = attrs.as(); - if (param == nullptr) { - reporter->GetDiagCtx().EmitFatal(Diagnostic::Error(reporter->GetSpan()) - << "the call attributes are not defined"); - return false; - } - - // types: [data, result] - ICHECK_EQ(types.size(), 2) << "the arity of einsum is 2, not " << types.size(); - - // Check input type is a tuple. - const auto* tensor_tuple = types[0].as(); - if (tensor_tuple == nullptr) { - reporter->GetDiagCtx().EmitFatal( - Diagnostic::Error(reporter->GetSpan()) - << "einsum requires a tuple of tensors as the first argument, found " - << PrettyPrint(types[0])); - return false; - } - - // Check the input tuple consists of tensors with consistent dtype. - if (tensor_tuple->fields[0].as()) { - return false; - } - ICHECK(tensor_tuple->fields[0].as()); - const auto& first = Downcast(tensor_tuple->fields[0]); - const DataType dtype = first->dtype; - std::vector> input_shapes; - for (const Type& ele : tensor_tuple->fields) { - if (ele.as()) { - return false; - } - - const auto& e = Downcast(ele); - - const DataType& e_dtype = e->dtype; - if (e_dtype != dtype) { - throw Error("relay.einsum requires all tensors have the same dtype"); - } - input_shapes.push_back(e->shape); - } - - // Calculate output shape - Array oshape = topi::InferEinsumShape(param->equation, input_shapes); - - auto rtype = TensorType(oshape, dtype); - reporter->Assign(types[1], rtype); - return true; -} - -Array EinsumCompute(const Attrs& attrs, const Array& inputs, - const Type& out_type) { - const EinsumAttrs* param = attrs.as(); - ICHECK(param != nullptr); - return Array{topi::einsum(param->equation, inputs)}; -} - -Expr MakeEinsum(Expr data, String equation) { - auto attrs = make_object(); - attrs->equation = std::move(equation); - static const Op& op = Op::Get("einsum"); - return Call(op, {data}, Attrs(attrs), {}); -} - -TVM_REGISTER_GLOBAL("relay.op._make.einsum").set_body_typed(MakeEinsum); - -RELAY_REGISTER_OP("einsum") - .describe(R"doc(Evaluates the Einstein summation convention -on the operands)doc" TVM_ADD_FILELINE) - .set_attrs_type() - .set_num_inputs(1) - .add_argument("data", "Tuple of Tensors", "The input list of tensors.") - .set_support_level(11) - .add_type_rel("Einsum", EinsumRel) - .set_attr("FTVMCompute", EinsumCompute) - .set_attr("TOpPattern", kInjective); - -} // namespace relay -} // namespace tvm diff --git a/src/relay/op/tensor/reduce.cc b/src/relay/op/tensor/reduce.cc deleted file mode 100644 index d82705e3fc55..000000000000 --- a/src/relay/op/tensor/reduce.cc +++ /dev/null @@ -1,746 +0,0 @@ -/* - * 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 reduce.cc - * \brief Reduction operators. - */ -#include -#include -#include -#include -#include - -#include -#include - -#include "../make_op.h" -#include "../op_common.h" -#include "../type_relations.h" - -namespace tvm { -namespace relay { - -TVM_REGISTER_NODE_TYPE(ReduceAttrs); -TVM_REGISTER_NODE_TYPE(ArgReduceAttrs); -TVM_REGISTER_NODE_TYPE(VarianceAttrs); - -/*! - * \brief GetReduceAxes, get the new axis from indim and other arguments - * \param indim Number of dimensions of input data. - * \param axis The input axis vector. - * \param exclude Whether 'axis' input given is the excluded axis. - * \return r_axes The new reduced axes of the output. - */ -inline std::vector GetReduceAxes(const uint32_t indim, const Array& inaxis, - bool exclude) { - if (!inaxis.defined() || inaxis.empty()) { - std::vector r_axes(indim); - std::iota(r_axes.begin(), r_axes.end(), 0); - return r_axes; - } - - std::vector in_axes; - for (auto i : inaxis) { - int64_t axis = i->value; - if (axis < 0) { - axis = axis + indim; - } - - // Check out of bounds error - ICHECK(axis >= 0) << "Axis out of bounds in reduce operator."; - ICHECK(axis < indim) << "Axis out of bounds in reduce operator."; - in_axes.push_back(axis); - } - - ICHECK(in_axes[in_axes.size() - 1] < indim) - << "Reduction axis " << in_axes[in_axes.size() - 1] << " exceeds input dimensions " << indim; - - std::sort(in_axes.begin(), in_axes.end()); - - if (!exclude) { - return in_axes; - } - - auto r_size = indim - in_axes.size(); - std::vector r_axes(r_size); - for (uint32_t i = 0, j = 0, k = 0; i < indim; ++i) { - if (j < in_axes.size() && in_axes[j] == i) { - ++j; - continue; - } - r_axes[k++] = i; - } - return r_axes; -} - -// Get axis under exclude condition. -Array GetExcludeAxes(size_t indim, const Array& inaxis) { - ICHECK(inaxis.defined()) << "Cannot set exclude when axis=None"; - std::vector axis_flag(indim, true); - for (auto i : inaxis) { - int64_t axis = i->value; - if (axis < 0) { - axis = axis + static_cast(indim); - } - // Check out of bounds error - ICHECK_GE(axis, 0) << "Axis out of bounds in reduce operator."; - ICHECK_LT(axis, static_cast(indim)) << "Axis out of bounds in reduce operator."; - axis_flag[axis] = false; - } - - Array r_axes; - - for (size_t i = 0; i < axis_flag.size(); ++i) { - if (axis_flag[i]) { - r_axes.push_back(static_cast(i)); - } - } - return r_axes; -} - -// Return the modified layout for AlterOpLayout pass. -template -InferCorrectLayoutOutput ReduceInferCorrectLayout(const Attrs& attrs, - const Array& new_in_layouts, - const Array& old_in_layouts, - const Array& old_in_types) { - const auto* attrs_ptr = attrs.as(); - ICHECK(attrs_ptr); - ObjectPtr params = make_object(*attrs_ptr); - - // Get the reduce axes. - Array> old_in_shapes; - for (auto old_in_t : old_in_types) { - ICHECK(old_in_t.as()); - old_in_shapes.push_back(old_in_t.as()->shape); - } - uint32_t indim = old_in_shapes[0].size(); - auto r_axes = GetReduceAxes(indim, params->axis, params->exclude); - - Layout inferred_in = Layout::Undef(); - Layout inferred_out = Layout::Undef(); - - // Infer [in_layout, out_layout, new_r_axes] from old_in_layout or new_in_layout - auto infer = [&](const Layout& layout) { - // 1) Collect the original axes - std::unordered_set old_r_dims; - for (auto r_axis : r_axes) { - old_r_dims.emplace(old_in_layouts[0][r_axis].name()); - } - - // 2) Collect the new axes by walking new_layout. - tvm::Array new_r_axes; - std::string inferred_in_string = ""; - std::string inferred_out_string = ""; - auto push_new_axis = [&](const std::string& layout_dim, int axis) { - if ((old_r_dims.count(layout_dim) && !params->exclude) || - (!old_r_dims.count(layout_dim) && params->exclude)) { - new_r_axes.push_back(tvm::Integer(axis)); - return true; - } - return false; - }; - for (size_t axis_index = 0; axis_index < layout->axes.size(); ++axis_index) { - const auto& layout_axis = LayoutAxis::Get(layout->axes[axis_index]); - const std::string& layout_dim = layout_axis.name(); - if (layout_axis.IsPrimal()) { - push_new_axis(layout_dim, axis_index); - inferred_in_string += layout_dim; - if (!old_r_dims.count(layout_dim) || params->keepdims) { - inferred_out_string += layout_dim; - } - } else { - // For example, if the original layout is NCHW, the new layout is NCHW8c, and the original - // reduce axes is [1], the new reduce axes become [1, 4]. - auto primal_dim = layout_axis.ToPrimal().name(); - auto packed_dim = std::to_string(layout.FactorOf(layout_axis)) + layout_dim; - inferred_in_string += packed_dim; - if (push_new_axis(primal_dim, axis_index)) { - if (params->exclude) { - // The primal axis is not reduced, so keep the input packed dim. - inferred_out_string += packed_dim; - } else if (params->keepdims) { - // If the primal axis is part of reduce axes in the original layout, the inner dim - // becomes 1 after reduction. - inferred_out_string += "1" + layout_dim; - } - } else { - inferred_out_string += packed_dim; - } - } - } - - // 3) Set the new axis and layout. - return std::make_tuple(Layout(inferred_in_string), Layout(inferred_out_string), new_r_axes); - }; - - std::string new_layout_string; - Array new_r_axes; - Array new_input_layouts; - - auto check_num_input_layouts = [](Array in_layouts) { - // The second case is for variance op - ICHECK(in_layouts.size() == 1 || in_layouts.size() == 2); - }; - - if (new_in_layouts.defined() && r_axes.size()) { - // Adapt to new layout. The axis has to change. Record original reduce axes. Convert to the - // modified layout axes. - check_num_input_layouts(new_in_layouts); - check_num_input_layouts(old_in_layouts); - - // Get inferred_in and inferred_out from new_in_layout. - std::tie(inferred_in, inferred_out, new_r_axes) = infer(new_in_layouts[0]); - params->axis = new_r_axes; - } else if (old_in_layouts.defined()) { - check_num_input_layouts(old_in_layouts); - - // If the new layout is undefined, get inferred_in and inferred_out from old_in_layout. - if (old_in_layouts[0].defined()) { - std::tie(inferred_in, inferred_out, std::ignore) = infer(old_in_layouts[0]); - } - } - - new_input_layouts.push_back(inferred_in); - - if (old_in_layouts.size() == 2) { - new_input_layouts.push_back(inferred_in); - } - - return InferCorrectLayoutOutput(new_input_layouts, {inferred_out}, Attrs(params)); -} - -template -Array ReduceCompute(const Attrs& attrs, const Array& inputs, - const Type& out_type, F f) { - const ReduceAttrs* param = attrs.as(); - ICHECK(param != nullptr); - if (inputs[0]->shape.size() == 0) { - return {topi::identity(inputs[0])}; - } - auto axes = param->axis; - if (param->exclude) { - axes = GetExcludeAxes(inputs[0]->shape.size(), param->axis); - if (axes.size() == 0) { - return {topi::identity(inputs[0])}; - } - } - - return {f(inputs[0], axes, param->keepdims, false)}; -} - -template -Array ArgReduceCompute(const Attrs& attrs, const Array& inputs, - const Type& out_type, F f) { - const ArgReduceAttrs* param = attrs.as(); - ICHECK(param != nullptr); - if (inputs[0]->shape.size() == 0) { - return {topi::identity(inputs[0])}; - } - auto axes = param->axis; - if (param->exclude) { - axes = GetExcludeAxes(inputs[0]->shape.size(), param->axis); - if (axes.size() == 0) { - return {topi::identity(inputs[0])}; - } - } - - return {f(inputs[0], axes, param->keepdims, false, param->select_last_index)}; -} - -/*! - * \brief ReduceShapeImpl get the outshape for the reduction operator - * \param in_shape Shape of input data. - * \param param Attrs details. - * \param reporter The reporter to report solution to. - * \return oshape Output shape inferred. - * \tparam AttrsType The attribute type. - */ -template -inline std::vector ReduceShapeImpl(const std::vector& in_shape, - const AttrsType* param, - const TypeReporter& reporter) { - uint32_t indim = in_shape.size(); - auto r_axes = GetReduceAxes(indim, param->axis, param->exclude); - if (!r_axes.size()) { - return in_shape; - } - - auto max_shape = tir::make_const(DataType::Int(64), 1); - bool is_dynamic_input = false; - for (int64_t axis : r_axes) { - if (in_shape[axis].as()) { - max_shape *= in_shape[axis]; - } else { - is_dynamic_input = true; - break; - } - } - - if (is_dynamic_input) { - ICHECK(reporter->Assert( - max_shape < tir::make_const(DataType::Int(64), std::numeric_limits::max()))) - << "The maximum possible index of reduced shape cannot be more than int32 max."; - } - - if (param->keepdims) { - std::vector oshape(in_shape); - for (unsigned i = 0, j = 0; i < indim; ++i) { - if (j >= r_axes.size() || !(r_axes[j] == i)) { - continue; - } - oshape[i] = 1; - ++j; - } - return oshape; - } else { - auto osize = indim - r_axes.size(); - std::vector oshape(osize); - for (unsigned i = 0, j = 0, k = 0; i < indim; ++i) { - if (j < r_axes.size() && (r_axes[j] == i)) { - ++j; - continue; - } - oshape[k++] = in_shape[i]; - } - return oshape; - } -} - -template -bool GenericReduceRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - ICHECK_EQ(types.size(), 2); - const auto* data = types[0].as(); - if (data == nullptr) return false; - ICHECK(static_cast(data->shape.size()) != 0); - std::vector in_shape(data->shape.begin(), data->shape.end()); - - const T* param = attrs.as(); - ICHECK(param != nullptr); - - // assign output type and shape - auto oshape = ReduceShapeImpl(in_shape, param, reporter); - reporter->Assign(types[1], TensorType(oshape, data->shape[0].dtype())); - return true; -} -/*! - * \brief ArgReduceRel Output type and shape relation evaluation function. - * \param num_inputs Number of input types in the args. - * \param attrs The additional attributes of the operator. - * \param reporter The reporter to report solution to. - * \return false if This relation cannot be resolved. true if this relation has been resolved. - */ -bool ArgReduceRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - return GenericReduceRel(types, num_inputs, attrs, reporter); -} - -/*! - * \brief ReduceRel Output type and shape relation evaluation function. - * \param num_inputs Number of input types in the args. - * \param attrs The additional attributes of the operator. - * \param reporter The reporter to report solution to. - * \return false if This relation cannot be resolved. true if this relation has been resolved. - */ -bool ReduceRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - ICHECK_EQ(types.size(), 2); - const auto* data = types[0].as(); - if (data == nullptr) return false; - std::vector in_shape(data->shape.begin(), data->shape.end()); - - const ReduceAttrs* param = attrs.as(); - ICHECK(param != nullptr); - - // assign output type and shape - auto oshape = ReduceShapeImpl(in_shape, param, reporter); - reporter->Assign(types[1], TensorType(oshape, data->dtype)); - return true; -} - -Expr MakeReduce(Expr data, Array axis, bool keepdims, bool exclude, String op_name) { - auto attrs = make_object(); - attrs->axis = std::move(axis); - attrs->keepdims = keepdims; - attrs->exclude = exclude; - return Call(Op::Get(op_name), {data}, Attrs(attrs), {}); -} - -Expr MakeOneElementReduce(Expr data, Array axis, bool keepdims, bool exclude, - bool select_last_index, String op_name) { - auto attrs = make_object(); - attrs->axis = std::move(axis); - attrs->keepdims = keepdims; - attrs->exclude = exclude; - attrs->select_last_index = select_last_index; - return Call(Op::Get(op_name), {data}, Attrs(attrs), {}); -} - -#define RELAY_REGISTER_REDUCE_OP(OpName) \ - TVM_REGISTER_GLOBAL("relay.op._make." OpName) \ - .set_body_typed([](Expr data, Array axis, bool keepdims, bool exclude) { \ - return MakeReduce(data, axis, keepdims, exclude, OpName); \ - }); \ - RELAY_REGISTER_OP(OpName).set_num_inputs(1).add_argument("data", "Tensor", "The input tensor.") - -#define RELAY_REGISTER_ONE_ELEMENT_REDUCE_OP(OpName) \ - TVM_REGISTER_GLOBAL("relay.op._make." OpName) \ - .set_body_typed([](Expr data, Array axis, bool keepdims, bool exclude, \ - bool select_last_index) { \ - return MakeOneElementReduce(data, axis, keepdims, exclude, select_last_index, OpName); \ - }); \ - RELAY_REGISTER_OP(OpName).set_num_inputs(1).add_argument("data", "Tensor", "The input tensor.") - -Array ArgMaxCompute(const Attrs& attrs, const Array& inputs, - const Type& out_type) { - return ArgReduceCompute(attrs, inputs, out_type, topi::argmax); -} - -RELAY_REGISTER_ONE_ELEMENT_REDUCE_OP("argmax") - .describe(R"code(Creates an operation that finds the indices of the maximum -values over a given axis. - -)code" TVM_ADD_FILELINE) - .set_attrs_type() - .set_support_level(4) - .add_type_rel("ArgReduce", GenericReduceRel) - .set_attr("FTVMCompute", ArgMaxCompute) - .set_attr("FInferCorrectLayout", ReduceInferCorrectLayout) - .set_attr("TOpPattern", kCommReduce); - -Array ArgMinCompute(const Attrs& attrs, const Array& inputs, - const Type& out_type) { - return ArgReduceCompute(attrs, inputs, out_type, topi::argmin); -} - -RELAY_REGISTER_ONE_ELEMENT_REDUCE_OP("argmin") - .describe(R"code(Creates an operation that finds the indices of the minimum -values over a given axis. - -)code" TVM_ADD_FILELINE) - .set_attrs_type() - .set_support_level(4) - .add_type_rel("ArgReduce", GenericReduceRel) - .set_attr("FTVMCompute", ArgMinCompute) - .set_attr("FInferCorrectLayout", ReduceInferCorrectLayout) - .set_attr("TOpPattern", kCommReduce); - -Array SumCompute(const Attrs& attrs, const Array& inputs, - const Type& out_type) { - return ReduceCompute(attrs, inputs, out_type, topi::sum); -} - -RELAY_REGISTER_REDUCE_OP("sum") - .describe(R"code(Computes the sum of array elements over given axes. - -Example:: - - data = [[[1,2],[2,3],[1,3]], - [[1,4],[4,3],[5,2]], - [[7,1],[7,2],[7,3]]] - - sum(data, axis=1) - [[ 4. 8.] - [ 10. 9.] - [ 21. 6.]] - - sum(data, axis=[1,2]) - [ 12. 19. 27.] - -)code" TVM_ADD_FILELINE) - .set_attrs_type() - .set_support_level(4) - .add_type_rel("Reduce", ReduceRel) - .set_attr("FInferCorrectLayout", ReduceInferCorrectLayout) - .set_attr("FTVMCompute", SumCompute) - .set_attr("TOpPattern", kCommReduce); - -Array AllCompute(const Attrs& attrs, const Array& inputs, - const Type& out_type) { - return ReduceCompute(attrs, inputs, out_type, topi::all); -} - -RELAY_REGISTER_REDUCE_OP("all") - .describe(R"code(Computes the logical AND of boolean array elements over given axes. - -Example:: - - data = [[[ True, True, True], - [ True, True, True], - [False, True, False]], - [[ True, False, False], - [ True, True, False], - [False, True, True]]] - - all(data, axis=1) - [[False, True, False], - [False, False, False]] - - all(data, axis=0) - [[ True, False, False], - [ True, True, False], - [False, True, False]] - -)code" TVM_ADD_FILELINE) - .set_attrs_type() - .set_support_level(4) - .add_type_rel("Reduce", ReduceRel) - .set_attr("FTVMCompute", AllCompute) - .set_attr("FInferCorrectLayout", ReduceInferCorrectLayout) - .set_attr("TOpPattern", kCommReduce); - -Array AnyCompute(const Attrs& attrs, const Array& inputs, - const Type& out_type) { - return ReduceCompute(attrs, inputs, out_type, topi::any); -} - -RELAY_REGISTER_REDUCE_OP("any") - .describe(R"code(Computes the logical OR of boolean array elements over given axes. - -Example:: - - data = [[[ True, True, True], - [ True, True, True], - [False, True, False]], - [[ True, False, False], - [ True, True, False], - [False, True, True]]] - - any(data, axis=1) - [[True, True, True], - [True, True, True]] - - any(data, axis=0) - [[ True, True, True], - [ True, True, True], - [False, True, True]] - -)code" TVM_ADD_FILELINE) - .set_attrs_type() - .set_support_level(4) - .add_type_rel("Reduce", ReduceRel) - .set_attr("FTVMCompute", AnyCompute) - .set_attr("TOpPattern", kCommReduce); - -Array MaxCompute(const Attrs& attrs, const Array& inputs, - const Type& out_type) { - return ReduceCompute(attrs, inputs, out_type, topi::max); -} - -RELAY_REGISTER_REDUCE_OP("max") - .describe(R"code(Computes the max of array elements over given axes. - -)code" TVM_ADD_FILELINE) - .set_attrs_type() - .set_support_level(4) - .add_type_rel("Reduce", ReduceRel) - .set_attr("FTVMCompute", MaxCompute) - .set_attr("FInferCorrectLayout", ReduceInferCorrectLayout) - .set_attr("TOpPattern", kCommReduce); - -Array MinCompute(const Attrs& attrs, const Array& inputs, - const Type& out_type) { - return ReduceCompute(attrs, inputs, out_type, topi::min); -} - -RELAY_REGISTER_REDUCE_OP("min") - .describe(R"code(Computes the min of array elements over given axes. - -)code" TVM_ADD_FILELINE) - .set_attrs_type() - .set_support_level(4) - .add_type_rel("Reduce", ReduceRel) - .set_attr("FTVMCompute", MinCompute) - .set_attr("FInferCorrectLayout", ReduceInferCorrectLayout) - .set_attr("TOpPattern", kCommReduce); - -Array ProdCompute(const Attrs& attrs, const Array& inputs, - const Type& out_type) { - return ReduceCompute(attrs, inputs, out_type, topi::prod); -} - -TVM_REGISTER_GLOBAL("relay.op._make.prod").set_body_typed(Prod); - -RELAY_REGISTER_OP("prod") - .set_num_inputs(1) - .add_argument("data", "Tensor", "The input tensor.") - .describe(R"code(Computes the products of array elements over given axes. - -Example:: - - data = [[[1,2],[2,3],[1,3]], - [[1,4],[4,3],[5,2]], - [[7,1],[7,2],[7,3]]] - - prod(data, axis=1) - [35562240] - - prod(data, axis=[1,2]) - [ 36 480 2058] - -)code" TVM_ADD_FILELINE) - .set_attrs_type() - .set_support_level(4) - .add_type_rel("Reduce", ReduceRel) - .set_attr("FTVMCompute", ProdCompute) - .set_attr("FInferCorrectLayout", ReduceInferCorrectLayout) - .set_attr("TOpPattern", kCommReduce); - -Array MeanCompute(const Attrs& attrs, const Array& inputs, - const Type& out_type) { - auto data = inputs[0]; - IndexExpr count = tir::make_const(DataType::Int(64), 1); - const ReduceAttrs* param = attrs.as(); - ICHECK(param != nullptr); - auto axes = param->axis; - for (int64_t i : GetReduceAxes(inputs[0]->shape.size(), param->axis, param->exclude)) { - count *= inputs[0]->shape[i]; - } - // Check the datatype of input data. If it's fp16, we'll have trouble representing all - // indices and summation needed so we instead just cast to fp32. - bool recast_fp16 = false; - if (data->dtype.is_float16()) { - recast_fp16 = true; - data = topi::cast(data, DataType::Float(32)); - } - count = cast(data->dtype, count); - auto res = ReduceCompute(attrs, {data}, out_type, topi::sum); - auto output = topi::divide(res[0], count); - // Set the output back to the appropriate fp16 type if needed. - if (recast_fp16) { - output = topi::cast(output, DataType::Float(16)); - } - return {output}; -} - -RELAY_REGISTER_REDUCE_OP("mean") - .describe(R"code(Computes the mean of array elements over given axes. - -Example:: - - data = [[[1,2],[2,3],[1,3]], - [[1,4],[4,3],[5,2]], - [[7,1],[7,2],[7,3]]] - - mean(data) - [3.22] - - mean(data, axis=[1,2]) - [ 2. 3.16666667 4.5] - -)code" TVM_ADD_FILELINE) - .set_attrs_type() - .set_support_level(4) - .add_type_rel("Reduce", ReduceRel) - .set_attr("FTVMCompute", MeanCompute) - .set_attr("FInferCorrectLayout", ReduceInferCorrectLayout) - .set_attr("TOpPattern", kCommReduce); - -bool VarianceRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - ICHECK_EQ(types.size(), 3); - const auto* data = types[0].as(); - if (data == nullptr) return false; - ICHECK(static_cast(data->shape.size()) != 0); - const auto* mean = types[1].as(); - if (mean == nullptr) return false; - - std::vector in_shape(data->shape.begin(), data->shape.end()); - std::vector mean_shape(mean->shape.begin(), mean->shape.end()); - ICHECK_EQ(in_shape.size(), mean_shape.size()); - - const VarianceAttrs* param = attrs.as(); - ICHECK(param != nullptr); - - // assign output type and shape - auto oshape = ReduceShapeImpl(in_shape, param, reporter); - reporter->Assign(types[2], TensorType(oshape, data->dtype)); - return true; -} - -Array VarianceCompute(const Attrs& attrs, const Array& inputs, - const Type& out_type) { - IndexExpr count = tir::make_const(DataType::Int(64), 1); - const VarianceAttrs* param = attrs.as(); - ICHECK(param != nullptr); - auto axes = param->axis; - bool unbiased = param->unbiased; - auto data = inputs[0]; - auto mean = inputs[1]; - for (int64_t i : GetReduceAxes(data->shape.size(), param->axis, param->exclude)) { - count *= data->shape[i]; - } - if (unbiased) { - count -= 1; - } - std::vector expand_shape; - auto diff = topi::subtract(data, mean); - auto sq_diff = topi::multiply(diff, diff); - if (param->exclude) { - axes = GetExcludeAxes(sq_diff->shape.size(), param->axis); - ICHECK_NE(axes.size(), 0); - } - // If the input is fp16, we might have trouble representing the full sum of - // indices or values. We recast to fp32 to avoid this issue. - bool recast_fp16 = false; - if (data->dtype.is_float16()) { - recast_fp16 = true; - sq_diff = topi::cast(sq_diff, DataType::Float(32)); - } - auto var = topi::divide(topi::sum(sq_diff, axes, param->keepdims, false), count); - - // Recast back to fp16 if needed. - if (recast_fp16) { - var = topi::cast(var, DataType::Float(16)); - } - - return {var}; -} - -Expr MakeVariance(Expr data, Expr mean, Array axis, bool keepdims, bool exclude, - bool unbiased = false) { - auto attrs = make_object(); - attrs->axis = std::move(axis); - attrs->keepdims = keepdims; - attrs->exclude = exclude; - attrs->unbiased = unbiased; - static const Op& op = Op::Get("variance"); - return Call(op, {data, mean}, Attrs(attrs), {}); -} - -TVM_REGISTER_GLOBAL("relay.op._make._variance").set_body_typed(MakeVariance); - -RELAY_REGISTER_OP("variance") - .describe(R"code(Computes the variance of array elements over given axes. - -)code" TVM_ADD_FILELINE) - .set_attrs_type() - .set_support_level(4) - .set_num_inputs(2) - .add_argument("data", "Tensor", "The input tensor.") - .add_argument("mean", "Tensor", "The mean tensor.") - .add_type_rel("Variance", VarianceRel) - .set_attr("FTVMCompute", VarianceCompute) - .set_attr("FInferCorrectLayout", ReduceInferCorrectLayout) - .set_attr("TOpPattern", kCommReduce); - -} // namespace relay -} // namespace tvm diff --git a/src/relay/op/tensor/transform.cc b/src/relay/op/tensor/transform.cc deleted file mode 100644 index 96f833d80505..000000000000 --- a/src/relay/op/tensor/transform.cc +++ /dev/null @@ -1,4438 +0,0 @@ -/* - * 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 transform.cc - * \brief Transform operators. - */ -#include "transform.h" - -#include -#include -#include -#include -#include -#include -#include -#include -#include -#include -#include -#include -#include -#include -#include - -#include -#include - -#include "../../transforms/infer_layout_utils.h" -#include "../../transforms/pass_utils.h" -#include "../../transforms/pattern_utils.h" -#include "../make_op.h" -#include "../op_common.h" -#include "../type_relations.h" - -namespace tvm { -namespace relay { -using tir::IntImmNode; - -TVM_REGISTER_NODE_TYPE(SlidingWindowAttrs); - -bool SlidingWindowRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - // `types` contains: [data, result] - ICHECK_EQ(types.size(), 2); - const auto* data = types[0].as(); - if (data == nullptr) { - reporter->GetDiagCtx().EmitFatal(Diagnostic::Error(reporter->GetSpan()) - << "SlidingWindow operator expects input to be of TensorType " - << "but got " << PrettyPrint(types[0])); - return false; - } - const auto* param = attrs.as(); - const int axis = param->axis; - - std::vector oshape; - - // Dimensions up until `axis` remain the same. - for (int i = 0; i < axis; ++i) { - oshape.emplace_back(data->shape[i]); - } - - // New dimensions which result from sliding the window in each dimension. One new dimension per - // window dimension. - for (size_t i = 0; i < param->window_shape.size(); ++i) { - // Length of the shape along this dimension. - auto dim_len = data->shape[axis + i]; - // Length of the window along this dimension. - auto window_len = param->window_shape[i]; - // Strides along this dimension. - auto stride = param->strides[i]; - - oshape.push_back(floordiv(dim_len - (window_len - 1) + stride - 1, stride)); - } - - // Dimensions comprising the window. - for (size_t i = 0; i < param->window_shape.size(); ++i) { - oshape.push_back(param->window_shape[i]); - } - - reporter->Assign(types[1], TensorType(oshape, data->dtype)); - return true; -} - -Array SlidingWindowCompute(const Attrs& attrs, const Array& inputs, - const Type& out_type) { - const SlidingWindowAttrs* param = attrs.as(); - ICHECK(param != nullptr); - return {topi::sliding_window(inputs[0], param->axis, param->window_shape, param->strides)}; -} - -Expr MakeSlidingWindow(Expr data, int axis, Array window_shape, Array strides) { - auto attrs = make_object(); - attrs->axis = axis; - attrs->window_shape = window_shape; - attrs->strides = strides; - static const Op& op = Op::Get("sliding_window"); - return Call(op, {data}, Attrs(attrs), {}); -} - -TVM_REGISTER_GLOBAL("relay.ir.sliding_window").set_body_typed(MakeSlidingWindow); - -RELAY_REGISTER_OP("sliding_window") - .describe(R"code(Slide window over a tensor.)code" TVM_ADD_FILELINE) - .set_num_inputs(1) - .set_attrs_type() - .add_argument("data", "Tensor", "The input tensor.") - .add_type_rel("SlidingWindow", SlidingWindowRel) - .set_attr("TOpPattern", kOpaque); - -// relay.cast -TVM_REGISTER_NODE_TYPE(CastAttrs); - -bool CastRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - ICHECK_EQ(types.size(), 2); - const auto* data = types[0].as(); - if (data == nullptr) { - ICHECK(types[0].as()) - << "cast: expect input type to be TensorType but get " << types[0]; - return false; - } - const auto* param = attrs.as(); - reporter->Assign(types[1], TensorType(data->shape, param->dtype)); - return true; -} - -Array CastCompute(const Attrs& attrs, const Array& inputs, - const Type& out_type) { - const CastAttrs* param = attrs.as(); - ICHECK(param != nullptr); - DataType dtype = param->dtype; - return {topi::cast(inputs[0], dtype)}; -} - -Expr MakeCast(Expr data, DataType dtype) { - auto attrs = make_object(); - attrs->dtype = dtype; - static const Op& op = Op::Get("cast"); - return Call(op, {data}, Attrs(attrs), {}); -} - -TVM_REGISTER_GLOBAL("relay.ir.cast").set_body_typed(MakeCast); - -RELAY_REGISTER_OP("cast") - .describe(R"code(Cast the data into a new data type. - -)code" TVM_ADD_FILELINE) - .set_num_inputs(1) - .set_attrs_type() - .add_argument("data", "Tensor", "The input tensor.") - .set_support_level(3) - .add_type_rel("Cast", CastRel) - .set_attr("FTVMCompute", CastCompute) - .set_attr("TOpPattern", kElemWise) - .set_attr("FInferCorrectLayout", ElemwiseArbitraryLayout); - -// relay.cast_like -bool CastLikeRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - ICHECK_EQ(types.size(), 3); - const auto* data = types[0].as(); - if (data == nullptr) { - ICHECK(types[0].as()) - << "cast: expect input type to be TensorType but get " << types[0]; - return false; - } - const auto* dtype_like = types[1].as(); - if (dtype_like == nullptr) { - ICHECK(types[1].as()) - << "cast: expect input type to be TensorType but get " << types[1]; - return false; - } - reporter->Assign(types[2], TensorType(data->shape, dtype_like->dtype)); - return true; -} - -Array CastLikeCompute(const Attrs& attrs, const Array& inputs, - const Type& out_type) { - return {topi::cast(inputs[0], inputs[1]->dtype)}; -} - -Expr MakeCastLike(Expr data, Expr dtype_like) { - static const Op& op = Op::Get("cast_like"); - return Call(op, {data, dtype_like}, Attrs(), {}); -} - -TVM_REGISTER_GLOBAL("relay.ir.cast_like").set_body_typed(MakeCastLike); - -RELAY_REGISTER_OP("cast_like") - .describe(R"code(Cast the data into the type of another tensor. -)code" TVM_ADD_FILELINE) - .set_num_inputs(2) - .add_argument("data", "Tensor", "The input tensor.") - .add_argument("dtype_like", "Tensor", "The tensor to cast to.") - .set_support_level(3) - .add_type_rel("CastLike", CastLikeRel) - .set_attr("FTVMCompute", CastLikeCompute) - .set_attr("TOpPattern", kElemWise) - .set_attr("FInferCorrectLayout", ElemwiseArbitraryLayout); - -Array ReinterpretCompute(const Attrs& attrs, const Array& inputs, - const Type& out_type) { - const CastAttrs* param = attrs.as(); - ICHECK(param != nullptr); - DataType dtype = param->dtype; - return {topi::reinterpret(inputs[0], dtype)}; -} - -Expr MakeReinterpret(Expr data, DataType dtype) { - auto attrs = make_object(); - attrs->dtype = dtype; - static const Op& op = Op::Get("reinterpret"); - return Call(op, {data}, Attrs(attrs), {}); -} - -TVM_REGISTER_GLOBAL("relay._make.reinterpret").set_body_typed(MakeReinterpret); - -RELAY_REGISTER_OP("reinterpret") - .describe(R"code(Reinterpret the data into a new data type. -)code" TVM_ADD_FILELINE) - .set_num_inputs(1) - .set_attrs_type() - .add_argument("data", "Tensor", "The input tensor.") - .set_support_level(3) - .add_type_rel("Reinterpret", CastRel) - .set_attr("FTVMCompute", ReinterpretCompute) - .set_attr("TOpPattern", kElemWise) - .set_attr("FInferCorrectLayout", ElemwiseArbitraryLayout); - -// relay.expand_dims -TVM_REGISTER_NODE_TYPE(ExpandDimsAttrs); - -bool ExpandDimsRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - // `types` contains: [data, result] - ICHECK_EQ(types.size(), 2); - const auto* data = types[0].as(); - if (data == nullptr) { - ICHECK(types[0].as()) - << "expand_dims: expect input type to be TensorType but get " << types[0]; - return false; - } - const auto* param = attrs.as(); - const int ndim = static_cast(data->shape.size()); - const int axis = param->axis; - const int num_newaxis = param->num_newaxis; - ICHECK(num_newaxis >= 0) << "expand_dims only accepts `num_newaxis >= 0`" - << ", but got num_newaxis = " << num_newaxis; - ICHECK(-ndim - 1 <= axis && axis <= ndim) - << "expand_dims only accepts `axis` in [-data.ndim - 1, data.ndim]" - << ", but got axis = " << axis << ", and data.ndim = " << ndim; - const int pivot = axis < 0 ? ndim + axis + 1 : axis; - std::vector oshape; - oshape.reserve(ndim + num_newaxis); - for (int i = 0; i < pivot; ++i) { - oshape.emplace_back(data->shape[i]); - } - for (int i = 0; i < num_newaxis; ++i) { - oshape.emplace_back(1); - } - for (int i = pivot; i < ndim; ++i) { - oshape.emplace_back(data->shape[i]); - } - reporter->Assign(types[1], TensorType(oshape, data->dtype)); - return true; -} - -Array ExpandDimsCompute(const Attrs& attrs, const Array& inputs, - const Type& out_type) { - const ExpandDimsAttrs* param = attrs.as(); - ICHECK(param != nullptr); - return {topi::expand_dims(inputs[0], param->axis, param->num_newaxis)}; -} - -Expr MakeExpandDims(Expr data, int axis, int num_newaxis) { - auto attrs = make_object(); - attrs->axis = axis; - attrs->num_newaxis = num_newaxis; - static const Op& op = Op::Get("expand_dims"); - return Call(op, {data}, Attrs(attrs), {}); -} - -TVM_REGISTER_GLOBAL("relay.op._make.expand_dims").set_body_typed(MakeExpandDims); - -RELAY_REGISTER_OP("expand_dims") - .describe(R"code(Insert `num_newaxis` axes at the position given by `axis` - -- **data**: The input data to the operator. - -)code" TVM_ADD_FILELINE) - .set_num_inputs(1) - .set_attrs_type() - .add_argument("data", "Tensor", "The input tensor.") - .set_support_level(1) - .add_type_rel("ExpandDims", ExpandDimsRel) - .set_attr("FTVMCompute", ExpandDimsCompute) - .set_attr("TOpPattern", kBroadcast) - .set_attr("TReshapeOp", true); - -// relay.concatenate -TVM_REGISTER_NODE_TYPE(ConcatenateAttrs); - -Array ConcatenateCompute(const Attrs& attrs, const Array& inputs, - const Type& out_type) { - const ConcatenateAttrs* param = attrs.as(); - ICHECK(param != nullptr); - return {topi::concatenate(inputs, param->axis)}; -} - -Expr MakeConcatenate(Expr data, int axis) { - auto attrs = make_object(); - attrs->axis = axis; - static const Op& op = Op::Get("concatenate"); - return Call(op, {data}, Attrs(attrs), {}); -} - -TVM_REGISTER_GLOBAL("relay.op._make.concatenate").set_body_typed(MakeConcatenate); - -RELAY_REGISTER_OP("concatenate") - .describe(R"code(Concatenate the input tensors along the given axis. - -- **data** : A list of tensors. - -- **axis** : The axis along which the tensors are concatenated. - -)code" TVM_ADD_FILELINE) - .set_attrs_type() - .set_num_inputs(1) - .add_argument("data", "Tensor", "The input list of tensors.") - .set_support_level(1) - .add_type_rel("Concatenate", ConcatenateRel) - .set_attr("FInferCorrectLayout", ConcatenateLayout) - .set_attr("TOpPattern", kInjective); - -TVM_REGISTER_NODE_TYPE(StackAttrs); - -bool StackRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - // types: [data, result] - ICHECK_EQ(types.size(), 2); - const auto* tensor_tuple = types[0].as(); - if (tensor_tuple == nullptr) { - ICHECK(types[0].as()) - << "cast: expect input type to be TupleType but get " << types[0]; - return false; - } - for (auto field : tensor_tuple->fields) { - if (field.as()) { - return false; - } - } - const auto* param = attrs.as(); - const auto& first = Downcast(tensor_tuple->fields[0]); - const int ndim = static_cast(first->shape.size()); - - // Sanity check: axis - int axis = param->axis.IntValue(); - ICHECK(-(ndim + 1) <= axis && axis < ndim + 1) - << "stack only accepts `axis` in [-(ndim+1), ndim+1)" - << ", but got axis = " << axis << ", and ndim = " << ndim; - axis = axis < 0 ? ndim + axis + 1 : axis; - - // Sanity check: ndim and dtype. - const DataType dtype = first->dtype; - for (const Type& ele : tensor_tuple->fields) { - const auto& e = Downcast(ele); - int e_ndim = static_cast(e->shape.size()); - const DataType& e_dtype = e->dtype; - ICHECK_EQ(e_ndim, ndim) << "relay.stack requires all tensors have the same ndim"; - ICHECK_EQ(e_dtype, dtype) << "relay.stack requires all tensors have the same dtype"; - for (size_t j = 0; j < first->shape.size(); ++j) { - if (j == static_cast(axis)) continue; - if (first->shape[j].as() || e->shape[j].as() || - reporter->AssertEQ(first->shape[j], e->shape[j])) - continue; - throw CompileError( - "relay.stack requires all tensors have the same shape " - "on non-stacking axes"); - } - } - - // Calculate shape - std::vector oshape; - oshape.reserve(ndim + 1); - const int stack_dim = static_cast(tensor_tuple->fields.size()); - for (int i = 0; i < axis; ++i) { - oshape.emplace_back(first->shape[i]); - } - oshape.emplace_back(stack_dim); - for (int i = axis; i < ndim; ++i) { - oshape.emplace_back(first->shape[i]); - } - reporter->Assign(types[1], TensorType(oshape, dtype)); - return true; -} - -Array StackCompute(const Attrs& attrs, const Array& inputs, - const Type& out_type) { - const StackAttrs* param = attrs.as(); - ICHECK(param != nullptr); - return {topi::stack(inputs, param->axis.IntValue())}; -} - -Expr MakeStack(Expr data, int axis) { - auto attrs = make_object(); - attrs->axis = axis; - static const Op& op = Op::Get("stack"); - return Call(op, {data}, Attrs(attrs), {}); -} - -TVM_REGISTER_GLOBAL("relay.op._make.stack").set_body_typed(MakeStack); - -RELAY_REGISTER_OP("stack") - .describe(R"code(Stack the input tensors along the given axis. - -- **data** : A list of tensors. - -- **axis** : The axis along which the tensors are stacked. - -)code" TVM_ADD_FILELINE) - .set_attrs_type() - .set_num_inputs(1) - .add_argument("data", "Tensor", "The input list of tensors.") - .set_support_level(3) - .add_type_rel("Stack", StackRel) - .set_attr("FTVMCompute", StackCompute) - .set_attr("TOpPattern", kInjective); - -/* relay.transpose */ -TVM_REGISTER_NODE_TYPE(TransposeAttrs); - -bool TransposeRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - // types: [data, result] - ICHECK_EQ(types.size(), 2); - const auto* data = types[0].as(); - if (data == nullptr) { - ICHECK(types[0].as()) - << "transpose: expect input type to be TensorType but get " << types[0]; - return false; - } - const auto* param = attrs.as(); - const int ndim = data->shape.size(); - const Array& axes = param->axes; - // check dimension match - ICHECK(!axes.defined() || static_cast(axes.size()) == ndim) - << "Dimension mismatch: axes has " << axes.size() << " elements" - << ", but data.ndim = " << ndim; - // construct int_axes - std::vector int_axes; - int_axes.reserve(ndim); - // used not defined to check if it is None. - if (!axes.defined()) { - for (int i = ndim - 1; i >= 0; --i) { - int_axes.push_back(i); - } - } else { - std::vector axis_used(ndim, 0); - for (const Integer& e : axes) { - int64_t axis = e.IntValue(); - // sanity check for axis and ndim - ICHECK(-ndim <= axis && axis < ndim) - << "transpose only allows each `axis` in `axes` in range [-data.ndim, data.ndim)" - << ", but got axis = " << axis << ", and data.ndim = " << ndim; - axis = axis < 0 ? axis + ndim : axis; - // sanity check for duplication - ICHECK(!axis_used[axis]) << "Duplicate axes in transpose: " << axis; - axis_used[axis] = 1; - int_axes.push_back(static_cast(axis)); - } - } - std::vector oshape; - oshape.reserve(ndim); - for (int axis : int_axes) { - oshape.push_back(data->shape[axis]); - } - reporter->Assign(types[1], TensorType(oshape, data->dtype)); - return true; -} - -InferCorrectLayoutOutput TransposeInferCorrectLayout(const Attrs& attrs, - const Array& new_in_layouts, - const Array& old_in_layouts, - const Array& old_in_types) { - const auto* attrs_ptr = attrs.as(); - ICHECK(attrs_ptr); - ObjectPtr params = make_object(*attrs_ptr); - - std::string in_layout_str = ""; - std::string out_layout_str = ""; - - // Infer the input layout string and update the axes. - if (old_in_layouts.defined() && old_in_layouts[0].defined()) { - ICHECK_EQ(old_in_layouts.size(), 1); - auto old_layout = old_in_layouts[0]; - Array old_axes = params->axes; - - // Deal with default axes and negative axes. - if (!old_axes.defined() || old_axes.size() == 0) { - for (int i = old_layout.ndim() - 1; i >= 0; --i) { - old_axes.push_back(i); - } - } - for (size_t i = 0; i < old_axes.size(); ++i) { - int axis = static_cast(old_axes[i]->value); - if (axis < 0) { - int pos_axis = static_cast(old_layout.ndim()) + axis; - old_axes.Set(i, pos_axis); - } - } - - if (new_in_layouts.defined() && new_in_layouts[0].defined()) { - ICHECK_EQ(new_in_layouts.size(), 1); - auto new_layout = new_in_layouts[0]; - - // Update the axes based on the new layout. - Array new_axes = Array(); - for (auto axis : old_axes) { - auto new_axis = new_layout.IndexOf(old_layout[axis->value]); - if (new_axis == -1) { // Cannot find the target axis in the new layout. - new_axes.clear(); - break; - } - new_axes.push_back(new_axis); - } - if (new_axes.defined() && new_axes.size() == new_layout.ndim()) { - params->axes = std::move(new_axes); - in_layout_str = new_layout.name(); - } - } - - // If the input layout string cannot be determined, propagate the old layout. - if (in_layout_str == "") { - params->axes = std::move(old_axes); - in_layout_str = old_layout.name(); - } - } - - // Infer the output layout string based on the input layout and the axes. - Attrs new_attrs(params); - if (in_layout_str != "") { - for (auto axis : params->axes) { - ICHECK_LT(axis->value, in_layout_str.length()); - out_layout_str += in_layout_str[axis->value]; - } - try { - return InferCorrectLayoutOutput({Layout(in_layout_str)}, {Layout(out_layout_str)}, new_attrs); - } catch (const tvm::Error& e) { - // If the layout string is invalid for any reason, give up. - return InferCorrectLayoutOutput({Layout::Undef()}, {Layout::Undef()}, attrs); - } - } - return InferCorrectLayoutOutput({Layout::Undef()}, {Layout::Undef()}, attrs); -} - -Array TransposeCompute(const Attrs& attrs, const Array& inputs, - const Type& out_type) { - const auto* param = attrs.as(); - ICHECK(param != nullptr); - return Array{topi::transpose(inputs[0], param->axes)}; -} - -Expr MakeTranspose(Expr data, Array axes) { - auto attrs = make_object(); - attrs->axes = std::move(axes); - static const Op& op = Op::Get("transpose"); - return Call(op, {data}, Attrs(attrs), {}); -} - -TVM_REGISTER_GLOBAL("relay.op._make.transpose").set_body_typed(MakeTranspose); - -RELAY_REGISTER_OP("transpose") - .describe(R"code(Permutes the dimensions of an array. - -- **data**: The input data to the operator. - -- **axes**: The target axes order, reverse order if not specified. - -)code" TVM_ADD_FILELINE) - .set_num_inputs(1) - .set_attrs_type() - .add_argument("data", "Tensor", "The input tensor.") - .set_support_level(3) - .add_type_rel("Transpose", TransposeRel) - .set_attr("FTVMCompute", TransposeCompute) - .set_attr("FInferCorrectLayout", TransposeInferCorrectLayout) - .set_attr("TOpPattern", kInjective); - -/* relay.reshape */ -TVM_REGISTER_NODE_TYPE(ReshapeAttrs); -TVM_REGISTER_NODE_TYPE(ReshapeLikeAttrs); - -Array InferNewShape(const Array& data_shape, const Attrs& attrs, - bool reverse) { - const auto* param = attrs.as(); - Array oshape; - Array ishape; - Array newshape; - - if (reverse) { - ishape.Assign(data_shape.rbegin(), data_shape.rend()); - newshape.Assign(param->newshape.rbegin(), param->newshape.rend()); - } else { - ishape = data_shape; - newshape = param->newshape; - } - - bool allowzero = param->allowzero; - - std::unordered_set used_input_dims; - std::unordered_set used_output_dims; - size_t src_idx = 0; - int infer_idx = -1; - - for (size_t i = 0; i < newshape.size(); ++i) { - int svalue = newshape[i]->value; - // special flag handling for shape inference. - if (svalue > 0) { - oshape.push_back(newshape[i]); - ++src_idx; - } else if (svalue == 0) { - if (allowzero) { - // 0 means empty tensor, thus default behavior - oshape.push_back(newshape[i]); - ++src_idx; - } else { - // 0 means to copy at equivilant position in data tensor - ICHECK_LT(src_idx, ishape.size()); - used_input_dims.insert(src_idx); - used_output_dims.insert(oshape.size()); - oshape.push_back(ishape[src_idx++]); - } - } else if (svalue == -1) { - // inference based on rest - ICHECK_LT(infer_idx, 0) << "One and only one dim can be inferred"; - infer_idx = i; - oshape.push_back(1); - ++src_idx; - } else if (svalue == -2) { - // copy all remaining dims from source - while (src_idx < ishape.size()) { - used_input_dims.insert(src_idx); - used_output_dims.insert(oshape.size()); - oshape.push_back(ishape[src_idx++]); - } - } else if (svalue == -3) { - // merge two dims from source - ICHECK_LT(src_idx + 1, ishape.size()); - used_input_dims.insert(src_idx); - IndexExpr d1 = ishape[src_idx++]; - used_input_dims.insert(src_idx); - IndexExpr d2 = ishape[src_idx++]; - used_output_dims.insert(oshape.size()); - if (d1.as() || d2.as()) { - oshape.push_back(Any()); - } else { - oshape.push_back(d1 * d2); - } - } else if (svalue == -4) { - // split the source dim s into two dims - // read the left dim and then the right dim (either can be -1) - ICHECK_LT(i + 2, newshape.size()); - ICHECK_LT(src_idx, ishape.size()); - used_input_dims.insert(src_idx); - IndexExpr d0 = ishape[src_idx++]; - Integer d1 = newshape[++i]; - Integer d2 = newshape[++i]; - if (d1->value == -1) { - ICHECK_NE(d2->value, -1) << "Split dims cannot both be -1."; - used_output_dims.insert(oshape.size()); - if (d0.as()) { - oshape.push_back(Any()); - } else { - oshape.push_back(indexdiv(d0, d2)); - } - used_output_dims.insert(oshape.size()); - oshape.push_back(d2); - } else { - used_output_dims.insert(oshape.size()); - oshape.push_back(d1); - used_output_dims.insert(oshape.size()); - if (d2->value == -1) { - if (d0.as()) { - oshape.push_back(Any()); - } else { - oshape.push_back(indexdiv(d0, d1)); - } - } else { - oshape.push_back(d2); - } - } - } else { - LOG(FATAL) << "Unsupported special value: " << svalue; - } - } - - if (infer_idx >= 0) { - IndexExpr infer_dim = 1; - for (size_t i = 0; i < ishape.size(); ++i) { - if (used_input_dims.count(i) != 0) { - continue; - } - if (ishape[i].as()) { - infer_dim = Any(); - break; - } - infer_dim *= ishape[i]; - } - if (!infer_dim.as()) { - for (size_t i = 0; i < oshape.size(); ++i) { - if (used_output_dims.count(i) != 0) { - continue; - } - if (oshape[i].as()) { - infer_dim = Any(); - break; - } - infer_dim = indexdiv(infer_dim, oshape[i]); - } - } - arith::Analyzer ana; - infer_dim = ana.Simplify(infer_dim); - oshape.Set(infer_idx, infer_dim); - } - - return oshape; -} - -bool ReshapeRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - // types: [data, result] - ICHECK_EQ(types.size(), 2); - const auto* data = types[0].as(); - if (data == nullptr) { - ICHECK(types[0].as()) - << "reshape: expect input type to be TensorType but get " << types[0]; - return false; - } - - const auto& oshape = InferNewShape(data->shape, attrs, false); - - // Verify that the sum of dimensions in the output shape is the sum of - // dimensions in the input shape - Array data_shape; - data_shape = data->shape; - - bool found_dynamic = false; - int64_t oshape_sum = 1; - for (auto& x : oshape) { - // Check if we have a dynamic shape. If we do, we can't verify if the - // reshape is valid. Dynamic shapes are marker by using Any, but can also - // occur from SizeVar's. In the case of SizeVar, the shape expression can - // be an AST. We can't easily check if we have an AST because of a ShapeVar - // or some other reason, so our check for dynamic shape is just if we can - // convert the shape to in integer or not. - if (!x->IsInstance()) { - found_dynamic = true; - break; - } - oshape_sum *= Downcast(x)->value; - } - int64_t data_shape_sum = 1; - for (auto& x : data_shape) { - if (!x->IsInstance()) { - found_dynamic = true; - break; - } - data_shape_sum *= Downcast(x)->value; - } - if (!found_dynamic && oshape_sum != data_shape_sum) { - std::ostringstream oshape_str, data_shape_str; - for (auto iter = oshape.begin(); iter != oshape.end(); iter++) { - oshape_str << (iter != oshape.begin() ? "," : "") << *iter; - } - for (auto iter = data_shape.begin(); iter != data_shape.end(); iter++) { - data_shape_str << (iter != data_shape.begin() ? "," : "") << *iter; - } - ICHECK_EQ(oshape_sum, data_shape_sum) - << "Input tensor shape(" << data_shape_str.str() << ") and reshaped shape(" - << oshape_str.str() << ") are not compatible!"; - } - - reporter->Assign(types[1], TensorType(oshape, data->dtype)); - return true; -} - -bool ReverseReshapeRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - // types: [data, result] - ICHECK_EQ(types.size(), 2); - const auto* data = types[0].as(); - if (data == nullptr) { - ICHECK(types[0].as()) - << "reshape: expect input type to be TensorType but get " << types[0]; - return false; - } - - const auto& oshape = InferNewShape(data->shape, attrs, true); - - // Verify that the sum of dimensions in the output shape is the sum of - // dimensions in the input shape - Array data_shape; - data_shape.Assign(data->shape.rbegin(), data->shape.rend()); - - bool found_dynamic = false; - int64_t oshape_sum = 1; - for (auto& x : oshape) { - // Check if we have a dynamic shape. If we do, we can't verify if the - // reshape is valid. Dynamic shapes are marker by using Any, but can also - // occur from SizeVar's. In the case of SizeVar, the shape expression can - // be an AST. We can't easily check if we have an AST because of a ShapeVar - // or some other reason, so our check for dynamic shape is just if we can - // convert the shape to in integer or not. - if (!x->IsInstance()) { - found_dynamic = true; - break; - } - oshape_sum *= Downcast(x)->value; - } - int64_t data_shape_sum = 1; - for (auto& x : data_shape) { - if (!x->IsInstance()) { - found_dynamic = true; - break; - } - data_shape_sum *= Downcast(x)->value; - } - if (!found_dynamic) { - ICHECK_EQ(oshape_sum, data_shape_sum) - << "Input tensor shape and reshaped shape are not compatible"; - } - - reporter->Assign(types[1], - TensorType(Array(oshape.rbegin(), oshape.rend()), data->dtype)); - return true; -} - -Array infer_reshape_like(const Array& lhs_shape, - const Array& rhs_shape, const Attrs& attrs) { - const auto* like_attrs = attrs.as(); - CHECK(!like_attrs->lhs_end.defined() || like_attrs->lhs_end.as()) - << "lhs_end must be a concrete integer or None"; - CHECK(!like_attrs->rhs_end.defined() || like_attrs->rhs_end.as()) - << "rhs_end must be a concrete integer or None"; - - int64_t lhs_shape_size = static_cast(lhs_shape.size()); - int64_t rhs_shape_size = static_cast(rhs_shape.size()); - int64_t lhs_begin = static_cast(like_attrs->lhs_begin); - int64_t lhs_end = - like_attrs->lhs_end.defined() ? like_attrs->lhs_end.as()->value : lhs_shape_size; - int64_t rhs_begin = static_cast(like_attrs->rhs_begin); - int64_t rhs_end = - like_attrs->rhs_end.defined() ? like_attrs->rhs_end.as()->value : rhs_shape_size; - - // handle negative axes - lhs_begin = lhs_begin < 0 ? lhs_begin + lhs_shape_size : lhs_begin; - lhs_end = lhs_end < 0 ? lhs_end + lhs_shape_size : lhs_end; - rhs_begin = rhs_begin < 0 ? rhs_begin + rhs_shape_size : rhs_begin; - rhs_end = rhs_end < 0 ? rhs_end + rhs_shape_size : rhs_end; - - Array shape_like; - for (auto i = 0; i < lhs_begin; i++) { - shape_like.push_back(lhs_shape[i]); - } - for (auto i = rhs_begin; i < rhs_end; i++) { - shape_like.push_back(rhs_shape[i]); - } - for (auto i = lhs_end; i < lhs_shape_size; i++) { - shape_like.push_back(lhs_shape[i]); - } - return shape_like; -} - -Array ReshapeCompute(const Attrs& attrs, const Array& inputs, - const Type& out_type) { - // Quick path for reshape_like - if (!attrs.as()) { - ICHECK(attrs.as() != nullptr); - auto shape_like = infer_reshape_like(inputs[0]->shape, inputs[1]->shape, attrs); - return {topi::reshape(inputs[0], shape_like)}; - } - - const auto* out_ttype = out_type.as(); - ICHECK(out_ttype != nullptr); - Array newshape; - bool newshape_has_any = false; - for (auto val : out_ttype->shape) { - if (val->IsInstance() || val->IsInstance()) { - newshape_has_any = true; - break; - } else { - newshape.push_back(val); - } - } - - if (newshape_has_any) { - newshape = InferNewShape(inputs[0]->shape, attrs, false); - } - return {topi::reshape(inputs[0], newshape)}; -} - -Expr MakeReshape(Expr data, Array newshape, bool allowzero) { - auto attrs = make_object(); - attrs->newshape = std::move(newshape); - attrs->allowzero = allowzero; - static const Op& op = Op::Get("reshape"); - return Call(op, {data}, Attrs(attrs), {}); -} - -TVM_REGISTER_GLOBAL("relay.op._make.reshape").set_body_typed(MakeReshape); - -RELAY_REGISTER_OP("reshape") - .describe(R"code(Reshapes the input array. - -Example:: - -To give user more convenience in without doing manual shape inference, -some dimensions of the shape can take special values from the set {0, -1, -2, -3, -4}. -The significance of each is explained below: - -- ``0`` copy this dimension from the input to the output shape. - -Example:: - -- data.shape = (2,3,4), newshape = (4,0,2), result.shape = (4,3,2) -- data.shape = (2,3,4), newshape = (2,0,0), result.shape = (2,3,4) - -- ``-1`` infers the dimension of the output shape by using the remainder of the input dimensions -keeping the size of the new array same as that of the input array. -At most one dimension of shape can be -1. - -Example:: - -- data.shape = (2,3,4), newshape = (6,1,-1), result.shape = (6,1,4) -- data.shape = (2,3,4), newshape = (3,-1,8), result.shape = (3,1,8) -- data.shape = (2,3,4), newshape = (-1,), result.shape = (24,) - -- ``-2`` copy all/remainder of the input dimensions to the output shape. - -Example:: - -- data.shape = (2,3,4), newshape = (-2,), result.shape = (2,3,4) -- data.shape = (2,3,4), newshape = (2,-2), result.shape = (2,3,4) -- data.shape = (2,3,4), newshape = (-2,1,1), result.shape = (2,3,4,1,1) - -- ``-3`` use the product of two consecutive dimensions of the input shape as the output dimension. - -Example:: - -- data.shape = (2,3,4), newshape = (-3,4), result.shape = (6,4) -- data.shape = (2,3,4,5), newshape = (-3,-3), result.shape = (6,20) -- data.shape = (2,3,4), newshape = (0,-3), result.shape = (2,12) -- data.shape = (2,3,4), newshape = (-3,-2), result.shape = (6,4) - -- ``-4`` split one dimension of the input into two dimensions passed subsequent to -4 in shape (can contain -1). - -Example:: - -- data.shape = (2,3,4), newshape = (-4,1,2,-2), result.shape =(1,2,3,4) -- data.shape = (2,3,4), newshape = (2,-4,-1,3,-2), result.shape = (2,1,3,4) - -)code" TVM_ADD_FILELINE) - .set_num_inputs(1) - .set_attrs_type() - .add_argument("data", "Tensor", "The input tensor.") - .set_support_level(3) - .add_type_rel("Reshape", ReshapeRel) - .set_attr("FTVMCompute", ReshapeCompute) - .set_attr("TOpPattern", kInjective) - .set_attr("TReshapeOp", true); - -/*! - * \brief ReshapeLikeRel User defined type constraint function. - * \param num_inputs Number of input types in the args. - * \param attrs The additional attributes of the operator. - * \param reporter The reporter to report solution to. - * \return False if the relation has not been resolved, it might be resolved later. - * True if this relation has been resolved. - */ -bool ReshapeLikeRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - ICHECK(attrs.as() != nullptr); - ICHECK_EQ(types.size(), 3); - const auto* data = types[0].as(); - if (data == nullptr) { - return false; - } - const auto* reshape_like = types[1].as(); - if (reshape_like == nullptr) { - return false; - } - auto shape_like = infer_reshape_like(data->shape, reshape_like->shape, attrs); - // Only check When input data has static shape. - bool is_static_shape = true; - for (size_t i = 0; i < data->shape.size(); ++i) { - if (!data->shape[i].as()) { - is_static_shape = false; - break; - } - } - auto output_type = TensorType(shape_like, data->dtype); - if (is_static_shape) { - ICHECK(reporter->AssertEQ(data->Size(), output_type->Size())) - << "Reshape inputs size should be compatible, " - << "but found data_shape " << data->shape << " not same as output_shape " - << output_type->shape; - } - reporter->Assign(types[2], output_type); - return true; -} - -Expr MakeReshapeLike(Expr lhs, Expr rhs, int lhs_begin, Integer lhs_end, int rhs_begin, - Integer rhs_end) { - auto attrs = make_object(); - attrs->lhs_begin = std::move(lhs_begin); - attrs->lhs_end = std::move(lhs_end); - attrs->rhs_begin = std::move(rhs_begin); - attrs->rhs_end = std::move(rhs_end); - static const Op& op = Op::Get("reshape_like"); - return Call(op, {lhs, rhs}, Attrs(attrs), {}); -} - -TVM_REGISTER_GLOBAL("relay.op._make.reshape_like").set_body_typed(MakeReshapeLike); - -RELAY_REGISTER_OP("reshape_like") - .describe(R"code(Reshapes the input array by the size of another array. -For an input array with shape ``(d1, d2, ..., dk)``, `reshape_like` operation reshapes -the input array into an output array with the same shape as the second input array. -.. note:: - Sizes for both array should be compatible. -Example:: - - data.shape == (1, 2, 3, 4) - shape_like.shape == (6, 2, 2, 3) - - ret = reshape_like(data, shape_like, lhs_begin=1, rhs_end=3) - ret.shape == (1, 6, 2, 2) -)code" TVM_ADD_FILELINE) - .set_attrs_type() - .set_num_inputs(2) - .add_argument("data", "Tensor", "The input tensor.") - .add_argument("shape_like", "Tensor", "Shape tensor.") - .set_support_level(3) - .add_type_rel("ReshapeLike", ReshapeLikeRel) - .set_attr("FTVMCompute", ReshapeCompute) - .set_attr("TOpPattern", kInjective); - -// ArgWhere -bool ArgWhereRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - ICHECK_EQ(num_inputs, 1); - auto tt = types[0].as(); - - if (tt == nullptr) { - return false; - } - - const auto& input_shape = tt->shape; - const auto& input_rank = input_shape.size(); - std::vector result_shape; - result_shape.push_back(Any()); - result_shape.push_back(IntImm(DataType::Int(32), input_rank)); - reporter->Assign(types[1], TensorType(result_shape, DataType::Int(32))); - return true; -} - -TVM_REGISTER_GLOBAL("relay.op._make.argwhere").set_body_typed([](Expr data) { - static const Op& op = Op::Get("argwhere"); - return Call(op, {data}, Attrs(), {}); -}); - -RELAY_REGISTER_OP("argwhere") - .describe(R"doc(Find the indices of elements of a tensor that are -non-zero)doc" TVM_ADD_FILELINE) - .set_num_inputs(1) - .add_argument("condition", "Tensor", "The input condition tensor.") - .add_type_rel("ArgWhere", ArgWhereRel) - .set_attr("TOpIsStateful", false) - .set_attr("TOpPattern", kOpaque) - .set_support_level(10); - -// scatter_elements operator -TVM_REGISTER_NODE_TYPE(ScatterElementsAttrs); - -bool ScatterElementsRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - // `types` contains: [data, indices, updates, output] - ICHECK_EQ(types.size(), 4); - const auto* data = types[0].as(); - const auto* indices = types[1].as(); - const auto* updates = types[2].as(); - if (data == nullptr) { - ICHECK(types[0].as()) - << "ScatterElements: expect input data type to be TensorType but got " << types[0]; - return false; - } - if (indices == nullptr) { - ICHECK(types[1].as()) - << "ScatterElements: expect indices type to be TensorType but got " << types[1]; - return false; - } - if (updates == nullptr) { - ICHECK(types[2].as()) - << "ScatterElements: expect updates type to be TensorType but got " << types[2]; - return false; - } - ICHECK(indices->dtype.is_int() || indices->dtype.is_uint()) - << "ScatterElements: indices must be a tensor of integers."; - - // Assign output - reporter->Assign(types[3], TensorType(data->shape, data->dtype)); - return true; -} - -Expr MakeScatterElements(Expr data, Expr indices, Expr updates, int axis, String reduction) { - auto attrs = make_object(); - attrs->axis = std::move(axis); - attrs->reduction = std::move(reduction); - static const Op& op = Op::Get("scatter_elements"); - return Call(op, {data, indices, updates}, Attrs(attrs), {}); -} - -TVM_REGISTER_GLOBAL("relay.op._make.scatter_elements").set_body_typed(MakeScatterElements); - -// scatter_elements op has extern schedules: convert to Opaque to prevent compilation failures -RELAY_REGISTER_OP("scatter_elements") - .describe(R"code(Scatter elements with updating data by reduction of values in updates -at positions defined by indices.)code" TVM_ADD_FILELINE) - .set_num_inputs(3) - .add_argument("data", "Tensor", "The input tensor.") - .add_argument("indices", "Tensor", "The indices tensor.") - .add_argument("updates", "Tensor", "The input tensor of updates.") - .set_support_level(3) - .add_type_rel("ScatterElements", ScatterElementsRel) - .set_attr("TOpPattern", kOpaque); - -// scatter_nd operator -TVM_REGISTER_NODE_TYPE(ScatterNDAttrs); - -bool ScatterNDRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - // `types` contains: [data, indices, updates, result] - ICHECK_EQ(types.size(), 4); - const auto* data = types[0].as(); - const auto* indices = types[1].as(); - const auto* updates = types[2].as(); - if (data == nullptr) { - ICHECK(types[0].as()) - << "ScatterND: expect input data type to be TensorType but got " << types[0]; - return false; - } - if (indices == nullptr) { - ICHECK(types[1].as()) - << "ScatterND: expect indices type to be TensorType but got " << types[1]; - return false; - } - if (updates == nullptr) { - ICHECK(types[2].as()) - << "ScatterND: expect updates type to be TensorType but got " << types[2]; - return false; - } - ICHECK(indices->dtype.is_int() || indices->dtype.is_uint()) - << "ScatterND: indices must be a tensor of integers."; - - const auto out_shape = data->shape; - const IntImmNode* mdim = indices->shape[0].as(); - ICHECK(mdim) << "ScatterND needs a static shape for the first axis of indices, got " - << indices->shape; - const size_t kdim = indices->shape.size() - 1; - const size_t ndim = out_shape.size(); - ICHECK_LE(size_t(mdim->value), ndim) - << "ScatterND: Given updates with shape (Y_0, ..., Y_{K-1}, X_M, ..., X_{N-1}), and indices " - "with shape (M, Y_0, ..., Y_{K-1}), M must be less than or equal to N."; - // Indices: (M, Y_0, .. Y_{K-1}) data: (Y_0, .. Y_{K-1}, ...), verify Y's. - for (size_t i = 0; i < kdim; i++) { - reporter->AssertEQ(indices->shape[i + 1], updates->shape[i]); - } - - std::vector oshape; - for (auto& x : out_shape) { - oshape.push_back(x); - } - - // updates: (Y_0, .. Y_{K-1}, X_M, .. X_{N-1}) out: (X_0, .. X_{N-1}), verify X_M to X_{N-1} - for (size_t i = mdim->value; i < ndim; i++) { - reporter->AssertEQ(updates->shape[i - mdim->value + kdim], oshape[i]); - } - - reporter->Assign(types[3], TensorType(data->shape, data->dtype)); - return true; -} - -Expr MakeScatterND(Expr data, Expr indices, Expr updates, String mode) { - auto attrs = make_object(); - attrs->mode = std::move(mode); - static const Op& op = Op::Get("scatter_nd"); - return Call(op, {data, indices, updates}, Attrs(attrs), {}); -} - -TVM_REGISTER_GLOBAL("relay.op._make.scatter_nd").set_body_typed(MakeScatterND); - -// scatter_nd operator has extern schedules for CPU and GPU devices. -// Fusing extern schedules with Injective schedules leads to errors. -// So, converting the scatter_nd to Opaque to prevent compilation failures -RELAY_REGISTER_OP("scatter_nd") - .describe(R"code(Scatter elements or slices from data and store to a tensor -whose shape is defined by indices. - -Given data with shape (Y_0, ..., Y_{K-1}, X_M, ..., X_{N-1}) and indices with shape -(M, Y_0, ..., Y_{K-1}), the output will have shape (X_0, X_1, ..., X_{N-1}). -)code" TVM_ADD_FILELINE) - .set_num_inputs(3) - .add_argument("data", "Tensor", "The input tensor.") - .add_argument("indices", "Tensor", "The indices tensor.") - .add_argument("updates", "Tensor", "The input tensor.") - .set_support_level(3) - .add_type_rel("ScatterND", ScatterNDRel) - .set_attr("TOpPattern", kOpaque); - -// Take -TVM_REGISTER_NODE_TYPE(TakeAttrs); - -bool TakeRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - // `types` contains: [data, indices, result] - ICHECK_EQ(types.size(), 3); - const auto* data = types[0].as(); - if (data == nullptr) { - return false; - } - const auto* indices = types[1].as(); - if (indices == nullptr) { - return false; - } - ICHECK(indices->dtype.is_int() || indices->dtype.is_uint()) - << "indices of take must be tensor of integer"; - const auto param = attrs.as(); - ICHECK(param != nullptr); - - if (!param->axis.defined()) { - std::vector oshape(indices->shape.begin(), indices->shape.end()); - reporter->Assign(types[2], TensorType(oshape, data->dtype)); - return true; - } - - std::vector oshape; - const auto ndim_data = static_cast(data->shape.size()); - const auto ndim_indices = static_cast(indices->shape.size()); - int axis = static_cast(param->axis->value); - int batch_dims = static_cast(param->batch_dims->value); - if (axis < 0) axis += ndim_data; - if (batch_dims < 0) axis += ndim_indices; - ICHECK_LE(axis, ndim_data) << "axis should be with in data shape" - << ", but got = " << axis; - ICHECK_LE(batch_dims, ndim_indices) << "batch_dims should be with in indices shape" - << ", but got = " << batch_dims; - ICHECK_LE(batch_dims, axis) << "batch_dims should be less than or equal to axis" - << ", but got = " << batch_dims; - - oshape.reserve(ndim_data - 1 + ndim_indices - batch_dims); - for (int i = 0; i < batch_dims; ++i) { - oshape.emplace_back(data->shape[i]); - } - for (int i = batch_dims; i < axis; ++i) { - oshape.emplace_back(data->shape[i]); - } - for (int i = batch_dims; i < ndim_indices; ++i) { - oshape.emplace_back(indices->shape[i]); - } - for (int i = axis + 1; i < ndim_data; ++i) { - oshape.emplace_back(data->shape[i]); - } - - reporter->Assign(types[2], TensorType(oshape, data->dtype)); - return true; -} - -Array TakeCompute(const Attrs& attrs, const Array& inputs, - const Type& out_type) { - const auto* param = attrs.as(); - ICHECK(param != nullptr); - if (!param->axis.defined()) { - return Array{ - topi::take(inputs[0], inputs[1], param->batch_dims.IntValue(), param->mode)}; - } else { - return Array{topi::take(inputs[0], inputs[1], param->batch_dims.IntValue(), - param->axis.IntValue(), param->mode)}; - } -} - -Expr MakeTake(Expr data, Expr indices, Integer batch_dims, Integer axis, String mode) { - auto attrs = make_object(); - attrs->batch_dims = std::move(batch_dims); - attrs->axis = std::move(axis); - attrs->mode = std::move(mode); - static const Op& op = Op::Get("take"); - return Call(op, {data, indices}, Attrs(attrs), {}); -} - -TVM_REGISTER_GLOBAL("relay.op._make.take").set_body_typed(MakeTake); - -RELAY_REGISTER_OP("take") - .describe(R"code(Take elements from an array along an axis. - -When axis is not None, this function does the same thing as 'fancy' indexing -(indexing arrays using arrays); however, it can be easier to use if you need -elements along a given axis. - -**Note** that when axis is none the flattened input array is used. - -Examples:: - - a = [[ 1, 2], - [ 3, 4]] - indices = [3, 0, 2] - take(a, indices) = [ 4, 1, 3] - - a = [[ 1., 2.], - [ 3., 4.]] - indices = [1, 0] - take(a, indices, axis=1) = [[ 2., 1.], - [ 4., 3.]] - -)code" TVM_ADD_FILELINE) - .set_attrs_type() - .set_num_inputs(2) - .add_argument("data", "Tensor", "The input tensor.") - .add_argument("indices", "Tensor", "The indices tensor.") - .set_support_level(3) - .add_type_rel("Take", TakeRel) - .set_attr("FTVMCompute", TakeCompute) - .set_attr("TOpPattern", kInjective); - -// Init ops -TVM_REGISTER_NODE_TYPE(InitOpAttrs); - -bool FullRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - ICHECK_EQ(types.size(), 2); - const InitOpAttrs* param = attrs.as(); - const auto* fill_value = types[0].as(); - if (fill_value == nullptr) { - return false; - } - - DataType out_dtype = param->dtype; - if (out_dtype.bits() == 0) { - out_dtype = fill_value->dtype; - } - - ICHECK_EQ(fill_value->shape.size(), 0) - << "Fill value should be a scalar but has dimension " << fill_value->shape.size() << "."; - - std::vector oshape; - const Array& cshape_array = param->shape.value(); - for (size_t i = 0; i < cshape_array.size(); ++i) { - oshape.push_back(cshape_array[i]); - } - reporter->Assign(types[1], TensorType(oshape, out_dtype)); - return true; -} - -Expr MakeFull(Expr fill_value, Array shape, DataType dtype) { - auto attrs = make_object(); - attrs->dtype = std::move(dtype); - attrs->shape = std::move(shape); - static const Op& op = Op::Get("full"); - return Call(op, {fill_value}, Attrs(attrs), {}); -} - -Array FullCompute(const Attrs& attrs, const Array& inputs, - const Type& out_type) { - const auto* out_ttype = out_type.as(); - return {topi::full(out_ttype->shape, out_ttype->dtype, inputs[0]())}; -} - -TVM_REGISTER_GLOBAL("relay.op._make.full").set_body_typed(MakeFull); - -RELAY_REGISTER_OP("full") - .describe(R"code(Fill array with scalar value. - -)code" TVM_ADD_FILELINE) - .set_attrs_type() - .set_num_inputs(1) - .add_argument("fill_value", "double", "The value to fill.") - .set_support_level(3) - .add_type_rel("Full", FullRel) - .set_attr("FTVMCompute", FullCompute) - .set_attr("TOpPattern", kElemWise); - -bool InitOpRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - // types = [ret_type] - ICHECK_EQ(types.size(), 1); - - const InitOpAttrs* param = attrs.as(); - ICHECK(param); - - DataType out_dtype = param->dtype; - std::vector oshape; - - const Array& cshape_array = param->shape.value(); - for (size_t i = 0; i < cshape_array.size(); ++i) { - oshape.push_back(cshape_array[i]); - } - reporter->Assign(types[0], TensorType(oshape, out_dtype)); - return true; -} - -Expr MakeZeros(Array shape, DataType dtype) { - auto attrs = make_object(); - attrs->shape = std::move(shape); - attrs->dtype = std::move(dtype); - static const Op& op = Op::Get("zeros"); - return Call(op, {}, Attrs(attrs), {}); -} - -TVM_REGISTER_GLOBAL("relay.op._make.zeros").set_body_typed(MakeZeros); - -RELAY_REGISTER_OP("zeros") - .describe(R"code(Fill array with zeros. - -)code" TVM_ADD_FILELINE) - .set_attrs_type() - .set_num_inputs(0) - .set_support_level(3) - .add_type_rel("InitOp", InitOpRel); - -Expr MakeOnes(Array shape, DataType dtype) { - auto attrs = make_object(); - attrs->shape = std::move(shape); - attrs->dtype = std::move(dtype); - static const Op& op = Op::Get("ones"); - return Call(op, {}, Attrs(attrs), {}); -} - -TVM_REGISTER_GLOBAL("relay.op._make.ones").set_body_typed(MakeOnes); - -RELAY_REGISTER_OP("ones") - .describe(R"code(Fill array with ones. - -)code" TVM_ADD_FILELINE) - .set_attrs_type() - .set_num_inputs(0) - .set_support_level(3) - .add_type_rel("InitOp", InitOpRel); - -bool FullLikeRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - ICHECK_EQ(types.size(), 3); - const auto* data = types[0].as(); - if (data == nullptr) { - return false; - } - const auto* fill_value = types[1].as(); - if (fill_value == nullptr) { - return false; - } - - ICHECK_EQ(fill_value->shape.size(), 0) - << "The fill value should be a scalar but here it has dimension " << fill_value->shape.size() - << "."; - - reporter->Assign(types[2], TensorType(data->shape, data->dtype)); - return true; -} - -Array FullLikeCompute(const Attrs& attrs, const Array& inputs, - const Type& out_type) { - return {topi::full_like(inputs[0], inputs[1]())}; -} - -Expr MakeFullLike(Expr data, Expr fill_value) { - static const Op& op = Op::Get("full_like"); - return Call(op, {data, fill_value}, Attrs(), {}); -} - -TVM_REGISTER_GLOBAL("relay.op._make.full_like").set_body_typed(MakeFullLike); - -RELAY_REGISTER_OP("full_like") - .describe(R"code(Return an scalar value array with the same shape -and type as the input array. - -)code" TVM_ADD_FILELINE) - .set_num_inputs(2) - .add_argument("data", "Tensor", "The input tensor.") - .add_argument("fill_value", "double", "Scalar value to fill.") - .set_support_level(3) - .add_type_rel("FullLike", FullLikeRel) - .set_attr("FTVMCompute", FullLikeCompute) - .set_attr("TOpPattern", kElemWise); - -// arange operator -TVM_REGISTER_NODE_TYPE(ArangeAttrs); - -bool ArangeRel(const Array& types, int num_inputs, const Attrs& raw_attrs, - const TypeReporter& reporter) { - ICHECK_EQ(types.size(), 4); - const ArangeAttrs* attrs = raw_attrs.as(); - const ConstantNode *cstart, *cstop, *cstep; - - reporter->Assign(types[0], types[1]); - reporter->Assign(types[1], types[2]); - reporter->Assign(types[2], TensorType({}, attrs->dtype)); - - if ((cstart = attrs->start.as()) && (cstop = attrs->stop.as()) && - (cstep = attrs->step.as())) { - double start = ToScalar(cstart->data); - double stop = ToScalar(cstop->data); - double step = ToScalar(cstep->data); - int32_t num_elem = static_cast(std::ceil((stop - start) / step)); - ICHECK_GT(num_elem, 0) << "Invalid arange attributes (start, stop, step): " << attrs->start - << ", " << attrs->stop << ", " << attrs->step; - reporter->Assign(types[3], TensorType({num_elem}, attrs->dtype)); - return true; - } else { - reporter->Assign(types[3], TensorType({Any()}, attrs->dtype)); - return true; - } -} - -inline te::Tensor DynamicArange(const te::Tensor& start, const te::Tensor& stop, - const te::Tensor& step, tvm::DataType dtype, - std::string name = "T_arange_dynamic", - std::string tag = topi::kInjective) { - ICHECK_EQ(start.ndim(), 0); - ICHECK_EQ(stop.ndim(), 0); - ICHECK_EQ(step.ndim(), 0); - tvm::PrimExpr num_elem = tvm::tir::Var("num_elem"); - return te::compute( - {num_elem}, - [&](const Array& indices) { - Array empty_indices; - return tvm::cast(dtype, start(empty_indices) + step(empty_indices) * indices[0]); - }, - name, tag); -} - -Array ArangeCompute(const Attrs& attrs, const Array& inputs, - const Type& out_type) { - const ArangeAttrs* param = attrs.as(); - ICHECK(param != nullptr); - te::Tensor start = inputs[0]; - te::Tensor stop = inputs[1]; - te::Tensor step = inputs[2]; - return {DynamicArange(start, stop, step, param->dtype)}; -} - -Expr MakeArange(Expr start, Expr stop, Expr step, DataType dtype) { - auto attrs = make_object(); - attrs->start = start; - attrs->stop = stop; - attrs->step = step; - attrs->dtype = dtype; - static const Op& op = Op::Get("arange"); - return Call(op, {start, stop, step}, Attrs(attrs), {}); -} - -TVM_REGISTER_GLOBAL("relay.op._make.arange").set_body_typed(MakeArange); - -// An issue with the existing design is that we require dependency -// to type the operator precisely. -// -// Supporting this in general is challenging so we duplicate the -// secondary arguments as args and attributes. -// -// In this way reify the arguments at both the value and type level. -// -// In the case our arguments are constant we can immediately recover -// the type of arange. -// -// In general I think we should avoid this pattern, and introduce -// a secondary shape analysis to recover more precise information. -RELAY_REGISTER_OP("arange") - .describe(R"code(Returns evenly spaced values within a given interval. - -)code" TVM_ADD_FILELINE) - .set_attrs_type() - .set_num_inputs(3) - .add_argument("start", "Expr", "Start of interval. The interval includes this value.") - .add_argument("end", "Expr", "Stop of interval. The interval does not include this value.") - .add_argument("step", "Expr", "Spacing between values.") - .set_support_level(3) - .add_type_rel("Arange", ArangeRel) - .set_attr("FTVMCompute", ArangeCompute) - // TODO(@icemelon): Change arange to kOpaque because FuseOps doesn't consider dynamic shape - .set_attr("TOpPattern", kOpaque) - .set_attr("AnyCodegenStrategy", kVariableDimensions); - -// repeat operator -TVM_REGISTER_NODE_TYPE(RepeatAttrs); - -bool RepeatRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - // `types` contains: [data, result] - ICHECK_EQ(types.size(), 2); - const auto* data = types[0].as(); - if (data == nullptr) { - ICHECK(types[0].as()) - << "repeat: expect input type to be TensorType but get " << types[0]; - return false; - } - const auto* param = attrs.as(); - const int ndim = static_cast(data->shape.size()); - const int repeats = param->repeats.IntValue(); - const int axis = param->axis.IntValue(); - ICHECK(repeats >= 1) << "repeat only accepts `repeats >= 1`" - << ", but got repeats = " << repeats; - ICHECK(-ndim - 1 <= axis && axis <= ndim) - << "repeat only accepts `axis` in [-data.ndim - 1, data.ndim]" - << ", but got axis = " << axis << ", and data.ndim = " << ndim; - const int pivot = axis < 0 ? ndim + axis : axis; - std::vector oshape; - oshape.reserve(ndim + repeats); - for (int i = 0; i < pivot; ++i) { - oshape.emplace_back(data->shape[i]); - } - if (data->shape[pivot].as()) { - oshape.emplace_back(Any()); - } else { - oshape.emplace_back(data->shape[pivot] * repeats); - } - for (int i = pivot + 1; i < ndim; ++i) { - oshape.emplace_back(data->shape[i]); - } - reporter->Assign(types[1], TensorType(oshape, data->dtype)); - return true; -} - -Array RepeatCompute(const Attrs& attrs, const Array& inputs, - const Type& out_type) { - const RepeatAttrs* param = attrs.as(); - ICHECK(param != nullptr); - return {topi::repeat(inputs[0], param->repeats.IntValue(), param->axis.IntValue())}; -} - -Expr MakeRepeat(Expr data, int repeats, int axis) { - auto attrs = make_object(); - attrs->repeats = repeats; - attrs->axis = axis; - static const Op& op = Op::Get("repeat"); - return Call(op, {data}, Attrs(attrs), {}); -} - -TVM_REGISTER_GLOBAL("relay.op._make.repeat").set_body_typed(MakeRepeat); - -RELAY_REGISTER_OP("repeat") - .describe(R"code(Repeat elements of an array `repeats` times along axis `axis` - -- **data**: The input data to the operator. - -)code" TVM_ADD_FILELINE) - .set_num_inputs(1) - .set_attrs_type() - .add_argument("data", "Tensor", "The input tensor.") - .set_support_level(3) - .add_type_rel("Repeat", RepeatRel) - .set_attr("FTVMCompute", RepeatCompute) - .set_attr("TOpPattern", kBroadcast); - -bool SparseFillEmptyRowsRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - // types: [sparse_indices, sparse_values, dense_shape, default_value, result] - ICHECK_EQ(types.size(), 5) << "SparseFillEmptyRowsRel expects 5 inputs but " << types.size() - << "provided"; - std::vector fields; - auto sparse_indices = types[0].as(); - auto ndims = sparse_indices->shape[1]; - fields.push_back(TensorType(Array{Any(), ndims}, tvm::DataType::Int(64))); - fields.push_back(TensorType(Array{Any()}, tvm::DataType::Int(64))); - fields.push_back(TensorType(Array{Any()}, tvm::DataType::Int(64))); - reporter->Assign(types[types.size() - 1], TupleType(Array(fields))); - return true; -} - -Expr MakeSparseFillEmptyRows(Expr sparse_indices, Expr sparse_values, Expr dense_shape, - Expr default_value) { - static const Op& op = Op::Get("sparse_fill_empty_rows"); - return Call(op, {sparse_indices, sparse_values, dense_shape, default_value}, Attrs(), {}); -} - -TVM_REGISTER_GLOBAL("relay.op._make.sparse_fill_empty_rows") - .set_body_typed(MakeSparseFillEmptyRows); - -RELAY_REGISTER_OP("sparse_fill_empty_rows") - .describe( - R"code(Fill empty rows of a sparse tensor with a default value.)code" TVM_ADD_FILELINE) - .set_num_inputs(4) - .add_argument("sparse_indices", "Tensor", - "A 2-D int64 tensor of shape [N, ndims], which specifies the indices of the" - "elements in the sparse tensor that contain nonzero values. COO Format") - .add_argument( - "sparse_values", "Tensor", - "A 1-D tensor[N] which supplies the values for each element in indices. COO Format") - .add_argument("dense_shape", "Tensor", - "A 1-D int64 tensor of shape [ndims], which specifies the dense_shape of the" - "sparse tensor. Takes a list indicating the number of elements in each " - "dimension") - .add_argument("default_value", "Tensor", - "The value to fill for empty rows, with the same type as sparse_values") - .add_type_rel("sparse_fill_empty_rows", SparseFillEmptyRowsRel) - .set_support_level(3) - .set_attr("TOpPattern", kOpaque); - -bool SparseReshapeRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - // types: [sparse_indices, prev_shape, new_shape, result] - ICHECK_EQ(types.size(), 4) << "SparseReshapeRel expects 4 types but " << types.size() - << " provided"; - ICHECK_EQ(num_inputs, 3) << "SparseReshapeRel expects 4 inputs but " << num_inputs << " provided"; - auto sparse_indices = types[0].as(); - auto prev_shape = types[1].as(); - auto new_shape = types[2].as(); - if (sparse_indices == nullptr || prev_shape == nullptr || new_shape == nullptr) { - return false; - } - CHECK(sparse_indices->dtype.is_int()) << "sparse_indices must be tensor of integers"; - CHECK(prev_shape->dtype.is_int()) << "prev_shape must be tensor of integers"; - CHECK(new_shape->dtype.is_int()) << "new_shape must be tensor of integers"; - ICHECK_EQ(sparse_indices->shape.size(), 2) << "sparse_indices must be 2-D tensor"; - ICHECK_EQ(prev_shape->shape.size(), 1) << "prev_shape must be 1-D tensor"; - ICHECK_EQ(new_shape->shape.size(), 1) << "new_shape must be 1-D tensor"; - std::vector fields; - Array new_sparse_indices_shape{sparse_indices->shape[0], new_shape->shape[0]}; - fields.push_back(TensorType(new_sparse_indices_shape, sparse_indices->dtype)); - fields.push_back(TensorType(new_shape->shape, new_shape->dtype)); - reporter->Assign(types[3], TupleType(Array(fields))); - return true; -} - -Expr MakeSparseReshape(Expr sparse_indices, Expr prev_shape, Expr new_shape) { - static const Op& op = Op::Get("sparse_reshape"); - return Call(op, {sparse_indices, prev_shape, new_shape}, Attrs(), {}); -} - -TVM_REGISTER_GLOBAL("relay.op._make.sparse_reshape").set_body_typed(MakeSparseReshape); - -RELAY_REGISTER_OP("sparse_reshape") - .describe(R"code(Return new sparse indices of the reshaped tensor -)code" TVM_ADD_FILELINE) - .set_num_inputs(3) - .add_argument("sparse_indices", "Tensor", - "A 2-D tensor of shape [N, ndims], which specifies the indices of the" - "elements in the sparse tensor that contain nonzero values. COO Format") - .add_argument("prev_shape", "Tensor", - "A 1-D tensor of shape [ndims], which specifies the previous dense shape of the" - "sparse tensor") - .add_argument("new_shape", "Tensor", - "A 1-D tensor of shape [ndims], which specifies the desired dense shape of the" - "sparse tensor") - .add_type_rel("sparse_reshape", SparseReshapeRel) - .set_attr("TOpPattern", kInjective) - .set_support_level(3); - -TVM_REGISTER_NODE_TYPE(StftAttrs); - -bool STFTRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - // types: [data, window, result] - ICHECK_EQ(types.size(), 3) << "STFTRel expects 3 types but " << types.size() << "provided"; - ICHECK_EQ(num_inputs, 2) << "Unique: expect 2 inputs but " << num_inputs << " provided"; - auto data = types[0].as(); - if (data == nullptr) { - ICHECK(types[0].as()) - << "Unique: expect input type to be TensorType but get " << types[0]; - return false; - } - const auto* param = attrs.as(); - const int ndim = static_cast(data->shape.size()); - std::vector oshape; - int dim = 0; - if (ndim == 2) { - oshape.push_back(data->shape[0]); // batch dimension - dim += 1; - } - oshape.push_back(param->onesided ? param->n_fft / 2 + 1 : param->n_fft); - if (data->shape[dim].as()) - oshape.push_back(Any()); - else - oshape.push_back(indexdiv((data->shape[dim] - param->n_fft), param->hop_length) + - 1); // n_frames - oshape.push_back(2); - reporter->Assign(types[2], TensorType(oshape, data->dtype)); - return true; -} - -Expr MakeSTFT(Expr data, int n_fft, int hop_length, int win_length, Expr window, bool normalized, - bool onesided) { - auto attrs = make_object(); - attrs->n_fft = n_fft; - attrs->hop_length = hop_length; - attrs->win_length = win_length; - attrs->normalized = normalized; - attrs->onesided = onesided; - static const Op& op = Op::Get("stft"); - return Call(op, {data, window}, Attrs(attrs), {}); -} - -TVM_REGISTER_GLOBAL("relay.op._make.stft").set_body_typed(MakeSTFT); - -RELAY_REGISTER_OP("stft") - .describe( - R"code(The STFT computes the Fourier transform of short overlapping windows of the input. -)code" TVM_ADD_FILELINE) - .set_num_inputs(2) - .add_argument("data", "Tensor", "the input tensor") - .add_argument("window", "Tensor", "the optional window function") - .add_type_rel("stft", STFTRel) - .set_support_level(3) - .set_attr("TOpPattern", kOpaque); - -// DFT -TVM_REGISTER_NODE_TYPE(DFTAttrs); -bool DFTRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - // types: [re_data, im_data, output] - ICHECK_EQ(types.size(), 3) - << "DFT: expects three types, two for the input and one for the output"; - ICHECK_EQ(num_inputs, 2) << "DFT: expect 2 inputs but " << num_inputs << " provided"; - const auto* re_data = types[0].as(); - const auto* im_data = types[1].as(); - - if (re_data == nullptr) { - ICHECK(types[0].as()) - << "DFT: expect re_data type to be TensorType but get " << types[0]; - return false; - } - if (im_data == nullptr) { - ICHECK(types[1].as()) - << "DFT: expect im_data type to be TensorType but get " << types[1]; - return false; - } - - std::vector shapes; - shapes.push_back(TensorType(re_data->shape, re_data->dtype)); - shapes.push_back(TensorType(im_data->shape, im_data->dtype)); - - reporter->Assign(types[2], TupleType(Array(shapes))); - - return true; -} - -Expr MakeDFT(Expr re_data, Expr im_data, Bool inverse) { - auto attrs = make_object(); - attrs->inverse = inverse; - static const Op& op = Op::Get("dft"); - return Call(op, {re_data, im_data}, Attrs(attrs), {}); -} - -TVM_REGISTER_GLOBAL("relay.op._make.dft").set_body_typed(MakeDFT); - -RELAY_REGISTER_OP("dft") - .describe(R"doc(Computes the discrete Fourier transform of input.)doc" TVM_ADD_FILELINE) - .set_num_inputs(2) - .add_argument("re_data", "Tensor", "Real part of input tensor.") - .add_argument("im_data", "Tensor", "Imaginary part of input tensor.") - .set_support_level(3) - .set_attr("TOpPattern", kOpaque) - .add_type_rel("DFT", DFTRel); - -// meshgrid operator -TVM_REGISTER_NODE_TYPE(MeshgridAttrs); - -bool MeshgridRel(const Array& types, int num_inputs, const Attrs& raw_attrs, - const TypeReporter& reporter) { - // types: [data, result] - ICHECK_EQ(types.size(), 2); - const MeshgridAttrs* attrs = raw_attrs.as(); - const auto* tensor_tuple = types[0].as(); - if (tensor_tuple == nullptr) { - throw CompileError(ErrorBuilder() - << "meshgrid requires a tuple of tensors as the first argument, found " - << PrettyPrint(types[0])); - } else if (types[0].as() != nullptr) { - return false; - } - const int data_length = static_cast(tensor_tuple->fields.size()); - - // Get first dtype. - const auto& first = Downcast(tensor_tuple->fields[0]); - const DataType dtype = first->dtype; - - // Get size of output grid. - std::vector grid_shape; - grid_shape.reserve(data_length); - for (const Type& ele : tensor_tuple->fields) { - if (ele.as()) { - return false; - } - const auto& e = Downcast(ele); - int e_ndim = static_cast(e->shape.size()); - const DataType& e_dtype = e->dtype; - if (e_dtype != dtype) { - throw CompileError("relay.meshgrid requires all tensors have the same dtype"); - } - if (e_ndim == 0) { - grid_shape.emplace_back(1); - } else if (e_ndim == 1) { - grid_shape.emplace_back(e->shape[0]); - } else { - throw CompileError("relay.meshgrid requires all tensors be either scalars or 1-D vectors."); - } - } - - // "xy" mode swaps first two dimensions - if (attrs->indexing == "xy" && grid_shape.size() >= 2) { - std::swap(grid_shape[0], grid_shape[1]); - } - - // There is one output grid for each input, all with same shape. - std::vector grids; - grids.reserve(data_length); - for (int i = 0; i < data_length; i++) { - grids.emplace_back(TensorType(grid_shape, dtype)); - } - reporter->Assign(types[1], TupleType(Array(grids))); - return true; -} - -Array MeshgridCompute(const Attrs& attrs, const Array& inputs, - const Type& out_type) { - const MeshgridAttrs* param = attrs.as(); - ICHECK(param != nullptr); - return {topi::meshgrid(inputs, param->indexing)}; -} - -Expr MakeMeshgrid(Expr data, String indexing) { - auto attrs = make_object(); - attrs->indexing = std::move(indexing); - static const Op& op = Op::Get("meshgrid"); - return Call(op, {data}, Attrs(attrs), {}); -} - -TVM_REGISTER_GLOBAL("relay.op._make.meshgrid").set_body_typed(MakeMeshgrid); - -RELAY_REGISTER_OP("meshgrid") - .describe(R"code(Create coordinate matrices from coordinate vectors. - -)code" TVM_ADD_FILELINE) - .set_attrs_type() - .set_num_inputs(1) - .add_argument("data", "Tensor", "The input list of tensors.") - .set_support_level(3) - .add_type_rel("Meshgrid", MeshgridRel) - .set_attr("FTVMCompute", MeshgridCompute) - .set_attr("TOpPattern", kInjective); - -// tile operator -TVM_REGISTER_NODE_TYPE(TileAttrs); - -bool TileRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - // `types` contains: [data, result] - ICHECK_EQ(types.size(), 2); - const auto* data = types[0].as(); - if (data == nullptr) { - ICHECK(types[0].as()) - << "tile: expect input type to be TensorType but get " << types[0]; - return false; - } - const auto* param = attrs.as(); - const size_t ndim = data->shape.size(); - const Array& reps = param->reps; - // check dimension match - ICHECK(reps.defined()) << "repetition array is not defined. data.ndim = " << ndim; - const size_t rndim = reps.size(); - for (size_t i = 0; i < rndim; ++i) { - if (const tvm::tir::IntImmNode* val = reps[i].as()) { - ICHECK_GT(val->value, 0) << "Tile reps value should always be larger than 0, but get: " - << val->value; - } - } - size_t tndim = (ndim > rndim) ? ndim : rndim; - // re-construct data shape or reps shape - std::vector data_shape; - std::vector reps_shape; - data_shape.reserve(tndim); - reps_shape.reserve(tndim); - if (ndim == rndim) { - for (size_t i = 0; i < tndim; ++i) { - data_shape.emplace_back(data->shape[i]); - reps_shape.emplace_back(reps[i]); - } - } else if (ndim > rndim) { - for (size_t i = 0; i < ndim; ++i) { - data_shape.emplace_back(data->shape[i]); - } - for (size_t i = 0; i < (ndim - rndim); ++i) { - reps_shape.emplace_back(1); - } - for (size_t i = 0; i < rndim; ++i) { - reps_shape.emplace_back(reps[i]); - } - } else { - for (size_t i = 0; i < rndim; ++i) { - reps_shape.emplace_back(reps[i]); - } - for (size_t i = 0; i < (rndim - ndim); ++i) { - data_shape.emplace_back(1); - } - for (size_t i = 0; i < ndim; ++i) { - data_shape.emplace_back(data->shape[i]); - } - } - std::vector oshape; - oshape.reserve(tndim); - for (size_t i = 0; i < tndim; ++i) { - // Save Any if it is dynamic shape - if (!data_shape[i].as()) { - oshape.emplace_back(Any()); - } else { - oshape.emplace_back(data_shape[i] * reps_shape[i]); - } - } - reporter->Assign(types[1], TensorType(oshape, data->dtype)); - return true; -} - -Array TileCompute(const Attrs& attrs, const Array& inputs, - const Type& out_type) { - const TileAttrs* param = attrs.as(); - ICHECK(param != nullptr); - return {topi::tile(inputs[0], param->reps)}; -} - -Expr MakeTile(Expr data, Array reps) { - auto attrs = make_object(); - attrs->reps = reps; - static const Op& op = Op::Get("tile"); - return Call(op, {data}, Attrs(attrs), {}); -} - -TVM_REGISTER_GLOBAL("relay.op._make.tile").set_body_typed(MakeTile); - -RELAY_REGISTER_OP("tile") - .describe(R"code(Repeat the whole array multiple times. - -- **data**: The input data to the operator. - -)code" TVM_ADD_FILELINE) - .set_num_inputs(1) - .set_attrs_type() - .add_argument("data", "Tensor", "The input tensor.") - .set_support_level(3) - .add_type_rel("Tile", TileRel) - .set_attr("FTVMCompute", TileCompute) - .set_attr("TOpPattern", kBroadcast); - -// reverse operator -TVM_REGISTER_NODE_TYPE(ReverseAttrs); - -bool ReverseRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - // `types` contains: [data, result] - ICHECK_EQ(types.size(), 2); - const auto* data = types[0].as(); - if (data == nullptr) { - ICHECK(types[0].as()) - << "reverse: expect input type to be TensorType but get " << types[0]; - return false; - } - const auto* param = attrs.as(); - const int ndim = static_cast(data->shape.size()); - const int axis = param->axis.IntValue(); - ICHECK(-ndim <= axis && axis < ndim) - << "reverse only accepts `axis` in [-data.ndim, data.ndim - 1]" - << ", but got axis = " << axis << ", and data.ndim = " << ndim; - reporter->Assign(types[1], types[0]); - return true; -} - -Array ReverseCompute(const Attrs& attrs, const Array& inputs, - const Type& out_type) { - const ReverseAttrs* param = attrs.as(); - ICHECK(param != nullptr); - // pass empty seq_length tensor to reverse_sequence - return {topi::reverse_sequence(inputs[0], te::Tensor(), param->axis.IntValue())}; -} - -Expr MakeReverse(Expr data, int axis) { - auto attrs = make_object(); - attrs->axis = axis; - static const Op& op = Op::Get("reverse"); - return Call(op, {data}, Attrs(attrs), {}); -} - -TVM_REGISTER_GLOBAL("relay.op._make.reverse").set_body_typed(MakeReverse); - -RELAY_REGISTER_OP("reverse") - .describe(R"code(Reverses the order of elements along given `axis` while preserving array shape. - -- **data**: The input data to the operator. - -)code" TVM_ADD_FILELINE) - .set_num_inputs(1) - .set_attrs_type() - .add_argument("data", "Tensor", "The input tensor.") - .set_support_level(3) - .add_type_rel("Reverse", ReverseRel) - .set_attr("FTVMCompute", ReverseCompute) - .set_attr("TOpPattern", kInjective); - -// reverse sequence operator -TVM_REGISTER_NODE_TYPE(ReverseSequenceAttrs); - -bool ReverseSequenceRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - // `types` contains: [data, seq_lengths, result] - ICHECK_EQ(types.size(), 3); - const auto* data = types[0].as(); - - if (data == nullptr) { - ICHECK(types[0].as()) - << "reverse_sequence: expect input type to be TensorType but get " << types[0]; - return false; - } - - const auto* seq_lengths = types[1].as(); - if (seq_lengths == nullptr) { - ICHECK(types[1].as()) - << "reverse_sequence: expect input type to be TensorType but get " << types[1]; - return false; - } - - const int seq_lengths_dim = static_cast(seq_lengths->shape.size()); - ICHECK(seq_lengths_dim == 1) << "For reverse_sequnece, seq_lengths must be a 1D vector"; - ICHECK(seq_lengths->dtype.is_int()) - << "For reverse_sequnece, seq_lengths must be tensor of integer"; - - const auto* param = attrs.as(); - const int ndim = static_cast(data->shape.size()); - int batch_axis = param->batch_axis.IntValue(); - ICHECK(-ndim <= batch_axis && batch_axis < ndim) - << "reverse_sequence only accepts `batch_axis` in [-data.ndim, data.ndim - 1]" - << ", but got batch_axis = " << batch_axis << ", and data.ndim = " << ndim; - - if (batch_axis < 0) { - batch_axis = static_cast(data->shape.size()) + batch_axis; - } - ICHECK(reporter->Assert(seq_lengths->shape[0] == data->shape[batch_axis])) - << "For reverse_sequnece seq_lengths size should match with dimension of batch axis" - << ", but got dimension of batch_axis = " << data->shape[batch_axis] - << ", and seq_length size = " << seq_lengths->shape[0]; - - const int seq_axis = param->seq_axis.IntValue(); - ICHECK(-ndim <= seq_axis && seq_axis < ndim) - << "reverse_sequnece only accepts `seq_axis` in [-data.ndim, data.ndim - 1]" - << ", but got seq_axis = " << seq_axis << ", and data.ndim = " << ndim; - - reporter->Assign(types[2], types[0]); - return true; -} - -Array ReverseSequenceCompute(const Attrs& attrs, const Array& inputs, - const Type& out_type) { - const ReverseSequenceAttrs* param = attrs.as(); - ICHECK(param != nullptr); - return {topi::reverse_sequence(inputs[0], inputs[1], param->seq_axis.IntValue(), - param->batch_axis.IntValue())}; -} - -Expr MakeReverseSequence(Expr data, Expr seq_lengths, int seq_axis, int batch_axis) { - auto attrs = make_object(); - attrs->seq_axis = seq_axis; - attrs->batch_axis = batch_axis; - static const Op& op = Op::Get("reverse_sequence"); - return Call(op, {data, seq_lengths}, Attrs(attrs), {}); -} - -TVM_REGISTER_GLOBAL("relay.op._make.reverse_sequence").set_body_typed(MakeReverseSequence); - -RELAY_REGISTER_OP("reverse_sequence") - .describe(R"code(Reverses the tensor for variable length slices. -Input is first sliced along batch axis and then elements are reversed along seq axis. - -- **data**: The input data to the operator. - -- **seq_lengths**: A 1D Tensor with length data.dims[batch_axis]. - -- **seq_axis**: The axis along which the elements will be reversed. Default is 1. - -- **batch_axis**: The axis along which the tensor will be sliced. Default is 0. - -)code" TVM_ADD_FILELINE) - .set_num_inputs(2) - .set_attrs_type() - .add_argument("data", "Tensor", "The input tensor.") - .add_argument("seq_lengths", "Tensor", "A 1D Tensor with length data.dims[batch_axis]") - .set_support_level(3) - .add_type_rel("ReverseSequence", ReverseSequenceRel) - .set_attr("FTVMCompute", ReverseSequenceCompute) - .set_attr("TOpPattern", kInjective); - -// where operator -bool WhereRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - ICHECK_EQ(types.size(), 4U); - const auto* condition = types[0].as(); - const auto* x = types[1].as(); - const auto* y = types[2].as(); - - if (condition == nullptr || x == nullptr || y == nullptr) { - return false; - } - - ICHECK_EQ(x->dtype, y->dtype) << "x and y must have the same dtype: " << x->dtype << " vs " - << y->dtype; - - auto tensor_ty_condition = GetRef(condition); - auto tensor_ty_x = GetRef(x); - auto tensor_ty_y = GetRef(y); - - auto b_ty = ConcreteBroadcast(tensor_ty_x, tensor_ty_y, x->dtype); - auto ret_ty = ConcreteBroadcast(tensor_ty_condition, b_ty, b_ty->dtype); - - reporter->Assign(types[3], ret_ty); - return true; -} - -// Positional relay function to create where operator. -Expr MakeWhere(const Expr& condition, const Expr& x, const Expr& y) { - static const Op& op = Op::Get("where"); - return Call(op, {condition, x, y}); -} - -Array WhereCompute(const Attrs& attrs, const Array& inputs, - const Type& out_type) { - return {topi::where(inputs[0], inputs[1], inputs[2])}; -} - -TVM_REGISTER_GLOBAL("relay.op._make.where").set_body_typed(MakeWhere); - -RELAY_REGISTER_OP("where") - .describe(R"code( -Return the elements, either from x or y, depending on the condition. - -Given three ndarrays, condition, x, and y, return an ndarray with the elements -from x or y, depending on the elements from condition are true or false. - -Shapes of condition, x, and y must be broadcastable to a common shape, which -is the output shape of this op. Semantics follow numpy where function. -https://numpy.org/doc/stable/reference/generated/numpy.where.html - -Note that all non-zero values are interpreted as True in condition. - -Examples:: - - x = [[1, 2], [3, 4]] - y = [[5, 6], [7, 8]] - cond = [[0, 1], [-1, 0]] - where(cond, x, y) = [[5, 2], [3, 8]] - - - cond = [[1], [0]] - where(cond, x, y) = [[1, 2], [7, 8]] - - cond = [0, 1] - where(cond, 1, -1) = [-1, 1] - -)code" TVM_ADD_FILELINE) - .add_argument("condition", "Tensor", "Condition array") - .add_argument("x", "Tensor", "First array to be selected") - .add_argument("y", "Tensor", "Second array to be selected") - .set_num_inputs(3) - .set_support_level(4) - .add_type_rel("Where", WhereRel) - .set_attr("FTVMCompute", WhereCompute) - .set_attr("TOpPattern", kBroadcast); - -// Squeeze -TVM_REGISTER_NODE_TYPE(SqueezeAttrs); - -Expr MakeSqueeze(Expr data, Array axis) { - auto attrs = make_object(); - attrs->axis = std::move(axis); - static const Op& op = Op::Get("squeeze"); - return Call(op, {data}, Attrs(attrs), {}); -} - -TVM_REGISTER_GLOBAL("relay.op._make.squeeze").set_body_typed(MakeSqueeze); - -bool SqueezeRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - ICHECK_EQ(types.size(), 2); - const auto* data = types[0].as(); - if (data == nullptr) { - return false; - } - const auto* param = attrs.as(); - ICHECK(param != nullptr); - std::vector result_shape; - // if axes is None, squeeze all axes of dimension 1 - if (!param->axis.defined()) { - for (const auto& e : data->shape) { - if (!e.as()) { - LOG(FATAL) << "axis needs to be defined for dynamic input."; - } - const int64_t* axis_ptr = tir::as_const_int(e); - ICHECK(axis_ptr != nullptr) << "the axes attribute must be concrete"; - if (*axis_ptr != 1) { - result_shape.push_back(e); - } - } - } else { - // pair up original shape with a boolean which control whether it will be in the final shape. - std::vector> original_shape; - for (const auto& e : data->shape) { - original_shape.push_back(std::pair(e, true)); - } - for (const auto& e : param->axis) { - int64_t axis_val = e->value; - if (axis_val < 0) { - axis_val += static_cast(original_shape.size()); - } - ICHECK_GE(axis_val, 0); - ICHECK_LT(axis_val, original_shape.size()); - original_shape.at(axis_val).second = false; - } - for (const auto& p : original_shape) { - if (p.second) { - result_shape.push_back(p.first); - } else { - if (const int64_t* axis_ptr = tir::as_const_int(p.first)) { - ICHECK_EQ(*axis_ptr, 1) << "cannot squeeze axis with dimension not equal to 1"; - } - } - } - } - reporter->Assign(types[1], TensorType(result_shape, data->dtype)); - return true; -} - -Array SqueezeCompute(const Attrs& attrs, const Array& inputs, - const Type& out_type) { - const SqueezeAttrs* param = attrs.as(); - ICHECK(param != nullptr); - return {topi::squeeze(inputs[0], param->axis)}; -} - -InferCorrectLayoutOutput SqueezeInferCorrectLayout(const Attrs& attrs, - const Array& new_in_layouts, - const Array& old_in_layouts, - const Array& old_in_types) { - const auto* attrs_ptr = attrs.as(); - ICHECK(attrs_ptr); - ObjectPtr params = make_object(*attrs_ptr); - - Layout inferred_input = new_in_layouts.defined() ? new_in_layouts[0] : old_in_layouts[0]; - Layout inferred_output = inferred_input; - - ICHECK(old_in_types[0].as()); - const auto& shape = old_in_types[0].as()->shape; - - // axis to squeeze - Array axis; - if (params->axis.defined()) { - axis = params->axis; - } else { - // if axes is None, squeeze all axes of dimension 1 - for (size_t i = 0; i < shape.size(); i++) { - if (topi::detail::GetConstInt(shape[i]) == 1) { - axis.push_back(i); - } - } - } - - // If new_in_layouts are defined, this code tries to modify the layout - if (new_in_layouts.defined() && old_in_layouts.defined()) { - Array new_axis; - for (const auto& e : axis) { - const auto& dim = old_in_layouts[0][e.IntValue()]; - new_axis.push_back((new_in_layouts[0]).IndexOf(dim)); - } - params->axis = new_axis; - axis = new_axis; - } - - // Infer output layout - Array kept_axes; - for (size_t i = 0; i < inferred_input.ndim(); i++) { - bool is_dim_kept = true; - - // Check whether the dim should be kept - for (const auto& e : axis) { - int64_t axis_val = e->value; - if (axis_val < 0) { - axis_val += inferred_input.ndim(); - } - if (static_cast(i) == axis_val) { - is_dim_kept = false; - break; - } - } - - if (is_dim_kept) { - kept_axes.push_back(inferred_input->axes[i]); - } - } - inferred_output = Layout(kept_axes); - - return InferCorrectLayoutOutput({inferred_input}, {inferred_output}, Attrs(params)); -} - -RELAY_REGISTER_OP("squeeze") - .describe(R"code(Squeeze the input tensor at the dimensions given by axes - -- **data**: The input data to the operator. - -)code" TVM_ADD_FILELINE) - .set_num_inputs(1) - .set_attrs_type() - .add_argument("data", "Tensor", "The input tensor.") - .set_support_level(3) - .add_type_rel("Squeeze", SqueezeRel) - .set_attr("FTVMCompute", SqueezeCompute) - .set_attr("TOpPattern", kInjective) - .set_attr("FInferCorrectLayout", SqueezeInferCorrectLayout) - .set_attr("TReshapeOp", true); - -// CollapseSumLike: -> B where BroadCast(A, B) = A -bool CollapseSumLikeRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - ICHECK_EQ(types.size(), 3); - reporter->Assign(types[2], types[1]); - return BroadcastRel({types[0], types[1], types[0]}, 2, Attrs(), reporter); -} - -Expr MakeCollapseSumLike(Expr data, Expr collapse_type) { - static const Op& op = Op::Get("collapse_sum_like"); - return Call(op, {data, collapse_type}, Attrs(), {}); -} - -Array CollapseSumLikeCompute(const Attrs& attrs, const Array& inputs, - const Type& out_type) { - const auto* out_ttype = out_type.as(); - ICHECK(out_ttype != nullptr); - return {topi::collapse_sum(inputs[0], out_ttype->shape)}; -} - -TVM_REGISTER_GLOBAL("relay.op._make.collapse_sum_like").set_body_typed(MakeCollapseSumLike); - -RELAY_REGISTER_OP("collapse_sum_like") - .describe(R"code(Collapse the first input to match the shape of the second input. -)code" TVM_ADD_FILELINE) - .set_num_inputs(2) - .add_argument("data", "Tensor", "The input tensor.") - .add_argument("collapse_type", "Tensor", "Provide the type to collapse to.") - .set_support_level(10) - .add_type_rel("CollapseSumLike", CollapseSumLikeRel) - .set_attr("FTVMCompute", CollapseSumLikeCompute) - .set_attr("TOpPattern", kCommReduce); - -// CollapseSumTo: -> B where Broadcast(A, B) = A -bool CollapseSumToRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - ICHECK_EQ(types.size(), 3); - const InitOpAttrs* param = attrs.as(); - - const auto* target_shape = types[1].as(); - DataType out_dtype = types[0].as()->dtype; - - const IntImmNode* rank = target_shape->shape[0].as(); - ICHECK(rank) << "Parameter must have static rank"; - - std::vector oshape; - if (param->shape) { - const Array& cshape_array = param->shape.value(); - for (size_t i = 0; i < cshape_array.size(); i++) { - oshape.push_back(cshape_array[i]); - } - } else { - for (int i = 0; i < rank->value; i++) { - oshape.push_back(Any()); - } - } - reporter->Assign(types[2], TensorType(oshape, out_dtype)); - return BroadcastRel({types[0], types[2], types[0]}, 2, Attrs(), reporter); -} - -Expr MakeCollapseSumTo(Expr data, Expr shape) { - static const Op& op = Op::Get("collapse_sum_to"); - auto attrs = make_object(); - if (const auto* cshape = shape.as()) { - attrs->shape = ToVector(cshape->data); - } - return Call(op, {data, shape}, Attrs(attrs), {}); -} - -TVM_REGISTER_GLOBAL("relay.op._make.collapse_sum_to").set_body_typed(MakeCollapseSumTo); - -RELAY_REGISTER_OP("collapse_sum_to") - .describe(R"code(Broadcast the first input to match the shape argument. -)code" TVM_ADD_FILELINE) - .set_num_inputs(2) - .add_argument("data", "Tensor", "The input tensor.") - .add_argument("shape", "Tensor", "Target shape.") - .set_support_level(4) - .add_type_rel("CollapseSumTo", CollapseSumToRel) - .set_attr("FTVMCompute", CollapseSumLikeCompute) - .set_attr("TOpPattern", kCommReduce); - -bool BroadCastToRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - // types = [data_type, ret_type], broadcast_to_type is in attrs bc static - ICHECK_EQ(types.size(), 2); - - const InitOpAttrs* param = attrs.as(); - ICHECK(param); - - DataType out_dtype; - if (auto ttype = types[0].as()) { - out_dtype = ttype->dtype; - } else { - ICHECK(types[0].as()) - << "Broadcast: expect to be TensorType but get " << types[0]; - return false; - } - - std::vector oshape; - - const Array& cshape_array = param->shape.value(); - for (size_t i = 0; i < cshape_array.size(); ++i) { - oshape.push_back(cshape_array[i]); - } - reporter->Assign(types[1], TensorType(oshape, out_dtype)); - return BroadcastRel({types[0], types[1], types[1]}, 2, Attrs(), reporter); -} - -Expr MakeBroadCastTo(Expr data, Array shape) { - static const Op& op = Op::Get("broadcast_to"); - auto attrs = make_object(); - - attrs->shape = std::move(shape); - return Call(op, {data}, Attrs(attrs), {}); -} - -Array BroadCastToCompute(const Attrs& attrs, const Array& inputs, - const Type& out_type) { - const auto* out_ttype = out_type.as(); - return {topi::broadcast_to(inputs[0], out_ttype->shape)}; -} - -TVM_REGISTER_GLOBAL("relay.op._make.broadcast_to").set_body_typed(MakeBroadCastTo); - -RELAY_REGISTER_OP("broadcast_to") - .describe(R"code(Broadcast the first input to match the shape argument. -)code" TVM_ADD_FILELINE) - .set_num_inputs(1) - .add_argument("data", "Tensor", "The input tensor.") - .set_support_level(4) - .add_type_rel("BroadCastTo", BroadCastToRel) - .set_attrs_type() - .set_attr("FTVMCompute", BroadCastToCompute) - .set_attr("TOpPattern", kBroadcast); - -// BroadCastToLike: -> B where BroadCast(A, B) = B -bool BroadCastToLikeRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - ICHECK_EQ(types.size(), 3); - reporter->Assign(types[2], types[1]); - return BroadcastRel({types[0], types[1], types[1]}, 2, Attrs(), reporter); -} - -Expr MakeBroadCastToLike(Expr data, Expr broadcast_type) { - static const Op& op = Op::Get("broadcast_to_like"); - return Call(op, {data, broadcast_type}, Attrs(), {}); -} - -Array BroadCastToLikeCompute(const Attrs& attrs, const Array& inputs, - const Type& out_type) { - const auto* out_ttype = out_type.as(); - ICHECK(out_ttype != nullptr); - return {topi::broadcast_to(inputs[0], out_ttype->shape)}; -} - -TVM_REGISTER_GLOBAL("relay.op._make.broadcast_to_like").set_body_typed(MakeBroadCastToLike); - -RELAY_REGISTER_OP("broadcast_to_like") - .describe(R"code(Broadcast the first input to match the shape of the second input. -)code" TVM_ADD_FILELINE) - .set_num_inputs(2) - .add_argument("data", "Tensor", "The input tensor.") - .add_argument("broadcast_type", "Tensor", "Provide the type to broadcast to.") - .set_support_level(10) - .add_type_rel("BroadCastToLike", BroadCastToLikeRel) - .set_attr("FTVMCompute", BroadCastToLikeCompute) - .set_attr("TOpPattern", kBroadcast); - -// Adapter function to make int array. -Array GetIntArray(Array arr) { - for (size_t i = 0; i < arr.size(); ++i) { - ICHECK(!arr[i].defined() || arr[i].as()) << "Expect an int array"; - } - return Downcast>(arr); -} - -// strided_slice -TVM_REGISTER_NODE_TYPE(StridedSliceAttrs); - -bool StridedSliceRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - ICHECK_EQ(types.size(), 2); - const StridedSliceAttrs* param = attrs.as(); - if (param == nullptr) { - return false; - } - const auto* data = types[0].as(); - - if (data == nullptr) { - return false; - } - - ICHECK(param->begin) << "strided_slice received invalid begin " << param->begin; - ICHECK(param->end) << "strided_slice received invalid end " << param->end; - ICHECK(param->strides) << "strided_slice received invalid strides " << param->strides; - - auto begin = param->begin.value(); - auto end = param->end.value(); - auto strides = param->strides.value(); - - const size_t src_tensor_dim = static_cast(data->shape.size()); - Array axes; - if (param->axes) { - axes = param->axes.value(); - ICHECK(axes.size() == begin.size() && axes.size() == end.size() && - axes.size() == strides.size()) - << "axes, begin, end, and strides must have the same length"; - } else { - for (size_t i = 0; i < src_tensor_dim; ++i) axes.push_back(i); - - const IntImm one = IntImm(DataType::Int(64), 1); - const IntImm zero = IntImm(DataType::Int(64), 0); - const IntImm max_range = IntImm(DataType::Int(64), std::numeric_limits::max()); - - for (size_t i = strides.size(); i < src_tensor_dim; ++i) { - strides.push_back(one); - } - for (size_t i = begin.size(); i < src_tensor_dim; ++i) { - begin.push_back(topi::GetConstInt(strides[i]) > 0 ? zero : max_range); - } - for (size_t i = end.size(); i < src_tensor_dim; ++i) { - end.push_back(topi::GetConstInt(strides[i]) < 0 ? zero : max_range); - } - } - auto oshape = - topi::StridedSliceOutputShape(data->shape, begin, end, strides, axes, param->slice_mode); - reporter->Assign(types[1], TensorType(oshape, data->dtype)); - return true; -} - -InferCorrectLayoutOutput StridedSliceInferCorrectLayout( - const Attrs& attrs, const Array& new_in_layouts, const Array& old_in_layouts, - const Array& old_in_types) { - Array> old_in_shapes; - for (auto old_in_t : old_in_types) { - ICHECK(old_in_t.as()); - old_in_shapes.push_back(old_in_t.as()->shape); - } - - ICHECK(old_in_layouts.defined()); - ICHECK_GE(old_in_layouts.size(), 1); - ICHECK(old_in_shapes.defined()); - ICHECK_GE(old_in_shapes.size(), 1); - - auto layout = old_in_layouts[0]; - InferCorrectLayoutOutput out_default{{Layout::Undef()}, {Layout::Undef()}, attrs}; - - if (layout.defined() && new_in_layouts.defined()) { - ICHECK_GE(new_in_layouts.size(), 1); - auto new_layout = new_in_layouts[0]; - auto shape = old_in_shapes[0]; - - const auto* attrs_ptr = attrs.as(); - ICHECK(attrs_ptr); - ObjectPtr params = make_object(*attrs_ptr); - - Array begin, end, strides; - if (params->begin && params->end && params->strides) { - for (Integer i : params->strides.value()) { - ICHECK(i.defined()); - auto slice_val = Integer(IntImm(i->dtype, i->value)); - strides.push_back(params->slice_mode == "size" ? Integer(IntImm(i->dtype, 1)) : slice_val); - } - - for (Integer i : params->begin.value()) { - ICHECK(i.defined()); - begin.push_back(IntImm(i->dtype, i->value)); - } - for (Integer i : params->end.value()) { - ICHECK(i.defined()); - end.push_back(IntImm(i->dtype, i->value)); - } - } - - Array new_begin, new_end, new_strides; - - // Handles layout conversion like NHWC -> NCHW - auto old_layout_name = layout.name(); - auto new_layout_name = new_layout.name(); - - if (old_layout_name.rfind(new_layout_name, 0) != 0 && - new_layout_name.rfind(old_layout_name, 0) != 0) { - if (old_layout_name.size() != new_layout_name.size()) { - // Not support NHW4c -> NCHW - return out_default; - } else { - if (params->axes) { - auto axes = params->axes.value(); - Array new_axes; - - for (size_t i = 0; i < axes.size(); ++i) { - auto old_idx = axes[i].IntValue(); - auto new_idx = new_layout.IndexOf(layout[old_idx]); - new_begin.push_back(begin[i]); - new_end.push_back(end[i]); - new_strides.push_back(strides[i]); - new_axes.push_back(new_idx); - } - params->axes = new_axes; - - } else { - for (size_t i = 0; i < new_layout_name.size(); ++i) { - auto index = layout.IndexOf(new_layout[i]); - if (index == -1) { - return out_default; - } - - size_t new_index = static_cast(index); - int64_t bg, ed, st; - if (strides.defined() && new_index < strides.size() && strides[new_index].defined()) { - st = strides[new_index]->value; - } else { - st = 1; - } - if (new_index < begin.size() && begin[new_index].defined()) { - bg = begin[new_index]->value; - } else { - bg = 0; - } - if (new_index < end.size() && end[new_index].defined()) { - ed = end[new_index]->value; - } else { - ed = shape[new_index].as()->value; - } - - new_begin.push_back(IntImm(begin[0]->dtype, bg)); - new_end.push_back(IntImm(end[0]->dtype, ed)); - new_strides.push_back(IntImm(strides[0]->dtype, st)); - } - } - - params->begin = new_begin; - params->end = new_end; - params->strides = new_strides; - layout = new_layout; - } - } else if (old_layout_name.size() < - new_layout_name.size()) { // prohibit transforms such as NCHW4c -> NCHW - if (params->axes) { - auto axes = params->axes.value(); - Array new_axes; - for (size_t i = 0; i < axes.size(); ++i) { - auto old_idx = axes[i].IntValue(); - auto new_idx = new_layout.IndexOf(layout[old_idx]); - new_axes.push_back(new_idx); - - const LayoutAxis& axis = layout[old_idx]; - ICHECK(axis.IsPrimal()); - auto factor = new_layout.FactorOf(axis); - if (factor == -1) { - new_begin.push_back(begin[i]); - new_end.push_back(end[i]); - } else { - if (strides.defined() && i < strides.size()) { - auto stride = strides[i]; - // arbitrary stride is not supported - if (stride.defined() && stride->value != 1) { - return out_default; - } - } - int64_t bg = begin[i].IntValue(); - int64_t ed = end[i].IntValue(); - if (bg % factor || ed % factor) { - // transform to original layout - return out_default; - } - new_begin.push_back(IntImm(begin[0]->dtype, (bg / factor))); - new_end.push_back(IntImm(end[0]->dtype, (ed / factor))); - } - } - params->axes = new_axes; - - } else { - for (size_t i = 0; i < begin.size(); i++) { - const LayoutAxis& axis = layout[i]; - ICHECK(axis.IsPrimal()); - auto factor = new_layout.FactorOf(axis); - if (factor == -1) { - new_begin.push_back(IntImm(begin[i]->dtype, begin[i].IntValue())); - new_end.push_back(IntImm(end[i]->dtype, end[i].IntValue())); - } else { - if (strides.defined() && i < strides.size()) { - auto stride = strides[i]; - // arbitrary stride is not supported - if (stride.defined() && stride->value != 1) { - return out_default; - } - } - int64_t bg = begin[i].defined() ? begin[i]->value : 0; - int64_t ed; - if (!end[i].defined()) { - ed = shape[i].as()->value; - } else if (params->slice_mode == "size") { - if (end[i]->value < 0) { - ed = shape[i].as()->value; - } else { - ed = bg + end[i]->value; - } - } else { - ed = end[i]->value; - } - - if (bg % factor || ed % factor) { - // transform to original layout - return out_default; - } - new_begin.push_back(IntImm(begin[0]->dtype, (bg / factor))); - new_end.push_back(IntImm(end[0]->dtype, (ed / factor))); - } - } - } - - layout = new_layout; - params->begin = new_begin; - params->end = new_end; - } - return InferCorrectLayoutOutput({layout}, {layout}, Attrs(params)); - } - return InferCorrectLayoutOutput({layout}, {layout}, attrs); -} - -Array StridedSliceCompute(const Attrs& attrs, const Array& inputs, - const Type& out_type) { - const StridedSliceAttrs* param = attrs.as(); - ICHECK(param != nullptr); - ICHECK(param->begin && param->end && param->strides); - Array begin = param->begin.value(); - Array end = param->end.value(); - Array strides = param->strides.value(); - if (param->axes) { - auto axes = param->axes.value(); - return Array{ - topi::strided_slice_with_axes(inputs[0], begin, end, strides, axes, param->slice_mode)}; - } - return Array{topi::strided_slice(inputs[0], begin, end, strides, param->slice_mode)}; -} - -// Positional relay function to create StridedSlice operator used by frontend FFI. -Expr MakeStridedSlice(Expr data, Array begin, Array end, Array strides, - String slice_mode, Optional> axes) { - auto attrs = make_object(); - attrs->begin = std::move(begin); - attrs->end = std::move(end); - attrs->strides = std::move(strides); - attrs->slice_mode = slice_mode; - attrs->axes = std::move(axes); - static const Op& op = Op::Get("strided_slice"); - return Call(op, {data}, Attrs(attrs), {}); -} - -TVM_REGISTER_GLOBAL("relay.op._make.strided_slice").set_body_typed(MakeStridedSlice); - -RELAY_REGISTER_OP("strided_slice") - .describe(R"code(Strided slice of an array. - -Examples:: - - x = [[ 1., 4., 7., 10.], - [ 2., 5., 8., 11.], - [ 3., 6., 9., 12.]] - - strided_slice(x, begin=[0, 1], end=[2, 4], stride=[1, 1]) = [[ 4., 7., 10.], - [ 5., 8., 11.]] - - x = [[[ 1., 2.], - [ 3., 4.]], - - [[ 5., 6.], - [ 7., 8.]]] - - strided_slice(x, begin=[0, 0], end=[2, 2]) = [[[ 1., 2.], - [ 3., 4.]], - - [[ 5., 6.], - [ 7., 8.]]] -)code" TVM_ADD_FILELINE) - .set_num_inputs(1) - .add_argument("data", "Tensor", "The input tensor.") - .set_support_level(4) - .set_attrs_type() - .add_type_rel("StridedSlice", StridedSliceRel) - .set_attr("FTVMCompute", StridedSliceCompute) - .set_attr("TOpPattern", kInjective) - .set_attr("AnyCodegenStrategy", kVariableDimensions) - .set_attr("FInferCorrectLayout", StridedSliceInferCorrectLayout); - -// strided_set -bool StridedSetRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - ICHECK_EQ(types.size(), 6); - reporter->Assign(types[5], types[0]); - return true; -} - -Expr MakeStridedSet(Expr data, Expr v, Expr begin, Expr end, Expr strides) { - static const Op& op = Op::Get("strided_set"); - return Call(op, {data, v, begin, end, strides}, {}); -} - -TVM_REGISTER_GLOBAL("relay.op._make.strided_set").set_body_typed(MakeStridedSet); - -RELAY_REGISTER_OP("strided_set") - .describe(R"code(Strided set of an array. -Example:: - - x = [[ 1., 4., 7., 10.], - [ 2., 5., 8., 11.], - [ 3., 6., 9., 12.]] - - v = [[ 11., 22., 33.] - [ 44., 55., 66.]] - - strided_set(x, v, begin=[0, 1], end=[2, 4], stride=[1, 1]) = \ - [[ 1., 11., 22., 33.], - [ 2., 44., 55., 66.], - [ 3., 6., 9., 12.]] -)code" TVM_ADD_FILELINE) - .set_num_inputs(5) - .add_argument("data", "Tensor", "The input tensor.") - .add_argument("v", "Tensor", "The data to set.") - .add_argument("begin", "Tensor", "Indices for the start of the slice.") - .add_argument("end", "Tensor", "Indices indicating the end of the slice.") - .add_argument("strides", "Tensor", "The strides values.") - .set_support_level(4) - .set_attr("TOpPattern", kInjective) - .add_type_rel("StridedSet", StridedSetRel); - -// relay.split -TVM_REGISTER_NODE_TYPE(SplitAttrs); - -InferCorrectLayoutOutput SplitInferCorrectLayout(const Attrs& attrs, - const Array& new_in_layouts, - const Array& old_in_layouts, - const Array& old_in_types) { - const auto* attrs_ptr = attrs.as(); - ICHECK(attrs_ptr); - ObjectPtr param = make_object(*attrs_ptr); - - Array> old_in_shapes; - for (auto old_in_t : old_in_types) { - ICHECK(old_in_t.as()); - old_in_shapes.push_back(old_in_t.as()->shape); - } - - size_t axis = - param->axis < 0 ? param->axis + old_in_shapes[0].size() : static_cast(param->axis); - - Layout ret = Layout::Undef(); - size_t size = 0; - if (const auto* sections = param->indices_or_sections.as()) { - size = sections->value; - } else { - size = Downcast>(param->indices_or_sections).size() + 1; - } - - // If new_in_layouts are defined, this code tries to modify the layout. - if (new_in_layouts.defined() && old_in_layouts.defined()) { - bool divisible = true; - const auto& sp_dim = old_in_layouts[0][axis]; - auto new_index = new_in_layouts[0].IndexOf(sp_dim); - param->axis = new_index; - int factor = new_in_layouts[0].FactorOf(sp_dim); - if (factor > 1) { - if (!param->indices_or_sections.as()) { - auto ios = Downcast>(param->indices_or_sections); - Array new_ios; - for (const auto& v : ios) { - new_ios.push_back(runtime::Int(v->value / factor)); - if (v->value % factor) { - divisible = false; - } - } - if (divisible) { - param->indices_or_sections = new_ios; - } - } - } - if (divisible) { - ret = new_in_layouts[0]; - } else { - ret = old_in_layouts[0]; - } - } else if (old_in_layouts.defined()) { - ret = old_in_layouts[0]; - } - - return InferCorrectLayoutOutput({ret}, {Array(size, ret)}, Attrs(param)); -} - -bool SplitRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - // `types` contains: [data, result] - ICHECK_EQ(types.size(), 2); - const auto* data = types[0].as(); - if (data == nullptr) return false; - ICHECK_NE(data->shape.size(), 0) << "Input shape cannot be empty"; - const auto param = attrs.as(); - ICHECK(param != nullptr); - auto axis = param->axis; - if (axis < 0) { - axis += data->shape.size(); - } - ICHECK_LT(axis, data->shape.size()) << "axis should be within the input dimension range."; - ICHECK_GE(axis, 0) << "axis should be within the input dimension range."; - - if (const auto* sections = param->indices_or_sections.as()) { - if (!data->shape[axis].as()) { - ICHECK(reporter->Assert(indexmod(data->shape[axis], sections->value) == - tir::make_zero(DataType::Int(64)))) - << "indices_or_sections need to be able to divide input.shape[axis]"; - } - std::vector fields; - for (int i = 0; i < sections->value; ++i) { - std::vector oshape(data->shape.begin(), data->shape.end()); - if (data->shape[axis].as()) { - oshape[axis] = Any(); - } else { - oshape[axis] = indexdiv(oshape[axis], sections->value); - } - auto vec_type = TensorType(oshape, data->dtype); - fields.push_back(vec_type); - } - reporter->Assign(types[1], TupleType(Array(fields))); - } else { - Array indices; - for (auto index : Downcast>(param->indices_or_sections)) { - indices.push_back(IntImm(DataType::Int(32), index->value)); - } - auto begin = IndexExpr(tir::make_zero(DataType::Int(32))); - std::vector fields; - for (unsigned int i = 0; i < indices.size(); ++i) { - ICHECK(reporter->Assert(indices[i] > begin)) - << "indices_or_sections need to be a sorted ascending list"; - std::vector oshape(data->shape.begin(), data->shape.end()); - oshape[axis] = indices[i] - begin; - begin = indices[i]; - auto vec_type = TensorType(oshape, data->dtype); - fields.push_back(vec_type); - } - if (!data->shape[axis].as()) { - ICHECK(reporter->Assert(begin < data->shape[axis])) - << "The sum of sections must match the input.shape[axis]"; - } - std::vector oshape(data->shape.begin(), data->shape.end()); - if (data->shape[axis].as()) { - oshape[axis] = Any(); - } else { - oshape[axis] = data->shape[axis] - begin; - } - auto vec_type = TensorType(oshape, data->dtype); - fields.push_back(vec_type); - reporter->Assign(types[1], TupleType(Array(fields))); - } - return true; -} - -Array SplitCompute(const Attrs& attrs, const Array& inputs, - const Type& out_type) { - const auto param = attrs.as(); - ICHECK(param != nullptr); - - if (const auto* sections = param->indices_or_sections.as()) { - int64_t num_sections = sections->value; - return Array{topi::split_sections(inputs[0], num_sections, param->axis)}; - } else { - Array indices; - for (auto index : Downcast>(param->indices_or_sections)) { - indices.push_back(IntImm(DataType::Int(32), index->value)); - } - return Array{topi::split(inputs[0], indices, param->axis)}; - } -} - -Expr MakeSplit(Expr data, Variant> indices_or_sections, - int axis) { - auto attrs = make_object(); - attrs->axis = axis; - attrs->indices_or_sections = std::move(indices_or_sections); - static const Op& op = Op::Get("split"); - return Call(op, {data}, Attrs(attrs), {}); -} - -TVM_REGISTER_GLOBAL("relay.op._make.split").set_body_typed(MakeSplit); - -RELAY_REGISTER_OP("split") - .describe(R"code(Splits an array along a particular axis into multiple sub-arrays. - -Indices or sections to split into. Accepts an int or a tuple -If indices_or_sections is an integer, the input will be divided equally -along given axis. If such a split is not possible, an error is raised. - -If indices_or_sections is a tuple of sorted integers, -the entries indicate where along axis the array is split. - -)code" TVM_ADD_FILELINE) - .set_attrs_type() - .set_num_inputs(1) - .add_argument("data", "Tensor", "The input tensor.") - .set_support_level(3) - .add_type_rel("Split", SplitRel) - .set_attr("FTVMCompute", SplitCompute) - .set_attr("FInferCorrectLayout", SplitInferCorrectLayout) - .set_attr("TOpPattern", kInjective); - -// relay.slice_like -TVM_REGISTER_NODE_TYPE(SliceLikeAttrs); - -/*! - * \brief SliceLikeRel User defined type constraint function. - * \param num_inputs Number of input types in the args. - * \param attrs The additional attributes of the operator. - * \param reporter The reporter to report solution to. - * \return False if the relation has not been resolved, it might be resolved later. - * True if this relation has been resolved. - */ -bool SliceLikeRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - ICHECK_EQ(types.size(), 3); - const auto* data = types[0].as(); - if (data == nullptr) { - return false; - } - - const auto* target = types[1].as(); - if (target == nullptr) { - return false; - } - - const auto param = attrs.as(); - ICHECK(param != nullptr); - - const Array& dshape = data->shape; - const Array& target_shape = target->shape; - std::vector oshape(dshape.begin(), dshape.end()); - - if (!param->axes.defined()) { - for (size_t i = 0; i < dshape.size(); ++i) { - if (i < target_shape.size()) { - oshape[i] = target_shape[i]; - ICHECK(reporter->Assert(oshape[i] <= dshape[i])) - << "End index of axis " << i << " exceeds input shape: " << oshape[i] << " vs " - << dshape[i]; - } - } - } else { - ICHECK(param->axes.size() != 0) << "Axes cannot be empty."; - for (Integer val : param->axes) { - int axis = val->value; - if (axis < 0) { - axis += dshape.size(); - } - ICHECK(axis < static_cast(target_shape.size())) - << "Axis " << axis << " exceeds dimension " << target_shape.size() << " of target_shape."; - oshape[axis] = target_shape[axis]; - ICHECK(reporter->Assert(oshape[axis] <= dshape[axis])) - << "End index of axis " << axis << " exceeds input shape: " << oshape[axis] << " vs " - << dshape[axis]; - } - } - - reporter->Assign(types[2], TensorType(oshape, data->dtype)); - return true; -} - -Expr MakeSliceLike(Expr data, Expr shape_like, Array axes) { - auto attrs = make_object(); - attrs->axes = std::move(axes); - static const Op& op = Op::Get("slice_like"); - return Call(op, {data, shape_like}, Attrs(attrs), {}); -} - -InferCorrectLayoutOutput SliceLikeInferCorrectLayout(const Attrs& attrs, - const Array& new_in_layouts, - const Array& old_in_layouts, - const Array& old_in_types) { - Array new_axes; - if (old_in_layouts.defined() && new_in_layouts.defined()) { - ICHECK_EQ(new_in_layouts.size(), 2); - ICHECK_EQ(new_in_layouts[0]->name, new_in_layouts[1]->name); - ICHECK_EQ(old_in_layouts.size(), 2); - ICHECK_EQ(old_in_layouts[0]->name, old_in_layouts[1]->name); - - auto old_layout = old_in_layouts[0]; - auto new_layout = new_in_layouts[0]; - - const auto* attrs_ptr = attrs.as(); - ICHECK(attrs_ptr); - ObjectPtr params = make_object(*attrs_ptr); - - for (auto axis : params->axes) { - auto new_axis = new_layout.IndexOf(old_layout[axis->value]); - // Cannot find the target axis in the new layout. - if (new_axis == -1) { - new_axes.clear(); - break; - } - new_axes.push_back(new_axis); - } - if (!new_axes.empty()) { - params->axes = std::move(new_axes); - return InferCorrectLayoutOutput({new_layout, new_layout}, {new_layout}, Attrs(params)); - } - } - - if (old_in_layouts.defined()) { - ICHECK_EQ(old_in_layouts.size(), 2); - return InferCorrectLayoutOutput({old_in_layouts[0], old_in_layouts[1]}, {old_in_layouts[1]}, - attrs); - } - return InferCorrectLayoutOutput({Layout::Undef(), Layout::Undef()}, {Layout::Undef()}, attrs); -} - -Array SliceLikeCompute(const Attrs& attrs, const Array& inputs, - const Type& out_type) { - const auto* param = attrs.as(); - ICHECK(param != nullptr); - Array src_shape = inputs[0]->shape; - Array target_shape = inputs[1]->shape; - Array begin_idx, end_idx, strides; - for (size_t i = 0; i < src_shape.size(); ++i) { - begin_idx.push_back(0); - strides.push_back(1); - } - for (auto s : src_shape) { - ICHECK(s->IsInstance()) << "slice_like does not support dynamic input shape"; - end_idx.push_back(topi::GetConstInt(s)); - } - if (!param->axes.defined()) { - for (size_t i = 0; i < src_shape.size(); ++i) { - if (i < target_shape.size()) { - ICHECK(target_shape[i]->IsInstance()) - << "slice_like does not support dynamic output shape"; - end_idx.Set(i, topi::GetConstInt(target_shape[i])); - ICHECK_LE(topi::GetConstInt(end_idx[i]), topi::GetConstInt(src_shape[i])) - << "End index of axis " << i - << " exceeds input shape: " << topi::GetConstInt(end_idx[i]) << " vs " - << topi::GetConstInt(src_shape[i]); - } - } - } else { - for (Integer axis : param->axes) { - int a = axis.IntValue(); - if (a < 0) { - a = static_cast(src_shape.size()) + a; - } - ICHECK(target_shape[a]->IsInstance()) - << "slice_like does not support dynamic output shape"; - end_idx.Set(a, topi::GetConstInt(target_shape[a])); - ICHECK_LE(topi::GetConstInt(end_idx[a]), topi::GetConstInt(src_shape[a])) - << "End index of axis " << a << " exceeds input shape: " << topi::GetConstInt(end_idx[a]) - << " vs " << topi::GetConstInt(src_shape[a]); - } - } - return Array{topi::strided_slice(inputs[0], begin_idx, end_idx, strides, "end")}; -} - -TVM_REGISTER_GLOBAL("relay.op._make.slice_like").set_body_typed(MakeSliceLike); - -RELAY_REGISTER_OP("slice_like") - .describe(R"code(Slice the first input respect to the second input. -)code" TVM_ADD_FILELINE) - .set_attrs_type() - .set_num_inputs(2) - .add_argument("data", "Tensor", "The input tensor.") - .add_argument("shape_like", "Tensor", "Shape tensor.") - .set_support_level(10) - .add_type_rel("SliceLike", SliceLikeRel) - .set_attr("FTVMCompute", SliceLikeCompute) - .set_attr("FInferCorrectLayout", SliceLikeInferCorrectLayout) - .set_attr("TOpPattern", kInjective); - -// relay.layout_transform -TVM_REGISTER_NODE_TYPE(LayoutTransformAttrs); - -Array LayoutTransformCompute(const Attrs& attrs, const Array& inputs, - const Type& out_type) { - const auto* param = attrs.as(); - ICHECK(param != nullptr); - return Array{topi::layout_transform(inputs[0], param->src_layout, param->dst_layout)}; -} - -bool LayoutTransformRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - const auto* data = types[0].as(); - if (data == nullptr) { - ICHECK(types[0].as()) - << "LayoutTransform: expect input data type to be TensorType but get " << types[0]; - return false; - } - const LayoutTransformAttrs* params = attrs.as(); - - Layout src_layout(params->src_layout); - Layout dst_layout(params->dst_layout); - - ICHECK(src_layout.defined() && dst_layout.defined()) << "cannot convert from/to undefined layout"; - auto layout_converter = tir::BijectiveLayout(src_layout, dst_layout); - ICHECK(layout_converter.defined()) - << "cannot convert from " << params->src_layout << " to " << params->dst_layout; - - const auto& out_shape = layout_converter.ForwardShape(data->shape); - reporter->Assign(types[1], TensorType(out_shape, data->dtype)); - return true; -} - -Expr MakeLayoutTransform(Expr data, String src_layout, String dst_layout) { - auto attrs = make_object(); - attrs->src_layout = std::move(src_layout); - attrs->dst_layout = std::move(dst_layout); - static const Op& op = Op::Get("layout_transform"); - return Call(op, {data}, Attrs(attrs), {}); -} - -TVM_REGISTER_GLOBAL("relay.op._make.layout_transform").set_body_typed(MakeLayoutTransform); - -RELAY_REGISTER_OP("layout_transform") - .describe(R"code(Transform the input data layout. - -For transforming from NCHW to N16cHWC, the `__layout_transform__` operator reshapes -the input array by output[n, c, h, w, C] = data[n, C*16+c, h, w] - -)code" TVM_ADD_FILELINE) - .set_attrs_type() - .set_num_inputs(1) - .add_argument("data", "Tensor", "The input tensor.") - .add_type_rel("layout_transform", LayoutTransformRel) - .set_support_level(5) - .set_attr("FTVMCompute", LayoutTransformCompute); - -// relay.auto_scheduler_layout_transform -TVM_REGISTER_NODE_TYPE(AutoSchedulerLayoutTransformAttrs); - -Array AutoSchedulerLayoutTransformCompute(const Attrs& attrs, - const Array& inputs, - const Type& out_type) { - const auto* param = attrs.as(); - CHECK(param != nullptr); - return Array{ - topi::auto_scheduler_layout_transform(inputs[0], param->src_layout, param->dst_layout)}; -} - -bool AutoSchedulerLayoutTransformRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - const auto* data = types[0].as(); - CHECK(data != nullptr); - const AutoSchedulerLayoutTransformAttrs* params = attrs.as(); - - Array dst_shape; - std::vector dst_axes; - - topi::parse_auto_scheduler_layout(params->dst_layout, &dst_shape, &dst_axes); - - reporter->Assign(types[1], TensorType(dst_shape, data->dtype)); - return true; -} - -Expr MakeAutoSchedulerLayoutTransform(Expr data, String src_layout, String dst_layout) { - auto attrs = make_object(); - attrs->src_layout = std::move(src_layout); - attrs->dst_layout = std::move(dst_layout); - static const Op& op = Op::Get("auto_scheduler_layout_transform"); - return Call(op, {data}, Attrs(attrs), {}); -} - -TVM_REGISTER_GLOBAL("relay.op._make.auto_scheduler_layout_transform") - .set_body_typed(MakeAutoSchedulerLayoutTransform); - -RELAY_REGISTER_OP("auto_scheduler_layout_transform") - .describe(R"code(Transform the input kernel layout. -)code" TVM_ADD_FILELINE) - .set_attrs_type() - .set_num_inputs(1) - .add_argument("data", "Tensor", "The input tensor.") - .add_type_rel("auto_scheduler_layout_transform", AutoSchedulerLayoutTransformRel) - .set_support_level(5) - .set_attr("FTVMCompute", AutoSchedulerLayoutTransformCompute); - -// relay.meta_schedule_layout_transform -TVM_REGISTER_NODE_TYPE(MetaScheduleLayoutTransformAttrs); - -Array MetaScheduleLayoutTransformCompute(const Attrs& attrs, - const Array& inputs, - const Type& out_type) { - const auto* param = attrs.as(); - CHECK(param != nullptr); - return Array{topi::meta_schedule_layout_transform(inputs[0], param->index_map)}; -} - -bool MetaScheduleLayoutTransformRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - TensorType data_type = Downcast(types[0]); - arith::Analyzer analyzer; - const MetaScheduleLayoutTransformAttrs* params = attrs.as(); - ICHECK(params); - Array new_shape = params->index_map->MapShape(data_type->shape, &analyzer); - reporter->Assign(types[1], TensorType(new_shape, data_type->dtype)); - return true; -} - -Expr MakeMetaScheduleLayoutTransform(Expr data, tir::IndexMap index_map) { - static const Op& op = Op::Get("meta_schedule_layout_transform"); - auto attrs = make_object(); - attrs->index_map = index_map; - return Call(op, {data}, Attrs(attrs), {}); -} - -TVM_REGISTER_GLOBAL("relay.op._make.meta_schedule_layout_transform") - .set_body_typed(MakeMetaScheduleLayoutTransform); - -RELAY_REGISTER_OP("meta_schedule_layout_transform") - .describe(R"code(Transform the input kernel layout. -)code" TVM_ADD_FILELINE) - .set_attrs_type() - .set_num_inputs(1) - .add_argument("data", "Tensor", "The input tensor.") - .add_type_rel("meta_schedule_layout_transform", MetaScheduleLayoutTransformRel) - .set_support_level(5) - .set_attr("FTVMCompute", MetaScheduleLayoutTransformCompute); - -// relay._contrib_reverse_reshape -Expr MakeReverseReshape(Expr data, Array newshape) { - auto attrs = make_object(); - attrs->newshape = std::move(newshape); - static const Op& op = Op::Get("contrib_reverse_reshape"); - return Call(op, {data}, Attrs(attrs), {}); -} - -TVM_REGISTER_GLOBAL("relay.op._make.contrib_reverse_reshape").set_body_typed(MakeReverseReshape); - -RELAY_REGISTER_OP("contrib_reverse_reshape") - .describe(R"code(Reshapes the input array where the special values are inferred from -right to left. - -Example:: - -The special values have the same semantics as reshape. The difference is that -special values are inferred from right to left. It can be explained in the -example below:: - -- data.shape = (10,5,4), newshape = (-1,0), reshape results in (40,5) -- data.shape = (10,5,4), newshape = (-1,0), reverse_reshape results in (40,5) - -)code" TVM_ADD_FILELINE) - .set_num_inputs(1) - .set_attrs_type() - .add_argument("data", "Tensor", "The input tensor.") - .set_support_level(10) - .add_type_rel("ReverseReshape", ReverseReshapeRel) - .set_attr("FTVMCompute", ReshapeCompute) - .set_attr("TOpPattern", kInjective) - .set_attr("TReshapeOp", true); - -// gather operator -TVM_REGISTER_NODE_TYPE(GatherAttrs); - -bool GatherRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - // `types` contains: [data, indices, result] - ICHECK_EQ(types.size(), 3); - const auto* data = types[0].as(); - const auto* indices = types[1].as(); - if (data == nullptr) { - ICHECK(types[0].as()) - << "Gather: expect input data type to be TensorType but get " << types[0]; - return false; - } - if (indices == nullptr) { - ICHECK(types[1].as()) - << "Gather: expect indices type to be TensorType but get " << types[1]; - return false; - } - ICHECK(indices->dtype.is_int()) << "indices of take must be tensor of integer"; - const auto param = attrs.as(); - ICHECK(param != nullptr); - ICHECK(param->axis.defined()); - - const auto ndim_data = data->shape.size(); - const auto ndim_indices = indices->shape.size(); - int axis = param->axis->value; - ICHECK_EQ(ndim_data, ndim_indices); - if (axis < 0) { - axis += ndim_data; - } - ICHECK_GE(axis, 0); - ICHECK_LT(axis, ndim_data); - - std::vector oshape; - oshape.reserve(ndim_data); - for (size_t i = 0; i < ndim_data; ++i) { - if (i == static_cast(axis)) { - if (indices->shape[i].as()) { - const int64_t* indice_shape_i = tir::as_const_int(indices->shape[i]); - ICHECK_GE(*indice_shape_i, 1); - } - } else { - ICHECK(reporter->AssertEQ(indices->shape[i], data->shape[i])); - } - oshape.emplace_back(indices->shape[i]); - } - reporter->Assign(types[2], TensorType(oshape, data->dtype)); - return true; -} - -Array GatherCompute(const Attrs& attrs, const Array& inputs, - const Type& out_type) { - const auto* param = attrs.as(); - return {topi::gather(inputs[0], param->axis.IntValue(), inputs[1])}; -} - -Expr MakeGather(Expr data, Integer axis, Expr indices) { - auto attrs = make_object(); - attrs->axis = std::move(axis); - static const Op& op = Op::Get("gather"); - return Call(op, {data, indices}, Attrs(attrs), {}); -} - -TVM_REGISTER_GLOBAL("relay.op._make.gather").set_body_typed(MakeGather); - -RELAY_REGISTER_OP("gather") - .describe(R"code(Gather values along given axis from given indices. - -E.g. for a 3D tensor, output is computed as: - - out[i][j][k] = data[indices[i][j][k]][j][k] # if axis == 0 - out[i][j][k] = data[i][indices[i][j][k]][k] # if axis == 1 - out[i][j][k] = data[i][j][indices[i][j][k]] # if axis == 2 - -``indices`` must have same shape as ``data``, except at dimension ``axis`` -which must just be not null. Output will have same shape as ``indices``. -)code" TVM_ADD_FILELINE) - .set_attrs_type() - .set_num_inputs(2) - .add_argument("data", "Tensor", "The input data to the operator.") - .add_argument("indices", "Tensor", "The indices of values to gather.") - .set_support_level(3) - .add_type_rel("Gather", GatherRel) - .set_attr("FTVMCompute", GatherCompute) - .set_attr("TOpPattern", kInjective); - -TVM_REGISTER_NODE_TYPE(GatherNDAttrs); - -// gather_nd operator -bool GatherNDRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - // `types` contains: [data, indices, result] - ICHECK_EQ(types.size(), 3); - const auto* data = types[0].as(); - const auto* indices = types[1].as(); - if (data == nullptr) { - ICHECK(types[0].as()) - << "GatherND: expect input data type to be TensorType but get " << types[0]; - return false; - } - if (indices == nullptr) { - ICHECK(types[1].as()) - << "GatherND: expect indices type to be TensorType but get " << types[1]; - return false; - } - const size_t ndim = data->shape.size(); - const IntImmNode* mdim = indices->shape[0].as(); - ICHECK(mdim) << "GatherND needs a static shape for the first axis of indices, got " - << indices->shape; - const size_t kdim = indices->shape.size() - 1; - ICHECK(size_t(mdim->value) <= ndim) << "GatherND: indices shape does satisfy."; - - const auto param = attrs.as(); - ICHECK(param != nullptr); - - for (int i = 0; i < param->batch_dims->value; ++i) { - ICHECK(reporter->AssertEQ( - data->shape[i], indices->shape[i + 1])); // +1 since the first axis is the index tuple - } - - Array oshape; - for (size_t i = 1; i < kdim + 1; ++i) oshape.push_back(indices->shape[i]); - for (size_t i = mdim->value + param->batch_dims->value; i < ndim; ++i) - oshape.push_back(data->shape[i]); - reporter->Assign(types[2], TensorType(oshape, data->dtype)); - return true; -} - -Array GatherNDCompute(const Attrs& attrs, const Array& inputs, - const Type& out_type) { - const auto* param = attrs.as(); - ICHECK(param); - return {topi::gather_nd(inputs[0], inputs[1], param->batch_dims.IntValue())}; -} - -Expr MakeGatherND(Expr data, Expr indices, int batch_dims = 0, - Optional index_rank = NullValue()) { - static const Op& op = Op::Get("gather_nd"); - auto attrs = make_object(); - attrs->batch_dims = batch_dims; - attrs->index_rank = index_rank; - return Call(op, {data, indices}, Attrs(attrs)); -} - -TVM_REGISTER_GLOBAL("relay.op._make.gather_nd").set_body_typed(MakeGatherND); - -RELAY_REGISTER_OP("gather_nd") - .describe(R"code(Gather elements or slices from data and store to - a tensor whose shape is defined by indices. - -Optionally, batch_dims, the number of batch dimensions, can be given, whose -default value is 0. - -Let B denote batch_dims, and data, indices shape be (X_0, X_1, ..., X_{N-1}), -(M, Y_0, ..., Y_{K-1}) respectively. - -When B > 0, indexing will start from the B-th axis, and it must be the case that -X_0, ... X_{B-1} == Y_0, ... Y_{B-1}. The output will have a shape -(X_0, ..., X_{B-1}, Y_B, ..., Y_{K-1}, X_{M+B}, ..., X_{N-1}), where M + B <= N. - -When B == 0 (the default case), the output shape will be (Y_0, ..., Y_{K-1}, X_M, ..., X_{N-1}). - -In both cases, if M + B == N, the output shape will simply be (Y_0, ..., Y_{K-1}). -)code" TVM_ADD_FILELINE) - .set_num_inputs(2) - .set_attrs_type() - .add_argument("data", "Tensor", "The input tensor.") - .add_argument("indices", "Tensor", "The indices of values to gather.") - .set_support_level(3) - .add_type_rel("GatherND", GatherNDRel) - .set_attr("FTVMCompute", GatherNDCompute) - .set_attr("TOpPattern", kInjective); - -// relay.sequence_mask -TVM_REGISTER_NODE_TYPE(SequenceMaskAttrs); - -bool SequenceMaskRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - // `types` contains: [data, valid_length, result] - ICHECK_EQ(types.size(), 3); - const auto* data = types[0].as(); - const auto* valid_length = types[1].as(); - ICHECK(data); - ICHECK(valid_length); - const auto param = attrs.as(); - Array valid_length_shape; - ICHECK(param->axis == 0 || param->axis == 1); - valid_length_shape.push_back(data->shape[1 - param->axis]); - reporter->Assign(types[1], TensorType(valid_length_shape, valid_length->dtype)); - reporter->Assign(types[2], types[0]); - return true; -} - -Array SequenceMaskCompute(const Attrs& attrs, const Array& inputs, - const Type& out_type) { - const auto* param = attrs.as(); - ICHECK(param != nullptr); - return Array{ - topi::sequence_mask(inputs[0], inputs[1], param->mask_value, param->axis)}; -} - -Expr MakeSequenceMask(Expr data, Expr valid_length, double mask_value, int axis) { - auto attrs = make_object(); - attrs->mask_value = std::move(mask_value); - attrs->axis = std::move(axis); - static const Op& op = Op::Get("sequence_mask"); - return Call(op, {data, valid_length}, Attrs(attrs), {}); -} - -TVM_REGISTER_GLOBAL("relay.op._make.sequence_mask").set_body_typed(MakeSequenceMask); - -RELAY_REGISTER_OP("sequence_mask") - .describe( - R"code(Sets all elements outside the expected length of the sequence to a constant value. - -This function takes an n-dimensional input array of the form [MAX_LENGTH, batch_size, ...] or -[batch_size, MAX_LENGTH, ...] and returns an array of the same shape. - -`axis` means the axis of the length dimension and can only be 0 or 1. If axis is 0, -the data must have shape [MAX_LENGTH, batch_size, ...]. Otherwise (axis=1), the data must have -shape [batch_size, MAX_LENGTH, ...]. - -`valid_length` gives the length of each sequence. `valid_length` should be -a 1D int array with positive ints and has dimension [batch_size,]. - -Examples:: - - x = [[[ 1., 2., 3.], - [ 4., 5., 6.]], - - [[ 7., 8., 9.], - [ 10., 11., 12.]], - - [[ 13., 14., 15.], - [ 16., 17., 18.]]] - - // valid_length [1, 1] means only the first block of each batch will be kept - // and other blocks are masked with default mask value = 0 - sequence_mask(x, valid_length=[1, 1]) = - [[[ 1., 2., 3.], - [ 4., 5., 6.]], - - [[ 0., 0., 0.], - [ 0., 0., 0.]], - - [[ 0., 0., 0.], - [ 0., 0., 0.]]] - - // valid_length [2, 3] means the first 2 blocks of the 1st batch will be kept - // and the first 3 blocks of the 2nd batch will be kept - // the masked values are set to be the specified mask value = 0.1 - sequence_mask(x, valid_length=[2, 3], mask_value=0.1) = - [[[ 1., 2., 3.], - [ 4., 5., 6.]], - - [[ 7., 8., 9.], - [ 10., 11., 12.]], - - [[ 0.1, 0.1, 0.1], - [ 16., 17., 18.]]] -)code" TVM_ADD_FILELINE) - .set_attrs_type() - .set_num_inputs(2) - .add_argument("data", "Tensor", "The input tensor.") - .add_argument("valid_length", "Tensor", "The real (valid) length of each sequence.") - .set_support_level(10) - .add_type_rel("SequenceMask", SequenceMaskRel) - .set_attr("FTVMCompute", SequenceMaskCompute) - .set_attr("TOpPattern", kInjective); - -// relay.one_hot -TVM_REGISTER_NODE_TYPE(OneHotAttrs); - -bool OneHotRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - // `types` contains: [indices, on_value, off_value, result] - ICHECK_EQ(types.size(), 4); - const auto* indices = types[0].as(); - ICHECK(indices); - - const auto param = attrs.as(); - ICHECK_GT(param->depth, 0); - - Array oshape; - int ndim = indices->shape.size() + 1; - int indices_index = 0; - int true_axis = (param->axis == -1) ? indices->shape.size() : param->axis; - for (int i = 0; i < ndim; i++) { - if (i == true_axis) { - oshape.push_back(Integer(param->depth)); - } else { - oshape.push_back(indices->shape[indices_index++]); - } - } - - reporter->Assign(types[3], TensorType(oshape, param->dtype)); - return true; -} - -Array OneHotCompute(const Attrs& attrs, const Array& inputs, - const Type& out_type) { - const auto* param = attrs.as(); - ICHECK(param != nullptr); - return Array{ - topi::one_hot(inputs[0], inputs[1](), inputs[2](), param->depth, param->axis, param->dtype)}; -} - -Expr MakeOneHot(Expr indices, Expr on_value, Expr off_value, int depth, int axis, DataType dtype) { - auto attrs = make_object(); - attrs->depth = std::move(depth); - attrs->axis = axis; - attrs->dtype = dtype; - static const Op& op = Op::Get("one_hot"); - return Call(op, {indices, on_value, off_value}, Attrs(attrs), {}); -} - -TVM_REGISTER_GLOBAL("relay.op._make.one_hot").set_body_typed(MakeOneHot); - -RELAY_REGISTER_OP("one_hot") - .describe(R"code(Returns a one-hot tensor where the locations repsented by indices take value 1, - other locations take value 0. Final dimension is x depth. - - **indices** Locations to set to 1. - - **on_value** Value to fill at indices. - - **off_value** Value to fill at all other positions besides indices. - - **depth** Depth of the one-hot dimension. - - **axis** Axis to fill. - - **dtype**)code" TVM_ADD_FILELINE) - .set_attrs_type() - .set_num_inputs(3) - .add_argument("indices", "Tensor", "Locations to set to on_value.") - .add_argument("on_value", "Expr", "Value to fill at indices.") - .add_argument("off_value", "Expr", "Value to fill at all other positions besides indices.") - .set_support_level(10) - .add_type_rel("OneHot", OneHotRel) - .set_attr("FTVMCompute", OneHotCompute) - .set_attr("TOpPattern", kOutEWiseFusable); - -/* relay.unravel_index */ -bool UnRavelIndexRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - ICHECK_EQ(types.size(), 3); - - const auto* indices = types[0].as(); - if (indices == nullptr) { - ICHECK(types[0].as()) - << "unravel_index: expect input type to be TensorType but get " << types[0]; - return false; - } - ICHECK(indices->dtype.is_int() || indices->dtype.is_uint()) - << "indices of unravel_index must be tensor of integer"; - - const auto* shape = types[1].as(); - if (shape == nullptr) { - ICHECK(types[1].as()) - << "unravel_index: expect input type to be TensorType but get " << types[1]; - return false; - } - ICHECK(shape->dtype.is_int() || shape->dtype.is_uint()) - << "shape of unravel_index must be tensor of integer"; - - Array indices_shape; - Array shape_shape; - indices_shape = indices->shape; - shape_shape = shape->shape; - - Array oshape; - oshape.push_back(shape_shape[0]); - if (indices_shape.size() != 0) { - oshape.push_back(indices_shape[0]); - } - reporter->Assign(types[2], TensorType(oshape, indices->dtype)); - return true; -} - -Array UnRavelIndexCompute(const Attrs& attrs, const Array& inputs, - const Type& out_type) { - return Array{topi::unravel_index(inputs[0], inputs[1])}; -} - -Expr MakeUnRavelIndex(Expr data, Expr shape) { - static const Op& op = Op::Get("unravel_index"); - return Call(op, {data, shape}, Attrs(), {}); -} - -TVM_REGISTER_GLOBAL("relay.op._make.unravel_index").set_body_typed(MakeUnRavelIndex); - -RELAY_REGISTER_OP("unravel_index") - .describe( - R"code(Converts a flat index or array of flat indices into a tuple of coordinate arrays. - -Example:: - - unravel_index([22, 41, 37], (7, 6)) = [[3, 6, 6], [4, 5, 1]] -)code" TVM_ADD_FILELINE) - .set_num_inputs(2) - .add_argument("data", "Tensor", "The input tensor.") - .add_argument("shape", "Tensor", "The shape tensor.") - .set_support_level(3) - .add_type_rel("UnRavelIndexRel", UnRavelIndexRel) - .set_attr("FTVMCompute", UnRavelIndexCompute) - .set_attr("TOpPattern", kInjective); - -// sparse_to_dense -TVM_REGISTER_NODE_TYPE(SparseToDenseAttrs); - -bool SparseToDenseRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - ICHECK_EQ(num_inputs, 3); - auto sparse_indices = types[0].as(); - auto sparse_values = types[1].as(); - auto default_value = types[2].as(); - - if (sparse_indices == nullptr || sparse_values == nullptr || default_value == nullptr) { - return false; - } - - ICHECK(sparse_indices->dtype.is_int()) << "sparse_indices must be tensor of integers"; - - ICHECK_LE(sparse_indices->shape.size(), 3) - << "sparse_indices must be a tensor of either 0D, 1D or 2D"; - - ICHECK_LE(sparse_values->shape.size(), 2) << "sparse_values must be a tensor of either 0D, 1D"; - - ICHECK_EQ(default_value->shape.size(), 0) << "default_value should be a scalar"; - - const auto* param = attrs.as(); - ICHECK(param != nullptr); - - Array oshape; - for (auto i : param->output_shape) { - oshape.push_back(i); - } - reporter->Assign(types[3], TensorType(oshape, sparse_values->dtype)); - return true; -} - -Array SparseToDenseCompute(const Attrs& attrs, const Array& inputs, - const Type& out_type) { - ICHECK_EQ(inputs.size(), 3); - const auto* param = attrs.as(); - ICHECK(param != nullptr); - Array output_shape; - for (auto val : param->output_shape) { - output_shape.push_back(val); - } - return {topi::sparse_to_dense(inputs[0], output_shape, inputs[1], inputs[2]())}; -} - -Expr MakeSparseToDense(Expr indices, Array output_shape, Expr values, Expr default_value) { - auto attrs = make_object(); - attrs->output_shape = std::move(output_shape); - static const Op& op = Op::Get("sparse_to_dense"); - return Call(op, {indices, values, default_value}, Attrs(attrs)); -} - -TVM_REGISTER_GLOBAL("relay.op._make.sparse_to_dense").set_body_typed(MakeSparseToDense); - -RELAY_REGISTER_OP("sparse_to_dense") - .describe(R"code(A dense tensor from a sparse representation. - - - **sparse_indices**: A 0-D, 1-D, or 2-D tensor of integers containing location of sparse values - - - **output_shape**: A list of integers. Shape of the dense output tensor. - - - **sparse_values**: A 0-D or 1-D tensor containing the sparse values for the sparse indices. - - - **default_value**: A 0-D tensor containing the default value for the remaining locations. Defaults to 0. - - Example:: - - sparse_to_dense([0, 0], [1, 2]], [3, 4], [1, 2], 0) = [[1, 0, 0, 0], [0, 0, 2, 0], [0, 0, 0, 0]] - - )code" TVM_ADD_FILELINE) - .set_num_inputs(3) - .set_support_level(3) - .set_attrs_type() - .add_argument("sparse_indices", "Tensor", "Contains sparse indices.") - .add_argument("sparse_values", "Tensor", "Contains values for sparse indices.") - .add_argument("default_value", "Tensor", "Value to set for non-sparse indices. Defaults to 0.") - .add_type_rel("SparseToDense", SparseToDenseRel) - .set_attr("TOpIsStateful", false) - .set_attr("TOpPattern", kOpaque) - .set_attr("FInferCorrectLayout", ElemwiseArbitraryLayout) - .set_attr("FTVMCompute", SparseToDenseCompute); - -// relay.matrix_set_diag -TVM_REGISTER_NODE_TYPE(MatrixSetDiagAttrs); - -bool MatrixSetDiagRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - // `types` contains: [input, diagonal, result] - ICHECK_EQ(types.size(), 3); - - const auto* input = types[0].as(); - ICHECK(input); - - const auto* diagonal = types[1].as(); - ICHECK(diagonal); - - const auto param = attrs.as(); - ICHECK_GE(param->k2, param->k1); - - int d_ndims = diagonal->shape.size(); - int i_ndims = input->shape.size(); - - reporter->Assert(input->shape[i_ndims - 2] > -param->k1); - reporter->Assert(input->shape[i_ndims - 1] > param->k2); - - for (int i = 0; i < d_ndims - 2; i++) { - reporter->AssertEQ(input->shape[i], diagonal->shape[i]); - } - if (param->k1 != param->k2) { - reporter->AssertEQ(diagonal->shape[d_ndims - 2], param->k2 - param->k1 + 1); - } else if (d_ndims >= 2) { - reporter->AssertEQ(input->shape[d_ndims - 2], diagonal->shape[d_ndims - 2]); - } - auto max_diag_len = if_then_else(input->shape[i_ndims - 2] + (param->k2 > 0 ? param->k2 : 0) <= - input->shape[i_ndims - 1] + (param->k1 < 0 ? -param->k1 : 0), - input->shape[i_ndims - 2] + (param->k2 > 0 ? param->k2 : 0), - input->shape[i_ndims - 1] + (param->k1 < 0 ? -param->k1 : 0)); - reporter->AssertEQ(diagonal->shape[d_ndims - 1], max_diag_len); - - reporter->Assign(types[2], TensorType(input->shape, input->dtype)); - return true; -} - -Array MatrixSetDiagCompute(const Attrs& attrs, const Array& inputs, - const Type& out_type) { - const auto* param = attrs.as(); - ICHECK(param != nullptr); - return Array{topi::matrix_set_diag(inputs[0], inputs[1], param->k1, param->k2, - param->super_diag_right_align, - param->sub_diag_right_align)}; -} - -Expr MakeMatrixSetDiag(Expr input, Expr diagonal, int k1, int k2, bool super_diag_right_align, - bool sub_diag_right_align) { - auto attrs = make_object(); - attrs->k1 = k1; - attrs->k2 = k2; - attrs->super_diag_right_align = super_diag_right_align; - attrs->sub_diag_right_align = sub_diag_right_align; - static const Op& op = Op::Get("matrix_set_diag"); - return Call(op, {input, diagonal}, Attrs(attrs), {}); -} - -TVM_REGISTER_GLOBAL("relay.op._make.matrix_set_diag").set_body_typed(MakeMatrixSetDiag); - -RELAY_REGISTER_OP("matrix_set_diag") - .describe( - R"code(Returns a tensor with the diagonals of input tensor replaced with the provided diagonal values. - **input** Input tensor. - **diagonal** Values to be filled in the diagonal. - **k1** Lower limit (included) of the range of diagonals. - **k2** Upper limit (included) of the range of diagonals. - **super_diag_right_align** Bool, true iff super-diagonal is right aligned (left-padded). - **sub_diag_right_align** Bool, true iff sub-diagonal is right aligned (left-padded). - )code" TVM_ADD_FILELINE) - .set_attrs_type() - .set_num_inputs(2) - .add_argument("input", "Tensor", "Input Tensor.") - .add_argument("diagonal", "Tensor", "Values to be filled in the diagonal.") - .set_support_level(10) - .add_type_rel("MatrixSetDiag", MatrixSetDiagRel) - .set_attr("FTVMCompute", MatrixSetDiagCompute) - .set_attr("TOpPattern", kInjective); - -// adv_index -bool AdvIndexRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - ICHECK_EQ(num_inputs, 1); - auto inputs = types[0].as(); - auto data = inputs->fields[0].as(); - - if (inputs == nullptr || data == nullptr) { - return false; - } - ICHECK_LE(inputs->fields.size() - 1, data->shape.size()) << "too many indices for data!"; - - Array oshape; - TensorType broadcast_type = Downcast(inputs->fields[1]); - for (size_t i = 2; i < inputs->fields.size(); ++i) { - broadcast_type = - ConcreteBroadcast(broadcast_type, Downcast(inputs->fields[i]), data->dtype); - } - - for (const auto& dim : broadcast_type->shape) { - oshape.push_back(dim); - } - for (size_t i = inputs->fields.size() - 1; i < data->shape.size(); ++i) { - oshape.push_back(data->shape[i]); - } - reporter->Assign(types[1], TensorType(oshape, data->dtype)); - return true; -} - -Array AdvIndexCompute(const Attrs& attrs, const Array& inputs, - const Type& out_type) { - Array indices; - for (size_t i = 1; i < inputs.size(); ++i) { - indices.push_back(inputs[i]); - } - return {topi::adv_index(inputs[0], indices)}; -} - -Expr MakeAdvIndex(Expr inputs) { - static const Op& op = Op::Get("adv_index"); - return Call(op, {inputs}, Attrs(), {}); -} - -TVM_REGISTER_GLOBAL("relay.op._make.adv_index").set_body_typed(MakeAdvIndex); - -RELAY_REGISTER_OP("adv_index") - .describe(R"code(Numpy style advanced indexing. Index with a list of tensors. - )code" TVM_ADD_FILELINE) - .set_num_inputs(1) - .set_support_level(3) - .add_argument("inputs", "Tuple of Tensors", "Input tensor and indices.") - .add_type_rel("AdvIndex", AdvIndexRel) - .set_attr("TOpIsStateful", false) - .set_attr("TOpPattern", kInjective) - .set_attr("FTVMCompute", AdvIndexCompute); - -TVM_REGISTER_NODE_TYPE(ScanopAttrs); - -bool ScanopRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - // types: [data, output] - ICHECK_EQ(types.size(), 2) << "Expects two types, one for the input and another for the output"; - const auto* data = types[0].as(); - if (data == nullptr) { - ICHECK(types[0].as()) - << "Scanop: expect input type to be TensorType but get " << types[0]; - return false; - } - - const auto* param = attrs.as(); - - auto dtype = param->dtype; - if (dtype.is_void()) { - dtype = data->dtype; - } - - if (param->axis.defined()) { - reporter->Assign(types[1], TensorType(data->shape, dtype)); - } else { - auto prod = data->shape[0]; - for (size_t i = 1; i < data->shape.size(); ++i) { - prod = prod * data->shape[i]; - } - reporter->Assign(types[1], TensorType({prod}, dtype)); - } - - return true; -} - -Expr MakeCumsum(Expr data, Integer axis, DataType dtype, Optional exclusive) { - auto attrs = make_object(); - attrs->dtype = dtype; - attrs->axis = axis; - if (exclusive.defined()) { - attrs->exclusive = exclusive.value(); - } - static const Op& op = Op::Get("cumsum"); - return Call(op, {data}, Attrs(attrs), {}); -} - -TVM_REGISTER_GLOBAL("relay.op._make.cumsum").set_body_typed(MakeCumsum); - -RELAY_REGISTER_OP("cumsum") - .describe( - R"doc(Return the cumulative sum of the elements along a given axis.)doc" TVM_ADD_FILELINE) - .set_num_inputs(1) - .add_argument("data", "Tensor", "The input tensor.") - .set_support_level(3) - .add_type_rel("Cumsum", ScanopRel) - .set_attr("TOpPattern", kOpaque); - -Expr MakeCumprod(Expr data, Integer axis, DataType dtype, Bool exclusive) { - auto attrs = make_object(); - attrs->dtype = dtype; - attrs->axis = axis; - attrs->exclusive = exclusive; - static const Op& op = Op::Get("cumprod"); - return Call(op, {data}, Attrs(attrs), {}); -} - -TVM_REGISTER_GLOBAL("relay.op._make.cumprod").set_body_typed(MakeCumprod); - -RELAY_REGISTER_OP("cumprod") - .describe( - R"doc(Return the cumulative product of the elements along a given axis.)doc" TVM_ADD_FILELINE) - .set_num_inputs(1) - .add_argument("data", "Tensor", "The input tensor.") - .set_support_level(3) - .add_type_rel("Cumprod", ScanopRel) - .set_attr("TOpPattern", kOpaque); - -TVM_REGISTER_NODE_TYPE(UniqueAttrs); - -bool UniqueRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - // types: [data, result] - ICHECK_EQ(types.size(), 2) << "Unique: expect 2 types but " << types.size() << " provided"; - ICHECK_EQ(num_inputs, 1) << "Unique: expect 1 inputs but " << num_inputs << " provided"; - auto data = types[0].as(); - if (data == nullptr) { - ICHECK(types[0].as()) - << "Unique: expect input type to be TensorType but get " << types[0]; - return false; - } - const int ndim = static_cast(data->shape.size()); - ICHECK_EQ(ndim, 1) << "Unique: input must be 1-D tensor"; - - std::vector fields; - fields.push_back(TensorType(data->shape, data->dtype)); // unique - fields.push_back(TensorType(data->shape, DataType::Int(32))); // indices - fields.push_back(TensorType(data->shape, DataType::Int(32))); // inverse_indices - fields.push_back(TensorType(Array{1}, DataType::Int(32))); // num_unique - const auto* param = attrs.as(); - if (param->return_counts) { - fields.push_back(TensorType(data->shape, DataType::Int(32))); // counts - } - reporter->Assign(types[1], TupleType(Array(fields))); - return true; -} - -Expr MakeUnique(Expr data, bool sorted, bool return_counts) { - auto attrs = make_object(); - attrs->sorted = sorted; - attrs->return_counts = return_counts; - static const Op& op = Op::Get("unique"); - return Call(op, {data}, Attrs(attrs), {}); -} - -TVM_REGISTER_GLOBAL("relay.op._make.unique").set_body_typed(MakeUnique); - -RELAY_REGISTER_OP("unique") - .describe( - R"code(This operation returns the unique elements and the new index of each item in a given 1-D array. - )code" TVM_ADD_FILELINE) - .set_num_inputs(1) - .add_argument("data", "Tensor", "The input tensor") - .add_type_rel("unique", UniqueRel) - .set_support_level(3) - .set_attr("TOpPattern", kOpaque); - -// invert_permutation -Expr MakeInvertPermutation(Expr data) { - static const Op& op = Op::Get("invert_permutation"); - return Call(op, {data}, Attrs(), {}); -} - -TVM_REGISTER_GLOBAL("relay.op._make.invert_permutation").set_body_typed(MakeInvertPermutation); - -RELAY_REGISTER_OP("invert_permutation") - .describe(R"doc(Computes the inverse permutation of a tensor.)doc" TVM_ADD_FILELINE) - .set_num_inputs(1) - .add_argument("data", "Tensor", "The input tensor.") - .add_type_rel("Identity", IdentityRel) - .set_support_level(1) - .set_attr("TOpPattern", kInjective) - .set_attr("TOpIsStateful", false); - -// Trilu - -TVM_REGISTER_NODE_TYPE(TriluAttrs); - -bool TriluRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - // types: [data, k, result] - ICHECK_EQ(types.size(), 3) << "Trilu: expect 3 types but " << types.size() << " provided"; - ICHECK_EQ(num_inputs, 2) << "Trilu: expect 2 inputs but " << num_inputs << " provided"; - auto data = types[0].as(); - if (data == nullptr) { - ICHECK(types[0].as()) - << "Trilu: expect input type to be TensorType but get " << types[0]; - return false; - } - - auto k = types[1].as(); - if (k == nullptr) { - ICHECK(types[1].as()) - << "Trilu: expect k type to be TensorType but get " << types[1]; - return false; - } - - ICHECK(k->shape.size() == 0) << "Trilu: k must be a 0-D tensor but get " << k; - - // Output shape is the same as input shape. - reporter->Assign(types[2], TensorType(data->shape, data->dtype)); - return true; -} - -Expr MakeTrilu(Expr data, Expr k, bool upper) { - auto attrs = make_object(); - attrs->upper = upper; - static const Op& op = Op::Get("trilu"); - return Call(op, {data, k}, Attrs(attrs), {}); -} - -TVM_REGISTER_GLOBAL("relay.op._make.trilu").set_body_typed(MakeTrilu); - -RELAY_REGISTER_OP("trilu") - .describe( - R"code(Filters out the upper or lower portion of an input tensor on one side of a diagonal. - )code" TVM_ADD_FILELINE) - .set_num_inputs(2) - .add_argument("data", "Tensor", "The input tensor") - .add_argument("k", "Tensor", "The number of diagonals above or below the main to exclude.") - .add_type_rel("trilu", TriluRel) - .set_support_level(3) - .set_attr("TOpPattern", kElemWise); - -// FixedPointMultiplyPerAxis - -TVM_REGISTER_NODE_TYPE(FixedPointMultiplyPerAxisAttrs); - -bool FixedPointMultiplyPerAxisRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - ICHECK_EQ(types.size(), 5) << "FixedPointMultiplyPerAxis: expect 5 types but " << types.size() - << " provided"; - ICHECK_EQ(num_inputs, 4) << "FixedPointMultiplyPerAxis: expect 4 inputs but " << num_inputs - << " provided"; - - for (int i = 0; i < num_inputs; i++) { - auto data = types[i].as(); - if (data == nullptr) { - ICHECK(types[i].as()) - << "FixedPointMultiplyPerAxis: expect input type to be TensorType but get " << types[i]; - return false; - } - } - - return IdentityRel({types[0], types[4]}, 1, attrs, reporter); -} - -InferCorrectLayoutOutput FixedPointMultiplyPerAxisInferCorrectLayout( - const Attrs& attrs, const Array& new_in_layouts, const Array& old_in_layouts, - const Array& old_in_types) { - const auto* attrs_ptr = attrs.as(); - ICHECK(attrs_ptr); - ObjectPtr param = - make_object(*attrs_ptr); - - Array> old_in_shapes; - for (auto old_in_t : old_in_types) { - ICHECK(old_in_t.as()); - old_in_shapes.push_back(old_in_t.as()->shape); - } - - Array input_layouts, output_layouts; - - if (new_in_layouts.defined()) { - const Layout& new_layout = new_in_layouts[0]; - const Layout& old_layout = old_in_layouts[0]; - - std::unordered_set old_dims; - for (auto axis : param->axes) { - ICHECK_GE(axis->value, 0) << "Axis out of bounds in FixedPointMultiplyPerAxis operator."; - ICHECK_LT(axis->value, old_in_shapes[0].size()) - << "Axis out of bounds in FixedPointMultiplyPerAxis operator."; - old_dims.emplace(old_layout[axis->value].name()); - } - - Array new_axes; - std::string new_layout_string = ""; - for (size_t axis_index = 0; axis_index < new_layout->axes.size(); ++axis_index) { - const auto& layout_axis = LayoutAxis::Get(new_layout->axes[axis_index]); - const std::string& layout_dim = layout_axis.name(); - if (layout_axis.IsPrimal()) { - if (old_dims.count(layout_dim)) { - new_axes.push_back(tvm::Integer(axis_index)); - new_layout_string += layout_dim; - } - } else { - auto primal_dim = layout_axis.ToPrimal().name(); - if (old_dims.count(primal_dim)) { - new_axes.push_back(tvm::Integer(axis_index)); - new_layout_string += std::to_string(new_layout.FactorOf(layout_axis)) + layout_dim; - } - } - } - - Layout channel_layout = Layout(new_layout_string); - - input_layouts = {new_layout, channel_layout, channel_layout, channel_layout}; - output_layouts = {new_layout}; - param->axes = std::move(new_axes); - } else if (old_in_layouts.defined()) { - ICHECK_EQ(old_in_layouts.size(), 4); - ICHECK_EQ(param->axes.size(), 1); // Not tested other cases - const Layout& old_layout = old_in_layouts[0]; - if (old_layout.defined()) { - std::string layout_string = old_layout[param->axes[0]->value].name(); - Layout channel_layout = Layout(layout_string); - - input_layouts = {old_layout, channel_layout, channel_layout, channel_layout}; - output_layouts = {old_layout}; - } else { - // Set the layouts to undef. - Layout undef = Layout::Undef(); - input_layouts = Array(4, undef); - output_layouts = {undef}; - } - } else { - // Set the layouts to undef. - Layout undef = Layout::Undef(); - input_layouts = Array(4, undef); - output_layouts = {undef}; - } - - return InferCorrectLayoutOutput(input_layouts, output_layouts, Attrs(param)); -} - -Expr MakeFixedPointMultiplyPerAxis(Expr x, Expr m, Expr lshift, Expr rshift, - bool is_lshift_required, bool is_rshift_required, - Array axes) { - auto attrs = make_object(); - attrs->is_lshift_required = is_lshift_required; - attrs->is_rshift_required = is_rshift_required; - attrs->axes = std::move(axes); - static const Op& op = Op::Get("fixed_point_multiply_per_axis"); - return Call(op, {x, m, lshift, rshift}, Attrs(attrs), {}); -} - -TVM_REGISTER_GLOBAL("relay.op._make.fixed_point_multiply_per_axis") - .set_body_typed(MakeFixedPointMultiplyPerAxis); - -RELAY_REGISTER_OP("fixed_point_multiply_per_axis") - .describe(R"code(per channel fixed point multiplication)code" TVM_ADD_FILELINE) - .set_num_inputs(4) - .add_argument("data", "Tensor", "The input tensor.") - .add_argument("fp_multiplier", "Tensor", "The multipliers tensor.") - .add_argument("left_shift", "Tensor", "The left shifts tensor.") - .add_argument("right_shift", "Tensor", "The right shifts tensor.") - .add_type_rel("FixedPointMultiplyPerAxis", FixedPointMultiplyPerAxisRel) - .set_attr("TOpPattern", kBroadcast) - .set_attr("FInferCorrectLayout", - FixedPointMultiplyPerAxisInferCorrectLayout) - .set_attrs_type() - .set_support_level(10); - -} // namespace relay -} // namespace tvm diff --git a/src/relay/op/tensor/transform.h b/src/relay/op/tensor/transform.h deleted file mode 100644 index 6c88aec8b957..000000000000 --- a/src/relay/op/tensor/transform.h +++ /dev/null @@ -1,241 +0,0 @@ -/* - * 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 src/relay/op/tensor/transform.h - * \brief Transform op attributes that can be shared among Relay and its dialects. - */ -#ifndef TVM_RELAY_OP_TENSOR_TRANSFORM_H_ -#define TVM_RELAY_OP_TENSOR_TRANSFORM_H_ - -#include -#include -#include - -#include -#include -#include -#include -#include -#include - -#include "../../transforms/infer_layout_utils.h" -#include "../make_op.h" - -namespace tvm { -namespace relay { - -template -bool ConcatenateRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - // types: [data, result] - ICHECK_EQ(types.size(), 2) << "the arity of concatenate is 2, not " << types.size(); - /* If we receive a tuple we can continue, if we receive - * anything but an incomplete type we should signal an - * error. - */ - const auto* tensor_tuple = types[0].as(); - if (tensor_tuple == nullptr) { - reporter->GetDiagCtx().EmitFatal( - Diagnostic::Error(reporter->GetSpan()) - << "concatenate requires a tuple of tensors as the first argument, found " - << PrettyPrint(types[0])); - return false; - } else if (types[0].as() != nullptr) { - return false; - } - - const auto* param = attrs.as(); - if (param == nullptr) { - reporter->GetDiagCtx().EmitFatal(Diagnostic::Error(reporter->GetSpan()) - << "the call attributes are not defined"); - return false; - } - - if (tensor_tuple->fields[0].as()) { - return false; - } - const auto& first = Downcast(tensor_tuple->fields[0]); - // Sanity check: ndim and dtype. - const int ndim = static_cast(first->shape.size()); - const DataType dtype = first->dtype; - - // Sanity check: axis - int axis = param->axis; - if (!(-ndim <= axis && axis < ndim)) { - throw CompileError(ErrorBuilder() << "concatenate only accepts `axis` in [-ndim, ndim)" - << ", but got axis = " << axis << ", and ndim = " << ndim); - } - axis = axis < 0 ? ndim + axis : axis; - - for (const Type& ele : tensor_tuple->fields) { - if (ele.as()) { - return false; - } - - const auto& e = Downcast(ele); - - int e_ndim = static_cast(e->shape.size()); - const DataType& e_dtype = e->dtype; - if (e_ndim != ndim) { - throw Error("relay.concatenate requires all tensors have the same ndim"); - } - if (e_dtype != dtype) { - throw Error("relay.concatenate requires all tensors have the same dtype"); - } - } - - // Calculate shape - std::vector oshape(ndim); - const size_t data_length = tensor_tuple->fields.size(); - - // Accumulate the concat axis output dim or decide if this is dynamic concat - bool is_dynamic_concat = false; - std::vector input_tensors; - IndexExpr concat_output_dim = first->shape[axis]; - for (size_t i = 0; i < data_length; ++i) { - const auto& e = Downcast(tensor_tuple->fields[i]); - input_tensors.push_back(e); - if (e->shape[axis].as()) { - is_dynamic_concat = true; - concat_output_dim = Any(); - } else if (i > 0 && !is_dynamic_concat) { - // accumulate axis dimension - concat_output_dim += e->shape[axis]; - } - } - - oshape[axis] = concat_output_dim; - - for (int i = 0; i < ndim; ++i) { - if (i == axis) { - // The concat axis is already handled above. - // The rest of the body sets the output shape for non-concat axes - continue; - } - std::vector non_any; - for (size_t j = 0; j < data_length; ++j) { - const auto& e = input_tensors[j]; - if (!e->shape[i].as()) { - non_any.push_back(e->shape[i]); - } - } - size_t non_any_size = non_any.size(); - for (size_t k = 1; k < non_any_size; k++) { - if (reporter->AssertEQ(non_any[0], non_any[k])) continue; - throw Error( - "relay.concatenate requires all tensors have the same shape " - "on non-concatenating axes"); - } - - if (non_any_size == data_length) { - // All static case - oshape[i] = non_any[0]; - } else if (non_any_size > 0 && is_dynamic_concat) { - // For non-concat axes, we want to enforce static shape constraint. - // However, if the concat axis is static, the output shape would become static while - // the input could be partially static/dynamic. To prevent runtime segfaults due to the lack - // of runtime input shape checking for such cases, static shape constraint is only enforced - // when the output concat axis is dynamic. - // - // Examples (both concat on the first axis): - // * [(?, 3), (?, ?)] -> (?, 3) - // * [(1, 3), (1, ?)] -> (2, ?) - oshape[i] = non_any[0]; - } else { - oshape[i] = Any(); - } - } - - auto rtype = TensorType(oshape, dtype); - reporter->Assign(types[1], rtype); - return true; -} - -static inline InferCorrectLayoutOutput ConcatenateLayout( - const Attrs& attrs, const Array& new_in_layouts, const Array& old_in_layouts, - const Array& old_in_types) { - const auto* attrs_ptr = attrs.as(); - ICHECK(attrs_ptr); - ObjectPtr param = make_object(*attrs_ptr); - - Array> old_in_shapes; - ICHECK_EQ(old_in_types.size(), 1); - for (auto old_in_tuple_t : old_in_types) { - ICHECK(old_in_tuple_t.as()); - for (auto old_in_t : old_in_tuple_t.as()->fields) { - old_in_shapes.push_back(old_in_t.as()->shape); - } - } - - size_t axis = - param->axis < 0 ? param->axis + old_in_shapes[0].size() : static_cast(param->axis); - - Layout ret; - bool is_new_layout_selected = false; - if (new_in_layouts.defined()) { // this function is called after some operators are alternated. - // If all the new input layouts are same, the new in layout gets selected. For axis, the new - // axis in the new layout is identified. The param->axis is then modified on the fly to conform - // to the new input layout. - const auto& concate_dim = old_in_layouts[0][axis]; - bool all_input_layouts_same = true; - for (auto new_layout : new_in_layouts) { - if (!new_layout.Equals(new_in_layouts[0])) { - all_input_layouts_same = false; - } - } - if (all_input_layouts_same) { - auto new_index = new_in_layouts[0].IndexOf(concate_dim); - ret = new_in_layouts[0]; - param->axis = new_index; - is_new_layout_selected = true; - } - } - - if (!is_new_layout_selected) { - // this function is called on the original correct relay ir - for (size_t i = 0; i < old_in_layouts.size(); ++i) { - if (old_in_layouts[i].defined()) { - ret = old_in_layouts[i]; - break; - } - } - - if (ret.ndim() <= axis || !ret[axis].IsPrimal()) { - return InferCorrectLayoutOutput({Layout::Undef()}, {Layout::Undef()}, attrs); - } - } - - return InferCorrectLayoutOutput(Array(old_in_layouts.size(), ret), {ret}, Attrs(param)); -} - -/*! - * \brief Infer output shape for reshape. - * - * \param data_shape The input data shape. - * \param attrs The attributes. - * \param reverse Whether to reverse the indices. - * \return Output shape. - */ -Array InferNewShape(const Array& data_shape, const Attrs& attrs, - bool reverse); - -} // namespace relay -} // namespace tvm -#endif // TVM_RELAY_OP_TENSOR_TRANSFORM_H_ diff --git a/src/relay/op/tensor/unary.cc b/src/relay/op/tensor/unary.cc deleted file mode 100644 index c6d149846e56..000000000000 --- a/src/relay/op/tensor/unary.cc +++ /dev/null @@ -1,528 +0,0 @@ -/* - * 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 unary.cc - * \brief Unary operators. - */ -#include -#include -#include -#include -#include - -#include "../make_op.h" -#include "../op_common.h" -#include "../type_relations.h" - -namespace tvm { -namespace relay { - -#define RELAY_UNARY_COMPUTE(FTOPI) \ - [](const Attrs& attrs, const Array& inputs, \ - const Type& out_type) -> Array { return {FTOPI(inputs[0])}; } - -RELAY_REGISTER_UNARY_OP("log") - .describe(R"code(Returns the log input array, computed element-wise. - -.. math:: - log(x) - -)code" TVM_ADD_FILELINE) - .set_support_level(1) - .set_attr("FTVMCompute", RELAY_UNARY_COMPUTE(topi::log)); - -RELAY_REGISTER_UNARY_OP("log2") - .describe(R"code(Returns the log to base 2 of input array, computed element-wise. - -.. math:: - log2(x) - -)code" TVM_ADD_FILELINE) - .set_support_level(1) - .set_attr("FTVMCompute", RELAY_UNARY_COMPUTE(topi::log2)); - -RELAY_REGISTER_UNARY_OP("log10") - .describe(R"code(Returns the log to base 10 of input array, computed element-wise. - -.. math:: - log10(x) - -)code" TVM_ADD_FILELINE) - .set_support_level(1) - .set_attr("FTVMCompute", RELAY_UNARY_COMPUTE(topi::log10)); - -RELAY_REGISTER_UNARY_OP("tan") - .describe(R"code(Returns the tan of input array, computed element-wise. - -.. math:: - Y = tan(X) - -)code" TVM_ADD_FILELINE) - .set_support_level(1) - .set_attr("FTVMCompute", RELAY_UNARY_COMPUTE(topi::tan)); - -RELAY_REGISTER_UNARY_OP("cos") - .describe(R"code(Returns the cos of input array, computed element-wise. - -.. math:: - Y = cos(X) - -)code" TVM_ADD_FILELINE) - .set_support_level(1) - .set_attr("FTVMCompute", RELAY_UNARY_COMPUTE(topi::cos)); - -RELAY_REGISTER_UNARY_OP("cosh") - .describe(R"code(Returns the cosh of input array, computed element-wise. - -.. math:: - Y = cosh(X) - -)code" TVM_ADD_FILELINE) - .set_support_level(1) - .set_attr("FTVMCompute", RELAY_UNARY_COMPUTE(topi::cosh)); - -RELAY_REGISTER_UNARY_OP("sin") - .describe(R"code(Returns the sin of input array, computed element-wise. - -.. math:: - Y = sin(X) - -)code" TVM_ADD_FILELINE) - .set_support_level(1) - .set_attr("FTVMCompute", RELAY_UNARY_COMPUTE(topi::sin)); - -RELAY_REGISTER_UNARY_OP("sinh") - .describe(R"code(Returns the sinh of input array, computed element-wise. - -.. math:: - Y = sinh(X) - -)code" TVM_ADD_FILELINE) - .set_support_level(1) - .set_attr("FTVMCompute", RELAY_UNARY_COMPUTE(topi::sinh)); - -RELAY_REGISTER_UNARY_OP("acos") - .describe(R"code(Returns the acos of input array, computed element-wise. - -.. math:: - Y = acos(X) - -)code" TVM_ADD_FILELINE) - .set_support_level(1) - .set_attr("FTVMCompute", RELAY_UNARY_COMPUTE(topi::acos)); - -RELAY_REGISTER_UNARY_OP("acosh") - .describe(R"code(Returns the acosh of input array, computed element-wise. - -.. math:: - Y = acosh(X) - -)code" TVM_ADD_FILELINE) - .set_support_level(1) - .set_attr("FTVMCompute", RELAY_UNARY_COMPUTE(topi::acosh)); - -RELAY_REGISTER_UNARY_OP("asin") - .describe(R"code(Returns the asin of input array, computed element-wise. - -.. math:: - Y = asin(X) - -)code" TVM_ADD_FILELINE) - .set_support_level(1) - .set_attr("FTVMCompute", RELAY_UNARY_COMPUTE(topi::asin)); - -RELAY_REGISTER_UNARY_OP("asinh") - .describe(R"code(Returns the asinh of input array, computed element-wise. - -.. math:: - Y = asinh(X) - -)code" TVM_ADD_FILELINE) - .set_support_level(1) - .set_attr("FTVMCompute", RELAY_UNARY_COMPUTE(topi::asinh)); - -RELAY_REGISTER_UNARY_OP("atan") - .describe(R"code(Returns the atan of input array, computed element-wise. - -.. math:: - Y = atan(X) - -)code" TVM_ADD_FILELINE) - .set_support_level(1) - .set_attr("FTVMCompute", RELAY_UNARY_COMPUTE(topi::atan)); - -RELAY_REGISTER_UNARY_OP("atanh") - .describe(R"code(Returns the atanh of input array, computed element-wise. - -.. math:: - Y = atanh(X) - -)code" TVM_ADD_FILELINE) - .set_support_level(1) - .set_attr("FTVMCompute", RELAY_UNARY_COMPUTE(topi::atanh)); - -RELAY_REGISTER_UNARY_OP("exp") - .describe(R"code(Returns the exp input array, computed element-wise. - -.. math:: - \exp(x) - -)code" TVM_ADD_FILELINE) - .set_support_level(1) - .set_attr("FTVMCompute", RELAY_UNARY_COMPUTE(topi::exp)); - -RELAY_REGISTER_UNARY_OP("fast_exp") - .describe(R"code(Returns the fast_exp input array, computed element-wise. - -.. math:: - \fast_exp(x) - -)code" TVM_ADD_FILELINE) - .set_support_level(1) - .set_attr("FTVMCompute", RELAY_UNARY_COMPUTE(topi::fast_exp)); - -RELAY_REGISTER_UNARY_OP("erf") - .describe(R"code(Returns the error function value for input array, computed element-wise. - -.. math:: - \erf(x) - -)code" TVM_ADD_FILELINE) - .set_support_level(1) - .set_attr("FTVMCompute", RELAY_UNARY_COMPUTE(topi::erf)); - -RELAY_REGISTER_UNARY_OP("fast_erf") - .describe(R"code(Returns the error function value for input array, computed element-wise. - -.. math:: - \fast_erf(x) - -)code" TVM_ADD_FILELINE) - .set_support_level(1) - .set_attr("FTVMCompute", RELAY_UNARY_COMPUTE(topi::fast_erf)); - -RELAY_REGISTER_UNARY_OP("sqrt") - .describe(R"code(Returns the sqrt input array, computed element-wise. - -.. math:: - sqrt(x) - -)code" TVM_ADD_FILELINE) - .set_support_level(1) - .set_attr("FTVMCompute", RELAY_UNARY_COMPUTE(topi::sqrt)); - -RELAY_REGISTER_UNARY_OP("rsqrt") - .describe(R"code(Returns the rsqrt input array, computed element-wise. - -.. math:: - 1/sqrt(x) - -)code" TVM_ADD_FILELINE) - .set_support_level(1) - .set_attr("FTVMCompute", RELAY_UNARY_COMPUTE(topi::rsqrt)); - -RELAY_REGISTER_UNARY_OP("zeros_like") - .describe(R"code(Returns an array of zeros, with same type and shape as the input. -)code" TVM_ADD_FILELINE) - .set_support_level(4); - -RELAY_REGISTER_UNARY_OP("ones_like") - .describe(R"code(Returns an array of ones, with same type and shape as the input. -)code" TVM_ADD_FILELINE) - .set_support_level(4); - -RELAY_REGISTER_UNARY_OP("sigmoid") - .describe(R"code(Returns the sigmoid input array, computed element-wise. - -.. math:: - sigmoid(x) - -)code" TVM_ADD_FILELINE) - .set_support_level(1) - .set_attr("FTVMCompute", RELAY_UNARY_COMPUTE(topi::sigmoid)); - -RELAY_REGISTER_UNARY_OP("copy") - .describe(R"code(Copy a tensor. -)code" TVM_ADD_FILELINE) - .set_support_level(3) - .set_attr("FTVMCompute", RELAY_UNARY_COMPUTE(topi::identity)); - -// relay.clip -TVM_REGISTER_NODE_TYPE(ClipAttrs); - -Expr MakeClip(Expr a, double a_min, double a_max) { - auto attrs = make_object(); - attrs->a_min = a_min; - attrs->a_max = a_max; - static const Op& op = Op::Get("clip"); - return Call(op, {a}, Attrs(attrs), {}); -} - -TVM_REGISTER_GLOBAL("relay.op._make.clip").set_body_typed(MakeClip); - -RELAY_REGISTER_OP("clip") - .describe(R"code(Clip tensor values. -This function takes a tensor, a minimum value `a_min`, and a maximum value `a_max`, and returns a clipped tensor where all values below `a_min` are set to `a_min` and all values above `a_max` are set to `a_max`. `a_min` and `a_max` are cast to the tensor's dtype. -)code" TVM_ADD_FILELINE) - .set_num_inputs(1) - .add_argument("data", "Tensor", "The input tensor.") - .add_type_rel("Identity", IdentityRel) - .set_attr("TOpPattern", kElemWise) - .set_attr("TOpIsStateful", false) - .set_attr("FInferCorrectLayout", ElemwiseArbitraryLayout) - .set_attrs_type() - .set_support_level(3); - -// relay.fixed_point_multiply -TVM_REGISTER_NODE_TYPE(FixedPointMultiplyAttrs); - -TVM_REGISTER_GLOBAL("relay.op._make.fixed_point_multiply") - .set_body_typed([](Expr a, int32_t multiplier, int32_t shift) { - auto attrs = make_object(); - attrs->multiplier = multiplier; - attrs->shift = shift; - static const Op& op = Op::Get("fixed_point_multiply"); - return Call(op, {a}, Attrs(attrs), {}); - }); - -RELAY_REGISTER_OP("fixed_point_multiply") - .describe(R"code(fixed point multiplication)code" TVM_ADD_FILELINE) - .set_num_inputs(1) - .add_argument("data", "Tensor", "The input tensor.") - .add_type_rel("Identity", IdentityRel) - .set_attr("TOpPattern", kElemWise) - .set_attr("TOpIsStateful", false) - .set_attr("FInferCorrectLayout", ElemwiseArbitraryLayout) - .set_attrs_type() - .set_support_level(10); - -RELAY_REGISTER_UNARY_OP("floor") - .describe(R"code(Returns the floor of input array, computed element-wise. -)code" TVM_ADD_FILELINE) - .set_support_level(3) - .set_attr("FTVMCompute", RELAY_UNARY_COMPUTE(topi::floor)); - -RELAY_REGISTER_UNARY_OP("ceil") - .describe(R"code(Returns the ceil of input array, computed element-wise. - -.. math:: - ceil(x) - -)code" TVM_ADD_FILELINE) - .set_support_level(3) - .set_attr("FTVMCompute", RELAY_UNARY_COMPUTE(topi::ceil)); - -RELAY_REGISTER_UNARY_OP("trunc") - .describe(R"code(Returns the trunc of input array, computed element-wise. - -.. math:: - trunc(x) - -)code" TVM_ADD_FILELINE) - .set_support_level(3) - .set_attr("FTVMCompute", RELAY_UNARY_COMPUTE(topi::trunc)); - -RELAY_REGISTER_UNARY_OP("round") - .describe(R"code(Returns the round of input array, computed element-wise. - -.. math:: - round(x) - -)code" TVM_ADD_FILELINE) - .set_support_level(3) - .set_attr("FTVMCompute", RELAY_UNARY_COMPUTE(topi::round)); - -RELAY_REGISTER_UNARY_OP("sign") - .describe(R"code(Returns the sign of input array, computed element-wise. - -.. numpy:: - sign(x) - -)code" TVM_ADD_FILELINE) - .set_support_level(3) - .set_attr("FTVMCompute", RELAY_UNARY_COMPUTE(topi::sign)); - -RELAY_REGISTER_UNARY_OP("abs") - .describe(R"code(Returns the abs of input array, computed element-wise. - -.. math:: - abs(x) - -)code" TVM_ADD_FILELINE) - .set_support_level(3) - .set_attr("FTVMCompute", RELAY_UNARY_COMPUTE(topi::abs)); - -RELAY_REGISTER_UNARY_OP("tanh") - .describe(R"code(Returns the tanh of input array, computed element-wise. - -.. math:: - Y = sinh(X) / cosh(X) - -)code" TVM_ADD_FILELINE) - .set_support_level(1) - .set_attr("FTVMCompute", RELAY_UNARY_COMPUTE(topi::tanh)); - -RELAY_REGISTER_UNARY_OP("fast_tanh") - .describe(R"code(Returns the fast_tanh of input array, computed element-wise. - -.. math:: - Y = sinh(X) / cosh(X) - -)code" TVM_ADD_FILELINE) - .set_support_level(1) - .set_attr("FTVMCompute", RELAY_UNARY_COMPUTE(topi::fast_tanh)); - -RELAY_REGISTER_UNARY_OP("negative") - .describe(R"code(Returns the numeric negative of input array, computed element-wise. - -.. math:: - -(x) - -)code" TVM_ADD_FILELINE) - .set_support_level(3) - .set_attr("FTVMCompute", RELAY_UNARY_COMPUTE(topi::negative)); - -RELAY_REGISTER_UNARY_OP("logical_not") - .describe(R"code(Returns the logical inverse of input array, computed element-wise. - -.. math:: - !(x) - -)code" TVM_ADD_FILELINE) - .set_support_level(4) - .set_attr("FTVMCompute", RELAY_UNARY_COMPUTE(topi::logical_not)); - -RELAY_REGISTER_UNARY_OP("bitwise_not") - .describe(R"code(Returns the bitwise inverse of input array, computed element-wise. - -.. math:: - ~(x) - -)code" TVM_ADD_FILELINE) - .set_support_level(4) - .set_attr("FTVMCompute", RELAY_UNARY_COMPUTE(topi::bitwise_not)); - -Array ShapeOfCompute(const Attrs& attrs, const Array& inputs, - const Type& out_type) { - ICHECK_EQ(inputs.size(), 1); - const auto* param = attrs.as(); - ICHECK(param != nullptr); - return {topi::shape(inputs[0], param->dtype)}; -} - -Expr MakeShapeOf(Expr data, DataType dtype) { - auto attrs = make_object(); - attrs->dtype = dtype; - static const Op& op = Op::Get("shape_of"); - return Call(op, {data}, Attrs(attrs), {}); -} - -TVM_REGISTER_GLOBAL("relay.op._make.shape_of").set_body_typed(MakeShapeOf); - -RELAY_REGISTER_OP("shape_of") - .describe(R"code(Returns a tensor representing the shape of a tensor. - -)code" TVM_ADD_FILELINE) - .set_num_inputs(1) - .set_attrs_type() - .add_argument("data", "Tensor", "The input tensor.") - .add_type_rel("ShapeOf", ShapeOfRel) - .set_attr("TOpIsStateful", false) - // Use kOpaque for shape_of op for now since it won't be performance critic, - // and it makes things easier for dynamic shape func - .set_attr("TOpPattern", kOpaque) - .set_support_level(10) - .set_attr("FTVMCompute", ShapeOfCompute); - -TVM_REGISTER_NODE_TYPE(NdarraySizeAttrs); - -bool NdarraySizeRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - ICHECK_EQ(num_inputs, 1); - auto tt = types[0].as(); - - if (tt == nullptr) { - return false; - } - - const auto* param = attrs.as(); - ICHECK(param != nullptr); - reporter->Assign(types[1], TensorType({}, param->dtype)); - return true; -} - -Array NdarraySizeCompute(const Attrs& attrs, const Array& inputs, - const Type& out_type) { - ICHECK_EQ(inputs.size(), 1); - const auto* param = attrs.as(); - ICHECK(param != nullptr); - return Array{topi::ndarray_size(inputs[0], param->dtype)}; -} - -TVM_REGISTER_GLOBAL("relay.op._make.ndarray_size").set_body_typed([](Expr data, DataType dtype) { - auto attrs = make_object(); - attrs->dtype = dtype; - static const Op& op = Op::Get("ndarray_size"); - return Call(op, {data}, Attrs(attrs), {}); -}); - -RELAY_REGISTER_OP("ndarray_size") - .describe(R"code(Returns a tensor representing the number of elements of input tensor. - -)code" TVM_ADD_FILELINE) - .set_num_inputs(1) - .set_attrs_type() - .add_argument("data", "Tensor", "The input tensor.") - .add_type_rel("NdarraySize", NdarraySizeRel) - .set_attr("TOpIsStateful", false) - .set_attr("TOpPattern", kInjective) - .set_attr("FInferCorrectLayout", ElemwiseArbitraryLayout) - .set_support_level(10) - .set_attr("FTVMCompute", NdarraySizeCompute); - -RELAY_REGISTER_UNARY_OP("isnan") - .describe(R"code(Returns whether the input contains any NaN, computed element-wise. -.. math:: - isnan(x) -)code" TVM_ADD_FILELINE) - .set_support_level(3) - .add_type_rel("IdentityCompRel", IdentityCompRel) - .set_attr("FTVMCompute", RELAY_UNARY_COMPUTE(topi::isnan)); - -RELAY_REGISTER_UNARY_OP("isfinite") - .describe(R"code(Returns the finiteness of input, computed element-wise. -.. math:: - isfinite(x) -)code" TVM_ADD_FILELINE) - .set_support_level(3) - .add_type_rel("IdentityCompRel", IdentityCompRel) - .set_attr("FTVMCompute", RELAY_UNARY_COMPUTE(topi::isfinite)); - -RELAY_REGISTER_UNARY_OP("isinf") - .describe(R"code(Returns the infiniteness of input, computed element-wise. -.. math:: - isinf(x) -)code" TVM_ADD_FILELINE) - .set_support_level(3) - .add_type_rel("IdentityCompRel", IdentityCompRel) - .set_attr("FTVMCompute", RELAY_UNARY_COMPUTE(topi::isinf)); - -} // namespace relay -} // namespace tvm diff --git a/src/relay/op/type_relations.cc b/src/relay/op/type_relations.cc deleted file mode 100644 index 71e58bb927f5..000000000000 --- a/src/relay/op/type_relations.cc +++ /dev/null @@ -1,173 +0,0 @@ -/* - * 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 type_relations.cc - * \brief A set of utilities and common functionality - * for type relations. - */ -#include "./type_relations.h" - -#include -#include -#include -#include -#include - -#include - -namespace tvm { -namespace relay { - -bool IdentityRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - for (size_t i = 1; i < types.size(); ++i) { - reporter->Assign(types[i], types[0]); - } - return true; -} - -bool EqualCheck(const IndexExpr& lhs, const IndexExpr& rhs) { - IndexExpr diff = lhs - rhs; - if (const int64_t* pdiff = tir::as_const_int(diff)) { - return pdiff[0] == 0; - } - // symbolic - tvm::arith::Analyzer ana; - diff = ana.Simplify(diff); - if (const int64_t* pdiff = tir::as_const_int(diff)) { - return pdiff[0] == 0; - } - return false; -} - -bool EqualConstInt(const IndexExpr& lhs, int64_t value) { - if (const int64_t* pvalue = tir::as_const_int(lhs)) { - return pvalue[0] == value; - } - return false; -} - -TensorType ConcreteBroadcast(const TensorType& t1, const TensorType& t2, DataType output_dtype) { - std::vector oshape; - size_t ndim1 = t1->shape.size(); - size_t ndim2 = t2->shape.size(); - size_t i = 1; - for (; i <= std::min(ndim1, ndim2); ++i) { - IndexExpr s1 = t1->shape[ndim1 - i]; - IndexExpr s2 = t2->shape[ndim2 - i]; - if (EqualConstInt(s1, 1)) { - oshape.push_back(s2); - } else if (EqualConstInt(s2, 1)) { - oshape.push_back(s1); - } else if (s1.as()) { - // s1 == 1 || s1 == s2 - oshape.push_back(s2); - } else if (s2.as()) { - // s2 == 1 || s2 == s1 - oshape.push_back(s1); - } else if (EqualCheck(s1, s2)) { - oshape.push_back(s1); - } else { - throw CompileError(ErrorBuilder() << "Incompatible broadcast type " << t1 << " and " << t2); - } - } - - size_t max_ndim = std::max(ndim1, ndim2); - auto& rshape = (ndim1 > ndim2) ? t1->shape : t2->shape; - for (; i <= max_ndim; ++i) { - oshape.push_back(rshape[max_ndim - i]); - } - return TensorType(Array(oshape.rbegin(), oshape.rend()), output_dtype); -} - -bool BroadcastRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - ICHECK_EQ(types.size(), 3); - // DLOG(INFO) << "In1:" << types[0] << ",In2:" << types[1] - // << ",Out:" << types[2] << std::endl; - if (auto* t0 = types[0].as()) { - if (auto* t1 = types[1].as()) { - if (t0->dtype != t1->dtype) { - reporter->GetDiagCtx().Emit(Diagnostic::Error(t0->span) - << "data types " << t0->dtype << " and " << t1->dtype - << " do not match in BroadcastRel"); - } - reporter->Assign( - types[2], ConcreteBroadcast(GetRef(t0), GetRef(t1), t0->dtype)); - return true; - } - } - return false; -} - -bool BroadcastCompRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - ICHECK_EQ(types.size(), 3); - // DLOG(INFO) << "In1:" << types[0] << ",In2:" << types[1] - // << ",Out:" << types[2] << std::endl; - if (auto* t0 = types[0].as()) { - if (auto* t1 = types[1].as()) { - if (t0->dtype != t1->dtype) { - reporter->GetDiagCtx().Emit(Diagnostic::Error(t0->span) - << "data types " << t0->dtype << " and " << t1->dtype - << " do not match in BroadcastCompRel"); - } - reporter->Assign(types[2], ConcreteBroadcast(GetRef(t0), GetRef(t1), - DataType::Bool())); - return true; - } - } - return false; -} - -bool IdentityCompRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - if (const auto* t0 = types[0].as()) { - Type out_type = TensorType(t0->shape, DataType::Bool()); - reporter->Assign(types[1], out_type); - return true; - } - return false; -} - -Array RankShape(const Array& shape) { - if (shape.size() == 0) { - return {}; - } else { - return {tvm::Integer(shape.size())}; - } -} - -bool ShapeOfRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - ICHECK_EQ(num_inputs, 1); - auto tt = types[0].as(); - if (tt == nullptr) { - return false; - } - const auto* param = attrs.as(); - ICHECK(param != nullptr); - auto rank_shape = RankShape(tt->shape); - reporter->Assign(types[1], TensorType(rank_shape, param->dtype)); - return true; -} - -} // namespace relay -} // namespace tvm diff --git a/src/relay/op/type_relations.h b/src/relay/op/type_relations.h deleted file mode 100644 index 740766172ddc..000000000000 --- a/src/relay/op/type_relations.h +++ /dev/null @@ -1,106 +0,0 @@ -/* - * 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 tvm/relay/op/type_relations.h - * \brief A set of utilities and common functionality - * for type relations. - */ -#ifndef TVM_RELAY_OP_TYPE_RELATIONS_H_ -#define TVM_RELAY_OP_TYPE_RELATIONS_H_ - -#include -#include - -#include - -namespace tvm { -namespace relay { -/*! - * \brief The identity type relation, all the types are equal. - * - * \param types The input and output types to the relation. - * \param num_inputs The number of input arguments. - * \param attrs The attributes - * \param reporter The reporter. - * \return true whether relation has been resolved. - */ -bool IdentityRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter); - -/*! - * \brief The broadcast type relation, implements the broadcasting - * rule over the two input types producing the broadcasted type. - * - * \param types The input and output types to the relation. - * \param num_inputs The number of input arguments. - * \param attrs The attributes - * \param reporter The reporter. - * \return true whether relation has been resolved. - */ -bool BroadcastRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter); - -/*! - * \brief Determine the broadcasted shape from two input shapes - * \param t1 One of two Tensortype whose shapes are broadcasted - * \param t2 One of two Tensortype whose shapes are broadcasted - * \param output_dtype dtype of the output TensorType - * \return A TensorType whose shape is broadcasted from two input TensorType. - */ -TensorType ConcreteBroadcast(const TensorType& t1, const TensorType& t2, DataType output_dtype); - -/*! - * \brief The broadcast type relation, implements the broadcasting - * rule over the two input types producing the broadcasted type. - * - * This differs from BroadcastRel in the return dtype, - * it instead returns bool(uint8), for use in comparsion operators - * such as equal, not_equal, lt, and so on. - * - * \param types The input and output types to the relation. - * \param num_inputs The number of input arguments. - * \param attrs The attributes - * \param reporter The reporter. - * \return true whether relation has been resolved. - */ -bool BroadcastCompRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter); - -bool IdentityCompRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter); - -Array RankShape(const Array& shape); - -/*! - * \brief The shape of type relation. - * - * \param types The input and output types to the relation. - * \param num_inputs The number of input arguments. - * \param attrs The attributes - * \param reporter The reporter. - * \return true whether relation has been resolved. - */ -bool ShapeOfRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter); - -} // namespace relay -} // namespace tvm - -#endif // TVM_RELAY_OP_TYPE_RELATIONS_H_ diff --git a/src/relay/op/vision/multibox_op.cc b/src/relay/op/vision/multibox_op.cc deleted file mode 100644 index c76316aad401..000000000000 --- a/src/relay/op/vision/multibox_op.cc +++ /dev/null @@ -1,144 +0,0 @@ -/* - * 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 multibox_op.cc - * \brief Multibox related operators - */ -#include -#include -#include - -namespace tvm { -namespace relay { - -TVM_REGISTER_NODE_TYPE(MultiBoxPriorAttrs); - -bool MultiboxPriorRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - ICHECK_EQ(types.size(), 2); - const auto* data = types[0].as(); - const MultiBoxPriorAttrs* param = attrs.as(); - const auto& dshape = data->shape; - ICHECK_EQ(dshape.size(), 4) << "Input data should be 4D: " - "[batch, channel, height, width]"; - IndexExpr in_height = dshape[2]; - IndexExpr in_width = dshape[3]; - int num_sizes = static_cast(param->sizes.size()); - int num_ratios = static_cast(param->ratios.size()); - - // since input sizes are same in each batch, we could share MultiBoxPrior - std::vector oshape({1, in_height * in_width * (num_sizes + num_ratios - 1), 4}); - - // assign output type - reporter->Assign(types[1], TensorType(oshape, data->dtype)); - return true; -} - -Expr MakeMultiBoxPrior(Expr data, Array sizes, Array ratios, - Array steps, Array offsets, bool clip) { - auto attrs = make_object(); - attrs->sizes = std::move(sizes); - attrs->ratios = std::move(ratios); - attrs->steps = std::move(steps); - attrs->offsets = std::move(offsets); - attrs->clip = clip; - static const Op& op = Op::Get("vision.multibox_prior"); - return Call(op, {data}, Attrs(attrs), {}); -} - -TVM_REGISTER_GLOBAL("relay.op.vision._make.multibox_prior").set_body_typed(MakeMultiBoxPrior); - -RELAY_REGISTER_OP("vision.multibox_prior") - .describe(R"doc("Generate prior(anchor) boxes from data, sizes and ratios." -)doc" TVM_ADD_FILELINE) - .set_attrs_type() - .set_num_inputs(1) - .add_argument("data", "Tensor", "The input tensor.") - .set_support_level(5) - .add_type_rel("MultiBoxPrior", MultiboxPriorRel); - -TVM_REGISTER_NODE_TYPE(MultiBoxTransformLocAttrs); - -bool MultiBoxTransformLocRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - ICHECK_EQ(types.size(), 4); - - const auto* cls_prob = types[0].as(); - const auto* loc_pred = types[1].as(); - const auto* anchor = types[2].as(); - - if (cls_prob == nullptr || loc_pred == nullptr || anchor == nullptr) { - return false; - } - - const auto& cls_shape = cls_prob->shape; - const auto& loc_shape = loc_pred->shape; - const auto& anchor_shape = anchor->shape; - - ICHECK_EQ(cls_shape.size(), 3U) << "The dimension of class probability should be 3, but received " - << cls_shape.size(); - ICHECK_EQ(loc_shape.size(), 2U) - << "The dimension of location prediction should be 2, but received " << loc_shape.size(); - ICHECK_EQ(anchor_shape.size(), 3U) - << "The dimension of anchor should be 3, but received " << anchor_shape.size(); - - ICHECK(reporter->AssertEQ(cls_shape[2], anchor_shape[1])) << "Number of anchors mismatch found"; - ICHECK(reporter->AssertEQ(cls_shape[2] * 4, loc_shape[1])) << "# anchors mismatch with # loc."; - ICHECK(reporter->Assert(anchor_shape[1] > 0)) << "Number of anchors must > 0."; - ICHECK(reporter->AssertEQ(anchor_shape[2], 4)); - - std::vector oshape0({cls_shape[0], anchor_shape[1], 6}); - std::vector oshape1({cls_shape[0]}); - std::vector fields; - fields.push_back(TensorType(oshape0, cls_prob->dtype)); - fields.push_back(TensorType(oshape1, DataType::Int(32))); - - // assign output type - reporter->Assign(types[3], TupleType(Array(fields))); - return true; -} - -Expr MakeMultiBoxTransformLoc(Expr cls_prob, Expr loc_pred, Expr anchor, bool clip, - double threshold, Array variances, bool keep_background) { - auto attrs = make_object(); - attrs->clip = std::move(clip); - attrs->threshold = std::move(threshold); - attrs->variances = std::move(variances); - attrs->keep_background = std::move(keep_background); - static const Op& op = Op::Get("vision.multibox_transform_loc"); - return Call(op, {cls_prob, loc_pred, anchor}, Attrs(attrs), {}); -} - -TVM_REGISTER_GLOBAL("relay.op.vision._make.multibox_transform_loc") - .set_body_typed(MakeMultiBoxTransformLoc); - -RELAY_REGISTER_OP("vision.multibox_transform_loc") - .describe(R"doc("Location transformation for multibox detection." -)doc" TVM_ADD_FILELINE) - .set_attrs_type() - .set_num_inputs(3) - .add_argument("cls_prob", "Tensor", "Class probabilities.") - .add_argument("loc_pred", "Tensor", "Location regression predictions.") - .add_argument("anchor", "Tensor", "Multibox prior anchor boxes") - .add_type_rel("MultiBoxTransformLoc", MultiBoxTransformLocRel) - .set_support_level(5); - -} // namespace relay -} // namespace tvm diff --git a/src/relay/op/vision/nms.cc b/src/relay/op/vision/nms.cc deleted file mode 100644 index 24873e468a41..000000000000 --- a/src/relay/op/vision/nms.cc +++ /dev/null @@ -1,275 +0,0 @@ -/* - * 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 nms.cc - * \brief Non-maximum suppression operators - */ -#include -#include -#include - -namespace tvm { -namespace relay { - -TVM_REGISTER_NODE_TYPE(GetValidCountsAttrs); - -bool GetValidCountRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - ICHECK_EQ(types.size(), 3); - const auto* data = types[0].as(); - if (data == nullptr) return false; - const auto& dshape = data->shape; - ICHECK_EQ(dshape.size(), 3) << "Input data should be 3-D."; - - std::vector oshape({data->shape[0]}); - std::vector oshape_indices({data->shape[0], data->shape[1]}); - std::vector fields; - fields.push_back(TensorType(oshape, DataType::Int(32))); - fields.push_back(TensorType(data->shape, data->dtype)); - fields.push_back(TensorType(oshape_indices, DataType::Int(32))); - - // assign output type - reporter->Assign(types[2], TupleType(Array(fields))); - return true; -} - -Expr MakeGetValidCounts(Expr data, Expr score_threshold, int id_index, int score_index) { - auto attrs = make_object(); - attrs->id_index = id_index; - attrs->score_index = score_index; - static const Op& op = Op::Get("vision.get_valid_counts"); - return Call(op, {data, score_threshold}, Attrs(attrs), {}); -} - -TVM_REGISTER_GLOBAL("relay.op.vision._make.get_valid_counts").set_body_typed(MakeGetValidCounts); - -RELAY_REGISTER_OP("vision.get_valid_counts") - .describe(R"doc(Get valid count of bounding boxes given -a score threshold. Also moves valid boxes to the top of -input data. -)doc" TVM_ADD_FILELINE) - .set_num_inputs(2) - .add_argument("data", "Tensor", "Input data.") - .add_argument("score_threshold", "Tensor", "Minimum Score.") - .set_support_level(5) - .add_type_rel("GetValidCount", GetValidCountRel); - -TVM_REGISTER_NODE_TYPE(NonMaximumSuppressionAttrs); - -bool NMSRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - ICHECK_EQ(types.size(), 6); - const auto* data = types[0].as(); - if (data == nullptr) return false; - const auto* valid_count = types[1].as(); - if (valid_count == nullptr) return false; - const NonMaximumSuppressionAttrs* param = attrs.as(); - const auto& dshape = data->shape; - const auto& vshape = valid_count->shape; - ICHECK_EQ(dshape.size(), 3) << "Input data should be 3-D."; - ICHECK_EQ(vshape.size(), 1) << "Input valid count should be 1-D."; - - // assign output type - if (param->return_indices) { - std::vector fields; - // dynamic happens for return_indices in TensorFlow & ONNX - std::vector oshape({dshape[0], dshape[1]}); - fields.push_back(TensorType(oshape, DataType::Int(32))); - std::vector countshape({dshape[0], 1}); - fields.push_back(TensorType(countshape, DataType::Int(32))); - reporter->Assign(types[5], TupleType(Array(fields))); - } else { - reporter->Assign(types[5], TensorType(dshape, data->dtype)); - } - return true; -} - -Expr MakeNMS(Expr data, Expr valid_count, Expr indices, Expr max_output_size, Expr iou_threshold, - bool force_suppress, int top_k, int coord_start, int score_index, int id_index, - bool return_indices, bool invalid_to_bottom) { - auto attrs = make_object(); - attrs->force_suppress = force_suppress; - attrs->top_k = top_k; - attrs->coord_start = coord_start; - attrs->score_index = score_index; - attrs->id_index = id_index; - attrs->return_indices = return_indices; - attrs->invalid_to_bottom = invalid_to_bottom; - static const Op& op = Op::Get("vision.non_max_suppression"); - return Call(op, {data, valid_count, indices, max_output_size, iou_threshold}, Attrs(attrs), {}); -} - -TVM_REGISTER_GLOBAL("relay.op.vision._make.non_max_suppression").set_body_typed(MakeNMS); - -RELAY_REGISTER_OP("vision.non_max_suppression") - .describe(R"doc(Non-maximum suppression. The input boxes should -be in the format of [class_id, score, left, top, right, bottom] -or [score, left, top, right, bottom]. Set id_index to be -1 to -ignore class_id axis. -)doc" TVM_ADD_FILELINE) - .set_num_inputs(5) - .add_argument("data", "Tensor", "Input data.") - .add_argument("valid_count", "Tensor", "Number of valid anchor boxes.") - .add_argument("indices", "Tensor", "Corresponding indices in original input tensor.") - .add_argument("max_output_size", "Tensor", "Max number of output valid boxes.") - .add_argument("iou_threshold", "Tensor", "Threshold for box overlap.") - .set_support_level(5) - .add_type_rel("NMS", NMSRel); - -TVM_REGISTER_NODE_TYPE(AllClassNonMaximumSuppressionAttrs); - -bool AllClassNMSRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - ICHECK_EQ(types.size(), 6); - const auto* boxes = types[0].as(); - if (boxes == nullptr) return false; - const auto* scores = types[1].as(); - if (scores == nullptr) return false; - - const auto& boxes_shape = boxes->shape; - const auto& scores_shape = scores->shape; - ICHECK_EQ(boxes_shape.size(), 3) << "Input boxes should be 3-D."; - ICHECK_EQ(scores_shape.size(), 3) << "Input scores count should be 3-D."; - - IndexExpr batch = boxes_shape[0]; - IndexExpr num_classes = scores_shape[1]; - IndexExpr num_boxes = boxes_shape[1]; - - const auto* param = attrs.as(); - CHECK(param); - - std::vector fields; - if (param->output_format == "onnx") { - IndexExpr num_total_boxes = Any(); - if (!batch.as() && !num_boxes.as()) { - num_total_boxes = batch * num_classes * num_boxes; - } - std::vector oshape{num_total_boxes, 3}; - std::vector counts_shape{1}; - fields.push_back(TensorType(oshape, DataType::Int(64))); - fields.push_back(TensorType(counts_shape, DataType::Int(64))); - } else { - IndexExpr num_total_boxes_per_batch = Any(); - if (!num_boxes.as()) { - num_total_boxes_per_batch = num_classes * num_boxes; - } - std::vector indices_shape{batch, num_total_boxes_per_batch, 2}; - std::vector scores_shape{batch, num_total_boxes_per_batch}; - std::vector counts_shape{batch}; - fields.push_back(TensorType(indices_shape, DataType::Int(64))); - fields.push_back(TensorType(scores_shape, DataType::Float(32))); - fields.push_back(TensorType(counts_shape, DataType::Int(64))); - } - reporter->Assign(types[5], TupleType(Array(fields))); - return true; -} - -Expr MakeAllClassNMS(Expr boxes, Expr scores, Expr max_output_boxes_per_class, Expr iou_threshold, - Expr score_threshold, std::string output_format = "onnx") { - auto attrs = make_object(); - attrs->output_format = std::move(output_format); - static const Op& op = Op::Get("vision.all_class_non_max_suppression"); - return Call(op, {boxes, scores, max_output_boxes_per_class, iou_threshold, score_threshold}, - Attrs(attrs), {}); -} - -TVM_REGISTER_GLOBAL("relay.op.vision._make.all_class_non_max_suppression") - .set_body_typed(MakeAllClassNMS); - -RELAY_REGISTER_OP("vision.all_class_non_max_suppression") - .describe(R"doc(Non-maximum suppression operator for object detection, corresponding to ONNX - NonMaxSuppression and TensorFlow combined_non_max_suppression. - NMS is performed for each class separately -)doc" TVM_ADD_FILELINE) - .set_num_inputs(5) - .add_argument("boxes", "Tensor", "The input boxes in the format [batch, num_boxes, 4].") - .add_argument("scores", "Tensor", - "Scores for each box and class in the format [batch, num_classes, num_boxes].") - .add_argument("max_output_boxes_per_class", "Tensor", - "The maximum number of output boxes per class.") - .add_argument("iou_threshold", "Tensor", "The IoU threshold for box the overlap test.") - .add_argument("score_threshold", "Tensor", - "The score threshold to filter out low score boxes early.") - .set_support_level(5) - .add_type_rel("AllClassNMS", AllClassNMSRel); - -TVM_REGISTER_NODE_TYPE(RegularNonMaximumSuppressionAttrs); - -bool RegularNMSRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - ICHECK_EQ(types.size(), 3); - const auto* boxes = types[0].as(); - if (boxes == nullptr) return false; - const auto* scores = types[1].as(); - if (scores == nullptr) return false; - - const auto& boxes_shape = boxes->shape; - const auto& scores_shape = scores->shape; - ICHECK_EQ(boxes_shape.size(), 3) << "Input boxes should be 3-D."; - ICHECK_EQ(scores_shape.size(), 3) << "Input scores count should be 3-D."; - - IndexExpr num_batches = boxes_shape[0]; - - const auto* param = attrs.as(); - CHECK(param); - - std::vector fields; - std::vector nmsed_boxes_shape{num_batches, param->max_detections, 4}; - std::vector nmsed_classes_shape{num_batches, param->max_detections}; - std::vector nmsed_scores_shape{num_batches, param->max_detections}; - std::vector nmsed_detections_number_shape{num_batches}; - fields.push_back(TensorType(nmsed_boxes_shape, DataType::Float(32))); - fields.push_back(TensorType(nmsed_classes_shape, DataType::Float(32))); - fields.push_back(TensorType(nmsed_scores_shape, DataType::Float(32))); - fields.push_back(TensorType(nmsed_detections_number_shape, DataType::Int(32))); - - reporter->Assign(types[2], TupleType(Array(fields))); - return true; -} - -Expr MakeRegularNMS(Expr boxes, Expr scores, int32_t max_detections_per_class, - int32_t max_detections, int32_t num_classes, double iou_threshold, - double score_threshold) { - auto attrs = make_object(); - attrs->max_detections_per_class = max_detections_per_class; - attrs->max_detections = max_detections; - attrs->num_classes = num_classes; - attrs->iou_threshold = iou_threshold; - attrs->score_threshold = score_threshold; - static const Op& op = Op::Get("vision.regular_non_max_suppression"); - return Call(op, {boxes, scores}, Attrs(attrs), {}); -} - -TVM_REGISTER_GLOBAL("relay.op.vision._make.regular_non_max_suppression") - .set_body_typed(MakeRegularNMS); - -RELAY_REGISTER_OP("vision.regular_non_max_suppression") - .describe(R"doc(TBD)doc" TVM_ADD_FILELINE) - .set_num_inputs(2) - .add_argument("boxes", "Tensor", - "3-D tensor with shape (batch_size, num_boxes, 4). The four values in boxes " - "encode (ymin, xmin, ymax, xmax) coordinates of a box.") - .add_argument("scores", "Tensor", - "3-D tensor with shape (batch_size, num_boxes, num_classes_with_background).") - .set_support_level(5) - .add_type_rel("RegularNMS", RegularNMSRel); - -} // namespace relay -} // namespace tvm diff --git a/src/relay/op/vision/rcnn_op.cc b/src/relay/op/vision/rcnn_op.cc deleted file mode 100644 index 5f1948dd951c..000000000000 --- a/src/relay/op/vision/rcnn_op.cc +++ /dev/null @@ -1,242 +0,0 @@ -/* - * 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 rcnn_op.cc - * \brief Faster RCNN and Mask RCNN operators - */ -#include -#include -#include - -#include "../../transforms/infer_layout_utils.h" - -namespace tvm { -namespace relay { - -TVM_REGISTER_NODE_TYPE(ROIAlignAttrs); - -bool ROIAlignRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - auto roi_align_attrs = attrs.as(); - ICHECK_EQ(types.size(), 3); - const auto* data = types[0].as(); - const auto* rois = types[1].as(); - ICHECK(data); - ICHECK(rois); - const auto& dshape = data->shape; - const auto& rshape = rois->shape; - ICHECK(roi_align_attrs); - ICHECK_EQ(dshape.size(), 4) << "Input data should be 4-D."; - ICHECK_EQ(rshape.size(), 2) << "Input rois should be 2-D."; - // assign output type - std::vector oshape; - if (roi_align_attrs->layout == "NCHW") { - oshape = {rshape[0], dshape[1], roi_align_attrs->pooled_size[0], - roi_align_attrs->pooled_size[1]}; - } else { - ICHECK_EQ(roi_align_attrs->layout, "NHWC") << "Unexpected ROI Align layout"; - oshape = {rshape[0], roi_align_attrs->pooled_size[0], roi_align_attrs->pooled_size[1], - dshape[3]}; - } - - reporter->Assign(types[2], TensorType(oshape, data->dtype)); - return true; -} - -template -InferCorrectLayoutOutput ROIAlignInferCorrectLayout(const Attrs& attrs, - const Array& new_in_layouts, - const Array& old_in_layouts, - const Array& old_in_types) { - const T* params = attrs.as(); - Layout data_layout = params->layout; - - // Layout inference needs to define the layout for all inputs and output data layouts. - // For roi_align, the second inputs is 2-D tensor with shape [num_roi, 5]. - // So, we set the layout as "N5". - return InferCorrectLayoutOutput({data_layout, Layout("N5")}, {data_layout}, attrs); -} - -Expr MakeROIAlign(Expr data, Expr rois, Array pooled_size, double spatial_scale, - int sample_ratio, String layout, String mode) { - auto attrs = make_object(); - attrs->pooled_size = pooled_size; - attrs->spatial_scale = spatial_scale; - attrs->sample_ratio = sample_ratio; - attrs->layout = layout; - attrs->mode = mode; - static const Op& op = Op::Get("vision.roi_align"); - return Call(op, {data, rois}, Attrs(attrs), {}); -} - -TVM_REGISTER_GLOBAL("relay.op.vision._make.roi_align").set_body_typed(MakeROIAlign); - -RELAY_REGISTER_OP("vision.roi_align") - .describe(R"doc(ROI Align operator. - - - **data**: This depends on the `layout` parameter. Input is 4D array of shape - (batch_size, channels, height, width) if `layout` is `NCHW`. - - **rois**: 2D array of shape (num_roi, 5). The last dimension should be in format of - [batch_index, w_start, h_start, w_end, h_end]. - - **out**: This depends on the `layout` parameter. Output is 4D array of shape - (num_roi, channels, pooled_height, pooled_width) if `layout` is `NCHW`. - )doc" TVM_ADD_FILELINE) - .set_num_inputs(2) - .add_argument("data", "Tensor", "The input tensor.") - .add_argument("rois", "Tensor", "The input rois") - .set_support_level(5) - .add_type_rel("ROIAlign", ROIAlignRel) - .set_attr("FInferCorrectLayout", - ROIAlignInferCorrectLayout); - -TVM_REGISTER_NODE_TYPE(ROIPoolAttrs); - -bool ROIPoolRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - auto roi_pool_attrs = attrs.as(); - ICHECK_EQ(types.size(), 3); - const auto* data = types[0].as(); - const auto* rois = types[1].as(); - const auto& dshape = data->shape; - const auto& rshape = rois->shape; - ICHECK(roi_pool_attrs); - ICHECK_EQ(dshape.size(), 4) << "Input data should be 4-D."; - ICHECK_EQ(rshape.size(), 2) << "Input rois should be 2-D."; - // assign output type - std::vector oshape; - if (roi_pool_attrs->layout == "NCHW") { - oshape = {rshape[0], dshape[1], roi_pool_attrs->pooled_size[0], roi_pool_attrs->pooled_size[1]}; - } else if (roi_pool_attrs->layout == "NHWC") { - oshape = {rshape[0], roi_pool_attrs->pooled_size[0], roi_pool_attrs->pooled_size[1], dshape[3]}; - } else { - LOG(FATAL) << "vision.roi_pool does not support " << roi_pool_attrs->layout << " layout"; - } - - reporter->Assign(types[2], TensorType(oshape, data->dtype)); - return true; -} - -template -InferCorrectLayoutOutput ROIPoolInferCorrectLayout(const Attrs& attrs, - const Array& new_in_layouts, - const Array& old_in_layouts, - const Array& old_in_types) { - const T* params = attrs.as(); - Layout data_layout = params->layout; - - // Layout inference needs to define the layout for all inputs and output data layouts. - // For roi_pool, the second inputs is 2-D tensor with shape [num_roi, 5]. - // So, we set the layout as "N5". - return InferCorrectLayoutOutput({data_layout, Layout("N5")}, {data_layout}, attrs); -} - -Expr MakeROIPool(Expr data, Expr rois, Array pooled_size, double spatial_scale, - String layout) { - auto attrs = make_object(); - attrs->pooled_size = pooled_size; - attrs->spatial_scale = spatial_scale; - attrs->layout = layout; - static const Op& op = Op::Get("vision.roi_pool"); - return Call(op, {data, rois}, Attrs(attrs), {}); -} - -TVM_REGISTER_GLOBAL("relay.op.vision._make.roi_pool").set_body_typed(MakeROIPool); - -RELAY_REGISTER_OP("vision.roi_pool") - .describe(R"doc(ROI Pool operator. - - - **data**: This depends on the `layout` parameter. Input is 4D array of shape - (batch_size, channels, height, width) if `layout` is `NCHW`. - - **rois**: 2D array of shape (num_roi, 5). The last dimension should be in format of - [batch_index, w_start, h_start, w_end, h_end]. - - **out**: This depends on the `layout` parameter. Output is 4D array of shape - (num_roi, channels, pooled_height, pooled_width) if `layout` is `NCHW`. - )doc" TVM_ADD_FILELINE) - .set_num_inputs(2) - .add_argument("data", "Tensor", "The input tensor.") - .add_argument("rois", "Tensor", "The input rois") - .set_support_level(5) - .add_type_rel("ROIPool", ROIPoolRel) - .set_attr("FInferCorrectLayout", ROIPoolInferCorrectLayout); - -TVM_REGISTER_NODE_TYPE(ProposalAttrs); - -bool ProposalRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - auto proposal_attrs = attrs.as(); - ICHECK_EQ(types.size(), 4); - const auto* cls_prob = types[0].as(); - const auto* bbox_pred = types[1].as(); - const auto* im_info = types[2].as(); - - if (!cls_prob || !bbox_pred || !im_info) { - return false; - } - - ICHECK_EQ(cls_prob->shape.size(), 4U) - << "The dimension of class probability should be 4, but received " << cls_prob->shape.size(); - ICHECK_EQ(bbox_pred->shape.size(), 4U) - << "The dimension of box prediction should be 4, but received " << bbox_pred->shape.size(); - ICHECK_EQ(im_info->shape.size(), 2U) - << "The dimension of image info should be 2, but received " << im_info->shape.size(); - ICHECK(reporter->AssertEQ(im_info->shape[1], 3)); - - auto batch = cls_prob->shape[0]; - - std::vector oshape({batch * proposal_attrs->rpn_post_nms_top_n, 5}); - reporter->Assign(types[3], TensorType(oshape, cls_prob->dtype)); - return true; -} - -Expr MakeProposal(Expr cls_prob, Expr bbox_pred, Expr im_info, Array scales, - Array ratios, int feature_stride, double threshold, - int rpn_pre_nms_top_n, int rpn_post_nms_top_n, int rpn_min_size, bool iou_loss) { - auto attrs = make_object(); - attrs->scales = scales; - attrs->ratios = ratios; - attrs->feature_stride = feature_stride; - attrs->threshold = threshold; - attrs->rpn_pre_nms_top_n = rpn_pre_nms_top_n; - attrs->rpn_post_nms_top_n = rpn_post_nms_top_n; - attrs->rpn_min_size = rpn_min_size; - attrs->iou_loss = iou_loss; - static const Op& op = Op::Get("vision.proposal"); - return Call(op, {cls_prob, bbox_pred, im_info}, Attrs(attrs), {}); -} - -TVM_REGISTER_GLOBAL("relay.op.vision._make.proposal").set_body_typed(MakeProposal); - -RELAY_REGISTER_OP("vision.proposal") - .describe(R"code(Generate region proposals via RPN. - - - **cls_prob**: 4-D with shape [batch, 2 * num_anchors, height, width]. - - **bbox_pred**: 4-D with shape [batch, 4 * num_anchors, height, width]. - - **im_info**: 2-D with shape [batch, 3]. - - **out**: 2-D with shape [batch * rpn_post_nms_top_n, 5]. - )code" TVM_ADD_FILELINE) - .set_num_inputs(3) - .add_argument("cls_prob", "Tensor", "Score of how likely proposal is object") - .add_argument("bbox_pred", "Tensor", "BBox predicted deltas from anchors for proposals") - .add_argument("im_info", "Tensor", "Image size and scale") - .set_support_level(5) - .add_type_rel("Proposal", ProposalRel); - -} // namespace relay -} // namespace tvm diff --git a/src/relay/op/vision/yolo.cc b/src/relay/op/vision/yolo.cc deleted file mode 100644 index 8979f939c32e..000000000000 --- a/src/relay/op/vision/yolo.cc +++ /dev/null @@ -1,88 +0,0 @@ -/* - * 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 yolo.cc - * \brief Yolo related operators - */ -#include -#include -#include - -#include - -#include "../op_common.h" -#include "../type_relations.h" - -namespace tvm { -namespace relay { - -TVM_REGISTER_NODE_TYPE(YoloReorgAttrs); - -/*! - * \brief YoloReorgRel Output type and shape relation evaluation function. - * \param num_inputs Number of input types in the args. - * \param attrs The additional attributes of the operator. - * \param reporter The reporter to report solution to. - * \return false if This relation cannot be resolved. true if this relation has been resolved. - */ -bool YoloReorgRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - ICHECK_EQ(types.size(), 2); - const auto* data = types[0].as(); - if (data == nullptr) return false; - - const YoloReorgAttrs* param = attrs.as(); - ICHECK(param != nullptr); - - ICHECK(data->shape.size() == 4) << "Yolo reorg supports only 4 dimension."; - std::vector oshape(data->shape.begin(), data->shape.end()); - oshape[1] = oshape[1] * param->stride * param->stride; - oshape[2] = indexdiv(oshape[2], param->stride); - oshape[3] = indexdiv(oshape[3], param->stride); - reporter->Assign(types[1], TensorType(oshape, data->dtype)); - return true; -} - -Expr MakeYoloReorg(Expr data, Integer stride) { - auto attrs = make_object(); - attrs->stride = stride; - static const Op& op = Op::Get("vision.yolo_reorg"); - return Call(op, {data}, Attrs(attrs), {}); -} - -TVM_REGISTER_GLOBAL("relay.op.vision._make.yolo_reorg").set_body_typed(MakeYoloReorg); - -RELAY_REGISTER_OP("vision.yolo_reorg") - .describe(R"doc("Yolo reorg operation. This layer reorganize the output. -Its function is mostly shape transform.")doc" TVM_ADD_FILELINE) - .add_argument("data", "Tensor", "The input tensor.") - .set_num_inputs(1) - .set_support_level(5) - .set_attrs_type() - .add_type_rel("YoloReorg", YoloReorgRel) - .set_attr("FTVMCompute", [](const Attrs& attrs, const Array& inputs, - const Type& out_type) { - const auto* params = attrs.as(); - ICHECK(params != nullptr); - return Array{topi::vision::reorg(inputs[0], params->stride.IntValue())}; - }); - -} // namespace relay -} // namespace tvm diff --git a/src/relay/op/vm/vm.cc b/src/relay/op/vm/vm.cc deleted file mode 100644 index cd54e345062b..000000000000 --- a/src/relay/op/vm/vm.cc +++ /dev/null @@ -1,163 +0,0 @@ -/* - * 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 src/relay/op/vm/vm.cc - * \brief Dialect operators for Relay VM. - */ - -#include "vm.h" - -#include -#include -#include -#include -#include -#include -#include - -#include - -#include "../../transforms/infer_layout_utils.h" -#include "../op_common.h" -#include "../type_relations.h" - -namespace tvm { -namespace relay { - -// shape_of -// register ShapeOfAttrs here to make sure it has been registered when vm.shape_of uses it -TVM_REGISTER_NODE_TYPE(ShapeOfAttrs); - -// vm.shape_func -TVM_REGISTER_NODE_TYPE(ShapeFuncAttrs); - -RELAY_REGISTER_OP("vm.shape_of") - .describe(R"code(Get the shape of an input tensor. -)code" TVM_ADD_FILELINE) - .set_num_inputs(1) - .add_argument("tensor", "Tensor", "The input tensor") - .add_type_rel("ShapeOf", ShapeOfRel) - .set_attrs_type_key("relay.attrs.ShapeOfAttrs") - .set_support_level(10) - .set_attr("TOpPattern", kOpaque) - .set_attr("TOpIsStateful", false) - .set_attr("TNonComputational", true) - .set_attr("FInferCorrectLayout", ElemwiseArbitraryLayout); - -Expr ShapeOf(Expr expr) { - auto attrs = make_object(); - attrs->dtype = DataType::Int(64); - static const Op& op = Op::Get("vm.shape_of"); - return Call(op, {std::move(expr)}, Attrs(std::move(attrs)), {}); -} - -TVM_REGISTER_GLOBAL("relay.op.vm.shape_of").set_body_typed(ShapeOf); - -// vm.invoke_tvm_op -bool InvokeTVMOpRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - ICHECK_EQ(types.size(), 4u); - auto func_type = types[0].as(); - ICHECK(func_type != nullptr) << "input must be operator with known type"; - auto input_type = types[1].as(); - auto output_type = types[2].as(); - ICHECK(input_type != nullptr) - << "internal invariant violated: invoke_tvm_op inputs must be a tuple"; - ICHECK(output_type != nullptr) - << "internal invariant violated: invoke_tvm_op outputs must be a tuple"; - Type ex_output; - if (func_type->ret_type.as()) { - ex_output = TupleType({func_type->ret_type}); - } else { - ICHECK(func_type->ret_type.as()) - << "expecting function result to be tuple type. Types:" << std::endl - << PrettyPrint(types); - ex_output = func_type->ret_type; - } - auto ex_input = TupleType(func_type->arg_types); - reporter->Assign(ex_input, GetRef(input_type)); - reporter->Assign(ex_output, GetRef(output_type)); - reporter->Assign(types[3], TupleType::Empty()); - return true; -} - -Expr InvokeTVMOp(Expr func, Expr inputs, Expr outputs, DictAttrs attrs) { - static const Op& op = Op::Get("vm.invoke_tvm_op"); - return Call(op, {std::move(func), std::move(inputs), std::move(outputs)}, std::move(attrs)); -} - -TVM_REGISTER_GLOBAL("relay.op.vm.invoke_tvm_op") - .set_body_typed([](Expr func, Expr inputs, Expr outputs, DictAttrs attrs) { - return InvokeTVMOp(std::move(func), std::move(inputs), std::move(outputs), std::move(attrs)); - }); - -RELAY_REGISTER_OP("vm.invoke_tvm_op") - .describe(R"code(Invoke an operation compiled by TVM.)code" TVM_ADD_FILELINE) - .set_num_inputs(3) - .add_argument("op", "Function", "The operation to call") - .add_argument("ins", "Tuple", "The input tensors.") - .add_argument("outs", "Tuple", "The output tensors.") - .add_type_rel("InvokeTVMOp", InvokeTVMOpRel) - .set_attrs_type_key("DictAttrs") - .set_support_level(10) - .set_attr("TOpPattern", kOpaque) - .set_attr("TOpIsStateful", true) - .set_attr("TNonComputational", true) - .set_attr("FInferCorrectLayout", ElemwiseArbitraryLayout); - -// vm.reshape -TVM_REGISTER_NODE_TYPE(ReshapeTensorAttrs); - -bool ReshapeTensorRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - ICHECK_EQ(types.size(), 3u); - auto reshape_attrs = attrs.as(); - ICHECK(reshape_attrs); - auto tt = types[0].as(); - ICHECK(tt) << "input must be tensor type"; - reporter->Assign(types[2], TensorType(reshape_attrs->newshape, tt->dtype)); - return true; -} - -RELAY_REGISTER_OP("vm.reshape_tensor") - .describe(R"code(Use VM reshape_tensor instruction to reshape the tensor. -)code" TVM_ADD_FILELINE) - .set_num_inputs(2) - .add_argument("data", "Tensor", "The input tensor") - .add_argument("shape", "Tensor", "The output shape tensor") - .add_type_rel("ReshapeTensor", ReshapeTensorRel) - .set_attrs_type_key("relay.attrs.ReshapeTensorAttrs") - .set_support_level(10) - .set_attr("TOpPattern", kOpaque) - .set_attr("TOpIsStateful", false) - .set_attr("TNonComputational", true) - .set_attr("FInferCorrectLayout", ElemwiseArbitraryLayout); - -Expr ReshapeTensor(Expr data, Expr shape, Array newshape) { - static const Op& op = Op::Get("vm.reshape_tensor"); - auto attrs = make_object(); - attrs->newshape = std::move(newshape); - return Call(op, {std::move(data), std::move(shape)}, Attrs(std::move(attrs)), {}); -} - -TVM_REGISTER_GLOBAL("relay.op.vm.reshape_tensor").set_body_typed(ReshapeTensor); - -} // namespace relay -} // namespace tvm diff --git a/src/relay/op/vm/vm.h b/src/relay/op/vm/vm.h deleted file mode 100644 index 77cc41e1dd83..000000000000 --- a/src/relay/op/vm/vm.h +++ /dev/null @@ -1,39 +0,0 @@ -/* - * 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 src/relay/op/vm/vm.h - * \brief Dialect operators for Relay VM. - */ -#ifndef TVM_RELAY_OP_VM_VM_H_ -#define TVM_RELAY_OP_VM_VM_H_ - -#include - -namespace tvm { -namespace relay { - -Expr InvokeTVMOp(Expr func, Expr inputs, Expr outputs, DictAttrs attrs); -Expr ShapeOf(Expr expr); -Expr ReshapeTensor(Expr data, Expr shape, Array newshape); - -} // namespace relay -} // namespace tvm - -#endif // TVM_RELAY_OP_VM_VM_H_ diff --git a/src/relay/parser/meta_ref.cc b/src/relay/parser/meta_ref.cc deleted file mode 100644 index cdc6929622dd..000000000000 --- a/src/relay/parser/meta_ref.cc +++ /dev/null @@ -1,99 +0,0 @@ -/* - * 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 src/parser/meta_ref.cc - * \brief An operator which allows forward referencing a yet-to-be parsed meta table reference. - */ - -#include "./meta_ref.h" - -#include -#include -#include -#include - -namespace tvm { -namespace relay { - -using tvm::relay::transform::CreateFunctionPass; -using tvm::transform::PassContext; - -/* Set to arbitrary high number, since we should never schedule in normal pass manager flow. */ -static int kMetaExpandOptLevel = 1337; - -TVM_REGISTER_NODE_TYPE(MetaRefAttrs); - -bool MetaRefRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - LOG(FATAL) << "need to expand before type checking"; -} - -RELAY_REGISTER_OP("parser.MetaRef") - .describe(R"code(A reference into the meta table.)code" TVM_ADD_FILELINE) - .set_attrs_type() - .set_num_inputs(0) - .set_support_level(10) - .add_type_rel("MetaRef", MetaRefRel) - .set_attr("TOpIsStateful", false) - .set_attr("TNonComputational", true); - -Expr MetaRef(std::string type_key, uint64_t node_index) { - static const Op& op = Op::Get("parser.MetaRef"); - auto attrs = make_object(); - attrs->node_type_key = tvm::String(type_key); - attrs->node_index = node_index; - return Call(op, {}, Attrs(attrs), {}); -} - -struct MetaRefExpander : public ExprMutator { - MetaTable table; - - explicit MetaRefExpander(const MetaTable& table) : table(table) {} - - Expr VisitExpr_(const CallNode* call) final { - if (auto op_node = call->op.as()) { - if (op_node->name == "parser.MetaRef") { - auto meta_attrs = call->attrs.as(); - ICHECK(meta_attrs) << "an internal error has occurred"; - auto nodes = table.at(meta_attrs->node_type_key); - ICHECK_LT(meta_attrs->node_index, nodes.size()); - return Downcast(nodes[meta_attrs->node_index]); - } - } - - return ExprMutator::VisitExpr_(call); - } -}; - -Function ExpandMetaRefs(const MetaTable& meta_table, const relay::Function& func) { - MetaRefExpander expander(meta_table); - return Downcast(expander.VisitExpr(func)); -} - -IRModule ExpandMetaRefs(const MetaTable& meta_table, const IRModule& mod) { - auto pass = CreateFunctionPass([&](Function func, IRModule module, - PassContext ctx) { return ExpandMetaRefs(meta_table, func); }, - kMetaExpandOptLevel, "ExpandMetaRefs", {}); - - return pass(mod, PassContext::Create()); -} - -} // namespace relay -} // namespace tvm diff --git a/src/relay/parser/meta_ref.h b/src/relay/parser/meta_ref.h deleted file mode 100644 index bed67bea05a4..000000000000 --- a/src/relay/parser/meta_ref.h +++ /dev/null @@ -1,82 +0,0 @@ -/* - * 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 meta_ref.h - * \brief A reference into the metadata section of the Relay text format. - */ - -#ifndef TVM_RELAY_PARSER_META_REF_H_ -#define TVM_RELAY_PARSER_META_REF_H_ - -#include -#include -#include -#include - -#include - -namespace tvm { -namespace relay { - -/*! - * \brief Options for allocating storage. - */ -struct MetaRefAttrs : public tvm::AttrsNode { - tvm::String node_type_key; - uint64_t node_index; - - TVM_DECLARE_ATTRS(MetaRefAttrs, "relay.attrs.MetaRefAttrs") { - TVM_ATTR_FIELD(node_type_key) - .describe("The type_key representing the type of the node referenced."); - TVM_ATTR_FIELD(node_index).describe("The index into the type specific node array."); - } -}; - -/*! \brief A reference to a "meta-expression". - * - * In the text format we allow referencing metadata which - * uses a compact serialization that proceeds the main - * program body. - * - * We can reference this table using an expression of - * the form `meta[Type][index]`. - * - * We must later resolve these references to actual in-memory - * AST nodes but this requires first parsing the full program - * then expanding these temporary AST nodes into their corresponding - * nodes. - * - * For example the nth large constant will be pretty-printed as meta[relay.Constant][n] - * with its compact binary serialization residing in the metadata section at the end - * of the program. - * - * \param type_key The type key of the object in the meta section. - * \param node_index The index into that subfield. - * \returns The meta table reference. - */ -Expr MetaRef(std::string type_key, uint64_t node_index); - -relay::Function ExpandMetaRefs(const MetaTable& meta_table, const relay::Function& func); -IRModule ExpandMetaRefs(const MetaTable& meta_table, const IRModule& mod); - -} // namespace relay -} // namespace tvm - -#endif // TVM_RELAY_PARSER_META_REF_H_ diff --git a/src/relay/parser/op_table.h b/src/relay/parser/op_table.h deleted file mode 100644 index 6ff2c05476f4..000000000000 --- a/src/relay/parser/op_table.h +++ /dev/null @@ -1,95 +0,0 @@ -/* - * 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 op_table.h - * \brief A operator table for parsing. - * Provides symbolic token sequences to map to TVM operators, with a given associativity and arity. - */ - -#ifndef TVM_RELAY_PARSER_OP_TABLE_H_ -#define TVM_RELAY_PARSER_OP_TABLE_H_ - -#include -#include - -#include -#include -#include -#include - -#include "./tokenizer.h" - -namespace tvm { -namespace relay { - -struct Rule { - std::vector tokens; - int precedence; - int arity; - tvm::Op op; - bool left_assoc; - - Rule() : tokens(), precedence(0), arity(0), op(tvm::Op()), left_assoc(false) {} - - Rule(std::vector tokens, tvm::Op op, int precedence, int arity = 2, - bool left_assoc = false) - : tokens(tokens), precedence(precedence), arity(arity), op(op), left_assoc(left_assoc) {} - - Rule(const Rule& rule) { - this->tokens = rule.tokens; - this->op = rule.op; - this->precedence = rule.precedence; - this->arity = rule.arity; - this->left_assoc = rule.left_assoc; - } -}; - -struct OperatorTable { - std::vector rules; - std::unordered_map this_is_a_hack; - - explicit OperatorTable(std::vector rules) : rules(rules), this_is_a_hack() { - for (auto rule : rules) { - std::stringstream key; - for (auto token : rule.tokens) { - key << ToString(token); - } - this->this_is_a_hack.insert({key.str(), rule}); - } - } -}; - -inline OperatorTable DefaultOpTable() { - return OperatorTable( - {Rule({TokenType::kStar}, Op::Get("multiply"), 12, 2, true), - Rule({TokenType::kDivision}, Op::Get("divide"), 12, 2, true), - Rule({TokenType::kPlus}, Op::Get("add"), 10, 2, true), - Rule({TokenType::kMinus}, Op::Get("subtract"), 10, 2, true), - Rule({TokenType::kLAngle}, Op::Get("less"), 8, 2, true), - Rule({TokenType::kLAngle, TokenType::kEqual}, Op::Get("less_equal"), 8, 2, true), - Rule({TokenType::kRAngle}, Op::Get("greater"), 8, 2, true), - Rule({TokenType::kRAngle, TokenType::kEqual}, Op::Get("greater_equal"), 8, 2, true), - Rule({TokenType::kEqual, TokenType::kEqual}, Op::Get("equal"), 7, 2, true), - Rule({TokenType::kBang, TokenType::kEqual}, Op::Get("not_equal"), 7, 2, true)}); -} - -} // namespace relay -} // namespace tvm -#endif // TVM_RELAY_PARSER_OP_TABLE_H_ diff --git a/src/relay/parser/parser.cc b/src/relay/parser/parser.cc deleted file mode 100644 index 233455bf89ba..000000000000 --- a/src/relay/parser/parser.cc +++ /dev/null @@ -1,1989 +0,0 @@ -/* - * 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 parser.cc - * \brief A parser for TVM IR. - */ -#include -#include -#include -#include -#include -#include -#include -#include -#include -#include -#include -#include - -#include - -#include "../../support/scalars.h" -#include "./meta_ref.h" -#include "./op_table.h" -#include "./span_check.h" -#include "./tokenizer.h" - -namespace tvm { -namespace relay { - -/*! \brief The meta table maps from type key to a sequence of objects. */ -using MetaTable = Map>; - -using tvm::runtime::NDArray; -using tvm::runtime::String2DLDataType; -using tvm::transform::CreateModulePass; -using tvm::transform::PassContext; - -/*! \brief A helper for passing around spans with data structures with - * no span field. - */ -template -struct Spanned { - T data; - Span span; - - Spanned() = default; - Spanned(const Spanned& other) = default; - Spanned(T data, Span span) : data(data), span(span) {} -}; - -/*! \brief A wrapper structure for capturing the result of parsing - * a global definition *before* we add it to the IRModule. - * - * This enables the parser to parse everything in one pass before - * constructing the IRModule. - */ -struct GlobalFunc { - GlobalVar global; - Function function; - GlobalFunc() : global(), function() {} - GlobalFunc(GlobalVar global, Function function) : global(global), function(function) {} - GlobalFunc(const GlobalFunc& gfunc) { - this->global = gfunc.global; - this->function = gfunc.function; - } -}; - -/*! \brief A wrapper structure for capturing all top-level definitions - * when parsing a module. - */ -struct Definitions { - /*! \brief The set of global functions. */ - std::vector funcs; - /*! \brief The set of type definitions. */ - std::vector types; - // TODO(@jroesch): contain meta-table below -}; - -/*! \brief A structure representing the semantic versioning information - * for a Relay program. - */ -class SemVer { - public: - int major_version; - int minor_version; - int patch_version; - - SemVer() : major_version(0), minor_version(0), patch_version(0) {} - SemVer(int major_version, int minor_version, int patch_version) - : major_version(major_version), minor_version(minor_version), patch_version(patch_version) {} - SemVer(const SemVer& other) - : major_version(other.major_version), - minor_version(other.minor_version), - patch_version(other.patch_version) {} -}; - -/*! \brief A simple wrapper around a mapping from raw string names - * to a TVM variable, type variable or other binder type. - */ -template -struct Scope { - /*! \brief The internal map. */ - std::unordered_map name_map; -}; - -/*! \brief A stack of scopes. - * - * In order to properly handle scoping we must maintain a stack of scopes. - * - * A stack allows users to write programs which contain repeated variable - * names and to properly handle both nested scopes and removal of variables - * when they go out of scope. - * - * This is the classic approach to lexical scoping. - */ -template -class ScopeStack { - private: - std::vector> scope_stack; - std::unordered_map free_vars; - - public: - /*! \brief Adds a variable binding to the current scope. */ - void Add(const std::string& name, const T& value) { - if (!this->scope_stack.size()) { - LOG(FATAL) << "internal issue"; - } - this->scope_stack.back().name_map.insert({name, value}); - } - - void AddFreeVar(const std::string& name, const T& value) { free_vars.insert({name, value}); } - - /*! \brief Looks up a variable name in the scope stack returning the matching variable - * in most recent scope. */ - T Lookup(const std::string& name) { - for (auto scope = this->scope_stack.rbegin(); scope != this->scope_stack.rend(); ++scope) { - auto it = scope->name_map.find(name); - if (it != scope->name_map.end()) { - return it->second; - } - } - - // Check if we bound a free variable declaration. - auto it = free_vars.find(name); - if (it != free_vars.end()) { - return it->second; - } - - return T(); - } - - /*! \brief Adds a fresh scope. */ - void PushStack() { this->scope_stack.push_back(Scope()); } - - /*! \brief Removes the most recent scope. */ - void PopStack() { this->scope_stack.pop_back(); } -}; - -struct DuplicateKeyError : public Error { - explicit DuplicateKeyError(const std::string& msg) : Error(msg) {} -}; - -/*! \brief A table of interning strings as global function and type names. */ -template -struct InternTable { - /*! \brief The internal table mapping strings to a unique allocation. */ - std::unordered_map table; - DiagnosticContext* ctx; - - /*! \brief Add the unique allocation. */ - void Add(const std::string& name, const T& t) { - auto it = table.find(name); - if (it != table.end()) { - throw DuplicateKeyError("duplicate key name in intern table"); - } else { - table.insert({name, t}); - } - } - - /*! \brief Return the unique allocation. */ - Optional Get(const std::string& name) const { - auto it = table.find(name); - if (it != table.end()) { - return Optional(it->second); - } else { - return Optional(); - } - } -}; - -GlobalVar AddOrGet(InternTable* table, const std::string& name) { - auto var = table->Get(name); - if (var) { - return var.value(); - } else { - auto gvar = GlobalVar(name); - table->Add(name, gvar); - return gvar; - } -} - -GlobalTypeVar AddOrGet(InternTable* table, const std::string& name, - TypeKind kind = TypeKind::kType) { - auto var = table->Get(name); - if (var) { - auto tvar = var.value(); - TypeKind& tvar_kind = const_cast(tvar->kind); - tvar_kind = kind; - return tvar; - } else { - auto gvar = GlobalTypeVar(name, kind); - table->Add(name, gvar); - return gvar; - } -} - -/*! \brief The parser class is the main interface to the parser. - * the parser is not currently exposed beyond this .cc file. - * - * The parser is initialized with a diagnostic context, an - * operator table, and a token stream. - * - * The rest of the internal state is used to map the human readable - * form to in-memory IR representation. - * - * The main entry point to the parser are a set of parsing methods - * such as `ParseModule` and `ParseExpr`. - * - * As with traditional recursive descent parsers the parsing methods - * are factored recursively just as one would do with a formal language - * grammar. - * - * You can view a recursive descent parser as a human friendly way to specify - * a state machine, and thus this factoring is necessary as the 'state' of this - * machine is the combination of the current parsing method and the next token. - * - * Parsing proceeds by matching a token and then dispatching to the appropriate - * method to parse the next tokens in the stream. - * - * For example if we are parsing a type and encounter a "Tensor" token we switch - * into a mode for parsing `[`, a shape, a comma, a data type and then a `]`. - * - * Certain matches like this are unambiguous and proceed in a straight line fashion - * once the initial token is found. Other parsing is more complex and requires some - * tricks to correctly parse. - * - * For example when we find a '(' in an expression context, it may be part of - * a tuple, the arguments to a call, or a parenthesized expression. The below code - * disambiguate these cases by factoring expression parsing into a series of methods - * which encode the parsing context and thus how to interpret the parenthesis. - * - * For more information one should be able to read the code in order starting with - * `ParseModule` or `ParseExpr`. - */ -class Parser { - public: - /*! \brief The version that the parser is parsing. */ - SemVer version; - - /*! \brief The IRModule we are building. */ - IRModule module; - - /*! \brief The diagnostic context used for error reporting. */ - DiagnosticContext diag_ctx; - - const Source& source; - - /*! \brief The current position in the token stream. */ - int pos; - - /*! \brief The token stream for the parser. */ - std::vector tokens; - - /*! \brief The configured operator table. */ - OperatorTable op_table; - - /*! \brief Configure the whitespace mode, right now we ignore all whitespace. */ - bool ignore_whitespace; - - /*! \brief A global mapping for GlobalVar. */ - InternTable global_names; - - /*! \brief A global mapping for type definitions. */ - InternTable type_names; - - /*! \brief A global mapping for constructor names. */ - InternTable ctors; - - /*! \brief A mapping from graph variable to expression, i.e., `%0 = expr`. */ - std::unordered_map graph_ctx; - - /*! \brief The set of type scopes used for generics. */ - ScopeStack type_scopes; - - /*! \brief The set of expression scopes used for lexical scope. */ - ScopeStack expr_scopes; - - /*! \brief The metadata section. */ - MetaTable meta_table; - - Parser(IRModule module, DiagnosticContext ctx, const Source& source, std::vector tokens, - OperatorTable op_table, MetaTable table) - : module(module), - diag_ctx(ctx), - source(source), - pos(0), - tokens(tokens), - op_table(op_table), - ignore_whitespace(true), - meta_table(table) { - InitializeGlobals(); - InitializeTypeDefs(); - } - - /*! If we are parsing into a module with previously loaded data types we need to - * map constructor names and variable names in the global tables. - */ - void InitializeTypeDefs() { - for (auto pair : this->module->type_definitions) { - type_names.Add(pair.first->name_hint, pair.first); - for (auto ctor : pair.second->constructors) { - ctors.Add(ctor->name_hint, ctor); - } - } - } - - void InitializeGlobals() { - for (auto pair : this->module->functions) { - global_names.Add(pair.first->name_hint, pair.first); - } - } - - /*! \brief Examine the next token in the stream, the current parser is configured to be - * whitespace insensitive so we will skip all whitespace or comment tokens. */ - Token Peek() { - // For now we ignore all whitespace tokens and comments. - // We can tweak this behavior later to enable white space sensitivity in the parser. - while (pos < static_cast(tokens.size()) && ignore_whitespace && - (tokens.at(pos)->token_type == TokenType::kWhitespace || - tokens.at(pos)->token_type == TokenType::kNewline || - tokens.at(pos)->token_type == TokenType::kLineComment || - tokens.at(pos)->token_type == TokenType::kComment)) { - pos++; - } - - if (pos < static_cast(tokens.size())) { - return Token(this->tokens.at(pos)); - } else { - return Token::Null(); - } - } - - /*! \brief Lookahead by N tokens. - * \param n The number of tokens to lookahead. - * \return The Nth token. - */ - Token Lookahead(int n) { - ICHECK_GE(n, 1) << "lookahead is only valid when n >= 1"; - - // We intend to skip n - 1 tokens, then return the nth. - auto old_pos = pos; - for (int i = 0; i < n - 1; i++) { - Peek(); - pos++; - } - - auto tok = Peek(); - pos = old_pos; - return tok; - } - - /*! \brief Consume a token, this method is the lowest level way to consume a token - * and will not ignore white space or look ahead in anyway. - * - * /param token_type The token type to match. - */ - void Consume(const TokenType& token_type) { - if (tokens[pos]->token_type != token_type) { - this->diag_ctx.EmitFatal(Diagnostic::Error(tokens[pos]->span) - << "expected a " << Pretty(token_type) << " found " - << Pretty(Peek()->token_type)); - } - pos++; - } - - /*! Match a token in the stream, this will first invoke Peek, ignoring tokens such - * as whitespace or comments returning the first meaningful token. - * - * We then try and consume the requested token, this will trigger an error if the - * current token does not match the token_type. - */ - Token Match(const TokenType& token_type) { - auto tok = Peek(); - Consume(token_type); - return tok; - } - - /*! Conditionally consume a token when it matches, this will never trigger an error - * as we guard against consuming the token before we do. - * - * Useful for matching optional tokens, effectively looksahead by one. - */ - bool WhenMatch(const TokenType& token_type) { - VLOG(9) << "Parser::WhenMatch: Peek() == " << Peek(); - if (Peek()->token_type == token_type) { - Consume(token_type); - return true; - } else { - return false; - } - } - - /* \brief Add a graph binding to the parsing context - * - * For example if we parse %0 = add(...), map 0 -> add(...), etc. - */ - void AddGraphBinding(const Token& token, const Expr& expr) { - auto graph_no = token.ToNumber(); - this->graph_ctx.insert({graph_no, expr}); - } - - /* \brief Lookup a previously bound graph variable. - * - * Note: we take tokens in all lookup methods so that we - * that we can do error reporting based on token location. - */ - Expr LookupGraphBinding(const Token& token) { - auto graph_no = token.ToNumber(); - auto it = this->graph_ctx.find(graph_no); - if (it != this->graph_ctx.end()) { - return it->second; - } else { - LOG(FATAL) << "Local variable %" << graph_no << " has not yet been defined"; - throw; - } - } - - /*! \brief Bind a local variable in the expression scope. - * - * "x" -> Var("x"), these are needed to map from the raw string names - * to unique variable nodes. - * If a virtual device is specified, sets the virtual device of the variable. - */ - Var BindVar(const std::string& name, const relay::Type& type_annotation, - Optional virtual_device = Optional()) { - auto var = Var(name, type_annotation); - var->virtual_device_ = virtual_device.value_or(VirtualDevice::FullyUnconstrained()); - VLOG(1) << "Binding var named " << name << " to variable node " << PrettyPrint(var); - this->expr_scopes.Add(name, var); - return var; - } - - /*! \brief Bind a local variable in the expression scope. - * - * "x" -> Var("x"), these are needed to map from the raw string names - * to unique variable nodes. - */ - Var BindFreeVar(const std::string& name, const relay::Type& type_annotation) { - auto var = Var(name, type_annotation); - this->expr_scopes.AddFreeVar(name, var); - return var; - } - - /*! \brief Bind a type variable in the type scope. - * - * "A" -> TypeVar("A", ...), these are needed to map from raw string names - * to unique type variable nodes. - */ - TypeVar BindTypeVar(const std::string& name, const TypeKind type_kind) { - auto type_var = TypeVar(name, type_kind); - this->type_scopes.Add(name, type_var); - return type_var; - } - - /*! \brief Lookup a variable in the expression scope. - * - * Note: all lookup methods take tokens intentionally for error reporting information. - */ - Var LookupLocal(const Token& local) { - auto var = this->expr_scopes.Lookup(local.ToString()); - if (!var.defined()) { - diag_ctx.Emit(Diagnostic::Error(local->span) - << "this local variable has not been previously declared"); - } - return var; - } - - /*! \brief Lookup a variable in the type scope. - * - * Note: all lookup methods take tokens intentionally for error reporting information. - */ - TypeVar LookupTypeVar(const Token& ident) { - auto var = this->type_scopes.Lookup(ident.ToString()); - return var; - } - - /*! \brief Add an expression scope to the scope stack. */ - void PushScope() { this->expr_scopes.PushStack(); } - - /*! \brief Remove N expression scopes from the scope stack. */ - void PopScopes(int n) { - for (int i = 0; i < n; i++) { - this->expr_scopes.PopStack(); - } - } - - /*! \brief Add an type scope to the scope stack. */ - void PushTypeScope() { this->type_scopes.PushStack(); } - - /*! \brief Remove N type scopes from the scope stack. */ - void PopTypeScopes(int n) { - for (int i = 0; i < n; i++) { - this->type_scopes.PopStack(); - } - } - - /*! \brief Convert a numeric token to an NDArray for embedding into the Relay program. */ - NDArray NumberToNDArray(const Token& token) { - if (token->token_type == TokenType::kInteger) { - return support::IntImmToNDArray(Downcast(token->data)); - } else if (token->token_type == TokenType::kFloat) { - return support::FloatImmToNDArray(Downcast(token->data)); - } else { - LOG(FATAL) << "internal error: should only call this function on numeric tokens"; - } - } - - [[noreturn]] void ParseError(const Token& token, const std::string& msg) { - throw std::runtime_error(msg); - } - - /*! \brief A parsing helper for a bracketed expression . */ - template - R Bracket(TokenType open, TokenType close, std::function parser) { - Match(open); - R result = parser(); - Match(close); - return result; - } - - /*! \brief Parse `(` parser() `)`. */ - template - R Parens(std::function parser) { - return Bracket(TokenType::kOpenParen, TokenType::kCloseParen, parser); - } - - /*! \brief Parse `{` parser() `}`. */ - template - R Block(std::function parser) { - return Bracket(TokenType::kLCurly, TokenType::kRCurly, parser); - } - - template - R WithSpan(std::function parser) { - auto start_span = Peek()->span; - VLOG(9) << "WithSpan: start_span = " << start_span; - R ast = parser(); - if (ast.defined()) { - // The token at the head of the stream is now 1 past where we parsed. So we find its start - // position as its start and end, so that when we merge we only grow the spanned region - // to the start of the current stream. - auto span_pos = pos - 1; - while ((tokens.at(span_pos)->token_type == TokenType::kWhitespace || - tokens.at(span_pos)->token_type == TokenType::kNewline || - tokens.at(span_pos)->token_type == TokenType::kLineComment || - tokens.at(span_pos)->token_type == TokenType::kComment)) { - span_pos--; - } - auto end_token = tokens.at(span_pos); - VLOG(9) << "WithSpan: end_span = " << end_token->span; - ast->span = start_span.Merge(end_token->span); - } - return ast; - } - - struct MetaRef { - std::string type_key; - uint64_t node_index; - Span span; - MetaRef(std::string type_key, uint64_t node_index, Span span) - : type_key(type_key), node_index(node_index), span(span) {} - }; - - MetaRef MetaRefFromToken(const Token& tok) { - Call ref = Downcast(tok->data); - auto attrs = ref->attrs.as(); - auto type_key = attrs->node_type_key; - auto index = attrs->node_index; - return MetaRef(type_key, index, ref->span); - } - - /*! \brief Parse a meta reference of the form `meta[type_key][node_index]`. - * For example `meta[relay.Constant][0]` references the first constant, `meta[relay.Constant][1]` - * the second, and so on. - */ - ObjectRef ParseMetaRef() { - auto meta_ref_tok = Match(TokenType::kMetaReference); - auto meta_ref = MetaRefFromToken(meta_ref_tok); - auto it = this->meta_table.find(meta_ref.type_key); - if (it != this->meta_table.end()) { - auto nodes = (*it).second; - if (meta_ref.node_index < nodes.size()) { - return nodes[meta_ref.node_index]; - } else { - this->diag_ctx.Emit(Diagnostic::Error(meta_ref.span) - << "the node index `" << meta_ref.node_index - << "` is out of bounds for `" << meta_ref.type_key << "`"); - return ObjectRef(); - } - } else { - this->diag_ctx.Emit(Diagnostic::Error(meta_ref.span) - << "no entry in the meta table for `" << meta_ref.type_key << "`"); - return ObjectRef(); - } - } - /*! \brief Parses a sequence beginning with a start token, separated by a seperator token, and - * ending with a stop token. - * - * The simple form being ( )* . - * - * This also provides a fourth argument which is allowed to run when the sequence which matches - * the inner sequence can not proceed. - * - * This is useful for parsing things like attributes which don't match the standard expression - * parsers but are contained within the stop token. - */ - template - Array ParseSequence(TokenType start, TokenType sep, TokenType stop, std::function parse, - std::function before_stop = nullptr) { - VLOG(9) << "Parser::ParseSequence: start=" << ToString(start) << " sep=" << ToString(sep) - << " stop=" << ToString(stop); - Match(start); - - // This is for the empty arguments list case, if we have token stream - // we must parse leftovers, then match a stop token. - if (before_stop) { - auto did_parse = before_stop(); - if (did_parse) { - Match(stop); - return {}; - } - } - - // This is the case in which we find an empty arguments lists and no leftovers. - if (WhenMatch(stop)) { - return Array(); - } else { - VLOG(9) << "Parser::ParseSequence: parse first"; - auto data = parse(); - Array elements = {data}; - - if (WhenMatch(stop)) { - return elements; - // parse '( expr ',' * ')' - } else if (WhenMatch(sep)) { - while (true) { - VLOG(9) << "Parser::ParseSequence: parse element"; - if (WhenMatch(stop)) { - break; - } else { - // If before stop is - if (before_stop) { - auto did_parse = before_stop(); - if (did_parse) { - Match(stop); - return elements; - } - } - auto data = parse(); - WhenMatch(sep); - elements.push_back(data); - } - } - return elements; - } else { - auto next = Peek(); - this->diag_ctx.EmitFatal(Diagnostic::Error(next->span) - << "expected a " << Pretty(stop) << " found " - << Pretty(next->token_type)); - return Array(nullptr); - } - } - } - - /*! \brief Parse a full IRModule. */ - IRModule ParseModule() { - // Parse the semver header at the top of the module. - this->version = ParseSemVer(); - // Parse the definitions. - auto defs = ParseDefinitions(); - // Parse the metadata section at the end. - auto metadata = ParseMetadata(); - - Match(TokenType::kEndOfFile); - - for (auto type_def : defs.types) { - module->AddTypeDef(type_def->header, type_def); - } - - for (auto func : defs.funcs) { - module->Add(func.global, func.function, true); - } - - return module; - } - - /*! \brief Parse the semantic versioning header. */ - SemVer ParseSemVer(bool required = true) { - if (Peek()->token_type == TokenType::kVersion) { - auto version = Match(TokenType::kVersion); - // TODO(@jroesch): we currently only support 0.0.5. - if (version.ToString() != "\"0.0.5\"") { - this->diag_ctx.Emit(Diagnostic::Error(version->span) - << "invalid semantic version `" << version.ToString() << "`"); - } - } else if (required) { - this->diag_ctx.Emit(Diagnostic::Error(Peek()->span) - << "expected text format semantic version, found a " - << PrettyPrint(Peek())); - - this->diag_ctx.Emit(Diagnostic::Help(Peek()->span) - << "you can annotate it as #[version = \"0.0.5\"]"); - } - return SemVer(0, 0, 5); - } - - /*! \brief Parse zero or more Relay definitions. */ - Definitions ParseDefinitions() { - Definitions defs; - - while (true) { - auto next = Peek(); - switch (next->token_type) { - case TokenType::kDefn: { - Consume(TokenType::kDefn); - auto global_tok = Match(TokenType::kGlobal); - auto global_name = global_tok.ToString(); - auto global = AddOrGet(&global_names, global_name); - auto func = WithSpan([&]() { return ParseFunctionDef(); }); - ICHECK(func->span.defined()) << "spans must be set in parser"; - defs.funcs.push_back(GlobalFunc(global, func)); - continue; - } - case TokenType::kTypeDef: { - defs.types.push_back(ParseTypeDef()); - continue; - } - case TokenType::kExtern: { - Consume(TokenType::kExtern); - auto type_def = ParseTypeDef(); - if (type_def->constructors.size()) { - diag_ctx.Emit(Diagnostic::Error(next->span) - << "an external type may not have any constructors"); - } - defs.types.push_back(type_def); - } - default: - return defs; - } - } - } - - /*! \brief Parse zero or more Relay type definitions. */ - TypeData ParseTypeDef() { - // Match the `type` keyword. - Match(TokenType::kTypeDef); - // Parse the type's identifier. - auto type_tok = Match(TokenType::kIdentifier); - auto type_id = type_tok.ToString(); - auto type_global = AddOrGet(&type_names, type_id, TypeKind::kAdtHandle); - - Array generics; - - bool should_pop = false; - if (Peek()->token_type == TokenType::kLSquare) { - // If we have generics we need to add a type scope. - PushTypeScope(); - should_pop = true; - generics = ParseSequence( - TokenType::kLSquare, TokenType::kComma, TokenType::kRSquare, [&]() { - auto type_var_name = Match(TokenType::kIdentifier).ToString(); - return BindTypeVar(type_var_name, TypeKind::kType); - }); - } - - Array ctors; - if (Peek()->token_type == TokenType::kLCurly) { - // Parse the list of constructors. - ctors = ParseSequence( - TokenType::kLCurly, TokenType::kComma, TokenType::kRCurly, [&]() { - // First match the name of the constructor. - auto ctor_tok = Match(TokenType::kIdentifier); - auto ctor_name = ctor_tok.ToString(); - - Constructor ctor; - // Match the optional field list. - if (Peek()->token_type != TokenType::kOpenParen) { - ctor = tvm::Constructor(ctor_name, {}, type_global); - } else { - auto arg_types = - ParseSequence(TokenType::kOpenParen, TokenType::kComma, - TokenType::kCloseParen, [&]() { return ParseType(); }); - ctor = tvm::Constructor(ctor_name, arg_types, type_global); - } - - ICHECK(ctor.defined()); - - try { - this->ctors.Add(ctor_name, ctor); - } catch (const DuplicateKeyError& e) { - this->diag_ctx.EmitFatal(Diagnostic::Error(ctor_tok->span) - << "a constructor with the name " - << "`" << ctor_name << "` " - << "was previously defined"); - } - - return ctor; - }); - } - - // Now pop the type scope. - if (should_pop) { - PopTypeScopes(1); - } - - return TypeData(type_global, generics, ctors); - } - - std::string HackTokensAsString(int n) { - std::stringstream key; - n = std::min(static_cast(tokens.size() - pos), n); - for (int i = 0; i < n; i++) { - key << ToString(tokens.at(pos + i)->token_type); - } - return key.str(); - } - - std::vector ParseOp() { - std::vector matched; - Peek(); - for (int i = 4; i > 0; i--) { - auto key = HackTokensAsString(i); - auto it = this->op_table.this_is_a_hack.find(key); - if (it != this->op_table.this_is_a_hack.end()) { - pos = pos + i; - matched.push_back(it->second); - } - } - - return matched; - } - - /*! \brief Parse a single Relay expression. */ - Expr ParseExpr() { - VLOG(9) << "Parser::ParseExpr"; - return WithSpan([this] { - std::vector exprs; - - while (true) { - VLOG(9) << "Parser::ParseExpr: parsing a single expression"; - auto next = Peek(); - switch (next->token_type) { - // For graph or let, match first rhs, then invoke ParseBindingExpr - // ParseBindingExpression then parse_lhs() parse_rhs() ';' continue - case TokenType::kLCurly: { - // NB: Might need to optimize to remove deep recursion. - // Stack should only grow proportionally to the number of - // nested scopes. - // Parses `{` expression `}`. - auto block = WithSpan([&]() { - return Bracket(TokenType::kLCurly, TokenType::kRCurly, [&]() { - PushScope(); - auto expr = ParseExpr(); - PopScopes(1); - return expr; - }); - }); - exprs.push_back(block); - break; - } - case TokenType::kFreeVar: { - Consume(TokenType::kFreeVar); - auto var_token = Match(TokenType::kLocal); - - Type type; - if (WhenMatch(TokenType::kColon)) { - type = ParseType(); - } else { - type = IncompleteType(); - } - - BindFreeVar(var_token.ToString(), type); - break; - } - // Parses `let ...`; - case TokenType::kLet: - exprs.push_back(ParseBindingExpr()); - break; - case TokenType::kMatch: - case TokenType::kPartialMatch: { - bool is_total = next->token_type == TokenType::kMatch; - Consume(next->token_type); - exprs.push_back(ParseMatch(is_total)); - break; - } - - // %x ... - case TokenType::kGraph: - if (Lookahead(2)->token_type == TokenType::kEqual) { - exprs.push_back(ParseBindingExpr()); - break; - } - // intentional fall through here. - default: { - exprs.push_back(ParseExprBinOp()); - break; - } - } - - if (!WhenMatch(TokenType::kSemicolon)) { - break; - } - } - - ICHECK_GE(exprs.size(), 1); - - if (exprs.size() == 1) { - // ICHECK(exprs[0].defined() && exprs[0]->span.defined()) - // << "parser must set expression spans.\n" - // << exprs[0]; - return exprs[0]; - } else { - auto body = exprs.back(); - exprs.pop_back(); - while (exprs.size()) { - auto value = exprs.back(); - ICHECK(value->span.defined()) << "parser must set expression spans."; - exprs.pop_back(); - body = relay::Let(Var("", IncompleteType()), value, body, value->span.Merge(body->span)); - } - ICHECK(body->span.defined()) << "parser must set expression spans."; - return body; - } - }); - } - - /*! \brief Parse a "binding expression"; an expression where - * a graph or let variable is bound. - * - * In order to avoid stack overflow this is implemented in a special - * iterative way to keep stack depth constant in a long chain of bindings. - */ - Expr ParseBindingExpr() { - // We use a loop here so that the stack depth - // does not grow linearly with a sequence of - // graph or let bindings. - // - // Assuming we start at call depth k, we will - // enter k + c call frames to parse the RHS - // of the bindings where `c` is the depth - // of recursion needed by RHS. - // - // If RHS is a call expresssion the c=1. - // - // Once we have parsed the RHS we will be - // back at depth K, and will return to - // this loop header to parse another - // graph or let binding. - // - // This ensures for n sequential bindings - // the call depth will be the same before - // and after parsing the n bindings. - VLOG(9) << "Parser::ParseBindingExpr"; - std::vector> bindings; - int scopes = 0; - - while (true) { - auto next = Peek(); - if (next->token_type == TokenType::kGraph && Lookahead(2)->token_type == TokenType::kEqual) { - Match(TokenType::kGraph); - Match(TokenType::kEqual); - auto val = this->ParseExprBinOp(); - Match(TokenType::kSemicolon); - AddGraphBinding(next, val); - } else if (next->token_type == TokenType::kLet) { - auto span = next->span; - // Parse the 'let'. - Consume(TokenType::kLet); - - // Parse the local '%'. - auto local_tok = Match(TokenType::kLocal); - auto string = local_tok.ToString(); - - // Parse the optional type annotation (':' ). - Type type; - if (WhenMatch(TokenType::kColon)) { - type = ParseType(); - } - - auto var = BindVar(string, type); - - // Parse the '='; - Match(TokenType::kEqual); - - // Parse the body, and the ';'. - auto val = this->ParseExprBinOp(); - Consume(TokenType::kSemicolon); - - // Add the bindings to the local data structure. - std::tuple tuple(var, val, span); - bindings.push_back(tuple); - scopes++; - PushScope(); - } else { - // This is the only case we will increase the stack - // depth. - // - // If we parse a program which is a sequence of N bindings - // followed by a single body expression we will end up with - // a call depth of 3, the first call to ParseExpr, then - // ParseBindingExpr, then finally ParseExpr once more. - - auto body = this->ParseExpr(); - - // Remove the same number of scopes we added. - PopScopes(scopes); - - if (bindings.size() == 0) { - return body; - } else { - // We can now build the let binding up backwards. - for (auto binding = bindings.rbegin(); binding != bindings.rend(); binding++) { - auto span = body->span.Merge(std::get<2>(*binding)); - body = relay::Let(std::get<0>(*binding), std::get<1>(*binding), body, span); - } - return body; - } - } - } - } - - /*! Parse a function definition without a leading keyword or identifier. - * - * Handles things of the form [T1, ..., TN](arg1: U1, ..., argN : UN) -> Ret { body }. - */ - Function ParseFunctionDef() { - VLOG(9) << "Parser::ParseFunctionDef"; - return WithSpan([&]() { - PushScope(); - PushTypeScope(); - - Array generics; - if (Peek()->token_type == TokenType::kLSquare) { - generics = ParseSequence( - TokenType::kLSquare, TokenType::kComma, TokenType::kRSquare, [&]() { - auto type_var_name = Match(TokenType::kIdentifier).ToString(); - return BindTypeVar(type_var_name, TypeKind::kType); - }); - } - - Map raw_attrs; - - auto params = ParseSequence( - TokenType::kOpenParen, TokenType::kComma, TokenType::kCloseParen, - [&]() { - auto token = Match(TokenType::kLocal); - auto string = token.ToString(); - - // The fake attributes where the virtual device is specified. - VirtualDevice virtual_device; - if (WhenMatch(TokenType::kLCurly)) { - Map fake_attrs = ParseAttrs(); - VLOG(9) << "Fake attributes for function parameter: " << fake_attrs; - Match(TokenType::kRCurly); - if (fake_attrs.size() == 1 && fake_attrs.count(kVirtualDevice)) { - ICHECK(fake_attrs[kVirtualDevice].as()) - << "Expected the " << kVirtualDevice - << " to have type VirtualDeviceNode, but got " << virtual_device->GetTypeKey(); - virtual_device = Downcast(fake_attrs[kVirtualDevice]); - } - } - - Type type; - if (WhenMatch(TokenType::kColon)) { - type = ParseType(); - } - return BindVar(string, type, virtual_device); - }, - [&] { - auto is_ident = Lookahead(1)->token_type == TokenType::kIdentifier; - auto next_is_equal = Lookahead(2)->token_type == TokenType::kEqual; - - if (is_ident && next_is_equal) { - raw_attrs = ParseAttrs(); - return true; - } - - return false; - }); - - Type ret_type; - if (WhenMatch(TokenType::kMinus)) { - Match(TokenType::kRAngle); - ret_type = ParseType(); - } - - auto body = Block([&]() { return ParseExpr(); }); - - PopTypeScopes(1); - PopScopes(1); - - // TODO(@jroesch): attributes should never be null, they should always be empty. - if (raw_attrs.size()) { - // Promote kVirtualDevice to first-class - if (raw_attrs.count(kVirtualDevice)) { - ObjectRef vid = raw_attrs.at(kVirtualDevice); - ICHECK(vid.as()) - << "Expected the " << kVirtualDevice << " to have type VirtualDeviceNode, but got " - << vid->GetTypeKey(); - - DictAttrs attrs; - // Don't fill the raw_attrs in if there's nothing other than kVirtualDevice in the - // attributes - if (raw_attrs.size() > 1) { - raw_attrs.erase(kVirtualDevice); - attrs = DictAttrs(raw_attrs); - } - Function func = relay::Function(params, body, ret_type, generics, attrs); - func->virtual_device_ = vid; - return func; - } else { - return relay::Function(params, body, ret_type, generics, DictAttrs(raw_attrs)); - } - } else { - return relay::Function(params, body, ret_type, generics, tvm::DictAttrs()); - } - }); - } - - /*! \brief Parse an if-expression. */ - Expr ParseIf() { - return WithSpan([&]() { - VLOG(9) << "Parser::ParseIf"; - Consume(TokenType::kIf); - - auto guard = WithSpan([&] { return Parens([&] { return ParseExpr(); }); }); - - auto true_branch = Block([&] { - this->PushScope(); - auto expr = ParseExpr(); - this->PopScopes(1); - return expr; - }); - - Match(TokenType::kElse); - - auto false_branch = Block([&] { - this->PushScope(); - auto expr = ParseExpr(); - this->PopScopes(1); - return expr; - }); - - return relay::If(guard, true_branch, false_branch); - }); - } - - /* This factors parsing a list of patterns for both tuples, and constructors. */ - Array ParsePatternList() { - return ParseSequence(TokenType::kOpenParen, TokenType::kComma, TokenType::kCloseParen, - [&] { return ParsePattern(); }); - } - - /*! \brief Parses a pattern for a match expression. - * - * A pattern is either a wildcard `_`, a local `%name`, - * a constructor `C(p1, ..., pn)` or tuple `(p1, ..., pn). - * - * This function recursively parses a pattern. - */ - Pattern ParsePattern() { - VLOG(9) << "Parser::ParsePattern"; - auto next = Peek(); - switch (next->token_type) { - case TokenType::kUnderscore: { - Match(TokenType::kUnderscore); - return PatternWildcard(); - } - case TokenType::kLocal: { - auto id = Match(TokenType::kLocal); - Type type_annotation; - if (WhenMatch(TokenType::kColon)) { - type_annotation = ParseType(); - } - auto var = BindVar(id.ToString(), type_annotation); - return PatternVar(var); - } - case TokenType::kIdentifier: { - auto id = Match(TokenType::kIdentifier); - auto ctor = ctors.Get(id.ToString()); - if (!ctor) { - diag_ctx.EmitFatal( - // TODO(@jroesch): split into error and help - // deal with multiple rendering - Diagnostic::Error(id->span) - << "undefined constructor name `" << id.ToString() - << "`, perhaps you intended to write a" - << "pattern variable, considering changing this to `%" << id.ToString() << "`"); - } - if (Peek()->token_type == TokenType::kOpenParen) { - auto fields = ParsePatternList(); - return PatternConstructor(ctor.value(), fields); - } else { - return PatternConstructor(ctor.value(), {}); - } - } - default: - return PatternTuple(ParsePatternList()); - } - } - - Clause ParseMatchArm() { - PushScope(); - auto pattern = ParsePattern(); - Match(TokenType::kEqual); - Consume(TokenType::kRAngle); - auto expr = ParseExpr(); - PopScopes(1); - return Clause(pattern, expr); - } - - Expr ParseMatch(bool is_total) { - return WithSpan([&]() { - Expr scrutinee = ParseAtomicExpr(); - - Array clauses = - ParseSequence(TokenType::kLCurly, TokenType::kComma, TokenType::kRCurly, - [&] { return ParseMatchArm(); }); - - return relay::Match(scrutinee, clauses, is_total); - }); - } - - Expr ParseExprBinOp() { - VLOG(9) << "Parser::ParseExprBinOp"; - return WithSpan([this] { - // We must parse at least one expression, the default - // case is that there is no operator and we will fall - // through. - std::vector exprs; - Expr expr = WithSpan([this] { return ParseCallExpr(); }); - - exprs.push_back(expr); - - // Now we parse an optional op. - std::vector ops; - - // We will now parse 0 or more operator occurrences. - while (true) { - auto opt_op = ParseOp(); - - // If we didn't parse one we done. - if (opt_op.size() == 0) { - break; - } - - // Read the operation we parsed; - auto op = opt_op[0]; - - Expr right = WithSpan([this] { return ParseCallExpr(); }); - ICHECK(right->span.defined()); - - // If the operator stack is empty - // we parse an operator and expression - // and push them to stacks, then - // continue. - if (ops.size() == 0) { - ops.push_back(op); - exprs.push_back(right); - continue; - } - - if (op.precedence > ops.back().precedence || - (op.precedence == ops.back().precedence && op.left_assoc == false)) { - ops.push_back(op); - exprs.push_back(right); - continue; - } - - while (ops.size() && (op.precedence < ops.back().precedence || - (op.precedence == ops.back().precedence && op.left_assoc == true))) { - Rule new_op = ops.back(); - ops.pop_back(); - Expr right = exprs.back(); - exprs.pop_back(); - Expr left = exprs.back(); - exprs.pop_back(); - ICHECK(new_op.op.defined()) << "a call op must be set " << new_op.op; - exprs.push_back( - relay::Call(new_op.op, {left, right}, Attrs(), {}, left->span.Merge(right->span))); - } - - exprs.push_back(right); - ops.push_back(op); - } - - while (ops.size()) { - Rule new_op = ops.back(); - ops.pop_back(); - Expr right = exprs.back(); - exprs.pop_back(); - Expr left = exprs.back(); - exprs.pop_back(); - ICHECK(new_op.op.defined()) << "a call op must be set " << new_op.op; - exprs.push_back( - relay::Call(new_op.op, {left, right}, Attrs(), {}, left->span.Merge(right->span))); - } - - ICHECK_EQ(ops.size(), 0) << "No operations should be left on the operation stack."; - - ICHECK_EQ(exprs.size(), 1) - << "Only a single expression should be left on the expression stack."; - - return exprs[0]; - }); - } - - ObjectRef ParseAttributeValue() { - VLOG(9) << "Parser::ParseAttributeValue"; - auto next = Peek(); - switch (next->token_type) { - case TokenType::kFloat: - case TokenType::kInteger: - case TokenType::kBoolean: - case TokenType::kStringLiteral: - return Match(next->token_type)->data; - case TokenType::kMetaReference: - return ParseMetaRef(); - case TokenType::kLSquare: { - return ParseSequence(TokenType::kLSquare, TokenType::kComma, TokenType::kRSquare, - [&]() { return ParseAttributeValue(); }); - } - case TokenType::kOpenParen: { - // TODO(@jroesch: need to figure out bracket vs. sequence) - // return ParseSequence(TokenType::kOpenParen, TokenType::kComma, - // TokenType::kCloseParen, - // [&]() { return ParseAttributeValue(); }); - return Bracket(TokenType::kOpenParen, TokenType::kCloseParen, - [&]() { return ParseAttributeValue(); }); - } - // TODO(@jroesch): not sure about this being the right way to handle nulls. - case TokenType::kIdentifier: { - if (auto text = next->data.as()) { - std::string id = text.value(); - if (id == "nullptr") { - Match(TokenType::kIdentifier); - return ObjectRef(); - } - if (id == "None") { - Match(TokenType::kIdentifier); - return Optional(); - } - } - } - default: - return ParseAtomicExpr(); - } - } - - Map ParseAttrs() { - VLOG(9) << "Parser::ParseAttrs"; - Map kwargs; - while (Peek()->token_type == TokenType::kIdentifier) { - auto key = GetHierarchicalName(ParseHierarchicalName().data); - Match(TokenType::kEqual); - // TOOD(@jroesch): syntactically what do we allow to appear in attribute right hand side. - auto value = ParseAttributeValue(); - // TODO(@jroesch): we need a robust way to handle this writing dtypes as strings in text - // format is bad. - kwargs.Set(key, value); - WhenMatch(TokenType::kComma); - } - VLOG(9) << "Parser::ParseAttrs: kwargs=" << kwargs; - return kwargs; - } - - Expr ParseCallArgs(Expr op) { - ICHECK(op.defined()) << "the operator must be defined"; - - VLOG(9) << "Parser::ParseCallArgs"; - Attrs attrs; - std::string op_key; - bool is_op = false; - - if (auto op_node = op.as()) { - is_op = true; - op_key = op_node->attrs_type_key; - } - - if (Peek()->token_type == TokenType::kOpenParen) { - Array args = ParseSequence( - TokenType::kOpenParen, TokenType::kComma, TokenType::kCloseParen, - [&] { return ParseExpr(); }, - [&] { - auto is_ident = Lookahead(1)->token_type == TokenType::kIdentifier; - auto next_is_equal = Lookahead(2)->token_type == TokenType::kEqual; - auto is_pretty_attrs = is_ident && next_is_equal; - auto is_meta_next = Lookahead(1)->token_type == TokenType::kMetaReference; - // TODO(@jroesch): might not handle trailing comma - auto last_meta = Lookahead(2)->token_type == TokenType::kCloseParen; - auto is_meta_attrs = is_meta_next && last_meta; - - if (is_pretty_attrs || is_meta_attrs) { - if (is_meta_attrs) { - auto meta_ref = ParseMetaRef(); - if (meta_ref.as()) { - attrs = Downcast(meta_ref); - } else { - // Not awesome parsing code here. - this->pos--; - return false; - } - } else { - auto raw_attrs = ParseAttrs(); - if (is_op && op_key.size()) { - auto attr_obj = tvm::ReflectionVTable::Global()->CreateObject(op_key, raw_attrs); - ICHECK(attr_obj.defined()); - attrs = Downcast(attr_obj); - } else if (raw_attrs.count("attrs_type_key")) { - String attr_key = Downcast(raw_attrs["attrs_type_key"]); - if (attr_key.size()) { - raw_attrs.erase("attrs_type_key"); - auto attr_obj = - tvm::ReflectionVTable::Global()->CreateObject(attr_key, raw_attrs); - ICHECK(attr_obj.defined()); - attrs = Downcast(attr_obj); - } - } else { - this->diag_ctx.EmitFatal(Diagnostic::Error(op->span) - << "unable to determine the 'attrs_type_key' with which " - "to represent the call attributes for this operator"); - } - } - return true; - } - return false; - }); - - if (!attrs.defined()) { - if (is_op && op_key.size()) { - auto attr_obj = tvm::ReflectionVTable::Global()->CreateObject(op_key, {}); - ICHECK(attr_obj.defined()); - attrs = Downcast(attr_obj); - } - } - - // TODO(@jroesch): in a secondary pass adjust spans. - return Expr(Call(op, args, attrs, {})); - } else { - return Expr(); - } - - return Expr(); - } - - Expr ParseCallExpr() { - VLOG(9) << "Parser::ParseCallExpr"; - return WithSpan([this] { - Expr expr = ParseAtomicExpr(); - // Parse as many call args as possible, building up expression - // - // NB(@jroesch): this seems like a hack but in order to parse curried functions - // and avoid complex grammar we will parse multiple call lists in a row. - while (Peek()->token_type == TokenType::kOpenParen) { - auto new_expr = ParseCallArgs(expr); - - if (new_expr.defined()) { - expr = new_expr; - } else { - break; - } - } - - // We need a zero-arity case for constructors. - if (auto ctor_node = expr.as()) { - if (ctor_node->inputs.size() == 0) { - return Expr(Call(expr, {})); - } - } - - return expr; - }); - } - - Expr GetOp(const std::string& op_name, const Span& span) { - VLOG(9) << "op_name=" << op_name << " span=" << span; - try { - return Op::Get(op_name); - } catch (const Error& e) { - // we can relax this, but probably need to relax checks or return non-null here. - this->diag_ctx.EmitFatal(Diagnostic::Error(span) - << "operator `" << op_name - << "` not found, perhaps you forgot to register it?"); - return Expr(); - } - } - - Expr ParseAtomicExpr() { - VLOG(9) << "Parser::ParseAtomicExpr"; - Expr expr = WithSpan([this] { - auto next = Peek(); - switch (next->token_type) { - case TokenType::kInteger: - case TokenType::kFloat: { - Consume(next->token_type); - auto number = NumberToNDArray(next); - Expr e = Constant(number, next->span); - ICHECK(e->span.defined()) << "constant spans must be defined"; - return e; - } - case TokenType::kBoolean: { - Consume(TokenType::kBoolean); - int64_t value = Downcast(next->data).IntValue(); - Expr e = Constant(support::BoolToNDArray(value), next->span); - ICHECK(e->span.defined()) << "constant spans must be defined"; - return e; - } - // Parse a local of the form `%x`. - case TokenType::kLocal: { - Consume(TokenType::kLocal); - return Expr(LookupLocal(next)); - } - // Parse a local of the form `@x`. - case TokenType::kGlobal: { - auto global_name = next.ToString(); - Consume(TokenType::kGlobal); - auto global = AddOrGet(&global_names, global_name); - return Expr(global); - } - // Parse a local of the form `x`. - // Right now we fail to parse `x.y`. - case TokenType::kIdentifier: { - auto ctor = ctors.Get(next.ToString()); - if (ctor) { - Consume(TokenType::kIdentifier); - return Expr(ctor.value()); - } else { - auto spanned_idents = ParseHierarchicalName(); - auto idents = spanned_idents.data; - auto span = spanned_idents.span; - return GetOp(GetHierarchicalName(idents), span); - } - } - case TokenType::kGraph: { - Consume(TokenType::kGraph); - return LookupGraphBinding(next); - } - case TokenType::kMetaReference: { - return Downcast(ParseMetaRef()); - } - case TokenType::kFn: { - Consume(TokenType::kFn); - Expr e = ParseFunctionDef(); - ICHECK(e->span.defined()) << "function spans must be defined.\n" << e; - return e; - } - case TokenType::kIf: { - Expr e = ParseIf(); - return e; - } - case TokenType::kRef: { - Consume(TokenType::kRef); - Match(TokenType::kOpenParen); - auto ref_value = ParseExpr(); - Match(TokenType::kCloseParen); - return static_cast(RefCreate(ref_value)); - } - case TokenType::kRefRead: { - return WithSpan([&]() { - Consume(TokenType::kRefRead); - Match(TokenType::kOpenParen); - auto ref = ParseExpr(); - Match(TokenType::kCloseParen); - return static_cast(RefRead(ref)); - }); - } - case TokenType::kRefWrite: { - return WithSpan([&]() { - Consume(TokenType::kRefWrite); - Match(TokenType::kOpenParen); - auto ref = ParseExpr(); - Match(TokenType::kComma); - auto value = ParseExpr(); - Match(TokenType::kCloseParen); - return static_cast(RefWrite(ref, value)); - }); - } - case TokenType::kOpenParen: { - Span sp = next->span; - Consume(TokenType::kOpenParen); - // parse '(' ')' - if (WhenMatch(TokenType::kCloseParen)) { - return Expr(Tuple(Array())); - } else { - Expr subexpr = ParseExpr(); - // parse '(' expr ')' - if (WhenMatch(TokenType::kCloseParen)) { - return subexpr; - // parse '( expr ',' * ')' - } else if (WhenMatch(TokenType::kComma)) { - Array exprs = {subexpr}; - while (true) { - if (WhenMatch(TokenType::kCloseParen)) { - break; - } else { - auto element = ParseExpr(); - auto comma = Peek(); - if (WhenMatch(TokenType::kComma)) { - sp = sp.Merge(element->span.Merge(comma->span)); - } else { - sp = sp.Merge(element->span); - } - exprs.push_back(element); - } - } - Expr tuple = Tuple(exprs, sp); - ICHECK(tuple->span.defined()) << "tuple span should be defined"; - return tuple; - } - } - } - default: { - this->diag_ctx.EmitFatal(Diagnostic::Error(next->span) - << "expected an expression found " << Pretty(next->token_type)); - return Expr(); - } - } - }); - - if (WhenMatch(TokenType::kPeriod)) { - auto token = Match(TokenType::kInteger); - auto index = token.ToNumber(); - auto span = token->span.Merge(expr->span); - VLOG(9) << "Parser::ParseAtomicExpr: tuple get item"; - return relay::TupleGetItem(expr, index, span); - } else { - return expr; - } - } - - /*! \brief Parse a hierarchical name. - * - * The tokenizer produces a token stream of . - * and so on for names of the form `nn.conv2d`. - * Currently we only use string names everywhere instead - * of a notion of a hierarchical name. - * - * The below utility reassembles a token stream into a - * single stream inserting the required periods needed - * to look up registered names. - */ - Spanned> ParseHierarchicalName() { - Array idents; - Span span; - while (Peek()->token_type == TokenType::kIdentifier) { - auto token = Peek(); - - if (span.defined()) { - span = span.Merge(token->span); - } else { - span = token->span; - } - - auto name = token.ToString(); - idents.push_back(name); - Consume(TokenType::kIdentifier); - - // Keep parsing while we see a trailing period. - if (Peek()->token_type == TokenType::kPeriod) { - Consume(TokenType::kPeriod); - continue; - } else { - // No more periods means we are done! - break; - } - } - - return Spanned>(idents, span); - } - - std::string GetHierarchicalName(Array idents) { - ICHECK_NE(idents.size(), 0); - std::stringstream hierarchical_name; - int i = 0; - int periods = idents.size() - 1; - for (auto ident : idents) { - hierarchical_name << ident; - if (i < periods) { - hierarchical_name << "."; - i++; - } - } - return hierarchical_name.str(); - } - - /*! \brief Parse a shape. */ - Array ParseShape() { - auto dims = ParseSequence( - TokenType::kOpenParen, TokenType::kComma, TokenType::kCloseParen, [&]() { - tvm::PrimExpr dim; - if (Peek()->token_type == TokenType::kMetaReference) { - dim = Downcast(ParseMetaRef()); - } else if (WhenMatch(TokenType::kQuestion)) { - dim = tvm::tir::Any(); - } else { - dim = Downcast(Match(TokenType::kInteger)->data); - } - - return dim; - }); - return dims; - } - - /*! \brief Parse a function type. */ - Type ParseFunctionType() { - auto ty_params = ParseSequence(TokenType::kOpenParen, TokenType::kComma, - TokenType::kCloseParen, [&]() { return ParseType(); }); - - Match(TokenType::kMinus); - Match(TokenType::kRAngle); - auto ret_type = ParseType(); - - return relay::FuncType(ty_params, ret_type, {}, {}); - } - - // Parses a user defined ADT or type variable. - Type ParseNonPrimitiveType(const Token& tok) { - return WithSpan([&]() { - auto name = tok.ToString(); - Type head_type = LookupTypeVar(tok); - - if (!head_type.defined()) { - // head_type = type_names.Get(name); - head_type = AddOrGet(&type_names, name, TypeKind::kAdtHandle); - } - - if (!head_type.defined()) { - diag_ctx.EmitFatal(Diagnostic::Error(tok->span) - << "the type constructor `" << name << "` is undefined"); - } - - Array arg_types; - if (Peek()->token_type == TokenType::kLSquare) { - arg_types = ParseSequence(TokenType::kLSquare, TokenType::kComma, TokenType::kRSquare, - [&]() { return ParseType(); }); - } - - if (arg_types.size()) { - return static_cast(TypeCall(head_type, arg_types)); - } else { - if (head_type.as()) { - return static_cast(TypeCall(head_type, {})); - } else { - return static_cast(head_type); - } - } - }); - } - - /*! \brief Parses a TVM type. - * - * This matches either a `Tensor[shape, dtype]`, a user defined ADT, a tuple type, - * a scalar type or an incomplete type `_`. - */ - Type ParseType() { - return WithSpan([&]() -> Type { - auto tok = Peek(); - - if (tok->token_type == TokenType::kOpenParen) { - auto tys = - ParseSequence(TokenType::kOpenParen, TokenType::kComma, - TokenType::kCloseParen, [&]() { return ParseType(); }); - return relay::TupleType(tys); - } else if (WhenMatch(TokenType::kFn)) { - return ParseFunctionType(); - } else if (WhenMatch(TokenType::kIdentifier)) { - auto id = tok.ToString(); - if (id == "Tensor") { - Match(TokenType::kLSquare); - auto shape = ParseShape(); - Match(TokenType::kComma); - auto dtype_tok = Match(TokenType::kIdentifier); - auto dtype = DataType(String2DLDataType(dtype_tok.ToString())); - Match(TokenType::kRSquare); - return TensorType(shape, dtype); - } else { - auto ty = tok.ToString(); - if (ty.rfind("int", 0) == 0 || ty.find("float", 0) == 0 || ty.find("uint", 0) == 0 || - ty.find("bool", 0) == 0) { - // Need to do better error handling here. - auto dtype = DataType(String2DLDataType(tok.ToString())); - return TensorType({}, dtype); - } else { - return ParseNonPrimitiveType(tok); - } - } - } else if (WhenMatch(TokenType::kUnderscore)) { - return IncompleteType(); - } else { - this->diag_ctx.EmitFatal(Diagnostic::Error(tok->span) - << "failed to parse type found " << tok); - return Type(); - } - }); - } - - template - R ConsumeWhitespace(std::function func) { - auto old = this->ignore_whitespace; - this->ignore_whitespace = true; - while (tokens[pos]->token_type == TokenType::kWhitespace) { - pos++; - } - auto res = func(); - this->ignore_whitespace = old; - return res; - } - - Map> ParseMetadata() { - if (Peek()->token_type == TokenType::kMetadata) { - return Match(TokenType::kMetadata).ToMetadata(); - } else { - return Map>(); - } - } - - /*! \brief A helper for debugging the parser, displays the next N tokens in the token stream. */ - void DisplayNextN(int n) { - std::cout << "remaining tokens: " << std::endl; - auto bound = std::min(pos + n, static_cast(tokens.size())); - for (int i = 0; i < bound - pos; i++) { - std::cout << tokens[pos + i] << std::endl; - } - } - - // A function for debugging the operator parser. - void DebugStack(const std::vector& exprs, const std::vector& rules) { - std::cout << "Expr Stack: "; - for (auto expr : exprs) { - std::cout << expr << ", "; - } - - std::cout << std::endl; - std::cout << "Op Stack: "; - for (auto rule : rules) { - std::cout << rule.op << ", "; - } - - std::cout << std::endl; - } -}; - -Parser InitParser(const std::string& file_name, const std::string& file_content, - const Optional& init_module, const MetaTable& init_meta_table) { - VLOG(9) << "InitParser: file_name: " << file_name << "file_content_size: " << file_content.size(); - SourceName src_name = SourceName::Get(file_name); - Source source(src_name, file_content); - - IRModule module; - if (!init_module) { - SourceMap source_map; - module = IRModule({}, {}, {}, source_map); - } else { - module = init_module.value(); - } - - module->source_map.Add(source); - - auto diag_ctx = DiagnosticContext::Default(module); - auto tokens_and_table = Tokenize(diag_ctx, source); - - auto tokens = tokens_and_table.first; - MetaTable meta_data_table = tokens_and_table.second.ToMetadata(); - - // Merge any entries in init_meta_table into anything captured in the #[metadata] section - // of the file_content. Metadata references within file_content must use indexes which account - // for this ordering. - for (const auto& pair : init_meta_table) { - Array items; - if (meta_data_table.count(pair.first)) { - items = meta_data_table[pair.first]; - } - for (const auto& obj : pair.second) { - items.push_back(obj); - } - meta_data_table.Set(pair.first, items); - } - - return Parser(module, diag_ctx, source, tokens, DefaultOpTable(), std::move(meta_data_table)); -} - -IRModule ParseModule(const std::string& file_name, const std::string& file_content, - const Optional& init_module, const MetaTable& init_meta_table) { - VLOG_CONTEXT << "ParseModule"; - VLOG(9) << "parsing and type-checking " << file_name; - auto parser = InitParser(file_name, file_content, init_module, init_meta_table); - auto mod = parser.ParseModule(); - ICHECK(mod.defined()) << "The parser must return a non-null module."; - // NB(@jroesch): it is very important that we render any errors before we proceed - // if there were any errors which allow the parser to proceed we must render them - // here. - parser.diag_ctx.Render(); - auto infer_type = tvm::relay::transform::InferType(); - ICHECK(infer_type.defined()) << "The type inferencer must be non-null."; - return infer_type(mod); -} - -Expr ParseExpr(const std::string& file_name, const std::string& file_content) { - VLOG(9) << "ParseExpr"; - auto parser = InitParser(file_name, file_content, Optional(), MetaTable()); - parser.ParseSemVer(false); - parser.PushScope(); - auto expr = parser.ParseExpr(); - parser.Match(TokenType::kEndOfFile); - // NB(@jroesch): it is very important that we render any errors before we proceed - // if there were any errors which allow the parser to proceed we must render them - // here. - parser.diag_ctx.Render(); - return expr; -} - -/*! - * \brief This pass pretty-prints mod then parses it back so as to establish spans and sources - * for all Relay sub-expressions. This improves error and debugging diagnostics downstream for - * modules constructed programaticaly rather than textually. - */ -Pass AnnotateSpans() { - auto pass_func = [](const IRModule& mod, const PassContext& ctx) { - String text = AsText(mod, /*show_meta_data=*/true); - VLOG(1) << "AnnotateSpans intermediate text:" << std::endl << text; - return ParseModule("GeneratedSource", text); - }; - return CreateModulePass(pass_func, 0, "AnnotateSpans", {}); -} - -TVM_REGISTER_GLOBAL("relay.parser.ParseModuleInContext") - .set_body_typed([](const std::string& file_name, const std::string& file_content, - const Optional& init_module, const MetaTable& init_meta_table) { - return ParseModule(file_name, file_content, init_module, init_meta_table); - }); - -TVM_REGISTER_GLOBAL("relay.parser.ParseModule").set_body([](TVMArgs args, TVMRetValue* ret) { - ICHECK(args.size() >= 2 && args.size() <= 4) << "Expected 2-4 arguments, but got " << args.size(); - if (args.size() == 2) { - *ret = ParseModule(args[0], args[1]); - } else if (args.size() == 3) { - *ret = ParseModule(args[0], args[1], args[2]); - } else { - *ret = ParseModule(args[0], args[1], args[2], args[3]); - } -}); - -TVM_REGISTER_GLOBAL("relay.parser.ParseExpr") - .set_body_typed([](tvm::String file_name, tvm::String file_content) { - return ParseExpr(file_name, file_content); - }); - -TVM_REGISTER_GLOBAL("relay._transform.AnnotateSpans").set_body_typed(AnnotateSpans); - -} // namespace relay -} // namespace tvm diff --git a/src/relay/parser/span_check.cc b/src/relay/parser/span_check.cc deleted file mode 100644 index 6bbf6317ad9f..000000000000 --- a/src/relay/parser/span_check.cc +++ /dev/null @@ -1,107 +0,0 @@ -/* - * 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 span_check.cc - * \brief A utility for checking and reporting malformed span information. - */ -#include "./span_check.h" - -#include - -namespace tvm { -namespace relay { - -using tvm::relay::transform::CreateFunctionPass; -using tvm::transform::PassContext; - -void SpanChecker::VisitExpr(const Expr& e) { - this->expression = e; - VisitSpan(e->span); - span_stack.push_back(e->span); - ExprVisitor::VisitExpr(e); - this->expression = e; - span_stack.pop_back(); -} - -// TODO(@jroesch, @junru): we need to deal with unique spans for global/var. -void SpanChecker::VisitExpr_(const VarNode* op) {} -void SpanChecker::VisitExpr_(const GlobalVarNode* op) {} -void SpanChecker::VisitExpr_(const ConstantNode* op) {} - -void SpanChecker::VisitExpr_(const TupleNode* op) { ExprVisitor::VisitExpr_(op); } - -void SpanChecker::VisitExpr_(const FunctionNode* op) { ExprVisitor::VisitExpr_(op); } - -void SpanChecker::VisitExpr_(const CallNode* op) { ExprVisitor::VisitExpr_(op); } - -void SpanChecker::VisitExpr_(const LetNode* op) { ExprVisitor::VisitExpr_(op); } - -void SpanChecker::VisitExpr_(const IfNode* op) { ExprVisitor::VisitExpr_(op); } - -void SpanChecker::VisitExpr_(const OpNode* op) {} - -void SpanChecker::VisitExpr_(const TupleGetItemNode* op) { ExprVisitor::VisitExpr_(op); } - -void SpanChecker::VisitExpr_(const RefCreateNode* op) { ExprVisitor::VisitExpr_(op); } - -void SpanChecker::VisitExpr_(const RefReadNode* op) { ExprVisitor::VisitExpr_(op); } - -void SpanChecker::VisitExpr_(const RefWriteNode* op) { ExprVisitor::VisitExpr_(op); } - -void SpanChecker::VisitExpr_(const ConstructorNode* op) {} // ExprVisitor::VisitExpr_(op); } - -void SpanChecker::VisitExpr_(const MatchNode* op) { ExprVisitor::VisitExpr_(op); } - -void SpanChecker::VisitSpan(const Span& sp) { - if (!sp.defined()) { - Span span; - for (auto spans = this->span_stack.rbegin(); spans != this->span_stack.rend(); spans++) { - span = this->span_stack.back(); - if (span.defined()) { - diag_ctx.Emit(Diagnostic::Warning(span) << "found null-span, i-nodes deep from this span."); - return; - } - } - auto warning = Diagnostic::Warning(span); - warning << "\tAll spans are null\n"; - warning << "\t" << this->expression; - diag_ctx.Emit(warning); - } -} - -void SpanChecker::VisitType(const Type& t) {} -void SpanChecker::VisitClause(const Clause& c) {} -void SpanChecker::VisitPattern(const Pattern& c) {} - -Pass SpanCheck() { - return CreateFunctionPass( - [](const Function& func, const IRModule& mod, const PassContext& ctx) { - ICHECK(ctx->diag_ctx) << "Diagnostic context must be set."; - SpanChecker checker(ctx->diag_ctx.value()); - checker.VisitExpr(func); - ctx->diag_ctx.value().Render(); - return func; - }, - 0, "SpanCheck", {}); -} - -TVM_REGISTER_GLOBAL("relay.parser.SpanCheck").set_body_typed([]() { return SpanCheck(); }); - -} // namespace relay -} // namespace tvm diff --git a/src/relay/parser/span_check.h b/src/relay/parser/span_check.h deleted file mode 100644 index b85b4a497965..000000000000 --- a/src/relay/parser/span_check.h +++ /dev/null @@ -1,78 +0,0 @@ -/* - * 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 span_check.h - * \brief Check that the Relay IR has correctly attached span information. - */ -#ifndef TVM_RELAY_PARSER_SPAN_CHECK_H_ -#define TVM_RELAY_PARSER_SPAN_CHECK_H_ - -#include -#include -#include -#include -#include -#include - -#include -#include -#include -#include - -namespace tvm { -namespace relay { - -using namespace tvm::relay; -using tvm::transform::Pass; - -struct SpanChecker : ExprVisitor { - Expr expression; - DiagnosticContext diag_ctx; - std::vector span_stack; - - explicit SpanChecker(DiagnosticContext diag_ctx) : diag_ctx(diag_ctx) {} - - void VisitExpr(const Expr& expr) override; - void VisitExpr_(const VarNode* op) override; - void VisitExpr_(const GlobalVarNode* op) override; - void VisitExpr_(const ConstantNode* op) override; - void VisitExpr_(const TupleNode* op) override; - void VisitExpr_(const FunctionNode* op) override; - void VisitExpr_(const CallNode* op) override; - void VisitExpr_(const LetNode* op) override; - void VisitExpr_(const IfNode* op) override; - void VisitExpr_(const OpNode* op) override; - void VisitExpr_(const TupleGetItemNode* op) override; - void VisitExpr_(const RefCreateNode* op) override; - void VisitExpr_(const RefReadNode* op) override; - void VisitExpr_(const RefWriteNode* op) override; - void VisitExpr_(const ConstructorNode* op) override; - void VisitExpr_(const MatchNode* op) override; - void VisitType(const Type& t) override; - void VisitClause(const Clause& c) override; - void VisitPattern(const Pattern& c) override; - void VisitSpan(const Span& span) override; -}; - -Pass SpanCheck(); - -} // namespace relay -} // namespace tvm -#endif // TVM_RELAY_PARSER_SPAN_CHECK_H_ diff --git a/src/relay/parser/token.h b/src/relay/parser/token.h deleted file mode 100644 index 13875cb09391..000000000000 --- a/src/relay/parser/token.h +++ /dev/null @@ -1,406 +0,0 @@ -/* - * 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 token.h - * \brief The definition of tokens for the TVM parser. - */ - -#ifndef TVM_RELAY_PARSER_TOKEN_H_ -#define TVM_RELAY_PARSER_TOKEN_H_ - -#include -#include -#include - -#include -#include -#include - -namespace tvm { -namespace relay { - -enum class TokenType { - kCommentStart, - kCommentEnd, - kLineComment, - kComment, - kWhitespace, - kNewline, - kStringLiteral, - kIdentifier, - kLocal, - kGlobal, - kOp, - kGraph, - kOpenParen, - kCloseParen, - kAtSymbol, - kPercent, - kComma, - kPeriod, - kEqual, - kSemicolon, - kColon, - kInteger, - kFloat, - kDivision, - kBoolean, - kPlus, - kStar, - kMinus, - kRAngle, - kLAngle, - kRCurly, - kLCurly, - kRSquare, - kLSquare, - kBang, - kAt, - kQuestion, - kIf, - kElse, - kUnderscore, - kLet, - kFn, - kDefn, - kTypeDef, - kExtern, - kMatch, - kPartialMatch, - kMetadata, - kMetaReference, - kFreeVar, - kRef, - kRefRead, - kRefWrite, - kVersion, - kUnknown, - kEndOfFile, - kNull, -}; - -inline std::string ToString(const TokenType& token_type) { - switch (token_type) { - case TokenType::kCommentStart: - return "CommentStart"; - case TokenType::kCommentEnd: - return "CommentEnd"; - case TokenType::kLineComment: - return "LineComment"; - case TokenType::kComment: - return "Comment"; - case TokenType::kWhitespace: - return "WhiteSpace"; - case TokenType::kNewline: - return "Newline"; - case TokenType::kStringLiteral: - return "StringLiteral"; - case TokenType::kIdentifier: - return "Identifier"; - case TokenType::kLocal: - return "Local"; - case TokenType::kGlobal: - return "Global"; - case TokenType::kGraph: - return "Graph"; - case TokenType::kOp: - return "Op"; - case TokenType::kOpenParen: - return "OpenParen"; - case TokenType::kCloseParen: - return "CloseParen"; - case TokenType::kAtSymbol: - return "AtSymbol"; - case TokenType::kPercent: - return "Percent"; - case TokenType::kComma: - return "Comma"; - case TokenType::kColon: - return "Colon"; - case TokenType::kSemicolon: - return "Semicolon"; - case TokenType::kPeriod: - return "Period"; - case TokenType::kEqual: - return "Equal"; - case TokenType::kInteger: - return "Integer"; - case TokenType::kFloat: - return "Float"; - case TokenType::kPlus: - return "Plus"; - case TokenType::kStar: - return "Star"; - case TokenType::kMinus: - return "Minus"; - case TokenType::kDivision: - return "Division"; - case TokenType::kRAngle: - return "RAngle"; - case TokenType::kLAngle: - return "LAngle"; - case TokenType::kRCurly: - return "RCurly"; - case TokenType::kLCurly: - return "LCurly"; - case TokenType::kRSquare: - return "RSquare"; - case TokenType::kLSquare: - return "LSquare"; - case TokenType::kBang: - return "Bang"; - case TokenType::kUnderscore: - return "Underscore"; - case TokenType::kAt: - return "At"; - case TokenType::kLet: - return "Let"; - case TokenType::kIf: - return "If"; - case TokenType::kElse: - return "Else"; - case TokenType::kFn: - return "Fn"; - case TokenType::kDefn: - return "Defn"; - case TokenType::kTypeDef: - return "TypeDef"; - case TokenType::kExtern: - return "Extern"; - case TokenType::kMatch: - return "Match"; - case TokenType::kPartialMatch: - return "PartialMatch"; - case TokenType::kQuestion: - return "Question"; - case TokenType::kBoolean: - return "Boolean"; - case TokenType::kMetadata: - return "Metadata"; - case TokenType::kMetaReference: - return "MetaReference"; - case TokenType::kFreeVar: - return "FreeVar"; - case TokenType::kVersion: - return "Version"; - case TokenType::kRef: - return "Ref"; - case TokenType::kRefRead: - return "RefRead"; - case TokenType::kRefWrite: - return "RefWrite"; - case TokenType::kUnknown: - return "Unknown"; - case TokenType::kEndOfFile: - return "EndOfFile"; - case TokenType::kNull: - return "Null"; - // Older compilers warn even though the above code is exhaustive. - default: - LOG(FATAL) << "unreachable code"; - } -} - -inline std::string Pretty(const TokenType& token_type) { - switch (token_type) { - case TokenType::kCommentStart: - return "`/*`"; - case TokenType::kCommentEnd: - return "`*/`"; - case TokenType::kLineComment: - return "`//`"; - case TokenType::kComment: - return "comment"; - case TokenType::kWhitespace: - return "whitespace"; - case TokenType::kNewline: - return "newline"; - case TokenType::kStringLiteral: - return "string literal"; - case TokenType::kIdentifier: - return "identifier"; - case TokenType::kLocal: - return "local variable"; - case TokenType::kGlobal: - return "global variable"; - case TokenType::kGraph: - return "graph variable"; - case TokenType::kOp: - return "operator"; - case TokenType::kOpenParen: - return "`(`"; - case TokenType::kCloseParen: - return "`)`"; - case TokenType::kAtSymbol: - return "`@`"; - case TokenType::kPercent: - return "`%`"; - case TokenType::kComma: - return "`,`"; - case TokenType::kColon: - return "`:`"; - case TokenType::kSemicolon: - return "`;`"; - case TokenType::kPeriod: - return "`.`"; - case TokenType::kEqual: - return "`=`"; - case TokenType::kInteger: - return "integer"; - case TokenType::kFloat: - return "float"; - case TokenType::kPlus: - return "`+`"; - case TokenType::kStar: - return "`*`"; - case TokenType::kMinus: - return "`-`"; - case TokenType::kDivision: - return "`/`"; - case TokenType::kRAngle: - return "`<`"; - case TokenType::kLAngle: - return "`>`"; - case TokenType::kRCurly: - return "`}`"; - case TokenType::kLCurly: - return "`{`"; - case TokenType::kRSquare: - return "`]`"; - case TokenType::kLSquare: - return "`[`"; - case TokenType::kBang: - return "`!`"; - case TokenType::kUnderscore: - return "`_`"; - case TokenType::kAt: - return "`@`"; - case TokenType::kLet: - return "`let`"; - case TokenType::kIf: - return "`if`"; - case TokenType::kElse: - return "`else`"; - case TokenType::kFn: - return "`fn`"; - case TokenType::kDefn: - return "`def`"; - case TokenType::kTypeDef: - return "`type`"; - case TokenType::kExtern: - return "`extern`"; - case TokenType::kBoolean: - return "boolean"; - case TokenType::kMetadata: - return "metadata section"; - case TokenType::kMetaReference: - return "`meta`"; - case TokenType::kFreeVar: - return "`free_var`"; - case TokenType::kMatch: - return "`match`"; - case TokenType::kPartialMatch: - return "`match?`"; - case TokenType::kQuestion: - return "`?`"; - case TokenType::kRef: - return "`ref`"; - case TokenType::kRefRead: - return "`ref_read`"; - case TokenType::kRefWrite: - return "`ref_write`"; - case TokenType::kUnknown: - return "unknown"; - case TokenType::kEndOfFile: - return "end of file"; - case TokenType::kNull: - return "null"; - case TokenType::kVersion: - return "version attribute"; - // Older compilers warn even though the above code is exhaustive. - default: - LOG(FATAL) << "unreachable code"; - } -} - -class Token; - -class TokenNode : public Object { - public: - Span span; - TokenType token_type; - mutable runtime::ObjectRef data; - - void VisitAttrs(AttrVisitor* v) {} - - static constexpr const char* _type_key = "parser.Token"; - TVM_DECLARE_FINAL_OBJECT_INFO(TokenNode, Object); -}; - -TVM_STATIC_IR_FUNCTOR(ReprPrinter, vtable) - .set_dispatch([](const ObjectRef& ref, ReprPrinter* p) { - auto* node = static_cast(ref.get()); - p->stream << "Token(span=" << node->span << ", token_type=" << ToString(node->token_type) - << ", data=" << node->data << ")"; - }); - -TVM_REGISTER_NODE_TYPE(TokenNode); - -class Token : public ObjectRef { - public: - TVM_DLL explicit Token(Span span, TokenType token_type, ObjectRef data = ObjectRef()); - - static Token Null(); - int64_t ToNumber() const; - std::string ToString() const; - Map> ToMetadata() const; - TVM_DEFINE_OBJECT_REF_METHODS(Token, ObjectRef, TokenNode); -}; - -inline Token::Token(Span span, TokenType token_type, ObjectRef data) { - ObjectPtr n = make_object(); - n->span = span; - n->token_type = token_type; - n->data = data; - data_ = std::move(n); -} - -inline Token Token::Null() { return Token(Span(SourceName(), 0, 0, 0, 0), TokenType::kNull); } - -inline int64_t Token::ToNumber() const { - return Downcast(this->operator->()->data).IntValue(); -} - -inline std::string Token::ToString() const { - return Downcast(this->operator->()->data); -} - -inline Map> Token::ToMetadata() const { - ObjectRef data = this->operator->()->data; - if (data.defined()) { - return Downcast>>(data); - } else { - return Map>({}); - } -} - -} // namespace relay -} // namespace tvm -#endif // TVM_RELAY_PARSER_TOKEN_H_ diff --git a/src/relay/parser/tokenizer.h b/src/relay/parser/tokenizer.h deleted file mode 100644 index 2b7ad4e5593e..000000000000 --- a/src/relay/parser/tokenizer.h +++ /dev/null @@ -1,696 +0,0 @@ -/* - * 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 tokenizer.h - * \brief A parser for TVM IR. - */ -#ifndef TVM_RELAY_PARSER_TOKENIZER_H_ -#define TVM_RELAY_PARSER_TOKENIZER_H_ - -#include -#include - -#include -#include -#include -#include -#include -#include - -#include "../../support/scalars.h" -#include "./meta_ref.h" -#include "./token.h" - -namespace tvm { -namespace relay { - -// trim from start (in place) -static inline void ltrim(std::string& s) { // NOLINT(*) - s.erase(s.begin(), std::find_if(s.begin(), s.end(), [](int ch) { return !std::isspace(ch); })); -} - -// trim from end (in place) -static inline void rtrim(std::string& s) { // NOLINT(*) - s.erase(std::find_if(s.rbegin(), s.rend(), [](int ch) { return !std::isspace(ch); }).base(), - s.end()); -} - -inline bool IsDigit(char c) { return '0' <= c && c <= '9'; } - -inline bool IsWhitespace(char c) { return ' ' == c || c == '\t' || c == '\n'; } - -inline bool IsNumeric(char c) { - return (IsDigit(c) || c == '.' || c == 'e' || c == '-' || c == '+' || c == 'E') && - !IsWhitespace(c); -} - -inline bool IsIdentLetter(char c) { - return '_' == c || c == '/' || ('a' <= c && c <= 'z') || ('A' <= c && c <= 'Z'); -} - -inline bool IsIdent(char c) { return IsIdentLetter(c) || IsDigit(c); } - -static std::unordered_map KEYWORD_TABLE = { - {"let", TokenType::kLet}, {"fn", TokenType::kFn}, - {"def", TokenType::kDefn}, {"if", TokenType::kIf}, - {"else", TokenType::kElse}, {"type", TokenType::kTypeDef}, - {"match", TokenType::kMatch}, {"extern", TokenType::kExtern}, - {"free_var", TokenType::kFreeVar}, {"ref", TokenType::kRef}, - {"ref_read", TokenType::kRefRead}, {"ref_write", TokenType::kRefWrite}}; - -struct Tokenizer { - DiagnosticContext diag_ctx; - const SourceName& source_name; - - size_t pos; - int col; - int line; - char next_char; - String source; - std::vector tokens; - - char Next() { - char c = this->source.at(this->pos); - if (c == '\n') { - this->line += 1; - this->col = 1; - } else { - this->col += 1; - } - pos += 1; - return c; - } - - bool More() { return this->pos < this->source.size(); } - - char Peek() { - ICHECK(pos < this->source.size()); - return this->source.at(this->pos); - } - - Token NewToken(TokenType token_type, ObjectRef data = ObjectRef(), int lines = 0, int cols = 1) { - auto span = - Span(this->source_name, this->line, this->line + lines, this->col, this->col + cols); - return Token(span, token_type, data); - } - - Span SpanFrom(int line, int column) { - int end_line = this->line; - int end_column = this->col; - return Span(this->source_name, line, end_line, column, end_column); - } - - enum CommentParserState { - Proceed, - Forward, - Backward, - }; - - void MatchComment(std::string* buffer) { - // We only invoke this after we have matched the first start - // token assume, we are proceeding the parse forward with - // nesting = 1. - // - // When we are done we should be at nesting zero and be - // in the stop state. - CommentParserState state = CommentParserState::Proceed; - int nesting = 1; - - while (More()) { - switch (state) { - case CommentParserState::Proceed: { - if (Peek() == '/') { - state = CommentParserState::Forward; - } else if (Peek() == '*') { - state = CommentParserState::Backward; - } - buffer->operator+=(Next()); - continue; - } - case CommentParserState::Forward: { - if (Peek() == '*') { - nesting += 1; - buffer->operator+=(Next()); - } - state = CommentParserState::Proceed; - continue; - } - case CommentParserState::Backward: { - if (Peek() == '/') { - nesting -= 1; - if (nesting == 0) { - Next(); - buffer->pop_back(); - return; - } - } - - buffer->operator+=(Next()); - state = CommentParserState::Proceed; - continue; - } - } - } - } - - Token ParseNumber(bool is_pos, bool is_float, std::string number) { - ICHECK(number.size() > 0) << "an empty string is an invalid number"; - - Token token = NewToken(is_float ? TokenType::kFloat : TokenType::kInteger); - size_t suffix_pos = number.rfind(is_float ? 'f' : 'i'); - if (suffix_pos == std::string::npos) { - suffix_pos = number.size(); - } - std::string literal_text = number.substr(0, suffix_pos); - std::string suffix; - if (suffix_pos < number.size()) { - suffix = number.substr(suffix_pos + 1, number.size() - suffix_pos); - } - int width = 32; - - if (suffix.size()) { - try { - width = std::stoi(suffix); - } catch (const std::invalid_argument& err) { - this->diag_ctx.Emit(Diagnostic::Error(token->span) - << "invalid numeric suffix `" << suffix << "`"); - } catch (const std::out_of_range& err) { - this->diag_ctx.Emit(Diagnostic::Error(token->span) - << "invalid numeric suffix `" << suffix << "`"); - } - } - - if (is_float) { - double value = 0.0; - size_t index = 0; - try { - value = stod(literal_text, &index); - } catch (const std::invalid_argument& err) { - this->diag_ctx.Emit(Diagnostic::Error(token->span) - << "invalid floating point number `" << literal_text << "`"); - } catch (const std::out_of_range& err) { - this->diag_ctx.Emit(Diagnostic::Error(token->span) - << "invalid floating point number `" << literal_text << "`"); - } - if (index < literal_text.size()) { - this->diag_ctx.Emit(Diagnostic::Error(token->span) - << "invalid floating point number `" << literal_text << "`"); - } - value = is_pos ? value : -value; - token->data = support::ValueToFloatImm(value, width); - if (!token->data.defined()) { - this->diag_ctx.Emit(Diagnostic::Error(token->span) - << "floating point number `" << literal_text - << "` unrepresentable in width " << width); - token->data = support::ValueToFloatImm(0.0, width); - } - } else { - int64_t value = 0; - size_t index = 0; - try { - value = std::stoll(literal_text, &index); - } catch (const std::invalid_argument& err) { - this->diag_ctx.Emit(Diagnostic::Error(token->span) - << "invalid integer number `" << literal_text << "`"); - } catch (const std::out_of_range& err) { - this->diag_ctx.Emit(Diagnostic::Error(token->span) - << "invalid integer number `" << literal_text << "`"); - } - if (index < literal_text.size()) { - this->diag_ctx.Emit(Diagnostic::Error(token->span) - << "invalid integer number `" << literal_text << "`"); - } - value = is_pos ? value : -value; - token->data = support::ValueToIntImm(value, width); - if (!token->data.defined() && suffix.empty()) { - // Without any i suffix the legacy behavior was to default to int64 if out of range - // for int32. - width = 64; - token->data = support::ValueToIntImm(value, width); - } - if (!token->data.defined()) { - this->diag_ctx.Emit(Diagnostic::Error(token->span) - << "integer number `" << literal_text << "` unrepresentable in width " - << width); - token->data = support::ValueToIntImm(0, width); - } - } - - return token; - } - - Token ParseNumber(bool is_pos) { - std::stringstream ss; - while (More() && IsNumeric(Peek())) { - ss << Next(); - } - - bool is_float = false; - if (More() && (Peek() == 'f' || Peek() == 'i')) { - is_float = Peek() == 'f'; - // Capture trailing width suffix - ss << Next(); - while (More() && IsNumeric(Peek())) { - ss << Next(); - } - } - return ParseNumber(is_pos, is_float, ss.str()); - } - - bool MatchString(const std::string& string) { - int start = this->pos; - - for (auto c : string) { - if (Peek() != c) { - this->pos = start; - return false; - } else { - Next(); - } - } - - return true; - } - - Token TokenizeMetaRef() { - int line = this->line; - int column = this->col; - - std::stringstream type_key; - while (More() && Peek() != ']') { - type_key << Next(); - } - ICHECK_EQ(Peek(), ']'); - Next(); - - ICHECK_EQ(Peek(), '['); - Next(); - std::stringstream str_index; - while (More() && Peek() != ']') { - str_index << Next(); - } - ICHECK_EQ(Peek(), ']'); - Next(); - // todo: add error handling around bad indices - auto index = ParseNumber(true, false, str_index.str()).ToNumber(); - auto span = SpanFrom(line, column); - return Token(span, TokenType::kMetaReference, MetaRef(type_key.str(), index)); - } - - Token TokenizeAttr() { - int line = this->line; - int column = this->col; - Next(); - if (Peek() == '[') { - Next(); - std::stringstream raw_attribute; - - while (More() && Peek() != ']') { - raw_attribute << Next(); - } - - ICHECK_EQ(Next(), ']'); - - auto attribute = raw_attribute.str(); - // Clean up the white-space on both sides. - ltrim(attribute); - rtrim(attribute); - - // Metadata can only appear at the bottom of a file and goes to EOF. - if (attribute == "metadata") { - std::stringstream metadata; - while (More()) { - metadata << Next(); - } - ObjectRef metadata_map = tvm::LoadJSON(metadata.str()); - auto span = SpanFrom(line, column); - return Token(span, TokenType::kMetadata, metadata_map); - } - if (attribute.rfind("version", 0) == 0) { - std::string version = attribute.substr(attribute.find("=") + 1); - ltrim(version); - rtrim(version); - auto span = SpanFrom(line, column); - return Token(span, TokenType::kVersion, tvm::String(version)); - } else { - // TOOD(@jroesch): maybe make this a warning an continue parsing? - auto span = SpanFrom(line, column); - this->diag_ctx.EmitFatal(Diagnostic::Error(span) << "unsupported attribute " << attribute); - return Token(); - } - } else { - auto span = SpanFrom(line, column); - this->diag_ctx - .EmitFatal(Diagnostic::Error(span) - << "`#` denotes the start of an attribute can only be followed by `[`" - << " found `" << Peek() << "`"); - return Token(); - } - } - - inline Token TokenizeOnce() { - int line = this->line; - int col = this->col; - auto next = Peek(); - VLOG(9) << "tvm::relay::TokenizeOnce: next=" << next; - if (next == '\n') { - auto token = NewToken(TokenType::kNewline); - Next(); - return token; - } else if (next == '\r') { - Next(); - if (More() && Peek() == '\n') { - auto token = NewToken(TokenType::kNewline); - return token; - } else { - auto span = SpanFrom(line, col); - this->diag_ctx.EmitFatal( - Diagnostic::Error(span) - << "\\r carriage returns must be followed by a \\n in the TVM text format"); - return Token(); - } - } else if (next == '"') { - // TODO(@jroesch): Properly tokenize escape sequences in strings. - // see https://github.com/apache/tvm/issues/6153. - Next(); - std::stringstream string_content; - while (More() && Peek() != '"') { - string_content << Next(); - } - Next(); - return NewToken(TokenType::kStringLiteral, tvm::String(string_content.str())); - } else if (IsWhitespace(next)) { - auto token = NewToken(TokenType::kWhitespace); - Next(); - return token; - } else if (next == '-') { - int negs = 0; - while (More() && Peek() == '-') { - Next(); - negs++; - } - bool is_neg = negs % 2 == 1; - if (More() && IsDigit(Peek())) { - return ParseNumber(!is_neg); - } else if (More() && MatchString("inff")) { - return ParseNumber(!is_neg, true, "inff"); - } else { - // If there isn't a number right after either, - // this is really slow for lexing, should replace - // with multi-token return or something. - pos = pos - (negs - 1); - return NewToken(TokenType::kMinus); - } - } else if (IsDigit(next)) { - return ParseNumber(true); - } else if (MatchString("inff")) { - return ParseNumber(true, true, "inff"); - } else if (next == '.') { - auto token = NewToken(TokenType::kPeriod); - Next(); - return token; - } else if (next == ',') { - auto token = NewToken(TokenType::kComma); - Next(); - return token; - } else if (next == '=') { - auto token = NewToken(TokenType::kEqual); - Next(); - return token; - } else if (next == ';') { - auto token = NewToken(TokenType::kSemicolon); - Next(); - return token; - } else if (next == ':') { - auto token = NewToken(TokenType::kColon); - Next(); - return token; - } else if (next == '(') { - auto token = NewToken(TokenType::kOpenParen); - Next(); - return token; - } else if (next == ')') { - auto token = NewToken(TokenType::kCloseParen); - Next(); - return token; - } else if (next == '+') { - auto token = NewToken(TokenType::kPlus); - Next(); - return token; - } else if (next == '*') { - auto token = NewToken(TokenType::kStar); - Next(); - return token; - } else if (next == '<') { - auto token = NewToken(TokenType::kLAngle); - Next(); - return token; - } else if (next == '>') { - auto token = NewToken(TokenType::kRAngle); - Next(); - return token; - } else if (next == '{') { - auto token = NewToken(TokenType::kLCurly); - Next(); - return token; - } else if (next == '}') { - auto token = NewToken(TokenType::kRCurly); - Next(); - return token; - } else if (next == '[') { - auto token = NewToken(TokenType::kLSquare); - Next(); - return token; - } else if (next == ']') { - auto token = NewToken(TokenType::kRSquare); - Next(); - return token; - } else if (next == '!') { - auto token = NewToken(TokenType::kBang); - Next(); - return token; - } else if (next == '@') { - auto token = NewToken(TokenType::kAt); - Next(); - return token; - } else if (next == '?') { - auto token = NewToken(TokenType::kQuestion); - Next(); - return token; - } else if (MatchString("meta[")) { - return TokenizeMetaRef(); - } else if (next == '#') { - return TokenizeAttr(); - } else if (next == '%') { - auto token = NewToken(TokenType::kPercent); - Next(); - - std::stringstream number; - while (More() && IsDigit(Peek())) { - number << Next(); - } - - auto number_str = number.str(); - if (number_str.size()) { - auto num_tok = ParseNumber(true, false, number_str); - auto span = SpanFrom(token->span->line, token->span->column); - token = Token(span, TokenType::kGraph, num_tok->data); - } - - return token; - } else if (next == '/') { - Next(); - if (Peek() == '/') { - auto token = NewToken(TokenType::kLineComment); - // Consume the / - Next(); - std::stringstream comment; - while (More() && Peek() != '\n') { - comment << Next(); - } - token->data = tvm::String(comment.str()); - return token; - } else if (Peek() == '*') { - // Eat the first /* pair before entering the state machine. - Next(); - std::string comment; - MatchComment(&comment); - auto token = NewToken(TokenType::kComment, tvm::String(comment)); - return token; - } else { - return NewToken(TokenType::kDivision); - } - } else if (IsIdentLetter(next)) { - std::stringstream ss; - // Due the below code we need to patch - // the line/col info to the start of - // token. - int line = this->line; - int col = this->col; - - while (More() && IsIdent(Peek())) { - ss << Next(); - } - - std::string keyword = ss.str(); - auto it = KEYWORD_TABLE.find(keyword); - - TokenType token_type; - if (it != KEYWORD_TABLE.end()) { - token_type = it->second; - - if (token_type == TokenType::kMatch) { - if (More() && Peek() == '?') { - Next(); - token_type = TokenType::kPartialMatch; - } - } - } else { - token_type = TokenType::kIdentifier; - } - - auto span = SpanFrom(line, col); - return Token(span, token_type, tvm::String(ss.str())); - } else { - std::stringstream ss; - while (More() && !IsWhitespace(Peek())) { - ss << Next(); - } - auto token = NewToken(TokenType::kUnknown); - token->data = tvm::String(ss.str()); - return token; - } - } - - void Tokenize() { - VLOG(9) << "tvm::relay::Tokenize"; - while (this->More()) { - auto token = TokenizeOnce(); - ICHECK(token.defined()); - this->tokens.push_back(token); - } - this->tokens.push_back(NewToken(TokenType::kEndOfFile)); - } - - explicit Tokenizer(const DiagnosticContext& ctx, const Source& source) - : diag_ctx(ctx), - source_name(source->source_name), - pos(0), - col(1), - line(1), - source(source->source), - tokens() {} -}; - -inline std::vector Condense(const std::vector& tokens, Token* table) { - std::vector out; - bool found_metadata = false; - - for (size_t i = 0; i < tokens.size(); i++) { - auto current = tokens.at(i); - switch (current->token_type) { - case TokenType::kMetadata: { - if (!found_metadata) { - found_metadata = true; - *table = current; - } else { - LOG(FATAL) << "duplicate metadata section"; - } - continue; - } - case TokenType::kPercent: { - auto next = tokens.at(i + 1); - if (next->token_type == TokenType::kIdentifier) { - // Match this token. - i += 1; - // TODO(@jroesch): merge spans - auto tok = Token(current->span, TokenType::kLocal, next->data); - ICHECK(tok.defined()); - out.push_back(tok); - } else if (next->token_type == TokenType::kInteger) { - i += 1; - auto tok = Token(current->span, TokenType::kGraph, next->data); - ICHECK(tok.defined()); - out.push_back(tok); - } else { - ICHECK(current.defined()); - out.push_back(current); - } - continue; - } - case TokenType::kAt: { - auto next = tokens.at(i + 1); - if (next->token_type == TokenType::kIdentifier) { - // Match this token. - i += 1; - // TODO(@jroesch): merge spans - auto tok = Token(current->span, TokenType::kGlobal, next->data); - ICHECK(tok.defined()); - out.push_back(tok); - } else { - ICHECK(current.defined()); - out.push_back(current); - } - continue; - } - case TokenType::kIdentifier: { - std::string str = Downcast(current->data); - Token tok; - // TODO(@jroesch): merge spans - if (str == "True") { - auto data = tvm::Integer(1); - tok = Token(current->span, TokenType::kBoolean, data); - } else if (str == "False") { - auto data = tvm::Integer(0); - tok = Token(current->span, TokenType::kBoolean, data); - } else if (str == "_") { - tok = Token(current->span, TokenType::kUnderscore); - } else { - tok = current; - } - out.push_back(tok); - continue; - } - default: { - out.push_back(current); - continue; - } - } - } - - return out; -} - -inline std::pair, Token> Tokenize(const DiagnosticContext& ctx, - const Source& source) { - auto tokenizer = Tokenizer(ctx, source); - tokenizer.Tokenize(); - Token meta_table(Span(), TokenType::kUnknown, ObjectRef()); - auto tokens = Condense(tokenizer.tokens, &meta_table); - for (auto token : tokens) { - ICHECK(token.defined()); - } - return {tokens, meta_table}; -} - -} // namespace relay -} // namespace tvm - -#endif // TVM_RELAY_PARSER_TOKENIZER_H_ diff --git a/src/relay/printer/doc.cc b/src/relay/printer/doc.cc deleted file mode 100644 index 79313c9a587f..000000000000 --- a/src/relay/printer/doc.cc +++ /dev/null @@ -1,162 +0,0 @@ -/* - * 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 src/relay/doc.cc - * \brief Doc ADT used for pretty printing. - * - * Reference: Philip Wadler. A Prettier Printer. Journal of Functional Programming'98 - */ -#include "doc.h" - -#include - -#include -#include - -#include "../../support/str_escape.h" - -namespace tvm { -namespace relay { - -/*! - * \brief Represent a piece of text in the doc. - */ -class DocTextNode : public DocAtomNode { - public: - /*! \brief The str content in the text. */ - std::string str; - - explicit DocTextNode(std::string str_val) : str(str_val) {} - - static constexpr const char* _type_key = "printer.DocText"; - TVM_DECLARE_FINAL_OBJECT_INFO(DocTextNode, DocAtomNode); -}; - -TVM_REGISTER_OBJECT_TYPE(DocTextNode); - -class DocText : public DocAtom { - public: - explicit DocText(std::string str) { data_ = runtime::make_object(str); } - - TVM_DEFINE_OBJECT_REF_METHODS(DocText, DocAtom, DocTextNode); -}; - -/*! - * \brief Represent a line breaker in the doc. - */ -class DocLineNode : public DocAtomNode { - public: - /*! \brief The amount of indent in newline. */ - int indent; - - explicit DocLineNode(int indent) : indent(indent) {} - - static constexpr const char* _type_key = "printer.DocLine"; - TVM_DECLARE_FINAL_OBJECT_INFO(DocLineNode, DocAtomNode); -}; - -TVM_REGISTER_OBJECT_TYPE(DocLineNode); - -class DocLine : public DocAtom { - public: - explicit DocLine(int indent) { data_ = runtime::make_object(indent); } - - TVM_DEFINE_OBJECT_REF_METHODS(DocLine, DocAtom, DocLineNode); -}; - -// DSL function implementations -Doc& Doc::operator<<(const Doc& right) { - ICHECK(this != &right); - this->stream_.insert(this->stream_.end(), right.stream_.begin(), right.stream_.end()); - return *this; -} - -Doc& Doc::operator<<(std::string right) { return *this << DocText(right); } - -Doc& Doc::operator<<(const DocAtom& right) { - this->stream_.push_back(right); - return *this; -} - -std::string Doc::str() { - std::ostringstream os; - for (auto atom : this->stream_) { - if (auto* text = atom.as()) { - os << text->str; - } else if (auto* line = atom.as()) { - os << "\n" << std::string(line->indent, ' '); - } else { - LOG(FATAL) << "do not expect type " << atom->GetTypeKey(); - } - } - return os.str(); -} - -Doc Doc::NewLine(int indent) { return Doc() << DocLine(indent); } - -Doc Doc::Text(std::string text) { return Doc() << DocText(text); } - -Doc Doc::RawText(std::string text) { - return Doc() << DocAtom(runtime::make_object(text)); -} - -Doc Doc::Indent(int indent, Doc doc) { - for (size_t i = 0; i < doc.stream_.size(); ++i) { - if (auto* line = doc.stream_[i].as()) { - doc.stream_[i] = DocLine(indent + line->indent); - } - } - return doc; -} - -Doc Doc::StrLiteral(const std::string& value, std::string quote) { - Doc doc; - return doc << quote << support::StrEscape(value) << quote; -} - -Doc Doc::PyBoolLiteral(bool value) { - if (value) { - return Doc::Text("True"); - } else { - return Doc::Text("False"); - } -} - -Doc Doc::Brace(std::string open, const Doc& body, std::string close, int indent) { - Doc doc; - doc << open; - doc << Indent(indent, NewLine() << body) << NewLine(); - doc << close; - return doc; -} - -Doc Doc::Concat(const std::vector& vec, const Doc& sep) { - Doc seq; - if (vec.size() != 0) { - if (vec.size() == 1) return vec[0]; - seq << vec[0]; - for (size_t i = 1; i < vec.size(); ++i) { - seq << sep << vec[i]; - } - } - return seq; -} -} // namespace relay -} // namespace tvm diff --git a/src/relay/printer/doc.h b/src/relay/printer/doc.h deleted file mode 100644 index 36f26d9bd24b..000000000000 --- a/src/relay/printer/doc.h +++ /dev/null @@ -1,168 +0,0 @@ -/* - * 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 tvm/printer/doc.h - * \brief Doc ADT used for pretty printing. - * - * Reference: Philip Wadler. A Prettier Printer. Journal of Functional Programming'98 - */ -#ifndef TVM_RELAY_PRINTER_DOC_H_ -#define TVM_RELAY_PRINTER_DOC_H_ - -#include -#include -#include - -#include -#include -#include - -namespace tvm { -namespace relay { - -/*! - * \brief Doc atom node for the ADT. - * \sa DocAtom - */ -class DocAtomNode : public Object { - public: - static constexpr const char* _type_key = "printer.DocAtom"; - TVM_DECLARE_BASE_OBJECT_INFO(DocAtomNode, Object); -}; - -/*! - * \brief Managed reference to DocAtomNode. - * \sa DocAtomNode. - */ -class DocAtom : public ObjectRef { - public: - TVM_DEFINE_OBJECT_REF_METHODS(DocAtom, ObjectRef, DocAtomNode); -}; - -/*! - * \brief Stream-like interface for Doc DSL. - * - * The Doc DSL de-couples the layout decision from the printing decision. - * - * The layout(code formating) decisions include: - * - Change indentation. - * - Break single line into multiple ones(subjected to future improvements). - */ -class Doc { - public: - /*! \brief default constructor */ - Doc() {} - /*! - * \brief Append right to the end of the current doc stream. - * \param right The doc to be appended. - * \return reference to self. - */ - Doc& operator<<(const Doc& right); - /*! - * \brief Append right to the end of the current doc stream. - * \param right The doc to be appended. - * \return reference to self. - * \note pass by value to allow copy elison optimization. - */ - Doc& operator<<(std::string right); - /*! - * \brief Append right to the end of the current doc stream. - * \param right The doc to be appended. - * \return reference to self. - */ - Doc& operator<<(const DocAtom& right); - /*! - * \brief Convert value to string via std::ostreamstream - * the append to the current doc stream. - * \param right The doc to be appended. - * \tparam T the type of the value. - * \return reference to self. - */ - template ::value>::type> - Doc& operator<<(const T& value) { - std::ostringstream os; - os << value; - return *this << os.str(); - } - /*! - * \brief Convert the doc stream into string. - * \return The string representation. - */ - std::string str(); - /*! - * \brief Create a doc that represents text content. - * \return The created doc. - */ - static Doc Text(std::string value); - /*! - * \brief Create a doc that represents raw text(can have new lines) - * \return The created doc. - */ - static Doc RawText(std::string value); - /*! - * \brief Create a doc that represents a new line. - * \return The created doc. - */ - static Doc NewLine(int indent = 0); - /*! - * \brief Create a new doc that adds indentation to everyline of the doc. - * \param indent The indent to be added. - * \param doc The doc to be indented. - * \return The created doc. - * \note pass by value to allow copy elison optimization. - */ - static Doc Indent(int indent, Doc doc); - /*! - * \brief Create a Doc that represents a string literal. - * \param value The content of the string literal. - * \param quote The quote in the literal. - * \return The created doc. - */ - static Doc StrLiteral(const std::string& value, std::string quote = "\""); - /*! - * \brief Create a Doc that represents a boolean literal in python syntax. - * \param value The bool value. - * \return The created doc. - */ - static Doc PyBoolLiteral(bool value); - /*! - * \brief Enclose body by brace and add indent. - * \param body The body - * \param open The open brace. - * \param close The close brace. - * \param indent amount of indentation. - * \return The created doc. - */ - static Doc Brace(std::string open, const Doc& body, std::string close, int indent = 2); - /*! - * \brief Create a doc by concatenating together with separator. - * \param vec The docs to be concatenated. - * \param sep The seperator. - * \return The created doc. - */ - static Doc Concat(const std::vector& vec, const Doc& sep = Text(", ")); - - private: - /*! \brief Internal doc stream. */ - std::vector stream_; -}; -} // namespace relay -} // namespace tvm -#endif // TVM_RELAY_PRINTER_DOC_H_ diff --git a/src/relay/printer/meta_data.h b/src/relay/printer/meta_data.h deleted file mode 100644 index 2dfd594de7eb..000000000000 --- a/src/relay/printer/meta_data.h +++ /dev/null @@ -1,141 +0,0 @@ -/* - * 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. - */ -#ifndef TVM_RELAY_PRINTER_META_DATA_H_ -#define TVM_RELAY_PRINTER_META_DATA_H_ - -#include - -#include -#include - -#include "doc.h" - -namespace tvm { -namespace relay { -/*! - * \brief Meta data context for Printers - * - * This is an important part to enable bi-directional serializability. - * We use tvm's Node system to build the current IR. - * It can be hard to design a text format for all the possible nodes - * as the set of nodes can grow when we do more extensions. - * - * Instead of trying to design readable text format for every node, - * we support a meta data section in the text format. - * We allow the text format to refer to a node in the meta data section. - * - * The meta data section is a json serialized string of an Map>. - * Each element in the meta data section can be referenced by the text format. - * Each meta data node is printed in the following format. - * - * meta[type-key-of-node>][] - * - * Specifically, consider the following IR(constructed by python). - * - * \code - * - * n = tvm.var("n") - * x = tvm.relay.var("x", shape=(n, 1)) - * f = tvm.relay.Function([x], x) - * print(f.astext()) - * - * \endcode - * - * The corresponding text format is shown in the following code block. - * - * \code - * - * fn (%x: Tensor[(meta[Variable][0],), float32]) { - * %x - * } - * # Meta data section is a json-serialized string - * # of the following array. - * # [tvm.var("n")] - * - * \endcode - * - * Note that we store tvm.var("n") in the meta data section. - * Since it is stored in the index-0 in the meta data section, - * we print it as meta[Variable][0]. - * - * The text parser can recover this object by loading from the corresponding - * location in the meta data section. - * - * This is a design trade-off. - * It allows us to embedded any meta data in the text format, - * while still being able to tweak the text part of the printed IR easily. - */ -class TextMetaDataContext { - public: - /*! - * \brief Get text representation of meta node. - * \param node The node to be converted to meta node. - * \return A string representation of the meta node. - */ - Doc GetMetaNode(const ObjectRef& node) { - auto it = meta_repr_.find(node); - if (it != meta_repr_.end()) { - return it->second; - } - std::string type_key = node->GetTypeKey(); - ICHECK(!type_key.empty()); - Array& mvector = meta_data_[type_key]; - int64_t index = static_cast(mvector.size()); - mvector.push_back(node); - Doc doc; - doc << "meta[" << type_key << "][" << index << "]"; - meta_repr_[node] = doc; - return meta_repr_[node]; - } - - /*! - * \brief Test whether a node has been put in meta - * \param node The query node - * \return whether the node has been put in meta - */ - bool InMeta(const ObjectRef& node) { return meta_repr_.find(node) != meta_repr_.end(); } - - /*! - * \brief Print a key value pair - */ - Doc PrintKeyValue(const std::string& str, const Doc& v) const { - return Doc() << "\"" << str << "\": " << v; - } - - /*! - * \brief Get the metadata section in json format. - * \return the meta data string. - */ - Doc GetMetaSection() const { - if (meta_data_.size() == 0) return Doc(); - return Doc::RawText(SaveJSON(Map(meta_data_.begin(), meta_data_.end()))); - } - - /*! \return whether the meta data context is empty. */ - bool empty() const { return meta_data_.empty(); } - - private: - /*! \brief additional metadata stored in TVM json format */ - std::unordered_map> meta_data_; - /*! \brief map from meta data into its string representation */ - std::unordered_map meta_repr_; -}; -} // namespace relay -} // namespace tvm -#endif // TVM_RELAY_PRINTER_META_DATA_H_ diff --git a/src/relay/printer/model_library_format_printer.cc b/src/relay/printer/model_library_format_printer.cc deleted file mode 100644 index aab70910f644..000000000000 --- a/src/relay/printer/model_library_format_printer.cc +++ /dev/null @@ -1,84 +0,0 @@ -/* - * 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 "text_printer.h" - -namespace tvm { -namespace relay { - -class ModelLibraryFormatPrinter : public ::tvm::runtime::ModuleNode { - public: - ModelLibraryFormatPrinter(bool show_meta_data, - const runtime::TypedPackedFunc& annotate, - bool show_warning) - : text_printer_{show_meta_data, annotate, show_warning} {} - - const char* type_key() const final { return "model_library_format_printer"; } - - /*! \brief Get the property of the runtime module .*/ - int GetPropertyMask() const final { return runtime::ModulePropertyMask::kRunnable; } - - std::string Print(const ObjectRef& node) { - std::ostringstream oss; - oss << node; - return oss.str(); - } - - TVMRetValue GetVarName(tir::Var var) { - TVMRetValue rv; - std::string var_name; - if (text_printer_.GetVarName(var, &var_name)) { - rv = var_name; - } - - return rv; - } - - PackedFunc GetFunction(const String& name, const ObjectPtr& sptr_to_self) override { - if (name == "print") { - return TypedPackedFunc( - [sptr_to_self, this](ObjectRef node) { return Print(node); }); - } else if (name == "get_var_name") { - return TypedPackedFunc( - [sptr_to_self, this](tir::Var var) { return GetVarName(var); }); - } else { - return PackedFunc(); - } - } - - private: - TextPrinter text_printer_; -}; - -TVM_REGISTER_GLOBAL("relay.ir.ModelLibraryFormatPrinter") - .set_body_typed([](bool show_meta_data, - const runtime::TypedPackedFunc& annotate, - bool show_warning) { - return ObjectRef( - make_object(show_meta_data, annotate, show_warning)); - }); - -} // namespace relay -} // namespace tvm diff --git a/src/relay/printer/relay_text_printer.cc b/src/relay/printer/relay_text_printer.cc deleted file mode 100644 index 618e8fe138d8..000000000000 --- a/src/relay/printer/relay_text_printer.cc +++ /dev/null @@ -1,974 +0,0 @@ -/* - * 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 relay_text_printer.cc - * \brief Printer to print out the IR text format - * that can be parsed by a parser. - * - * Supports ANF, GNF in relay and metadata. - * - * Inlining heuristics: - * - Always inline: - * - GlobalVar - * - Constant - * - Op - * - Var - * - Otherwise, inline if the node is at the end of a scope and is used at most once. - */ -#include -#include -#include -#include -#include -#include -#include - -#include "../../ir/attr_functor.h" -#include "../../support/scalars.h" -#include "../analysis/dependency_graph.h" -#include "../parser/meta_ref.h" -#include "doc.h" -#include "meta_data.h" -#include "text_printer.h" -#include "tvm/runtime/builtin_fp16.h" - -namespace tvm { -namespace relay { - -/*! - * \brief Print additional info about expr in comment. - * \param expr The expression. - */ -Doc RelayTextPrinter::PrintOptionalInfo(const Expr& expr) { - Doc doc; - if (!opt_info_memo_.insert(expr).second) { - return doc; - } - // default annotations - if (annotate_ == nullptr) { - if ((expr.as() || expr.as() || expr.as() || - expr.as() || expr.as() || expr.as()) && - (expr->checked_type_.defined() || expr->span.defined())) { - doc << " /*"; - if (expr->checked_type_.defined()) { - doc << " ty=" << Print(expr->checked_type()); - } - if (expr->span.defined()) { - doc << " span=" << PrintSpan(expr->span); - } - doc << " */"; - } - } else { - std::string annotated_expr = annotate_(expr); - if (annotated_expr != "") { - doc << annotated_expr; - } - } - return doc; -} - -// indent a new body -Doc RelayTextPrinter::PrintBody(const ObjectRef& node, int indent) { - Doc doc; - Doc body; - doc << "{"; - doc << Doc::Indent(indent, body << Doc::NewLine() << PrintScope(node)) << Doc::NewLine(); - doc << "}"; - return doc; -} - -// create a new scope by creating a new printer object. This allows temp var -// numbers to be reused and prevents hoisted vars from escaping too far -Doc RelayTextPrinter::PrintScope(const ObjectRef& node) { - // print in a new scope - doc_stack_.push_back(Doc()); - // must print first so doc_stack_.back() reference doesn't become stale - Doc doc = Print(node, false, true); - doc = doc_stack_.back() << doc; - doc_stack_.pop_back(); - return doc; -} - -Doc RelayTextPrinter::PrintFinal(const ObjectRef& node) { - if (node.defined() && node->IsInstance() && - !node->IsInstance()) { - // Temporarily skip non-relay functions. - // TODO(tvm-team) enhance the code to work for all functions - } else if (node.as()) { - Expr expr = Downcast(node); - dg_ = DependencyGraph::Create(&arena_, expr); - } - - Doc doc; - doc << PrintScope(node); - return doc; -} - -Doc RelayTextPrinter::Print(const ObjectRef& node, bool meta, bool try_inline) { - bool is_non_relay_func = node.defined() && node->IsInstance() && - !node->IsInstance(); - if (node.as() && !is_non_relay_func) { - return PrintExpr(Downcast(node), meta, try_inline); - } else if (node.as()) { - return PrintType(Downcast(node), meta); - } else if (node.as()) { - return PrintPattern(Downcast(node), meta); - } else if (node.as()) { - return PrintMod(Downcast(node)); - } else { - // default module. - std::ostringstream os; - os << node; - return Doc::RawText(os.str()); - } -} - -Doc RelayTextPrinter::TempVar(int n) { - Doc doc; - return doc << "%" << n; -} - -Doc RelayTextPrinter::AllocTemp() { return TempVar(temp_var_counter_++); } - -/*! - * \brief get a unique name with the corresponding prefix - * \param prefix The prefix of the name - * \return The returned name. - */ -Doc RelayTextPrinter::GetUniqueName(const std::string& prefix) { - std::string unique_prefix = prefix; - auto it = name_alloc_map_.find(prefix); - if (it != name_alloc_map_.end()) { - while (true) { - std::ostringstream os; - os << prefix << (++it->second); - std::string name = os.str(); - if (name_alloc_map_.count(name) == 0) { - unique_prefix = name; - break; - } - } - } - name_alloc_map_[unique_prefix] = 0; - return Doc::Text(unique_prefix); -} - -Doc RelayTextPrinter::Print(Kind k) { - switch (k) { - case kType: - return Doc::Text("Type"); - case kShapeVar: - return Doc::Text("Shape"); - case kBaseType: - return Doc::Text("BaseType"); - case kConstraint: - return Doc::Text("Constraint"); - case kAdtHandle: - return Doc::Text("AdtHandle"); - case kTypeData: - return Doc::Text("TypeData"); - default: - LOG(ERROR) << "Unknown Kind"; - throw; - } -} -/*! - * \brief Allocate name to a type variable. - * \param var The input type variable. - * \return The corresponding name. - */ -Doc RelayTextPrinter::AllocTypeVar(const TypeVar& var) { - if (memo_type_.count(var)) { - Doc val = memo_type_[var]; - val << "-malformed-ir"; - return val; - } - std::string name = var->name_hint; - if (name.length() == 0 || !std::isalpha(name[0])) { - name = "t" + name; - } - Doc val = GetUniqueName(name); - memo_type_[var] = val; - if (var->kind != kType) { - val << ": " << Print(var->kind); - } - return val; -} - -/*! - * \brief Allocate name to a variable. - * \param var The input variable. - * \return The corresponding name. - */ -Doc RelayTextPrinter::AllocVar(const Var& var) { - // still print if ir is malformed, but show the error. - if (memo_.count(var)) { - Doc val = memo_[var]; - val << "-malformed-ir"; - return val; - } - std::string name = var->name_hint(); - // always make sure first name is alpha - if (name.length() == 0 || !std::isalpha(name[0])) { - name = "v" + name; - } - Doc val = GetUniqueName("%" + name); - memo_[var] = val; // Referential occurrences will not include the following. - if (!var->virtual_device()->IsFullyUnconstrained()) { - val << " {" << kVirtualDevice << "=" << PrintAttributeValue(var->virtual_device()) << "}"; - } - if (var->type_annotation.defined()) { - val << ": " << Print(var->type_annotation); - } - - val << PrintOptionalInfo(var); - return val; -} - -bool RelayTextPrinter::IsUnique(const Expr& expr) { - auto it = dg_.expr_node.find(expr); - if (it == dg_.expr_node.end()) { - return true; - } else { - return !(it->second->parents.head && it->second->parents.head->next); - } -} - -bool RelayTextPrinter::AlwaysInline(const Expr& expr) { - return expr.as() || expr.as() || expr.as() || - expr.as() || expr.as(); -} - -Doc RelayTextPrinter::VisitLeaf(const Expr& expr) { - if (!CheckVisited(expr)) { - Doc result = ExprFunctor::VisitExpr(expr); - // Add if not added after visiting - if (!CheckVisited(expr)) { - memo_[expr] = result; - } else { - result_memo_[expr] = result; - } - return result; - } - return memo_[expr]; -} - -bool RelayTextPrinter::CheckVisited(const Expr& expr) { return (memo_.count(expr)); } - -Doc RelayTextPrinter::VisitExpr(const Expr& expr) { - auto fcheck_visited = [this](const Expr& expr) { return this->CheckVisited(expr); }; - auto fvisit_leaf = [this](const Expr& expr) { return this->VisitLeaf(expr); }; - - if (fcheck_visited(expr)) { - return memo_[expr]; - } else { - ExpandDataflow(expr, fcheck_visited, fvisit_leaf); - return memo_[expr]; - } -} - -//------------------------------------ -// Overload of Expr printing functions -//------------------------------------ -Doc RelayTextPrinter::PrintExpr(const Expr& expr, bool meta, bool try_inline, bool optional_info) { - // Exploit memoization to print GNF. - // The first time we visit an expression, we need to allocate a temp var - // for it. Every subsequent time we can just use its assigned variable. - // This works since hashing uses pointer equality. - - // determine whether to inline - bool inline_expr = AlwaysInline(expr); - - if (try_inline) { - inline_expr |= IsUnique(expr); - } - - Doc printed_expr; - - if (meta) { - printed_expr = meta_->GetMetaNode(GetRef(expr.get())); - } else if (!inline_expr && expr.as()) { - // wrap GNFed let in brackets - Doc body; - printed_expr << "("; - printed_expr << Doc::Indent(2, body << Doc::NewLine() << VisitExpr(expr)) << Doc::NewLine(); - printed_expr << ")"; - } else { - printed_expr = VisitExpr(expr); - } - - if (optional_info) { - printed_expr << PrintOptionalInfo(expr); - } - - // add expr to doc - if (expr.as()) { - // This is our first time visiting the var and we hit the VarNode case - // in the visitor. Thus the variable is free. - if (var_memo_.insert(expr).second && result_memo_.count(expr)) { - doc_stack_.back() << "free_var " << result_memo_[expr] << ";" << Doc::NewLine(); - } - // Memoization is done in AllocVar. - return memo_[expr]; - } else if (inline_expr) { - memo_[expr] = printed_expr; - return printed_expr; - } else { - // Already exists. Reuse - if (!var_memo_.insert(expr).second) { - return memo_[expr]; - } - Doc temp_var = AllocTemp(); - memo_[expr] = temp_var; - doc_stack_.back() << temp_var << " = " << printed_expr << ";" << Doc::NewLine(); - return temp_var; - } -} - -// Should only be triggered when op is a free variable being visited for the -// first time. -Doc RelayTextPrinter::VisitExpr_(const VarNode* op) { return AllocVar(GetRef(op)); } - -Doc RelayTextPrinter::VisitExpr_(const ConstantNode* op) { - // Print out simple scalars directly. - if (support::IsSimpleScalar(op)) { - return Doc::Text(support::NDArrayScalarToString(op->data)); - } - // Fallbock: record it as a meta node. - Doc doc; - // Don't append optional_info. Because the entry function is Print, - // and it will append the optional_info afterwards. - return doc << PrintExpr(GetRef(op), /*meta=*/true, /*try_inline=*/false, - /*optional_info=*/false); -} - -Doc RelayTextPrinter::VisitExpr_(const TupleNode* op) { - std::vector fields; - for (Expr field : op->fields) { - fields.push_back(Print(field)); - } - Doc doc; - doc << "(" << Doc::Concat(fields); - // conform to python tuple format (1,) - if (op->fields.size() == 1) { - doc << ","; - } - return doc << ")"; -} - -Doc RelayTextPrinter::VisitExpr_(const TupleGetItemNode* op) { - Doc doc; - return doc << Print(op->tuple) << "." << op->index; -} - -Doc RelayTextPrinter::VisitExpr_(const IfNode* op) { - Doc doc; - doc << "if (" << Print(op->cond) << ") "; - doc << PrintBody(op->true_branch); - doc << " else "; - doc << PrintBody(op->false_branch); - return doc; -} - -Doc RelayTextPrinter::VisitExpr_(const LetNode* op) { - int n = 0; - size_t l = doc_stack_.size(); - Expr let = GetRef(op); - while (auto let_node = let.as()) { - Doc doc; - doc << "let " << AllocVar(let_node->var) << " = " << Print(let_node->value, false, true) << ";" - << Doc::NewLine(); - doc_stack_.push_back(doc); - let = let_node->body; - ++n; - } - Doc doc = PrintScope(let); - Doc doc_last; - for (int i = 0; i < n; ++i) { - doc_last << doc_stack_[l + i]; - } - doc_last << doc; - for (int i = 0; i < n; ++i) { - doc_stack_.pop_back(); - } - return doc_last; -} - -Doc RelayTextPrinter::PrintFunc(const Doc& prefix, const relay::Function& fn) { - Doc doc; - doc << prefix; - if (fn->type_params.size() > 0) { - doc << "["; - std::vector type_params; - for (const TypeVar& tv : fn->type_params) { - type_params.push_back(Doc::Text(tv->name_hint)); - } - doc << Doc::Concat(type_params); - doc << "]"; - } - doc << "("; - std::vector params; - for (Var param : fn->params) { - params.push_back(AllocVar(param)); - } - for (const Doc& d : PrintDictAttrs(fn->attrs)) { - params.push_back(d); - } - if (!fn->virtual_device()->IsFullyUnconstrained()) { - Doc vid_doc; - vid_doc << kVirtualDevice << "=" << PrintAttributeValue(fn->virtual_device()); - params.push_back(vid_doc); - } - doc << Doc::Concat(params) << ") "; - if (fn->ret_type.defined()) { - doc << "-> " << Print(fn->ret_type) << " "; - } - doc << PrintBody(fn->body); - return doc; -} - -Doc RelayTextPrinter::PrintFunc(const Doc& prefix, const BaseFunc& base_func) { - if (auto func = base_func.as()) { - return PrintFunc(prefix, func.value()); - } else if (auto func = base_func.as()) { - std::ostringstream os; - os << func.value(); - return Doc::RawText(os.str()); - } else { - // def @xyz = meta['ExternalFunc'][id] - Doc doc; - doc << prefix << " = " << meta_->GetMetaNode(base_func); - return doc; - } -} - -Doc RelayTextPrinter::PrintMod(const IRModule& mod) { - Doc doc; - int counter = 0; - // type definitions - for (const auto& kv : mod->type_definitions) { - if (counter++ != 0) { - doc << Doc::NewLine(); - } - doc << Print(kv.second); - doc << Doc::NewLine(); - } - // functions - for (const auto& kv : mod->functions) { - if (kv.second.as()) { - dg_ = DependencyGraph::Create(&arena_, kv.second); - } - if (counter++ != 0) { - doc << Doc::NewLine(); - } - std::ostringstream os; - os << "def @" << kv.first->name_hint; - doc << PrintFunc(Doc::Text(os.str()), kv.second); - doc << Doc::NewLine(); - } - return doc; -} - -Doc RelayTextPrinter::VisitExpr_(const FunctionNode* op) { - return PrintFunc(Doc::Text("fn "), GetRef(op)); -} - -Doc RelayTextPrinter::VisitExpr_(const GlobalVarNode* op) { - Doc doc; - doc << "@" << op->name_hint; - return doc; -} - -Doc RelayTextPrinter::VisitExpr_(const OpNode* op) { return Doc::Text(op->name); } - -Doc RelayTextPrinter::VisitExpr_(const CallNode* op) { - Doc doc; - // visit args first so they are lifted before the op - // this places op closer to its call site - std::vector args; - for (const Expr& arg : op->args) { - args.push_back(Print(arg)); - } - - for (const Doc& d : PrintCallAttrs(op->attrs, op->op)) { - args.push_back(d); - } - const auto* cons_node = op->op.as(); - if (cons_node) { - doc << cons_node->name_hint; - } else { - doc << Print(op->op); - } - - if (cons_node && cons_node->inputs.size() == 0) { - // don't print as a call if it's a 0-arity cons - return doc; - } else { - doc << "(" << Doc::Concat(args) << ")"; - return doc; - } -} - -Doc RelayTextPrinter::VisitExpr_(const RefCreateNode* op) { - Doc doc; - return doc << "ref(" << Print(op->value) << ")"; -} - -Doc RelayTextPrinter::VisitExpr_(const RefReadNode* op) { - Doc doc; - return doc << "ref_read(" << Print(op->ref) << ")"; -} - -Doc RelayTextPrinter::VisitExpr_(const RefWriteNode* op) { - Doc doc; - return doc << "ref_write(" << Print(op->ref) << ", " << Print(op->value) << ")"; -} - -Doc RelayTextPrinter::VisitExpr_(const MatchNode* op) { - // TODO(jmp): Lots of code duplication here because PrintBody and PrintScope don't accept Docs. - Doc doc; - Doc body; - doc << "match"; - if (!op->complete) { - doc << "?"; - } - doc << " (" << Print(op->data) << ") {"; - std::vector clause_docs; - for (const auto& clause : op->clauses) { - Doc clause_doc; - clause_doc << PrintPattern(clause->lhs, false) << " => "; - Doc rhs_doc = PrintScope(clause->rhs); - // TODO(@jroesch): This is unsound right now, and we need to revisit it. - // if (clause->rhs.as()) { - // only add braces if there are multiple lines on the rhs - rhs_doc = Doc::Brace("{", rhs_doc, "}"); - // } - clause_doc << rhs_doc << ","; - clause_docs.push_back(clause_doc); - } - doc << Doc::Indent(2, body << Doc::NewLine() << Doc::Concat(clause_docs, Doc::NewLine())) - << Doc::NewLine() << "}"; - return doc; -} - -Doc RelayTextPrinter::PrintPattern(const Pattern& pattern, bool meta) { - auto it = memo_pattern_.find(pattern); - if (it != memo_pattern_.end()) return it->second; - Doc printed_pattern; - if (meta) { - printed_pattern = meta_->GetMetaNode(GetRef(pattern.get())); - } else { - printed_pattern = VisitPattern(pattern); - } - memo_pattern_[pattern] = printed_pattern; - return printed_pattern; -} - -Doc RelayTextPrinter::VisitPattern_(const PatternConstructorNode* p) { - Doc doc; - doc << p->constructor->name_hint; - if (!p->patterns.empty()) { - doc << "("; - std::vector pats; - for (const auto& pat : p->patterns) { - pats.push_back(Print(pat)); - } - doc << Doc::Concat(pats) << ")"; - } - return doc; -} - -Doc RelayTextPrinter::VisitPattern_(const PatternTupleNode* pt) { - Doc doc; - doc << "("; - std::vector pats; - for (const auto& pat : pt->patterns) { - pats.push_back(Print(pat)); - } - doc << Doc::Concat(pats) << ")"; - return doc; -} - -Doc RelayTextPrinter::VisitPattern_(const PatternWildcardNode* pw) { return Doc::Text("_"); } - -Doc RelayTextPrinter::VisitPattern_(const PatternVarNode* pv) { return AllocVar(pv->var); } - -Doc RelayTextPrinter::VisitExpr_(const ConstructorNode* n) { - Doc doc; - doc << n->name_hint; - if (in_adt_def_ && n->inputs.size() != 0) { - doc << "("; - std::vector inputs; - for (Type input : n->inputs) { - inputs.push_back(Print(input)); - } - doc << Doc::Concat(inputs) << ")"; - } - return doc; -} - -//------------------------------------ -// Overload of Type printing functions -//------------------------------------ -Doc RelayTextPrinter::PrintType(const Type& type, bool meta) { - auto it = memo_type_.find(type); - if (it != memo_type_.end()) return it->second; - Doc printed_type; - if (meta) { - printed_type = meta_->GetMetaNode(GetRef(type.get())); - } else { - printed_type = VisitType(type); - } - memo_type_[type] = printed_type; - return printed_type; -} - -Doc RelayTextPrinter::VisitTypeDefault_(const Object* node) { - // by default always print as meta data - return Print(GetRef(node), true); -} - -Doc RelayTextPrinter::VisitType_(const TypeVarNode* node) { return Doc::Text(node->name_hint); } - -Doc RelayTextPrinter::VisitType_(const GlobalTypeVarNode* node) { - return Doc::Text(node->name_hint); -} - -Doc RelayTextPrinter::VisitType_(const TypeCallNode* node) { - Doc doc = PrintType(node->func, false); - std::vector args; - for (const Type& t : node->args) { - args.push_back(PrintType(t, false)); - } - doc << "["; - doc << Doc::Concat(args); - doc << "]"; - return doc; -} - -Doc RelayTextPrinter::PrintDType(DataType dtype) { - return Doc::Text(runtime::DLDataType2String(dtype)); -} - -Doc RelayTextPrinter::VisitType_(const TensorTypeNode* node) { - // scalar type - if (node->shape.size() == 0) { - return PrintDType(node->dtype); - } - Doc doc; - doc << "Tensor[("; - std::vector shapes; - for (const PrimExpr& prim_expr : node->shape) { - // Though not bound within an attribute the attribute visitor will handle the PrimExprs we - // care about. - shapes.push_back(PrintAttributeValue(prim_expr)); - } - doc << Doc::Concat(shapes); - return doc << "), " << PrintDType(node->dtype) << "]"; -} - -Doc RelayTextPrinter::VisitType_(const TupleTypeNode* node) { - std::vector fields; - for (Type field : node->fields) { - fields.push_back(Print(field)); - } - Doc doc; - doc << "(" << Doc::Concat(fields); - // conform to python tuple format (1,) - if (node->fields.size() == 1) { - doc << ","; - } - return doc << ")"; -} - -Doc RelayTextPrinter::VisitType_(const FuncTypeNode* node) { - Doc doc; - doc << "fn "; - if (node->type_params.size() != 0) { - doc << "["; - std::vector type_params; - for (Type type_param : node->type_params) { - type_params.push_back(Print(type_param)); - } - doc << Doc::Concat(type_params); - doc << "]"; - } - std::vector arg_types; - for (Type arg_type : node->arg_types) { - arg_types.push_back(Print(arg_type)); - } - return doc << "(" << Doc::Concat(arg_types) << ") -> " << Print(node->ret_type); -} - -Doc RelayTextPrinter::VisitType_(const RelayRefTypeNode* node) { - Doc doc; - return doc << "ref(" << Print(node->value) << ")"; -} - -Doc RelayTextPrinter::VisitType_(const TypeDataNode* node) { - in_adt_def_ = true; - Doc doc; - doc << "type " << Print(node->header); - - // type vars - if (node->type_vars.size() != 0) { - doc << "["; - std::vector type_vars; - for (Type type_var : node->type_vars) { - type_vars.push_back(Print(type_var)); - } - doc << Doc::Concat(type_vars) << "]"; - } - doc << " "; - - std::vector constructor_docs; - for (Constructor constructor : node->constructors) { - constructor_docs.push_back(Print(constructor, /* meta */ false, /* try_inline */ true)); - } - Doc separator; - separator << "," << Doc::NewLine(); - Doc adt_body; - adt_body << Doc::Concat(constructor_docs, separator); - // add trailing comma if there are any constructors - if (!constructor_docs.empty()) { - adt_body << ","; - } - doc << Doc::Brace("{", adt_body, "}"); - in_adt_def_ = false; - return doc; -} - -//------------------------------------ -// Overload of Attr printing functions -//------------------------------------ - -Doc RelayTextPrinter::VisitAttrDefault_(const Object* op) { - // Since we don't have any overload for a specific attribute type we'll need to force - // the meta[...] representation to avoid infinite regress. - return PrintAttributeValue(GetRef(op), /*force_meta=*/true); -} - -Doc RelayTextPrinter::VisitAttr_(const ArrayNode* op) { - Doc doc; - doc << "["; - std::vector arr_vals; - for (const auto& val : *op) { - arr_vals.push_back(PrintAttributeValue(val)); - } - doc << Doc::Concat(arr_vals); - doc << "]"; - return doc; -} - -Doc RelayTextPrinter::VisitAttr_(const tir::IntImmNode* op) { - if (support::IsSimpleScalarDtype(op->dtype)) { - return Doc::Text(support::IntImmToString(GetRef(op))); - } else { - // Fallback: Print int64_t without width suffix. - return Doc::Text(std::to_string(op->value)); - } -} - -Doc RelayTextPrinter::VisitAttr_(const tir::FloatImmNode* op) { - if (support::IsSimpleScalarDtype(op->dtype)) { - return Doc::Text(support::FloatImmToString(GetRef(op))); - } else { - // Fallbock: Print double without width suffix. - return Doc::Text(std::to_string(op->value)); - } -} - -Doc RelayTextPrinter::VisitAttr_(const tir::StringImmNode* op) { - return Doc::StrLiteral(op->value); -} - -/*! - * \brief Attribute printer which prints the attributes in the call. - */ -class RelayTextPrinter::AttrPrinter : public AttrVisitor { - public: - AttrPrinter(std::vector* doc, RelayTextPrinter* parent) : docs(doc), parent_(parent) {} - - template - void PrintKV(const char* key, const T& value) { - Doc doc; - doc << key << "=" << value; - docs->push_back(doc); - } - - void Visit(const char* key, double* value) final { - Doc doc; - doc << key << "=" << *value << "f"; - docs->push_back(doc); - } - - void Visit(const char* key, int64_t* value) final { PrintKV(key, *value); } - void Visit(const char* key, uint64_t* value) final { PrintKV(key, *value); } - void Visit(const char* key, int* value) final { PrintKV(key, *value); } - void Visit(const char* key, bool* value) final { PrintKV(key, Doc::PyBoolLiteral(*value)); } - void Visit(const char* key, std::string* value) final { PrintKV(key, Doc::StrLiteral(*value)); } - void Visit(const char* key, void** value) final { LOG(FATAL) << "do not allow void as argument"; } - void Visit(const char* key, DataType* value) final { - PrintKV(key, Doc::StrLiteral(runtime::DLDataType2String(*value))); - } - void Visit(const char* key, runtime::NDArray* value) final { - LOG(FATAL) << "do not allow NDarray as argument"; - } - void Visit(const char* key, runtime::ObjectRef* obj) final { - PrintKV(key, parent_->PrintAttributeValue(*obj)); - } - - private: - std::vector* docs; - RelayTextPrinter* parent_; -}; - -void RelayTextPrinter::AppendGenericAttrs(std::vector* docs, const Attrs& attrs, - bool include_type_key) { - if (!attrs.defined()) { - return; - } - AttrPrinter printer(docs, this); - // Need to drop cost cast since in general VisitNonDefaultAttrs can mutate, but in this - // case we are read-only. - const_cast(attrs.get())->VisitNonDefaultAttrs(&printer); - if (include_type_key) { - std::string s = attrs->GetTypeKey(); - printer.Visit("attrs_type_key", &s); - } -} - -std::vector RelayTextPrinter::PrintCallAttrs(const Attrs& attrs, const Expr& op) { - std::vector docs; - if (!attrs.defined()) { - return docs; - } - const auto* op_node = op.as(); - if (show_meta_data_ && op_node && (attrs->type_index() != op_node->attrs_type_index)) { - // The parser can only understand calls with attributes if they match the operator's - // declared attribute type. If that's not the case fall back to the meta[...] representation. - docs.push_back(meta_->GetMetaNode(attrs)); - } else { - AppendGenericAttrs(&docs, attrs, /*include_type_key=*/!op_node); - } - return docs; -} - -std::vector RelayTextPrinter::PrintDictAttrs(const DictAttrs& dict_attrs) { - if (!dict_attrs.defined()) { - return {}; - } - return PrintDictAttrs(dict_attrs->dict); -} - -std::vector RelayTextPrinter::PrintDictAttrs(const Map& dict_attrs) { - std::vector docs; - if (!dict_attrs.defined()) { - return docs; - } - for (const auto& k : dict_attrs) { - Doc doc; - doc << k.first << "=" << PrintAttributeValue(k.second); - docs.push_back(doc); - } - return docs; -} - -Doc RelayTextPrinter::PrintAttributeValue(const ObjectRef& value, bool force_meta) { - if (value.defined()) { - Doc printed_attr; - if (value.as()) { - printed_attr << "?"; - } else if (auto str_obj = value.as()) { - printed_attr << Doc::StrLiteral(str_obj.value()); - } else if (force_meta) { - printed_attr = meta_->GetMetaNode(Downcast(value)); - } else if (auto virtual_device_node = value.as()) { - if (show_meta_data_) { - printed_attr = meta_->GetMetaNode(virtual_device_node.value()); - } else { - // Special case: The ReprPrinter for VirtualDeviceNodes is much easier to work with while - // debugging. - std::ostringstream os; - os << virtual_device_node.value(); - return Doc::Text(os.str()); - } - } else if (const auto* base_attr_node = value.as()) { - if (show_meta_data_) { - printed_attr = meta_->GetMetaNode(GetRef(base_attr_node)); - } else { - // Special case: The non-meta form for attributes are much easier to work with while - // debugging. - printed_attr = PrintAttrsAsAttributeValue(GetRef(base_attr_node)); - } - } else if (const auto* base_map_node = value.as()) { - if (show_meta_data_) { - printed_attr = meta_->GetMetaNode(GetRef(base_map_node)); - } else { - // Special case: Show maps fields as key=value pairs to help debugging. - printed_attr << PrintMapAsAttributeValue(GetRef>(base_map_node)); - } - } else if (auto global_var = value.as()) { - if (show_meta_data_) { - printed_attr = meta_->GetMetaNode(global_var.value()); - } else { - printed_attr << "'" << global_var.value()->name_hint << "'"; - } - } else { - printed_attr = VisitAttr(value); - } - return printed_attr; - } else { - return Doc::Text("None"); - } -} - -Doc RelayTextPrinter::PrintAttrsAsAttributeValue(const Attrs& attrs) { - std::vector docs; - AppendGenericAttrs(&docs, attrs, /*include_type_key=*/false); - Doc doc; - doc << "{" << Doc::Concat(docs) << "}"; - return doc; -} - -Doc RelayTextPrinter::PrintMapAsAttributeValue(const Map& map) { - std::vector docs; - for (const auto& k : map) { - Doc doc; - doc << PrintAttributeValue(k.first); - doc << "="; - doc << PrintAttributeValue(k.second); - docs.push_back(doc); - } - Doc doc; - doc << "{" << Doc::Concat(docs) << "}"; - return doc; -} - -Doc RelayTextPrinter::PrintSpan(const Span& span) { - Doc doc; - const auto* span_node = span.as(); - ICHECK(span_node); - doc << span_node->source_name->name << ":" << span_node->line << ":" << span_node->column; - return doc; -} - -} // namespace relay -} // namespace tvm diff --git a/src/relay/printer/text_printer.cc b/src/relay/printer/text_printer.cc deleted file mode 100644 index f51f7c3dfa57..000000000000 --- a/src/relay/printer/text_printer.cc +++ /dev/null @@ -1,132 +0,0 @@ -/* - * 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 text_printer.cc - * \brief Printer to print out the unified IR text format - * that can be parsed by a parser. - */ - -#include "./text_printer.h" - -#include - -#include -#include - -namespace tvm { -namespace relay { - -static const char* kSemVer = "0.0.5"; - -Doc TextPrinter::PrintMod(const IRModule& mod) { - Doc doc; - int counter = 0; - - // We'll print in alphabetical order to make a/b diffs easier to work with. - - // type definitions - std::vector tyvars; - for (const auto& kv : mod->type_definitions) { - tyvars.emplace_back(kv.first); - } - std::sort(tyvars.begin(), tyvars.end(), - [](const GlobalTypeVar& left, const GlobalTypeVar& right) { - return left->name_hint < right->name_hint; - }); - for (const auto& tyvar : tyvars) { - if (counter++ != 0) { - doc << Doc::NewLine(); - } - doc << relay_text_printer_.Print(mod->type_definitions[tyvar]); - doc << Doc::NewLine(); - } - - // functions - std::vector vars; - for (const auto& kv : mod->functions) { - vars.emplace_back(kv.first); - } - std::sort(vars.begin(), vars.end(), [](const GlobalVar& left, const GlobalVar& right) { - return left->name_hint < right->name_hint; - }); - for (const auto& var : vars) { - const BaseFunc& base_func = mod->functions[var]; - if (base_func.as()) { - relay_text_printer_.dg_ = - relay::DependencyGraph::Create(&relay_text_printer_.arena_, base_func); - } - if (counter++ != 0) { - doc << Doc::NewLine(); - } - if (base_func.as()) { - std::ostringstream os; - os << "def @" << var->name_hint; - doc << relay_text_printer_.PrintFunc(Doc::Text(os.str()), base_func); - } else if (base_func.as()) { - doc << "@" << var->name_hint; - doc << " = " << tir_text_printer_.PrintPrimFunc(Downcast(base_func)); - } - doc << Doc::NewLine(); - } - -#if TVM_LOG_DEBUG - // attributes - // TODO(mbs): Make this official, including support from parser. - if (mod->attrs.defined() && !mod->attrs->dict.empty()) { - std::vector keys; - for (const auto& kv : mod->attrs->dict) { - keys.emplace_back(kv.first); - } - std::sort(keys.begin(), keys.end()); - doc << "attributes {" << Doc::NewLine(); - for (const auto& key : keys) { - doc << " '" << key << "' = " << PrettyPrint(mod->attrs->dict[key]) << Doc::NewLine(); - } - doc << "}" << Doc::NewLine(); - } -#endif - - return doc; -} - -String PrettyPrint(const ObjectRef& node) { - Doc doc; - doc << TextPrinter(/*show_meta_data=*/false, nullptr, false).PrintFinal(node); - return doc.str(); -} - -String AsText(const ObjectRef& node, bool show_meta_data, - runtime::TypedPackedFunc annotate) { - Doc doc; - doc << "#[version = \"" << kSemVer << "\"]" << Doc::NewLine(); - runtime::TypedPackedFunc ftyped = nullptr; - if (annotate != nullptr) { - ftyped = runtime::TypedPackedFunc( - [&annotate](const ObjectRef& expr) -> std::string { return annotate(expr); }); - } - doc << TextPrinter(show_meta_data, ftyped).PrintFinal(node); - return doc.str(); -} - -TVM_REGISTER_GLOBAL("relay.ir.PrettyPrint").set_body_typed(PrettyPrint); -TVM_REGISTER_GLOBAL("relay.ir.AsText").set_body_typed(AsText); - -} // namespace relay -} // namespace tvm diff --git a/src/relay/printer/text_printer.h b/src/relay/printer/text_printer.h deleted file mode 100644 index a6684bf4e5ce..000000000000 --- a/src/relay/printer/text_printer.h +++ /dev/null @@ -1,464 +0,0 @@ -/* - * 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 text_printer.h - * \brief Printer to print out the unified IR text format - * that can be parsed by a parser. - */ - -#ifndef TVM_RELAY_PRINTER_TEXT_PRINTER_H_ -#define TVM_RELAY_PRINTER_TEXT_PRINTER_H_ - -#include -#include -#include -#include -#include -#include -#include -#include -#include - -#include -#include -#include -#include - -#include "../../ir/attr_functor.h" -#include "../analysis/dependency_graph.h" -#include "doc.h" -#include "meta_data.h" - -namespace tvm { -namespace relay { - -class TextPrinter; - -class RelayTextPrinter : public ExprFunctor, - public PatternFunctor, - public TypeFunctor, - public AttrFunctor { - public: - explicit RelayTextPrinter(bool show_meta_data, TextMetaDataContext* meta, - runtime::TypedPackedFunc annotate) - : show_meta_data_(show_meta_data), annotate_(annotate), meta_(meta) {} - Doc VisitExpr(const Expr& expr) override; - virtual Doc VisitLeaf(const Expr& expr); - virtual bool CheckVisited(const Expr& expr); - - /*! - * \brief Print additional info about expr in comment. - * \param expr The expression. - */ - Doc PrintOptionalInfo(const Expr& expr); - // indent a new body - Doc PrintBody(const ObjectRef& node, int indent = 2); - // create a new scope by creating a new printer object. This allows temp var - // numbers to be reused and prevents hoisted vars from escaping too far - Doc PrintScope(const ObjectRef& node); - Doc PrintFinal(const ObjectRef& node); - - /*! - * \brief Returns \p attrs printed using the generic attribute visitor, as a sequence - * of key=value entries, if any. - */ - void AppendGenericAttrs(std::vector* docs, const Attrs& attrs, bool include_type_key); - - /*! - * \brief Returns \p attrs printed as a sequence of key=value entries, if any. - * This is used for call attributes. - */ - std::vector PrintCallAttrs(const Attrs& attrs, const Expr& op); - - /*! - * \brief Returns \p dict_attrs printed as a sequence of key=value entries, if any. - * This is used for function definition attributes. - */ - std::vector PrintDictAttrs(const DictAttrs& dict_attrs); - std::vector PrintDictAttrs(const Map& dict_attrs); - - /*! - * \brief Returns \p value printed as the rhs of an attribute key=value entry. If \p force_meta - * is true then value is printed in meta[...] for irrespective of the show_meta_data_ flag. - */ - Doc PrintAttributeValue(const ObjectRef& value, bool force_meta = false); - - /*! - * \brief Returns \p attrs printed as a self-contained value, ie wrapped in braces. - */ - Doc PrintAttrsAsAttributeValue(const Attrs& attrs); - - /*! - * \brief Returns \p map printed as a self-contained value, ie wrapped in braces. - */ - Doc PrintMapAsAttributeValue(const Map& map); - - Doc PrintSpan(const Span& span); - - Doc Print(const ObjectRef& node, bool meta = false, bool try_inline = false); - - Doc TempVar(int n); - Doc AllocTemp(); - /*! - * \brief get a unique name with the corresponding prefix - * \param prefix The prefix of the name - * \return The returned name. - */ - Doc GetUniqueName(const std::string& prefix); - Doc Print(Kind k); - /*! - * \brief Allocate name to a type variable. - * \param var The input type variable. - * \return The corresponding name. - */ - Doc AllocTypeVar(const TypeVar& var); - /*! - * \brief Allocate name to a variable. - * \param var The input variable. - * \return The corresponding name. - */ - Doc AllocVar(const Var& var); - bool IsUnique(const Expr& expr); - bool AlwaysInline(const Expr& expr); - - Doc PrintFunc(const Doc& prefix, const relay::Function& fn); - Doc PrintFunc(const Doc& prefix, const BaseFunc& base_func); - Doc PrintMod(const IRModule& mod); - - //------------------------------------ - // Overload of Expr printing functions - //------------------------------------ - Doc PrintExpr(const Expr& expr, bool meta, bool try_inline, bool optional_info = true); - // Should only be triggered when op is a free variable being visited for the - // first time. - Doc VisitExpr_(const VarNode* op) final; - Doc VisitExpr_(const ConstantNode* op) final; - Doc VisitExpr_(const TupleNode* op) final; - Doc VisitExpr_(const TupleGetItemNode* op) final; - Doc VisitExpr_(const IfNode* op) final; - Doc VisitExpr_(const LetNode* op) final; - Doc VisitExpr_(const FunctionNode* op) final; - Doc VisitExpr_(const GlobalVarNode* op) final; - Doc VisitExpr_(const OpNode* op) final; - Doc VisitExpr_(const CallNode* op) final; - Doc VisitExpr_(const RefCreateNode* op) final; - Doc VisitExpr_(const RefReadNode* op) final; - Doc VisitExpr_(const RefWriteNode* op) final; - Doc VisitExpr_(const MatchNode* op) final; - Doc PrintPattern(const Pattern& pattern, bool meta); - Doc VisitPattern_(const PatternConstructorNode* p) final; - Doc VisitPattern_(const PatternTupleNode* pt) final; - Doc VisitPattern_(const PatternWildcardNode* pw) final; - Doc VisitPattern_(const PatternVarNode* pv) final; - Doc VisitExpr_(const ConstructorNode* n) final; - //------------------------------------ - // Overload of Type printing functions - //------------------------------------ - Doc PrintType(const Type& type, bool meta); - Doc VisitTypeDefault_(const Object* node) final; - Doc VisitType_(const TypeVarNode* node) final; - Doc VisitType_(const GlobalTypeVarNode* node) final; - Doc VisitType_(const TypeCallNode* node) final; - Doc PrintDType(DataType dtype); - Doc VisitType_(const TensorTypeNode* node) final; - Doc VisitType_(const TupleTypeNode* node) final; - Doc VisitType_(const FuncTypeNode* node) final; - Doc VisitType_(const RelayRefTypeNode* node) final; - Doc VisitType_(const TypeDataNode* node) final; - //------------------------------------ - // Overload of Attr printing functions - //------------------------------------ - Doc VisitAttrDefault_(const Object* op) final; - Doc VisitAttr_(const ArrayNode* op) final; - Doc VisitAttr_(const tir::IntImmNode* op) final; - Doc VisitAttr_(const tir::FloatImmNode* op) final; - Doc VisitAttr_(const tir::StringImmNode* op) final; - - private: - /*! \brief Whether to print meta data. */ - bool show_meta_data_; - /*! \brief additional comment function */ - runtime::TypedPackedFunc annotate_; - /*! \brief Stack of docs to implement scoped GNFing. */ - std::vector doc_stack_{}; - /*! \brief Set for introduced vars */ - std::unordered_set var_memo_; - /*! \brief Set for exprs have been printed optional information */ - std::unordered_set opt_info_memo_; - /*! \brief Map for result and memo_ diffs for visited expression */ - std::unordered_map result_memo_; - /*! \brief Map from Expr to Doc */ - std::unordered_map memo_; - /*! \brief Map from Type to Doc */ - std::unordered_map memo_type_; - /*! \brief Map from Type to Doc */ - std::unordered_map memo_pattern_; - /*! \brief name allocation map */ - std::unordered_map name_alloc_map_; - /*! \brief meta data context */ - TextMetaDataContext* meta_; - /*! \brief counter of temporary variable */ - size_t temp_var_counter_{0}; - /*! \brief whether the printer is currently in an ADT definition */ - bool in_adt_def_; - /*! \brief arena for dependency graph */ - support::Arena arena_; - /*! \brief dependency graph of the expr */ - DependencyGraph dg_; - class AttrPrinter; - friend class AttrPrinter; - friend class tvm::relay::TextPrinter; -}; - -using namespace ::tvm::tir; - -/*! - * \brief Meta node collector - * If we decide to put some node into meta, then all the sub-nodes inside - * it need to be put in meta as well, since when parsing we need to know - * whether two refs are the same - */ -class MetaCollector : public StmtExprVisitor { - public: - explicit MetaCollector(TextMetaDataContext* meta) : meta_(meta) {} - - void Collect(const ObjectRef& n) { - // these nodes can be print directly(StringLiteral or use identifier to identify) - if (!n.defined() || n.as() || n.as() || n.as() || - n.as() || n.as() || n.as()) { - return; - } - if (n->IsInstance()) { - VisitStmt(Downcast(n)); - } else if (n->IsInstance()) { - VisitExpr(Downcast(n)); - } - } - - void VisitStmt(const Stmt& n) override { - meta_->GetMetaNode(n); - StmtVisitor::VisitStmt(n); - } - - void VisitExpr(const PrimExpr& n) override { - meta_->GetMetaNode(n); - ExprVisitor::VisitExpr(n); - } - - private: - TextMetaDataContext* meta_; -}; - -class TIRTextPrinter : public StmtFunctor, - public tir::ExprFunctor, - public TypeFunctor { - public: - explicit TIRTextPrinter(bool show_meta, TextMetaDataContext* meta) - : show_meta_(show_meta), meta_(meta), meta_collector_(meta) {} - - /*! \brief Output a newline */ - virtual Doc NewLine(); - - /*! \brief Print the node */ - Doc Print(const ObjectRef& node); - - /*! \brief Place into `s` the name used in the preceding Print call for `v`. - * \param v Var instance to check. Must point to a VarNode visited by Print. - * \param s String to receive the name. - * \return true when a name re-mapping was found. - */ - bool GetVarName(::tvm::tir::Var v, std::string* s); - - protected: - Doc VisitExpr_(const IntImmNode* op) override; - Doc VisitExpr_(const FloatImmNode* op) override; - Doc VisitExpr_(const StringImmNode* op) override; - Doc VisitExpr_(const CastNode* op) override; - Doc VisitExpr_(const tir::VarNode* op) override; - Doc VisitExpr_(const AddNode* op) override; - Doc VisitExpr_(const SubNode* op) override; - Doc VisitExpr_(const MulNode* op) override; - Doc VisitExpr_(const DivNode* op) override; - Doc VisitExpr_(const ModNode* op) override; - Doc VisitExpr_(const FloorDivNode* op) override; - Doc VisitExpr_(const FloorModNode* op) override; - Doc VisitExpr_(const MinNode* op) override; - Doc VisitExpr_(const MaxNode* op) override; - Doc VisitExpr_(const EQNode* op) override; - Doc VisitExpr_(const NENode* op) override; - Doc VisitExpr_(const LTNode* op) override; - Doc VisitExpr_(const LENode* op) override; - Doc VisitExpr_(const GTNode* op) override; - Doc VisitExpr_(const GENode* op) override; - Doc VisitExpr_(const AndNode* op) override; - Doc VisitExpr_(const OrNode* op) override; - Doc VisitExpr_(const NotNode* op) override; - Doc VisitExpr_(const SelectNode* op) override; - Doc VisitExpr_(const BufferLoadNode* op) override; - Doc VisitExpr_(const ProducerLoadNode* op) override; - Doc VisitExpr_(const RampNode* op) override; - Doc VisitExpr_(const BroadcastNode* op) override; - Doc VisitExpr_(const tir::LetNode* op) override; - Doc VisitExpr_(const tir::CallNode* op) override; - Doc VisitExpr_(const ShuffleNode* op) override; - Doc VisitExpr_(const ReduceNode* op) override; - Doc VisitExprDefault_(const Object* op) override; - - Doc VisitStmt_(const LetStmtNode* op) override; - Doc VisitStmt_(const AttrStmtNode* op) override; - Doc VisitStmt_(const AssertStmtNode* op) override; - Doc VisitStmt_(const BufferStoreNode* op) override; - Doc VisitStmt_(const ProducerStoreNode* op) override; - Doc VisitStmt_(const BufferRealizeNode* op) override; - Doc VisitStmt_(const ProducerRealizeNode* op) override; - Doc VisitStmt_(const AllocateNode* op) override; - Doc VisitStmt_(const AllocateConstNode* op) override; - Doc VisitStmt_(const DeclBufferNode* op) override; - Doc VisitStmt_(const IfThenElseNode* op) override; - Doc VisitStmt_(const SeqStmtNode* op) override; - Doc VisitStmt_(const EvaluateNode* op) override; - Doc VisitStmt_(const ForNode* op) override; - Doc VisitStmt_(const WhileNode* op) override; - Doc VisitStmt_(const PrefetchNode* op) override; - Doc VisitStmt_(const BlockRealizeNode* op) override; - Doc VisitStmtDefault_(const Object* op) override; - - private: - /*! \brief whether show meta data */ - bool show_meta_; - /*! \brief meta data context */ - TextMetaDataContext* meta_; - /*! \brief meta collector */ - MetaCollector meta_collector_; - /*! \brief Map from Var to Doc */ - std::unordered_map memo_var_; - /*! \brief Map from Buffer to Doc */ - std::unordered_map memo_buf_; - /*! \brief Map from Buffer to Doc */ - std::unordered_map memo_producer_; - /*! \brief name allocation map */ - std::unordered_map name_alloc_map_; - - friend class TextPrinter; - - Doc VisitType_(const PrimTypeNode* node) override; - Doc VisitType_(const PointerTypeNode* node) override; - Doc VisitType_(const TupleTypeNode* node) override; - - Doc PrintIRModule(const IRModule& module); - Doc PrintPrimFunc(const PrimFunc& primFunc); - Doc PrintArray(const ArrayNode* op); - Doc PrintIterVar(const IterVarNode* op); - Doc PrintRange(const RangeNode* op); - Doc PrintBuffer(const BufferNode* op); - Doc PrintProducer(const DataProducerNode* op); - Doc BufferNode2Doc(const BufferNode* op, Doc doc); - Doc DataProducerNode2Doc(const DataProducerNode* op, Doc doc); - Doc PrintString(const StringObj* op) { return Doc::StrLiteral(op->data); } - Doc PrintBufferRegion(const BufferRegionNode* op); - - /*! - * \brief special method to print out data type - * \param dtype The data type - */ - static Doc PrintDType(DataType dtype); - /*! - * \brief special method to print out const scalar - * \param dtype The data type - * \param data The pointer to hold the data. - */ - template - static Doc PrintConstScalar(DataType dtype, const T& data); - Doc GetUniqueName(std::string prefix); - Doc AllocVar(const tir::Var& var); - Doc AllocConst(const AllocateConst& var); - Doc AllocBuf(const Buffer& buffer); - Doc AllocProducer(const DataProducer& buffer); - /*! - * \brief special method to render vectors of docs with a separator - * \param vec vector of docs - * \param sep separator - */ - static Doc PrintSep(const std::vector& vec, const Doc& sep); - Doc PrintBody(const Stmt& body, bool indent = true); -}; - -String AsTVMScriptWithDiagnostic(const ObjectRef& mod, const String& tir_prefix, bool show_meta, - runtime::TypedPackedFunc annotate); - -class TextPrinter { - public: - explicit TextPrinter(bool show_meta_data, - const runtime::TypedPackedFunc& annotate, - bool show_warning = true) - : show_meta_data_(show_meta_data), - show_warning_(show_warning), - annotate_(annotate), - relay_text_printer_(show_meta_data, &meta_, annotate), - tir_text_printer_(show_meta_data, &meta_) {} - - /*! \brief whether show meta data */ - bool show_meta_data_; - - /*! \brief whether show the meta data warning message */ - bool show_warning_; - - /*! \brief meta data context */ - TextMetaDataContext meta_; - /*! \brief additional comment function */ - runtime::TypedPackedFunc annotate_; - /*! \brief Relay Text Printer */ - relay::RelayTextPrinter relay_text_printer_; - /*! \brief TIR Text Printer */ - TIRTextPrinter tir_text_printer_; - - bool GetVarName(::tvm::tir::Var v, std::string* s) { return tir_text_printer_.GetVarName(v, s); } - - Doc PrintFinal(const ObjectRef& node) { - Doc doc; - if (node.defined() && node->IsInstance()) { - doc << PrintMod(Downcast(node)); - } else if (node.defined() && - (node->IsInstance() || node->IsInstance() || - node->IsInstance())) { - doc << tir_text_printer_.Print(node); - } else { - doc << relay_text_printer_.PrintFinal(node); - } - if (!meta_.empty()) { - doc << Doc::NewLine(); - if (show_meta_data_) { - doc << "#[metadata]" << Doc::NewLine() << meta_.GetMetaSection(); - } else if (show_warning_) { - doc << "/* For debugging purposes the metadata section has been omitted." << Doc::NewLine() - << " * If you would like to see the full metadata section you can set the " - << Doc::NewLine() << " * option to `True` when invoking `astext`. " << Doc::NewLine() - << " */"; - } - } - return doc; - } - - Doc PrintMod(const IRModule& mod); -}; -} // namespace relay -} // namespace tvm - -#endif // TVM_RELAY_PRINTER_TEXT_PRINTER_H_ diff --git a/src/relay/printer/tir_text_printer.cc b/src/relay/printer/tir_text_printer.cc deleted file mode 100644 index c34788be91b8..000000000000 --- a/src/relay/printer/tir_text_printer.cc +++ /dev/null @@ -1,826 +0,0 @@ -/* - * 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 tir_text_printer.cc - * \brief Printer to print out the IR text format - * that can be parsed by a parser. - */ - -#include -#include -#include -#include -#include -#include -#include -#include -#include - -#include -#include - -#include "../../tir/transforms/ir_utils.h" -#include "doc.h" -#include "meta_data.h" -#include "text_printer.h" - -namespace tvm { -namespace relay { - -Doc TIRTextPrinter::Print(const ObjectRef& node) { - if (!node.defined()) return Doc::Text("(nullptr)"); - if (node->IsInstance()) { - return VisitStmt(Downcast(node)); - } else if (node->IsInstance()) { - return Doc::Text("?"); - } else if (node->IsInstance()) { - return VisitExpr(Downcast(node)); - } else if (node->IsInstance()) { - return VisitType(Downcast(node)); - } else if (node->IsInstance()) { - return PrintPrimFunc(Downcast(node)); - } else if (node->IsInstance()) { - return PrintIRModule(Downcast(node)); - } else if (node->IsInstance()) { - return PrintArray(node.as()); - } else if (node->IsInstance()) { - return PrintIterVar(node.as()); - } else if (node->IsInstance()) { - return PrintRange(node.as()); - } else if (node->IsInstance()) { - return PrintBuffer(node.as()); - } else if (node->IsInstance()) { - return PrintProducer(node.as()); - } else if (node->IsInstance()) { - return PrintString(node.as()); - } else if (node->IsInstance()) { - return PrintBufferRegion(node.as()); - } else if (node->IsInstance()) { - return Doc::Text(node.as()->ToDebugString()); - } else { - return this->meta_->GetMetaNode(node); - } -} - -Doc TIRTextPrinter::PrintPrimFunc(const PrimFunc& prim_func) { - const auto* op = prim_func.operator->(); - const auto& signature = op->func_type_annotation(); - // collect Meta in DictAttr - if (prim_func->attrs.defined()) { - for (const auto& it : prim_func->attrs->dict) { - meta_collector_.Collect(it.second); - } - } - // collect buffers in buffer_map - memo_var_.clear(); - memo_buf_.clear(); - - // ordered vars associated with buffers, for consistent printing - std::vector buffer_vars_ordered; - - for (tir::Var v : op->params) { - auto buffer_map_find = op->buffer_map.find(v); - if (buffer_map_find != op->buffer_map.end()) { - auto map_data = *buffer_map_find; - buffer_vars_ordered.push_back(map_data.first); - memo_buf_[map_data.second] = AllocBuf(map_data.second); - } - } - - // print PrimFunc - Doc doc; - doc << "primfn" - << "("; - // print params and its type annotation - std::vector params; - for (const auto& param : op->params) { - params.push_back(Print(param)); - } - Doc sep; - doc << PrintSep(params, Doc::Indent(9, Doc::Text(", "))) << ")"; - // print return type - doc << " -> " << Print(signature->ret_type); - // print attr - Doc attr_doc; - std::vector attr_docs; - if (prim_func->attrs.defined()) { - for (const auto& it : op->attrs->dict) { - attr_docs.push_back(Doc::StrLiteral(it.first) << ": " << Print(it.second)); - } - attr_doc << NewLine() << "attr = {" << PrintSep(attr_docs, Doc::Text(", ")) << "}"; - doc << Doc::Indent(2, attr_doc); - } - - // print all the buffers in the tree - if (memo_buf_.size() != 0) { - Doc buffer_doc; - std::vector buffer_docs; - for (const tir::Var& v : buffer_vars_ordered) { - const Buffer buf = op->buffer_map[v]; - buffer_docs.push_back(BufferNode2Doc(buf.get(), Print(buf))); - } - buffer_doc << NewLine() << "buffers = {"; - buffer_doc << PrintSep(buffer_docs, Doc::Indent(11, Doc::Text(",") << NewLine())); - doc << Doc::Indent(2, buffer_doc) << "}"; - } - - if (op->buffer_map.size() != 0) { - // print buffer_map - std::vector buffer_map_doc; - for (const tir::Var& v : buffer_vars_ordered) { - const Buffer buf = op->buffer_map[v]; - buffer_map_doc.push_back(Print(v) << ": " << Print(buf)); - } - doc << Doc::Indent( - 2, NewLine() << "buffer_map = {" << PrintSep(buffer_map_doc, Doc::Text(", ")) << "}"); - } - - doc << PrintBody(op->body); - return doc; -} - -Doc TIRTextPrinter::NewLine() { return Doc::NewLine(); } - -Doc TIRTextPrinter::PrintIRModule(const IRModule& module) { - const auto* op = module.operator->(); - Doc doc; - - Doc body; - body << NewLine(); - std::vector functions; - for (auto it = op->functions.begin(); it != op->functions.end(); ++it) { - if ((*it).second.as()) { - functions.push_back(Print((*it).second)); - } - } - body << TIRTextPrinter::PrintSep(functions, NewLine() << NewLine()); - doc << Doc::Indent(0, body); - return doc; -} - -Doc TIRTextPrinter::PrintArray(const ArrayNode* op) { - Doc doc; - doc << '['; - for (size_t i = 0; i < op->size(); ++i) { - if (i != 0) { - doc << ", "; - } - doc << Print(op->at(i)); - } - doc << ']'; - return doc; -} - -Doc TIRTextPrinter::PrintIterVar(const IterVarNode* op) { - Doc doc; - doc << "IterVar(" << Print(op->var); - if (op->dom.defined()) { - doc << ", [" << Print(op->dom) << "], "; - } else { - doc << ", " << Print(op->dom) << ", "; - } - doc << Doc::StrLiteral(IterVarType2String(op->iter_type)) << ", "; - doc << Doc::StrLiteral(op->thread_tag) << ")"; - return doc; -} - -Doc TIRTextPrinter::PrintRange(const RangeNode* op) { - return Print(op->min) << ":" << Print(op->min + op->extent); -} - -Doc TIRTextPrinter::PrintBuffer(const BufferNode* op) { - const Buffer& buffer = GetRef(op); - - if (meta_->InMeta(buffer)) { - return meta_->GetMetaNode(buffer); - } else if (memo_buf_.count(buffer)) { - return memo_buf_[buffer]; - } else { - memo_buf_[buffer] = AllocBuf(buffer); - return BufferNode2Doc(op, memo_buf_[buffer]); - } -} - -Doc TIRTextPrinter::PrintProducer(const DataProducerNode* op) { - const DataProducer& prod = GetRef(op); - - if (meta_->InMeta(prod)) { - return meta_->GetMetaNode(prod); - } else if (memo_producer_.count(prod)) { - return memo_producer_[prod]; - } else { - memo_producer_[prod] = AllocProducer(prod); - return DataProducerNode2Doc(op, memo_producer_[prod]); - } -} - -Doc TIRTextPrinter::BufferNode2Doc(const BufferNode* buf, Doc doc) { - doc << Doc::Text(": Buffer(") << Print(buf->data) << ", " << PrintDType(buf->dtype) << ", " - << Print(buf->shape) << ", " << Print(buf->strides); - if (!is_zero(buf->elem_offset)) { - doc << ", elem_offset=" << Print(buf->elem_offset); - } - if (buf->axis_separators.size()) { - doc << ", axis_separators=" << Print(buf->axis_separators); - } - if (GetRef(buf).scope() != "global") { - doc << ", scope=" << Doc::StrLiteral(GetRef(buf).scope()); - } - if (buf->data_alignment != runtime::kAllocAlignment) { - doc << ", align=" << buf->data_alignment; - } - if (buf->offset_factor != 1) { - doc << ", offset_factor=" << buf->offset_factor; - } - if (buf->buffer_type != 1) { - doc << ", type=" << Doc::StrLiteral("auto"); - } - return doc << ")"; -} - -Doc TIRTextPrinter::DataProducerNode2Doc(const DataProducerNode* prod, Doc doc) { - return doc << Doc::Text(": DataProducer(") << Print(prod->GetNameHint()) << ", " - << PrintDType(prod->GetDataType()) << ", " << Print(prod->GetShape()) << ")"; -} - -Doc TIRTextPrinter::PrintBufferRegion(const BufferRegionNode* op) { - Doc doc; - doc << Print(op->buffer) << "["; - for (size_t i = 0; i < op->region.size(); ++i) { - if (i != 0) { - doc << ", "; - } - const auto& range = op->region[i]; - if (!is_one(range->extent)) { - doc << Print(range->min) << ":" << Print(range->min + range->extent); - } else { - doc << Print(range->min); - } - } - doc << "]"; - return doc; -} - -Doc TIRTextPrinter::VisitExprDefault_(const Object* op) { - return this->meta_->GetMetaNode(GetRef(op)); -} - -Doc TIRTextPrinter::VisitStmtDefault_(const Object* op) { - return this->meta_->GetMetaNode(GetRef(op)); -} - -Doc TIRTextPrinter::VisitExpr_(const IntImmNode* op) { - return PrintConstScalar(op->dtype, op->value); -} - -Doc TIRTextPrinter::VisitExpr_(const FloatImmNode* op) { - return PrintConstScalar(op->dtype, op->value); -} - -Doc TIRTextPrinter::VisitExpr_(const StringImmNode* op) { return Doc::StrLiteral(op->value); } - -Doc TIRTextPrinter::VisitExpr_(const CastNode* op) { - Doc doc; - doc << "cast(" << PrintDType(op->dtype) << ", " << Print(op->value) << ")"; - return doc; -} - -Doc TIRTextPrinter::VisitExpr_(const tir::VarNode* op) { - const tir::Var& var = GetRef(op); - return meta_->InMeta(var) ? meta_->GetMetaNode(var) : AllocVar(GetRef(op)); -} - -#define TVM_DECLARE_TIR_TEXT_PRINTER_BINOP(OpName, OpString) \ - Doc TIRTextPrinter::VisitExpr_(const OpName* op) { \ - Doc doc; \ - doc << "(" << Print(op->a) << OpString; \ - doc << Print(op->b) << ")"; \ - return doc; \ - } - -TVM_DECLARE_TIR_TEXT_PRINTER_BINOP(AddNode, " + ") -TVM_DECLARE_TIR_TEXT_PRINTER_BINOP(SubNode, " - ") -TVM_DECLARE_TIR_TEXT_PRINTER_BINOP(MulNode, "*") -TVM_DECLARE_TIR_TEXT_PRINTER_BINOP(DivNode, " / ") -TVM_DECLARE_TIR_TEXT_PRINTER_BINOP(ModNode, " % ") -TVM_DECLARE_TIR_TEXT_PRINTER_BINOP(EQNode, " == ") -TVM_DECLARE_TIR_TEXT_PRINTER_BINOP(NENode, " != ") -TVM_DECLARE_TIR_TEXT_PRINTER_BINOP(LTNode, " < ") -TVM_DECLARE_TIR_TEXT_PRINTER_BINOP(LENode, " <= ") -TVM_DECLARE_TIR_TEXT_PRINTER_BINOP(GTNode, " > ") -TVM_DECLARE_TIR_TEXT_PRINTER_BINOP(GENode, " >= ") -TVM_DECLARE_TIR_TEXT_PRINTER_BINOP(AndNode, " && ") -TVM_DECLARE_TIR_TEXT_PRINTER_BINOP(OrNode, " || ") - -Doc TIRTextPrinter::VisitExpr_(const FloorDivNode* op) { - Doc doc; - doc << "floordiv(" << Print(op->a) << ", " << Print(op->b) << ")"; - return doc; -} - -Doc TIRTextPrinter::VisitExpr_(const FloorModNode* op) { - Doc doc; - doc << "floormod(" << Print(op->a) << ", " << Print(op->b) << ")"; - return doc; -} - -Doc TIRTextPrinter::VisitExpr_(const MinNode* op) { - Doc doc; - doc << "min(" << Print(op->a) << ", " << Print(op->b) << ")"; - return doc; -} - -Doc TIRTextPrinter::VisitExpr_(const MaxNode* op) { - Doc doc; - doc << "max(" << Print(op->a) << ", " << Print(op->b) << ")"; - return doc; -} - -Doc TIRTextPrinter::VisitExpr_(const NotNode* op) { - Doc doc; - doc << "!" << Print(op->a); - return doc; -} - -Doc TIRTextPrinter::VisitExpr_(const SelectNode* op) { - Doc doc; - doc << "select(" << Print(op->condition) << ", " << Print(op->true_value) << ", " - << Print(op->false_value) << ")"; - return doc; -} - -Doc TIRTextPrinter::VisitExpr_(const BufferLoadNode* op) { - Doc doc; - doc << Print(op->buffer) << Print(op->indices); - return doc; -} - -Doc TIRTextPrinter::VisitExpr_(const ProducerLoadNode* op) { - // TODO(tvm-team): consider make a better text format for producer. - Doc doc; - doc << op->producer->GetNameHint() << Print(op->indices); - return doc; -} - -Doc TIRTextPrinter::VisitExpr_(const RampNode* op) { - Doc doc; - doc << "ramp(" << Print(op->base) << ", " << Print(op->stride) << ", " << Print(op->lanes) << ")"; - return doc; -} - -Doc TIRTextPrinter::VisitExpr_(const BroadcastNode* op) { - Doc doc; - doc << "broadcast(" << Print(op->value) << ", " << Print(op->lanes) << ")"; - return doc; -} - -Doc TIRTextPrinter::VisitExpr_(const tir::LetNode* op) { - Doc doc; - doc << "let " << Print(op->var) << " = " << Print(op->value) << " in " << Print(op->body); - return doc; -} - -Doc TIRTextPrinter::VisitExpr_(const tir::CallNode* op) { - Doc doc; - std::vector func_args; - if (auto* ptr_op = op->op.as()) { - doc << "@" << Doc::Text(ptr_op->name) << "("; - if (ptr_op->name == "tir.call_llvm_pure_intrin") { - auto f = tvm::runtime::Registry::Get("target.llvm_get_intrinsic_name"); - ICHECK(f != nullptr) - << "Cannot find target.llvm_get_intrinsic_name. Compile with USE_LLVM=On"; - func_args.push_back(Print((*f)(Downcast(op->args[0])->value))); - for (size_t i = 1; i < op->args.size(); i++) { - func_args.push_back(Print(op->args[i])); - } - } else { - for (const auto& arg : op->args) { - func_args.push_back(Print(arg)); - } - } - } else { - // TODO(bohan): Print out the name by he global var in the module. - auto* op_gvar = op->op.as(); - ICHECK(op_gvar != nullptr); - doc << "@" << Doc::Text(op_gvar->name_hint) << "("; - for (const auto& arg : op->args) { - func_args.push_back(Print(arg)); - } - } - doc << PrintSep(func_args, Doc::Text(", ")) << ", dtype=" << PrintDType(op->dtype) << ")"; - return doc; -} - -Doc TIRTextPrinter::VisitExpr_(const ShuffleNode* op) { - Doc doc; - doc << "shuffle(" << Print(op->vectors) << ", " << Print(op->indices) << ")"; - return doc; -} - -Doc TIRTextPrinter::VisitExpr_(const ReduceNode* op) { - Doc doc; - doc << "reduce(" << Print(op->combiner) << ", " << Print(op->source) << ", " << Print(op->axis) - << ", " << op->value_index << ", " << Print(op->init) << ")"; - return doc; -} - -Doc TIRTextPrinter::VisitStmt_(const LetStmtNode* op) { - Doc doc; - doc << "let " << Print(op->var) << " = " << Print(op->value) << NewLine() << Print(op->body); - return doc; -} - -Doc TIRTextPrinter::VisitStmt_(const AttrStmtNode* op) { - Doc doc; - meta_collector_.Collect(op->node); - doc << "attr [" << Print(op->node) << "] " << Doc::StrLiteral(op->attr_key) << " = " - << Print(op->value); - if (op->body->IsInstance()) { - doc << PrintBody(op->body); - } else { - doc << ";" << NewLine() << Print(op->body); - } - return doc; -} - -Doc TIRTextPrinter::VisitStmt_(const AssertStmtNode* op) { - Doc doc; - doc << "assert(" << Print(op->condition) << ", " << Print(op->message) << ")" << NewLine() - << Print(op->body); - return doc; -} - -Doc TIRTextPrinter::VisitStmt_(const BufferStoreNode* op) { - Doc doc; - doc << Print(op->buffer) << Print(op->indices) << " = " << Print(op->value); - return doc; -} - -Doc TIRTextPrinter::VisitStmt_(const ProducerStoreNode* op) { - Doc doc; - doc << Print(op->producer) << Print(op->indices) << " = " << Print(op->value); - return doc; -} - -Doc TIRTextPrinter::VisitStmt_(const BufferRealizeNode* op) { - Doc doc; - doc << "realize(" << Print(op->buffer) << ", " << Print(op->bounds) << ", " - << Print(op->condition) << PrintBody(op->body) << ")"; - return doc; -} - -Doc TIRTextPrinter::VisitStmt_(const ProducerRealizeNode* op) { - Doc doc; - doc << "producer_realize(" << Print(op->producer) << ", " << Print(op->bounds) << ", " - << Print(op->condition) << ", " << PrintBody(op->body) << ")"; - return doc; -} - -Doc TIRTextPrinter::VisitStmt_(const AllocateNode* op) { - Doc doc; - auto scope = GetPtrStorageScope(op->buffer_var); - doc << "allocate(" << Print(op->buffer_var) << ", "; - doc << PrintDType(op->dtype) << ", "; - doc << Print(op->extents) << "), storage_scope = " << scope; - if (!op->annotations.empty()) { - std::vector attr_docs; - for (const auto& it : op->annotations) { - attr_docs.push_back(Doc::StrLiteral(it.first) << ": " << Print(it.second)); - } - doc << ", annotations = {" << PrintSep(attr_docs, Doc::Text(", ")) << "})"; - } - if (!is_one(op->condition)) { - doc << " if " << Print(op->condition); - } - if (op->body->IsInstance()) { - doc << PrintBody(op->body); - } else { - doc << ";" << NewLine() << Print(op->body); - } - return doc; -} - -Doc TIRTextPrinter::VisitStmt_(const AllocateConstNode* op) { - Doc doc; - doc << "constant(" << Print(op->buffer_var) << ", " << PrintDType(op->dtype) << ", " - << Print(op->extents) << ")"; - - if (op->body->IsInstance()) { - doc << PrintBody(op->body); - } else { - doc << ";" << NewLine() << Print(op->body); - } - return doc; -} - -Doc TIRTextPrinter::VisitStmt_(const DeclBufferNode* op) { - Doc doc; - doc << AllocBuf(op->buffer) << " = decl_buffer(" << Print(op->buffer->data) << ", " - << PrintDType(op->buffer->dtype) << ", " << Print(op->buffer->shape) << ")" << NewLine(); - if (op->body->IsInstance()) { - doc << PrintBody(op->body); - } else { - doc << ";" << NewLine() << Print(op->body); - } - return doc; -} - -Doc TIRTextPrinter::VisitStmt_(const IfThenElseNode* op) { - Doc doc; - doc << "if " << Print(op->condition) << PrintBody(op->then_case); - if (!is_one(op->condition) && op->else_case) { - doc << " else" << PrintBody(op->else_case.value()); - } - return doc; -} - -Doc TIRTextPrinter::VisitStmt_(const SeqStmtNode* op) { - std::vector stmts; - Doc seq_doc, doc; - for (Stmt stmt : op->seq) { - seq_doc << NewLine() << Print(stmt); - } - doc << " {" << Doc::Indent(2, seq_doc) << NewLine() << "}"; - return doc; -} - -Doc TIRTextPrinter::VisitStmt_(const EvaluateNode* op) { - Doc doc; - doc << Print(op->value); - return doc; -} - -Doc TIRTextPrinter::VisitStmt_(const ForNode* op) { - Doc doc; - doc << "for (" << Print(op->loop_var) << ", " << Print(op->min) << ", " - << Print(op->min + op->extent) << ")"; - if (op->kind != ForKind::kSerial) { - doc << " " << Doc::StrLiteral(ForKind2String(op->kind)); - } - doc << PrintBody(op->body); - return doc; -} - -Doc TIRTextPrinter::VisitStmt_(const WhileNode* op) { - Doc doc; - doc << "while (" << Print(op->condition) << ")"; - doc << PrintBody(op->body); - return doc; -} - -Doc TIRTextPrinter::VisitStmt_(const PrefetchNode* op) { - Doc doc; - doc << "prefetch(" << Print(op->buffer) << ", " << Print(op->bounds) << ")"; - return doc; -} - -Doc TIRTextPrinter::VisitStmt_(const BlockRealizeNode* op) { - const auto* block_op = op->block.as(); - // print block name and block vars - Doc doc; - doc << "block(["; - std::vector block_var_docs; - for (const auto& iter_var : block_op->iter_vars) { - Doc block_var_doc; - if (is_zero(iter_var->dom->min) && iter_var->iter_type == kDataPar) { - block_var_doc << Print(iter_var->dom->extent); - } else { - block_var_doc << "tir."; - switch (iter_var->iter_type) { - case kDataPar: - block_var_doc << "range"; - break; - case kCommReduce: - block_var_doc << "reduce_axis"; - break; - case kOrdered: - block_var_doc << "scan_axis"; - break; - case kOpaque: - block_var_doc << "opaque_axis"; - break; - default: - LOG(FATAL) << "Unknown block var iter type"; - break; - } - block_var_doc << "(" << Print(iter_var->dom->min) << ", " - << Print(iter_var->dom->min + iter_var->dom->extent) << ")"; - } - block_var_docs.push_back(block_var_doc); - } - doc << PrintSep(block_var_docs, Doc::Text(", ")) << "], "; - doc << Doc::StrLiteral(block_op->name_hint) << ")"; - std::vector block_var_names; - for (const auto& iter_var : block_op->iter_vars) { - Doc block_var_name; - AllocVar(iter_var->var); - block_var_names.push_back(Print(iter_var->var)); - } - if (!block_var_names.empty()) { - doc << " as [" << PrintSep(block_var_names, Doc::Text(", ")) << "]"; - } - doc << " {"; - Doc block_attr_doc; - // print predicate, binding, read/write tensor region, annotations - if (!is_one(op->predicate)) { - block_attr_doc << NewLine() << "where(" << Print(op->predicate) << ")"; - } - for (size_t i = 0; i < block_op->iter_vars.size(); ++i) - block_attr_doc << NewLine() << "bind(" << Print(block_op->iter_vars[i]->var) << ", " - << Print(op->iter_values[i]) << ")"; - block_attr_doc << NewLine() << "tir.reads(" << Print(block_op->reads) << ")"; - block_attr_doc << NewLine() << "tir.writes(" << Print(block_op->writes) << ")"; - if (!block_op->annotations.empty()) { - std::vector attr_docs; - for (const auto& it : block_op->annotations) { - attr_docs.push_back(Doc::StrLiteral(it.first) << ": " << Print(it.second)); - } - block_attr_doc << NewLine() << "tir.attrs({" << PrintSep(attr_docs, Doc::Text(", ")) << "})"; - } - // print body - Doc body; - body << NewLine(); - for (const auto& alloc_buf : block_op->alloc_buffers) { - body << AllocBuf(alloc_buf) << " = alloc_buffer(" << PrintDType(alloc_buf->dtype) - << Print(alloc_buf->shape) << ")" << NewLine(); - } - for (const auto& match_buf : block_op->match_buffers) { - body << AllocBuf(match_buf->buffer) << " = match_buffer(" << Print(match_buf->source) << ")" - << NewLine(); - } - if (block_op->init.defined()) { - Doc init_block; - init_block << "with init()"; - init_block << PrintBody(block_op->init.value()); - body << init_block << NewLine(); - } - body << Print(block_op->body); - doc << Doc::Indent(2, block_attr_doc << body); - return doc; -} - -Doc TIRTextPrinter::VisitType_(const PrimTypeNode* node) { - Doc doc; - doc << PrintDType(node->dtype); - return doc; -} - -Doc TIRTextPrinter::VisitType_(const PointerTypeNode* node) { - Doc doc; - doc << "Pointer("; - if (!node->storage_scope.empty()) { - doc << node->storage_scope << " "; - } - doc << Print(node->element_type) << ")"; - return doc; -} - -Doc TIRTextPrinter::VisitType_(const TupleTypeNode* node) { - std::vector fields; - for (Type field : node->fields) { - fields.push_back(Print(field)); - } - Doc doc; - doc << "(" << Doc::Concat(fields); - // conform to python tuple format (1,) - if (node->fields.size() == 1) { - doc << ","; - } - return doc << ")"; -} - -Doc TIRTextPrinter::PrintDType(DataType dtype) { - return Doc::Text(runtime::DLDataType2String(dtype)); -} - -template -Doc TIRTextPrinter::PrintConstScalar(DataType dtype, const T& data) { - Doc doc; - std::ostringstream os; - os << data; - if (dtype == DataType::Int(32)) { - doc << Doc::Text(os.str()); - } else { - if (dtype.bits() == 1 && dtype.lanes() == 1 && dtype.code() == kDLUInt) { - doc << ((data == 1) ? "True" : "False"); - return doc; - } - doc << Doc::Text(os.str()); - switch (dtype.code()) { - case kDLInt: - doc << "i"; - break; - case kDLUInt: - doc << "u"; - break; - case kDLFloat: - doc << "f"; - break; - } - doc << Doc::Text(std::to_string(dtype.bits())); - if (dtype.lanes() != 1) doc << "x" << Doc::Text(std::to_string(dtype.lanes())); - } - return doc; -} - -Doc TIRTextPrinter::GetUniqueName(std::string prefix) { - // std::replace(prefix.begin(), prefix.end(), '.', '_'); - std::string unique_prefix = prefix; - auto it = name_alloc_map_.find(prefix); - if (it != name_alloc_map_.end()) { - while (name_alloc_map_.count(unique_prefix = prefix + "_" + std::to_string(++it->second)) > 0) { - } - } - name_alloc_map_[unique_prefix] = 0; - return Doc::Text(unique_prefix); -} - -Doc TIRTextPrinter::AllocVar(const tir::Var& var) { - const auto& it = memo_var_.find(var); - if (it != memo_var_.end()) { - return it->second; - } - std::string name = var->name_hint.operator std::string(); - if (name.length() == 0 || !std::isalpha(name[0])) { - name = "v" + name; - } - Doc val = GetUniqueName(name); - memo_var_[var] = val; - return val << ": " << Print(GetType(var)); -} - -Doc TIRTextPrinter::AllocBuf(const Buffer& buffer) { - const auto& it = memo_buf_.find(buffer); - if (it != memo_buf_.end()) { - return it->second; - } - std::string name = buffer->name; - if (name.length() == 0 || !std::isalpha(name[0])) { - name = "buf_" + name; - } - Doc val = GetUniqueName(name); - memo_buf_[buffer] = val; - return val; -} - -Doc TIRTextPrinter::AllocProducer(const DataProducer& producer) { - const auto& it = memo_producer_.find(producer); - if (it != memo_producer_.end()) { - return it->second; - } - std::string name = producer->GetNameHint(); - if (name.length() == 0 || !std::isalpha(name[0])) { - name = "tensor_" + name; - } - Doc val = GetUniqueName(name); - memo_producer_[producer] = val; - return val; -} - -Doc TIRTextPrinter::PrintSep(const std::vector& vec, const Doc& sep) { - Doc seq; - if (vec.size() != 0) { - seq = vec[0]; - for (size_t i = 1; i < vec.size(); i++) { - seq << sep << vec[i]; - } - } - return seq; -} - -Doc TIRTextPrinter::PrintBody(const Stmt& body, bool indent) { - Doc doc; - if (body->IsInstance()) return Print(body); - doc << " {" << Doc::Indent(2, NewLine() << Print(body)) << NewLine() << "}"; - return doc; -} - -bool TIRTextPrinter::GetVarName(tir::Var v, std::string* s) { - auto it = memo_var_.find(v); - if (it == memo_var_.end()) { - return false; - } - - *s = it->second.str(); - return true; -} - -} // namespace relay -} // namespace tvm diff --git a/src/relay/printer/tir_text_printer_debug.cc b/src/relay/printer/tir_text_printer_debug.cc deleted file mode 100644 index 914d8877d2f7..000000000000 --- a/src/relay/printer/tir_text_printer_debug.cc +++ /dev/null @@ -1,97 +0,0 @@ -/* - * 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 tir_text_printer.cc - * \brief Printer to print out the IR text format - * that can be parsed by a parser. - */ - -#include "tir_text_printer_debug.h" - -#include -#include - -namespace tvm { -namespace relay { - -std::optional span_text(const Span& span) { - if (!span.defined()) { - return std::nullopt; - } - - std::string source("main.tir"); - if (span->source_name.defined() && span->source_name->name.get()) { - source = span->source_name->name; - } - return source + ":" + std::to_string(span->line) + ":" + std::to_string(span->column); -} - -template -void add_all_relevant_lines(const std::vector>& data, - size_t current_line, Doc* output) { - ICHECK(output) << "output must be a valid Doc"; - for (const auto& item : data) { - if (std::get<1>(item) != current_line - 1) { - // Item is not relevant for this line, skip it - continue; - } - - // Print out the item's span info if present - auto text = span_text(std::get<0>(item)->span); - if (text.has_value()) { - *output << *text; - } else { - *output << "missing"; - } - *output << ", "; - } -} - -Doc TIRTextPrinterDebug::NewLine() { - current_line_ += 1; - - if (!show_spans_) { - return TIRTextPrinter::NewLine(); - } - - Doc output; - - output << " ["; - - add_all_relevant_lines(exprs_by_line_, current_line_, &output); - add_all_relevant_lines(stmts_by_line_, current_line_, &output); - - output << "]" << TIRTextPrinter::NewLine(); - - return output; -} - -Doc TIRTextPrinterDebug::VisitStmt(const tvm::tir::Stmt& n) { - stmts_by_line_.push_back(std::make_tuple(n.get(), current_line_)); - return TIRTextPrinter::VisitStmt(n); -} - -Doc TIRTextPrinterDebug::VisitExpr(const PrimExpr& e) { - exprs_by_line_.push_back(std::make_tuple(e.get(), current_line_)); - return TIRTextPrinter::VisitExpr(e); -} - -} // namespace relay -} // namespace tvm diff --git a/src/relay/printer/tir_text_printer_debug.h b/src/relay/printer/tir_text_printer_debug.h deleted file mode 100644 index f7cb7a6554ec..000000000000 --- a/src/relay/printer/tir_text_printer_debug.h +++ /dev/null @@ -1,70 +0,0 @@ -/* - * 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 text_printer.h - * \brief Printer to print out the unified IR text format - * that can be parsed by a parser. - */ - -#ifndef TVM_RELAY_PRINTER_TIR_TEXT_PRINTER_DEBUG_H_ -#define TVM_RELAY_PRINTER_TIR_TEXT_PRINTER_DEBUG_H_ - -#include -#include - -#include "text_printer.h" - -namespace tvm { -namespace relay { - -class TIRTextPrinterDebug : public TIRTextPrinter { - public: - explicit TIRTextPrinterDebug(bool show_spans) - : TIRTextPrinter(false, &meta_), current_line_(1), show_spans_(show_spans) {} - - std::vector> GetExprsByLine() const { - return exprs_by_line_; - } - - std::vector> GetStmtsByLine() const { return stmts_by_line_; } - - private: - Doc NewLine() override; - - Doc VisitStmt(const tvm::tir::Stmt& n) override; - Doc VisitExpr(const PrimExpr& e) override; - - TextMetaDataContext meta_; - - // Line that the printer is currently printing - size_t current_line_; - - // Whether to include spans relevant to each line before a newline or not - bool show_spans_; - - // Record of all stmts and exprs and their corresponding line - std::vector> stmts_by_line_; - std::vector> exprs_by_line_; -}; - -} // namespace relay -} // namespace tvm - -#endif // TVM_RELAY_PRINTER_TIR_TEXT_PRINTER_DEBUG_H_ diff --git a/src/relay/printer/tvmscript_printer.cc b/src/relay/printer/tvmscript_printer.cc deleted file mode 100644 index 1126e633d633..000000000000 --- a/src/relay/printer/tvmscript_printer.cc +++ /dev/null @@ -1,1996 +0,0 @@ -/* - * 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 printer/tvmscript_printer.cc - * \brief Printer class to print Tensor IR to python syntax script - */ - -#include -#include -#include -#include -#include -#include -#include -#include -#include -#include -#include -#include -#include - -#include -#include - -#include "../../tir/transforms/ir_utils.h" -#include "doc.h" -#include "meta_data.h" -#include "text_printer.h" - -namespace tvm { -namespace relay { - -using namespace tvm::tir; - -enum class ExprPrecedence : int { - /*! \brief Identity(e.g., IntImm, Var) and function call(e.g., floordiv, min) */ - kIdentity = 0, - /*! - * \brief Multiplication(*), division(/), and remainder(%) - * \note floorDiv, floorMod is marked as kIdentity since they are function calls. - */ - kMultiplicationDivision = 1, - /*! \brief Addition(+) and subtraction(-) */ - kAdditionSubtraction = 2, - /*! \brief For relational operators < and <= and > and >= respectively */ - kRelational = 3, - /*! \brief For equality operators = and != respectively */ - kEquality = 4, - /*! \brief And(&&) */ - kAnd = 5, - /*! \brief Or(||) */ - kOr = 6, - /*! \brief Unknown precedence */ - kUnknown = 7, -}; - -/*! \brief Utility used for identifying usage of a buffer_var - * - * \details Find the Buffer object that corresponds to a variable or - * allocation, based on the BufferLoad/BufferStore instances that - * occur within the allocation's body. - */ -class BufferUsageFinder : public StmtExprVisitor { - public: - static Map> FindUsage(Map> usage, Stmt body) { - BufferUsageFinder visitor(std::move(usage)); - visitor.VisitStmt(body); - return std::move(visitor.usage_); - } - - void VisitExpr_(const tir::VarNode* op) final { - tir::Var var = GetRef(op); - if (!usage_.count(var)) { - usage_.Set(var, {}); - } - } - - void VisitExpr_(const BufferLoadNode* op) final { - VisitBuffer(op->buffer); - StmtExprVisitor::VisitExpr_(op); - } - - void VisitStmt_(const BufferStoreNode* op) final { - VisitBuffer(op->buffer); - StmtExprVisitor::VisitStmt_(op); - } - - void VisitStmt_(const DeclBufferNode* op) final { - buffers_declared_.insert(op->buffer.get()); - StmtExprVisitor::VisitStmt_(op); - buffers_declared_.erase(op->buffer.get()); - } - - private: - explicit BufferUsageFinder(Map> usage) : usage_(usage) {} - - void VisitBuffer(const Buffer& buffer) { - if (buffers_visited_.count(buffer.get())) { - return; - } - if (buffers_declared_.count(buffer.get())) { - return; - } - buffers_visited_.insert(buffer.get()); - - Array arr = usage_.Get(buffer->data).value_or({}); - arr.push_back(buffer); - usage_.Set(buffer->data, arr); - } - - // The search result. - Map> usage_; - // The buffers that have been visited so far, to avoid duplicate - // entries in the search result. - std::unordered_set buffers_visited_; - // The buffers declared via `DeclBuffer`. These buffers are excluded from the result because - // T.buffer_decl shouldn't be printed for them. - std::unordered_set buffers_declared_; -}; - -/*! - * \brief The printer for TVMScript - * \details The printer obtain the precedence of the top-level operation when printing each - * subexpression to decide whether or not parentheses is needed. - */ -class TVMScriptPrinter : public StmtFunctor, - public tir::ExprFunctor, - public TypeFunctor { - public: - explicit TVMScriptPrinter(const String& tir_prefix, bool show_meta, - runtime::TypedPackedFunc annotate = nullptr) - : tir_prefix_(tir_prefix), - show_meta_(show_meta), - annotate_(std::move(annotate)), - meta_collector_(&meta_) {} - - /*! - * \brief Print the node. - * \param node The node to be printed. - * \param out_precedence The operator precedence of node if it's a PrimExpr, - * so we can simplify the bracket. - */ - TVM_DLL Doc Print(const ObjectRef& node); - - protected: - /*! \brief The tir prefix */ - String tir_prefix_; - /*! \brief whether show meta data */ - bool show_meta_; - /*! \brief additional comment function */ - runtime::TypedPackedFunc annotate_; - /*! \brief meta data context */ - TextMetaDataContext meta_; - /*! \brief meta collector */ - relay::MetaCollector meta_collector_; - /*! \brief map from Function to GlobalVar */ - std::unordered_map func2var_; - /*! \brief var collector (var defined by For/Loop/Block) */ - std::unordered_set var_not_in_headers_; - /*! - * \brief buffer collector - * (buffer defined in BufferMap, BufferAllocation and MatchBufferRegion) - */ - std::unordered_set buf_not_in_headers_; - /*! \brief Map from Var to thread env name */ - std::unordered_map var_env_map_; - /*! \brief Map from Var to Doc */ - std::unordered_map memo_var_; - /*! \brief Map from Buffer to Doc */ - std::unordered_map memo_buf_; - /*! \brief Map from Buffer to Declaration Doc */ - std::unordered_map memo_buf_decl_; - /*! \brief name allocation map */ - std::unordered_map name_alloc_map_; - /*! \brief number of children of current node's parent */ - int num_child_; - /*! \brief the number of current node */ - int current_num_; - /*! \brief loop stack without annotations */ - std::vector simple_loop_stack_; - /*! \brief the maps from loop_vars to the loops */ - std::unordered_map loop_var_map_; - /*! - * \brief simple block vars remap from loop vars - * simple_remap requires: - * 1. block var iter type is kDataPar or kCommReduce - * 2. value is a single Var, which is a loop_var outside the block - * 3. The iter range is equal to loop range - */ - std::vector> block_var_remaps_; - /*! - * \brief Map from variables to the buffers they are used in. - * - * Used for identifying buffers that should be declared after the - * LetStmt or Allocate that generates their data pointer, rather - * than in the header. - */ - Map> buffer_var_usage_; - /*! \brief Analyzer to simplify some expressions. */ - arith::Analyzer ana_; - - Doc VisitExpr_(const CastNode* op, ExprPrecedence* out_precedence) override; - Doc VisitExpr_(const tir::VarNode* op, ExprPrecedence* out_precedence) override; - Doc VisitExpr_(const AddNode* op, ExprPrecedence* out_precedence) override; - Doc VisitExpr_(const SubNode* op, ExprPrecedence* out_precedence) override; - Doc VisitExpr_(const MulNode* op, ExprPrecedence* out_precedence) override; - Doc VisitExpr_(const DivNode* op, ExprPrecedence* out_precedence) override; - Doc VisitExpr_(const ModNode* op, ExprPrecedence* out_precedence) override; - Doc VisitExpr_(const FloorDivNode* op, ExprPrecedence* out_precedence) override; - Doc VisitExpr_(const FloorModNode* op, ExprPrecedence* out_precedence) override; - Doc VisitExpr_(const MinNode* op, ExprPrecedence* out_precedence) override; - Doc VisitExpr_(const MaxNode* op, ExprPrecedence* out_precedence) override; - Doc VisitExpr_(const EQNode* op, ExprPrecedence* out_precedence) override; - Doc VisitExpr_(const NENode* op, ExprPrecedence* out_precedence) override; - Doc VisitExpr_(const LTNode* op, ExprPrecedence* out_precedence) override; - Doc VisitExpr_(const LENode* op, ExprPrecedence* out_precedence) override; - Doc VisitExpr_(const GTNode* op, ExprPrecedence* out_precedence) override; - Doc VisitExpr_(const GENode* op, ExprPrecedence* out_precedence) override; - Doc VisitExpr_(const AndNode* op, ExprPrecedence* out_precedence) override; - Doc VisitExpr_(const OrNode* op, ExprPrecedence* out_precedence) override; - Doc VisitExpr_(const NotNode* op, ExprPrecedence* out_precedence) override; - Doc VisitExpr_(const SelectNode* op, ExprPrecedence* out_precedence) override; - Doc VisitExpr_(const IntImmNode* op, ExprPrecedence* out_precedence) override; - Doc VisitExpr_(const FloatImmNode* op, ExprPrecedence* out_precedence) override; - Doc VisitExpr_(const StringImmNode* op, ExprPrecedence* out_precedence) override; - Doc VisitExpr_(const ProducerLoadNode* op, ExprPrecedence* out_precedence) override; - Doc VisitExpr_(const BufferLoadNode* op, ExprPrecedence* out_precedence) override; - Doc VisitExpr_(const RampNode* op, ExprPrecedence* out_precedence) override; - Doc VisitExpr_(const BroadcastNode* op, ExprPrecedence* out_precedence) override; - Doc VisitExpr_(const tir::LetNode* op, ExprPrecedence* out_precedence) override; - Doc VisitExpr_(const tir::CallNode* op, ExprPrecedence* out_precedence) override; - Doc VisitExpr_(const ShuffleNode* op, ExprPrecedence* out_precedence) override; - Doc VisitExpr_(const ReduceNode* op, ExprPrecedence* out_precedence) override; - Doc VisitExprDefault_(const Object* op, ExprPrecedence* out_precedence) override; - - Doc VisitStmt_(const LetStmtNode* op) override; - Doc VisitStmt_(const AttrStmtNode* op) override; - Doc VisitStmt_(const AssertStmtNode* op) override; - Doc VisitStmt_(const BufferStoreNode* op) override; - Doc VisitStmt_(const BufferRealizeNode* op) override; - Doc VisitStmt_(const AllocateNode* op) override; - Doc VisitStmt_(const AllocateConstNode* op) override; - Doc VisitStmt_(const DeclBufferNode* op) override; - Doc VisitStmt_(const IfThenElseNode* op) override; - Doc VisitStmt_(const SeqStmtNode* op) override; - Doc VisitStmt_(const ForNode* op) override; - Doc VisitStmt_(const WhileNode* op) override; - Doc VisitStmt_(const PrefetchNode* op) override; - Doc VisitStmt_(const EvaluateNode* op) override; - Doc VisitStmt_(const BlockRealizeNode* op) override; - Doc VisitStmtDefault_(const Object* op) override; - - Doc VisitType_(const PrimTypeNode* node) override; - Doc VisitType_(const PointerTypeNode* node) override; - Doc VisitType_(const TupleTypeNode* node) override; - - Doc PrintBody(const Stmt& body); - Doc PrintIRModule(const IRModule& module); - Doc PrintPrimFunc(const PrimFunc& primFunc); - Doc PrintIterVar(const IterVarNode* op); - Doc PrintRange(const RangeNode* op); - Doc PrintArray(const ArrayNode* op); - Doc PrintBuffer(const BufferNode* op); - Doc PrintBufferIndices(const Array& indices); - Doc PrintNonHeaderBufferDeclarations(const Array& aliasing_buffers); - Doc AllocBufferDeclaration(const Buffer& buf); - Doc PrintBlockVar(const IterVar& iter_var, const PrimExpr& value); - Doc PrintBlockVarRemaps(); - Doc PrintBlockPredicate(const BlockRealizeNode* op); - Doc PrintBlockVars(const BlockRealizeNode* op); - Doc PrintBlockAttr(const BlockRealizeNode* op); - Doc PrintExpandedArray(const ArrayNode* op); - Doc PrintBlockBody(const BlockNode* op); - virtual Doc PrintBlockName(const BlockNode* block_op); - Doc PrintBufferRegion(const BufferRegionNode* op); - Doc PrintMatchBufferRegion(const MatchBufferRegionNode* op); - Doc PrintCommReducer(const CommReducerNode* op); - Doc PrintAnnotations(const Map& annotations); - Doc PrintTarget(const TargetNode* target); - static Doc PrintString(const StringObj* op) { return Doc::StrLiteral(op->data); } - - Doc GetUniqueName(std::string prefix); - Doc AllocVar(const tir::Var& var); - Doc AllocBuf(const Buffer& buffer); - void TryDeallocVar(const tir::Var& var); - bool ContainsOptionalInfo(const Stmt& stmt); - /*! - * \brief Check if a buffer declaration satisfies: - * 1. has only 'shape' and 'dtype' arguments specified, - * 2. the shape and strides are not dynamic. - * \param buffer The match buffer to be checked - */ - bool IsSimpleBuffer(const Buffer& buffer); - Doc PrintInlineBufferBind(const Buffer& buffer); - Doc PrintTuple(const ArrayNode* op); - - /*! Helper functions for loop printing. */ - /*! - * \brief Print a single for loop - * \param loop The for loop to be printed - */ - virtual Doc PrintLoop(const For& loop); - /*! \brief Print all simple loops in stack into one line using tir_prefix_.grid(). */ - Doc PrintLoopStack(); - /*! - * \brief Check whether a loop satisfies: - * 1. the loop is serial; - * 2. the loop has no annotation; - * 3. the loop starts from 0; - * 4. there is no optional information. - * \param for_op the for node to be checked - * \return A boolean indicating whether the input loop satisfies the above conditions - */ - bool IsSimpleLoop(const ForNode* for_op) { - return for_op->kind == ForKind::kSerial && for_op->annotations.empty() && - is_zero(for_op->min) && !ContainsOptionalInfo(GetRef(for_op)); - } - /*! - * \brief Check whether the `min` or `extent` of a loop depends on previous loops - * \param for_op The loop to be checked - * \return A boolean indicating whether the input loop depends on previous loops - */ - bool DependOnPrevLoops(const ForNode* for_op) { - auto f_check = [&var_map = this->loop_var_map_](const tir::VarNode* v) { - return var_map.count(v); - }; - return UsesVar(for_op->min, f_check) || UsesVar(for_op->extent, f_check); - } - - /*! - * \brief Print additional info about expr in comment. - * \param expr The expression. - */ - Doc PrintOptionalInfo(const Stmt& stmt) { - Doc doc; - // default annotations - if (ContainsOptionalInfo(stmt)) { - std::string annotated_stmt = annotate_(stmt); - doc << "# " << annotated_stmt << Doc::NewLine(); - } - return doc; - } - - /*! - * \brief special method to render vectors of docs with a separator - * \param vec vector of docs - * \param sep separator - */ - static Doc PrintSep(const std::vector& vec, const Doc& sep) { - Doc seq; - if (vec.size() != 0) { - seq = vec[0]; - for (size_t i = 1; i < vec.size(); i++) { - seq << sep << vec[i]; - } - } - return seq; - } - - /*! - * \brief dump meta info - * \return Doc with meta info - */ - Doc DumpMeta() { - if (show_meta_) { - return Doc::Text("__tvm_meta__ = ") - << (meta_.empty() ? Doc::Text("None") : meta_.GetMetaSection()); - } else { - return Doc::Text(""); - } - } - - /*! - * \brief special method to print out data type - * \param dtype The data type - */ - static Doc PrintDType(DataType dtype) { - return Doc::StrLiteral(runtime::DLDataType2String(dtype)); - } - - /*! - * \brief special method to print out const int64_t scalar - * \param dtype The data type - * \param data The pointer to hold the data. - */ - Doc PrintConstScalar(DataType dtype, const int64_t* data) const { - Doc doc; - std::ostringstream os; - - os << data[0]; - - if (dtype == DataType::Int(32)) { - doc << Doc::Text(os.str()); - } else if (dtype == DataType::Bool()) { - doc << Doc::Text(data[0] ? "True" : "False"); - } else { - doc << tir_prefix_ << "." << runtime::DLDataType2String(dtype) << "(" << Doc::Text(os.str()) - << ")"; - } - return doc; - } - - /*! - * \brief special method to print out const double scalar - * \param dtype The data type - * \param data The pointer to hold the data. - * \note this overriden function is created as std::isnan of msvc will complain about int64_t - */ - Doc PrintConstScalar(DataType dtype, const double* data) const { - Doc doc; - std::ostringstream os; - - os.precision(17); - if (std::isinf(data[0]) || std::isnan(data[0])) { - os << "\"" << data[0] << "\""; - } else { - os << data[0]; - } - - doc << tir_prefix_ << "." << runtime::DLDataType2String(dtype) << "(" << Doc::Text(os.str()) - << ")"; - - return doc; - } - - public: - static Doc PrintHeader(const std::string& tir_prefix) { - Doc header; - if (tir_prefix != "tir") { - header << "# from tvm.script import tir as " << tir_prefix << Doc::NewLine(); - } else { - header << "# from tvm.script import tir" << Doc::NewLine(); - } - return header; - } -}; - -/*! - * \brief special method to print NDArray in TIR - * \param arr the NDArray to be printed - * \param os the output stream where the NDArray will be printed to - */ -template -void NDArrayToTIR(::tvm::runtime::NDArray arr, std::ostream& os) { - if ((arr.DataType().code() == runtime::DataType::kInt || - arr.DataType().code() == runtime::DataType::kUInt) && - arr.DataType().bits() == 8) { - // Printing int8 NDArrays causes "UnicodeDecodeError: 'utf-8' codec can't decode byte" - // error during MetaSchedule tuning on int8 models. - return; - } - int ndim = arr->ndim; - int tot_dim = 1; - for (int i = 0; i < ndim; i++) { - tot_dim *= arr->shape[i]; - } - T* data_ptr = reinterpret_cast(arr->data); - constexpr int NUM_PRINT = 20; - os << "["; - for (int i = 0; i < tot_dim; i++) { - os << (i != 0 ? ", " : "") << data_ptr[i]; - if (i == NUM_PRINT) { - os << "..."; - break; - } - } - os << "]"; -} - -Doc TVMScriptPrinter::GetUniqueName(std::string prefix) { - std::replace(prefix.begin(), prefix.end(), '.', '_'); - std::string unique_prefix = prefix; - auto it = name_alloc_map_.find(prefix); - if (it != name_alloc_map_.end() && it->second >= 0) { - while (name_alloc_map_.count(unique_prefix = prefix + "_" + std::to_string(++it->second)) > 0) { - } - } - name_alloc_map_[unique_prefix] = 0; - return Doc::Text(unique_prefix); -} - -Doc TVMScriptPrinter::AllocVar(const tir::Var& var) { - const auto& it = memo_var_.find(var); - if (it != memo_var_.end()) { - return it->second; - } - std::string name = var->name_hint.operator std::string(); - if (name.length() == 0 || !std::isalpha(name[0])) { - name = "v" + name; - } - Doc val = GetUniqueName(name); - memo_var_[var] = val; - return val; -} - -Doc TVMScriptPrinter::AllocBufferDeclaration(const Buffer& buf) { - Doc doc = Print(buf->shape); - bool print_factor_explicitly = false; - doc << ", dtype=" << PrintDType(buf->dtype); - if (memo_var_.find(buf->data) != memo_var_.end()) { - doc << ", data=" << Print(buf->data); - } else { - // implicitly define data - memo_var_[buf->data] = Doc::Text(memo_buf_[buf].str() + ".data"); - var_not_in_headers_.insert(buf->data.get()); - } - if (!buf->strides.empty()) { - doc << ", strides=" << Print(buf->strides); - } - if (buf->elem_offset->IsInstance()) { - tir::Var elem_offset = Downcast(buf->elem_offset); - if (memo_var_.find(elem_offset) != memo_var_.end()) { - doc << ", elem_offset=" << Print(buf->elem_offset); - } else { - // implicitly define elem_offset - memo_var_[elem_offset] = Doc::Text(memo_buf_[buf].str() + ".elem_offset"); - var_not_in_headers_.insert(elem_offset.get()); - print_factor_explicitly = true; - } - } else if (buf->elem_offset->IsInstance()) { - IntImm elem_offset = Downcast(buf->elem_offset); - if (elem_offset->value != 0) { - doc << ", elem_offset=" << Print(buf->elem_offset); - } - } - if (buf.scope() != "global") { - doc << ", scope=" << Doc::StrLiteral(buf.scope()); - } - if (buf->data_alignment != runtime::kAllocAlignment) { - doc << ", align=" << buf->data_alignment; - } - if (buf->offset_factor != 1 || print_factor_explicitly) { - doc << ", offset_factor=" << buf->offset_factor; - } - if (buf->buffer_type != BufferType::kDefault) { - doc << ", type=" << Doc::StrLiteral("auto"); - } - if (buf->axis_separators.size()) { - doc << ", axis_separators=" << Print(buf->axis_separators); - } - return doc; -} - -Doc TVMScriptPrinter::AllocBuf(const Buffer& buffer) { - const auto& it = memo_buf_.find(buffer); - if (it != memo_buf_.end()) { - return it->second; - } - std::string name = buffer->name; - if (name.length() == 0 || !std::isalpha(name[0])) { - name = "buf_" + name; - } - Doc val = GetUniqueName(name); - memo_buf_[buffer] = val; - memo_buf_decl_[buffer] = AllocBufferDeclaration(buffer); - return val; -} - -/*! - * \brief Check if any optional information exists in annotate_ for - * a given Stmt. - * \param stmt The statement. - */ -bool TVMScriptPrinter::ContainsOptionalInfo(const Stmt& stmt) { - if (annotate_ == nullptr) return false; - return !annotate_(stmt).empty(); -} - -/*! - * \brief Try to dealloc vars out of space and leave the index to coming vars. - * \note It is not a necessary step. - */ -void TVMScriptPrinter::TryDeallocVar(const tir::Var& var) { - auto it = memo_var_.find(var); - ICHECK(it != memo_var_.end()); - std::string print_name = it->second.str(); - - std::string name_hint = var->name_hint.operator std::string(); - if (name_hint.length() == 0 || !std::isalpha(name_hint[0])) { - name_hint = "v" + name_hint; - } - std::replace(name_hint.begin(), name_hint.end(), '.', '_'); - - auto it2 = name_alloc_map_.find(name_hint); - // Skip it if we can not find the name_hint in name_alloc_map_. - if (it2 == name_alloc_map_.end()) return; - if (it2->second > 0) { - name_hint = name_hint + '_' + std::to_string(it2->second); - } - // Skip it if the name_hint is not equal to how it should be printed. - if (name_hint != print_name) return; - // Free the conresponding name_alloc_map_ index - --it2->second; -} - -Doc TVMScriptPrinter::PrintMatchBufferRegion(const MatchBufferRegionNode* op) { - const Buffer& buf = op->buffer; - buf_not_in_headers_.insert(buf.get()); - - Doc doc = Print(op->buffer) << " = " << tir_prefix_ << ".match_buffer(" << Print(op->source) - << ", " << memo_buf_decl_[op->buffer] << ")"; - return doc; -} - -// check if all arguments, except the first two, are specified for T.match_buffer -// if not, then this match buffer is printed out as T.buffer in prim_func arguments -// and check whether there are undefined variables in the shape/strides. -bool TVMScriptPrinter::IsSimpleBuffer(const Buffer& buf) { - if (memo_var_.find(buf->data) != memo_var_.end()) { - return false; - } - if (!buf->strides.empty()) { - return false; - } - for (const PrimExpr& shp_i : buf->shape) { - if (!UndefinedVars(shp_i).empty()) { - return false; - } - } - for (const PrimExpr& stride_i : buf->strides) { - if (!UndefinedVars(stride_i).empty()) { - return false; - } - } - if (!UndefinedVars(buf->elem_offset).empty()) { - return false; - } else if (buf->elem_offset->IsInstance()) { - IntImm elem_offset = Downcast(buf->elem_offset); - if (elem_offset->value != 0) { - return false; - } - } - if (buf.scope() != "global") { - return false; - } - if (buf->data_alignment != runtime::kAllocAlignment) { - return false; - } - if (buf->offset_factor != 1) { - return false; - } - if (buf->buffer_type != BufferType::kDefault) { - return false; - } - if (buf->axis_separators.size()) { - return false; - } - return true; -} - -Doc TVMScriptPrinter::PrintInlineBufferBind(const Buffer& buffer) { - Doc doc; - doc << tir_prefix_ << ".Buffer["; - if (buffer->shape.size() == 1) { - doc << Print(buffer->shape[0]); - } else { - doc << PrintTuple(buffer->shape.as()); - } - doc << ", " << PrintDType(buffer->dtype) << "]"; - return doc; -} - -// print array out as tuple with parentheses -Doc TVMScriptPrinter::PrintTuple(const ArrayNode* op) { - Doc doc; - doc << '('; - for (size_t i = 0; i < op->size(); ++i) { - if (i != 0) { - doc << ", "; - } - doc << Print(op->at(i)); - } - if (op->size() == 1) doc << ","; - doc << ')'; - return doc; -} - -Doc TVMScriptPrinter::PrintCommReducer(const CommReducerNode* op) { - Doc doc; - int n_var = static_cast(op->rhs.size()); - - doc << tir_prefix_ << ".comm_reducer(lambda "; - for (const tir::Var& v_lhs : op->lhs) { - doc << Print(v_lhs) << ", "; - } - for (int i = 0; i < n_var; ++i) { - doc << Print(op->rhs[i]) << (i == n_var - 1 ? ": " : ", "); - } - if (n_var == 1) { - doc << Print(op->result[0]) << ", "; - } else { - doc << "("; - for (int i = 0; i < n_var; ++i) { - doc << Print(op->result[i]); - if (i != n_var - 1) { - doc << ", "; - } - } - doc << "), "; - } - doc << Print(op->identity_element) << ")"; - - // Remove the vars in `lhs` and `rhs`, because they are the parameters of the printed lambda. - for (int i = 0; i < n_var; ++i) { - memo_var_.erase(op->lhs[i]); - memo_var_.erase(op->rhs[i]); - } - return doc; -} - -Doc TVMScriptPrinter::Print(const ObjectRef& node) { - if (!node.defined()) return Doc::Text("None"); - if (node->IsInstance()) { - return PrintOptionalInfo(Downcast(node)) << VisitStmt(Downcast(node)); - } else if (node->IsInstance()) { - ExprPrecedence t = ExprPrecedence::kUnknown; - return VisitExpr(Downcast(node), &t); - } else if (node->IsInstance()) { - return VisitType(Downcast(node)); - } else if (node->IsInstance()) { - return PrintPrimFunc(Downcast(node)); - } else if (node->IsInstance()) { - return PrintIRModule(Downcast(node)); - } else if (node->IsInstance()) { - return PrintArray(node.as()); - } else if (node->IsInstance()) { - return PrintBuffer(node.as()); - } else if (node->IsInstance()) { - return PrintString(node.as()); - } else if (node->IsInstance()) { - return PrintIterVar(node.as()); - } else if (node->IsInstance()) { - return PrintRange(node.as()); - } else if (node->IsInstance()) { - return PrintBufferRegion(node.as()); - } else if (node->IsInstance()) { - return PrintMatchBufferRegion(node.as()); - } else if (node->IsInstance()) { - return PrintCommReducer(node.as()); - } else if (node->IsInstance()) { - return PrintTarget(node.as()); - } else { - LOG(FATAL) << "Do not know how to print " << node->GetTypeKey(); - } -} - -Doc TVMScriptPrinter::VisitExprDefault_(const Object* op, ExprPrecedence* out_precedence) { - LOG(FATAL) << "Do not know how to print " << op->GetTypeKey(); -} - -Doc TVMScriptPrinter::VisitStmtDefault_(const Object* op) { - LOG(FATAL) << "Do not know how to print " << op->GetTypeKey(); -} - -Doc TVMScriptPrinter::VisitExpr_(const IntImmNode* op, ExprPrecedence* out_precedence) { - *out_precedence = ExprPrecedence::kIdentity; - return PrintConstScalar(op->dtype, &(op->value)); -} - -Doc TVMScriptPrinter::VisitExpr_(const FloatImmNode* op, ExprPrecedence* out_precedence) { - *out_precedence = ExprPrecedence::kIdentity; - return PrintConstScalar(op->dtype, &(op->value)); -} - -Doc TVMScriptPrinter::VisitExpr_(const StringImmNode* op, ExprPrecedence* out_precedence) { - *out_precedence = ExprPrecedence::kIdentity; - return Doc::StrLiteral(op->value); -} - -Doc TVMScriptPrinter::VisitExpr_(const CastNode* op, ExprPrecedence* out_precedence) { - *out_precedence = ExprPrecedence::kIdentity; - Doc doc; - doc << tir_prefix_ << ".Cast(" << PrintDType(op->dtype) << ", " << Print(op->value) << ")"; - return doc; -} - -Doc TVMScriptPrinter::VisitExpr_(const tir::VarNode* op, ExprPrecedence* out_precedence) { - *out_precedence = ExprPrecedence::kIdentity; - const tir::Var& var = GetRef(op); - return meta_.InMeta(var) ? meta_.GetMetaNode(var) : AllocVar(GetRef(op)); -} - -bool WillPrintConstScalar(const PrimExpr& expr) { - if (const auto* imm = expr.as()) { - DataType dtype = imm->dtype; - return dtype == DataType::Int(32) || dtype == DataType::Bool(); - } - return false; -} - -#define TVM_DECLARE_TVMSCRIPT_PRINTER_BINOP(OpName, OpString, OpClass, OpPrecedence) \ - Doc TVMScriptPrinter::VisitExpr_(const OpName* op, ExprPrecedence* out_precedence) { \ - Doc doc; \ - if (WillPrintConstScalar(op->a) && WillPrintConstScalar(op->b)) { \ - *out_precedence = ExprPrecedence::kIdentity; \ - doc << tir_prefix_ << "." << OpClass << "(" << Print(op->a) << ", " << Print(op->b) << ")"; \ - return doc; \ - } \ - ExprPrecedence lhs_precedence = ExprPrecedence::kUnknown; \ - ExprPrecedence rhs_precedence = ExprPrecedence::kUnknown; \ - /* Get children expr out_precedence */ \ - Doc lhs_doc = VisitExpr(op->a, &lhs_precedence); \ - Doc rhs_doc = VisitExpr(op->b, &rhs_precedence); \ - ICHECK(lhs_precedence != ExprPrecedence::kUnknown); \ - ICHECK(rhs_precedence != ExprPrecedence::kUnknown); \ - /* Update out_precedence of current node. */ \ - *out_precedence = OpPrecedence; \ - if (lhs_precedence > OpPrecedence || \ - (lhs_precedence == ExprPrecedence::kAnd && OpPrecedence == ExprPrecedence::kOr)) { \ - doc << "(" << lhs_doc << ")"; \ - } else { \ - doc << lhs_doc; \ - } \ - doc << OpString; \ - if (rhs_precedence >= OpPrecedence || \ - (rhs_precedence == ExprPrecedence::kAnd && OpPrecedence == ExprPrecedence::kOr)) { \ - doc << "(" << rhs_doc << ")"; \ - } else { \ - doc << rhs_doc; \ - } \ - return doc; \ - } - -TVM_DECLARE_TVMSCRIPT_PRINTER_BINOP(MulNode, " * ", "Mul", ExprPrecedence::kMultiplicationDivision) -TVM_DECLARE_TVMSCRIPT_PRINTER_BINOP(DivNode, " / ", "Div", ExprPrecedence::kMultiplicationDivision) -TVM_DECLARE_TVMSCRIPT_PRINTER_BINOP(FloorDivNode, " // ", "FloorDiv", - ExprPrecedence::kMultiplicationDivision) -TVM_DECLARE_TVMSCRIPT_PRINTER_BINOP(FloorModNode, " % ", "FloorMod", - ExprPrecedence::kMultiplicationDivision) -TVM_DECLARE_TVMSCRIPT_PRINTER_BINOP(AddNode, " + ", "Add", ExprPrecedence::kAdditionSubtraction) -TVM_DECLARE_TVMSCRIPT_PRINTER_BINOP(SubNode, " - ", "Sub", ExprPrecedence::kAdditionSubtraction) -TVM_DECLARE_TVMSCRIPT_PRINTER_BINOP(LTNode, " < ", "LT", ExprPrecedence::kRelational) -TVM_DECLARE_TVMSCRIPT_PRINTER_BINOP(LENode, " <= ", "LE", ExprPrecedence::kRelational) -TVM_DECLARE_TVMSCRIPT_PRINTER_BINOP(GTNode, " > ", "GT", ExprPrecedence::kRelational) -TVM_DECLARE_TVMSCRIPT_PRINTER_BINOP(GENode, " >= ", "GE", ExprPrecedence::kRelational) -TVM_DECLARE_TVMSCRIPT_PRINTER_BINOP(EQNode, " == ", "EQ", ExprPrecedence::kEquality) -TVM_DECLARE_TVMSCRIPT_PRINTER_BINOP(NENode, " != ", "NE", ExprPrecedence::kEquality) -TVM_DECLARE_TVMSCRIPT_PRINTER_BINOP(AndNode, " and ", "And", ExprPrecedence::kAnd) -TVM_DECLARE_TVMSCRIPT_PRINTER_BINOP(OrNode, " or ", "Or", ExprPrecedence::kOr) - -Doc TVMScriptPrinter::VisitExpr_(const ModNode* op, ExprPrecedence* out_precedence) { - *out_precedence = ExprPrecedence::kIdentity; - Doc doc; - doc << tir_prefix_ << ".truncmod(" << Print(op->a) << ", " << Print(op->b) << ")"; - return doc; -} - -Doc TVMScriptPrinter::VisitExpr_(const MinNode* op, ExprPrecedence* out_precedence) { - *out_precedence = ExprPrecedence::kIdentity; - Doc doc; - doc << tir_prefix_ << ".min(" << Print(op->a) << ", " << Print(op->b) << ")"; - return doc; -} - -Doc TVMScriptPrinter::VisitExpr_(const MaxNode* op, ExprPrecedence* out_precedence) { - *out_precedence = ExprPrecedence::kIdentity; - Doc doc; - doc << tir_prefix_ << ".max(" << Print(op->a) << ", " << Print(op->b) << ")"; - return doc; -} - -Doc TVMScriptPrinter::VisitExpr_(const NotNode* op, ExprPrecedence* out_precedence) { - *out_precedence = ExprPrecedence::kIdentity; - Doc doc; - doc << "not(" << Print(op->a) << ")"; - return doc; -} - -Doc TVMScriptPrinter::VisitExpr_(const SelectNode* op, ExprPrecedence* out_precedence) { - *out_precedence = ExprPrecedence::kIdentity; - Doc doc; - doc << tir_prefix_ << ".Select(" << Print(op->condition) << ", " << Print(op->true_value) << ", " - << Print(op->false_value) << ")"; - return doc; -} - -Doc TVMScriptPrinter::VisitExpr_(const ProducerLoadNode* op, ExprPrecedence* out_precedence) { - LOG(FATAL) << "Cannot print a tir.ProducerLoad as it is not valid in TIR Primfuncs. You need to " - "lower this function first."; - return Doc(); -} - -Doc TVMScriptPrinter::VisitExpr_(const BufferLoadNode* op, ExprPrecedence* out_precedence) { - *out_precedence = ExprPrecedence::kIdentity; - Doc doc; - if (op->indices.size() == 0) { - doc << Print(op->buffer) << "[()]"; - } else { - doc << Print(op->buffer) << PrintBufferIndices(op->indices); - } - return doc; -} - -Doc TVMScriptPrinter::VisitExpr_(const RampNode* op, ExprPrecedence* out_precedence) { - *out_precedence = ExprPrecedence::kIdentity; - Doc doc; - doc << tir_prefix_ << ".ramp(" << Print(op->base) << ", " << Print(op->stride) << ", " - << Print(op->lanes) << ")"; - return doc; -} - -Doc TVMScriptPrinter::VisitExpr_(const BroadcastNode* op, ExprPrecedence* out_precedence) { - *out_precedence = ExprPrecedence::kIdentity; - Doc doc; - doc << tir_prefix_ << ".broadcast(" << Print(op->value) << ", " << Print(op->lanes) << ")"; - return doc; -} - -Doc TVMScriptPrinter::VisitExpr_(const tir::LetNode* op, ExprPrecedence* out_precedence) { - *out_precedence = ExprPrecedence::kIdentity; - Doc doc; - doc << tir_prefix_ << ".let(" << Print(op->var) << ", " << Print(op->value) << ", " - << Print(op->body) << ")"; - return doc; -} - -Doc TVMScriptPrinter::VisitExpr_(const tir::CallNode* op, ExprPrecedence* out_precedence) { - *out_precedence = ExprPrecedence::kIdentity; - Doc doc; - if (auto* ptr_op = op->op.as()) { - std::string name = ptr_op->name; - if (name.find("tir.") == 0) { - name = tir_prefix_ + "." + name.substr(4); - } - doc << name << "("; - } else { - auto* op_gvar = op->op.as(); - ICHECK(op_gvar != nullptr); - doc << Doc::Text(op_gvar->name_hint) << "("; - } - std::vector args; - for (const auto& arg : op->args) { - args.push_back(Print(arg)); - } - args.push_back(Doc::Text("dtype=") << PrintDType(op->dtype)); - doc << PrintSep(args, Doc::Text(", ")) << ")"; - return doc; -} - -Doc TVMScriptPrinter::VisitExpr_(const ShuffleNode* op, ExprPrecedence* out_precedence) { - *out_precedence = ExprPrecedence::kIdentity; - Doc doc; - doc << tir_prefix_ << ".shuffle(" << Print(op->vectors) << ", " << Print(op->indices) << ")"; - return doc; -} - -Doc TVMScriptPrinter::VisitExpr_(const ReduceNode* op, ExprPrecedence* out_precedence) { - *out_precedence = ExprPrecedence::kIdentity; - Doc doc; - doc << tir_prefix_ << ".reduce(" << Print(op->combiner) << ", " << Print(op->source) << ", " - << Print(op->axis) << ", " << op->value_index << ")"; - return doc; -} - -Doc TVMScriptPrinter::VisitStmt_(const LetStmtNode* op) { - if (!buffer_var_usage_.count(op->var)) { - buffer_var_usage_ = BufferUsageFinder::FindUsage(std::move(buffer_var_usage_), op->body); - } - Array buffer_usage = buffer_var_usage_.Get(op->var).value_or({}); - - Doc doc; - if (current_num_ != num_child_ - 1) { - doc << "with " << tir_prefix_ << ".let(" << Print(op->var) << ", " << Print(op->value) << "):"; - doc << Doc::Indent( - 4, Doc::NewLine() << PrintNonHeaderBufferDeclarations(buffer_usage) << PrintBody(op->body)); - } else { - if (memo_var_.find(op->var) == memo_var_.end()) var_not_in_headers_.insert(op->var.get()); - doc << Print(op->var) << ": " << Print(GetType(op->var)) << " = " << Print(op->value) - << Doc::NewLine(); - doc << PrintNonHeaderBufferDeclarations(buffer_usage) << PrintBody(op->body); - } - return doc; -} - -Doc TVMScriptPrinter::VisitStmt_(const AttrStmtNode* op) { - Doc doc; - if (op->node.defined()) { - // merge attr with realize when possible - if (op->node->IsInstance() && op->attr_key == "realize_scope" && - op->body->IsInstance()) { - const auto* realize = Downcast(op->body).get(); - if (realize->buffer.same_as(op->node)) { - if (current_num_ != num_child_ - 1) { - doc << "with " << tir_prefix_ << ".realize(" << Print(realize->buffer) - << Print(realize->bounds) << ", " << Print(op->value); - if (!is_one(realize->condition)) { - doc << ", " << Print(realize->condition); - } - doc << "):" << Doc::Indent(4, Doc::NewLine() << PrintBody(realize->body)); - } else { - doc << tir_prefix_ << ".realize(" << Print(realize->buffer) << Print(realize->bounds) - << ", " << Print(op->value); - if (!is_one(realize->condition)) { - doc << ", " << Print(realize->condition); - } - doc << ")" << Doc::NewLine() << PrintBody(realize->body); - } - return doc; - } - } - // concise thread env - if (op->node->IsInstance() && - (op->attr_key == "thread_extent" || op->attr_key == "virtual_thread")) { - const auto* iter_var = Downcast(op->node).get(); - var_not_in_headers_.insert(iter_var->var.get()); - var_env_map_[iter_var->var] = iter_var->thread_tag; - if (current_num_ != num_child_ - 1) { - doc << "with " << tir_prefix_ << ".launch_thread(" << Print(iter_var->var) << ", " - << Print(op->value) << "):"; - doc << Doc::Indent(4, Doc::NewLine() << PrintBody(op->body)); - } else { - doc << tir_prefix_ << ".launch_thread(" << Print(iter_var->var) << ", " << Print(op->value) - << ")"; - doc << Doc::NewLine() << PrintBody(op->body); - } - return doc; - } - } - // default - if (current_num_ != num_child_ - 1) { - doc << "with " << tir_prefix_ << ".attr(" << Print(op->node) << ", " - << Doc::StrLiteral(op->attr_key) << ", " << Print(op->value) << "):"; - doc << Doc::Indent(4, Doc::NewLine() << PrintBody(op->body)); - } else { - doc << tir_prefix_ << ".attr(" << Print(op->node) << ", " << Doc::StrLiteral(op->attr_key) - << ", " << Print(op->value) << ")"; - doc << Doc::NewLine() << PrintBody(op->body); - } - return doc; -} - -Doc TVMScriptPrinter::VisitStmt_(const AssertStmtNode* op) { - Doc doc; - if (current_num_ != num_child_ - 1) { - doc << "with " << tir_prefix_ << ".Assert(" << Print(op->condition) << ", " - << Print(op->message) << "):"; - doc << Doc::Indent(4, Doc::NewLine() << PrintBody(op->body)); - } else { - doc << "assert " << Print(op->condition) << ", " << Print(op->message); - doc << Doc::NewLine() << PrintBody(op->body); - } - return doc; -} - -Doc TVMScriptPrinter::VisitStmt_(const BufferRealizeNode* op) { - LOG(FATAL) - << "TVM Script Printer Internal Error: All the BufferRealize should be folded with Attr"; - return Doc(); -} - -namespace { - -bool IsAllocateDeclBufferPattern(const AllocateNode* allocate) { - const tir::Var& buffer_var = allocate->buffer_var; - const DeclBufferNode* decl_buffer = allocate->body.as(); - if (!decl_buffer) { - return false; - } - const Buffer& buffer = decl_buffer->buffer; - if (!buffer_var.same_as(buffer->data)) { - return false; - } - if (allocate->dtype != buffer->dtype) { - return false; - } - if (!is_one(allocate->condition)) { - return false; - } - if (allocate->annotations.size()) { - return false; - } - if (allocate->extents.size() != buffer->shape.size()) { - return false; - } - tir::ExprDeepEqual expr_equal; - for (size_t i = 0, n = allocate->extents.size(); i < n; ++i) { - if (!expr_equal(allocate->extents[i], buffer->shape[i])) { - return false; - } - } - return true; -} - -} // namespace - -Doc TVMScriptPrinter::VisitStmt_(const AllocateNode* op) { - var_not_in_headers_.insert(op->buffer_var.get()); - - if (!buffer_var_usage_.count(op->buffer_var)) { - buffer_var_usage_ = BufferUsageFinder::FindUsage(std::move(buffer_var_usage_), op->body); - } - Array buffer_usage = buffer_var_usage_.Get(op->buffer_var).value_or({}); - - if (buffer_usage.empty()) { - if (IsAllocateDeclBufferPattern(op)) { - // As a syntax sugar, we identify the pattern of Allocate and DeclBuffer and print a single - // DeclBuffer statement. It is intentionally to call `Print` instead of `PrintBody` here to - // delegate the printing of the current node to `DeclBufferNode` while maintaining the - // same value of `current_num_` and `num_child_`. - return Print(op->body); - } - } - - auto storage_scope = GetPtrStorageScope(op->buffer_var); - Doc func_call; - func_call << tir_prefix_ << ".allocate(" << Print(op->extents) << ", " << PrintDType(op->dtype) - << ", " << Print(storage_scope); - if (!is_one(op->condition)) { - func_call << ", " << Print(op->condition); - } - if (!op->annotations.empty()) { - func_call << ", annotations={"; - func_call << PrintAnnotations(op->annotations); - func_call << "}"; - } - func_call << ")"; - - Doc doc; - if (current_num_ != num_child_ - 1) { - doc << "with " << func_call << " as " << Print(op->buffer_var) << ":"; - doc << Doc::Indent( - 4, Doc::NewLine() << PrintNonHeaderBufferDeclarations(buffer_usage) << PrintBody(op->body)); - } else { - doc << Print(op->buffer_var) << " = " << func_call << Doc::NewLine(); - doc << PrintNonHeaderBufferDeclarations(buffer_usage) << PrintBody(op->body); - } - TryDeallocVar(op->buffer_var); - return doc; -} - -Doc TVMScriptPrinter::VisitStmt_(const AllocateConstNode* alloc) { - std::stringstream ss; - ICHECK(alloc->data) << "Should be presented"; - const auto& data = alloc->data.value(); - - if (alloc->dtype.is_int()) { - if (alloc->dtype.bits() == 8) { - NDArrayToTIR(data, ss); - } else if (alloc->dtype.bits() == 16) { - NDArrayToTIR(data, ss); - } else if (alloc->dtype.bits() == 32) { - NDArrayToTIR(data, ss); - } else if (alloc->dtype.bits() == 64) { - NDArrayToTIR(data, ss); - } else { - LOG(FATAL) << "DataType not supported"; - } - } else if (alloc->dtype.is_uint()) { - if (alloc->dtype.bits() == 8) { - NDArrayToTIR(data, ss); - } else if (alloc->dtype.bits() == 16) { - NDArrayToTIR(data, ss); - } else if (alloc->dtype.bits() == 32) { - NDArrayToTIR(data, ss); - } else if (alloc->dtype.bits() == 64) { - NDArrayToTIR(data, ss); - } else { - LOG(FATAL) << "DataType not supported"; - } - } else if (alloc->dtype.is_float()) { - if (alloc->dtype.bits() == 16) { - NDArrayToTIR(data, ss); - } else if (alloc->dtype.bits() == 32) { - NDArrayToTIR(data, ss); - } else if (alloc->dtype.bits() == 64) { - NDArrayToTIR(data, ss); - } else { - LOG(FATAL) << "DataType not supported"; - } - } else { - LOG(FATAL) << "DataType not supported"; - } - auto ndarray_str = ss.str(); - - var_not_in_headers_.insert(alloc->buffer_var.get()); - - if (!buffer_var_usage_.count(alloc->buffer_var)) { - buffer_var_usage_ = BufferUsageFinder::FindUsage(std::move(buffer_var_usage_), alloc->body); - } - Array buffer_usage = buffer_var_usage_.Get(alloc->buffer_var).value_or({}); - - Doc func_call; - func_call << tir_prefix_ << ".allocate_const(" << ndarray_str << ", " << PrintDType(alloc->dtype) - << ", " << Print(alloc->extents) << ")"; - - Doc doc; - var_not_in_headers_.insert(alloc->buffer_var.get()); - if (current_num_ != num_child_ - 1) { - doc << "with " << func_call << " as " << Print(alloc->buffer_var) << ":"; - doc << Doc::Indent(4, Doc::NewLine() << PrintNonHeaderBufferDeclarations(buffer_usage) - << PrintBody(alloc->body)); - } else { - doc << Print(alloc->buffer_var) << " = " << func_call << Doc::NewLine(); - doc << PrintNonHeaderBufferDeclarations(buffer_usage) << PrintBody(alloc->body); - } - return doc; -} - -Doc TVMScriptPrinter::VisitStmt_(const DeclBufferNode* op) { - const Buffer& buffer = op->buffer; - buf_not_in_headers_.insert(buffer.get()); - Doc buffer_name = Print(op->buffer); - Doc func_call; - func_call << tir_prefix_ << ".decl_buffer(" << memo_buf_decl_.at(buffer) << ")"; - - Doc doc; - if (current_num_ != num_child_ - 1) { - doc << "with " << func_call << " as " << buffer_name << ":"; - doc << Doc::Indent(4, Doc::NewLine() << PrintBody(op->body)); - } else { - doc << buffer_name << " = " << func_call << Doc::NewLine(); - doc << PrintBody(op->body); - } - return doc; -} - -Doc TVMScriptPrinter::VisitStmt_(const IfThenElseNode* op) { - Doc doc; - doc << "if " << Print(op->condition) << ":"; - doc << Doc::Indent(4, Doc::NewLine() << PrintBody(op->then_case)); - - Optional else_case = op->else_case; - while (else_case) { - if (auto* else_if = else_case.value().as()) { - doc << Doc::NewLine(); - doc << "elif " << Print(else_if->condition) << ":"; - doc << Doc::Indent(4, Doc::NewLine() << PrintBody(else_if->then_case)); - - else_case = else_if->else_case; - } else { - doc << Doc::NewLine(); - doc << "else:" << Doc::Indent(4, Doc::NewLine() << PrintBody(else_case.value())); - break; - } - } - - return doc; -} - -Doc TVMScriptPrinter::VisitStmt_(const SeqStmtNode* op) { - std::vector stmts; - for (Stmt stmt : op->seq) { - stmts.push_back(Print(stmt)); - } - return PrintSep(stmts, Doc::NewLine()); -} - -Doc TVMScriptPrinter::VisitStmt_(const EvaluateNode* op) { - // When parsing TVMScript, a PrimExpr that occurs as a statement is - // automatically wrapped in `tir::Evaluate`. Therefore, when - // printing, it's only necessary to print the value. For - // readability, though, we still print T.evaluate() when the - // expression is something other than a call node. - Doc doc; - if (op->value.as()) { - doc << Print(op->value); - } else { - doc << tir_prefix_ << ".evaluate(" << Print(op->value) << ")"; - } - return doc; -} - -Doc TVMScriptPrinter::VisitStmt_(const ForNode* op) { - Doc doc; - var_not_in_headers_.insert(op->loop_var.get()); - loop_var_map_[op->loop_var.get()] = GetRef(op); - const auto* body = op->body.as(); - bool simple_loop = IsSimpleLoop(op); - if (simple_loop) simple_loop_stack_.push_back(GetRef(op)); - // It is a loop that can be compressed, let the loops below print it out - if (simple_loop && body != nullptr && IsSimpleLoop(body) && !DependOnPrevLoops(body)) { - doc << Print(GetRef(body)); - TryDeallocVar(op->loop_var); - loop_var_map_.erase(op->loop_var.get()); - return doc; - } - // It is a loop that can not be compressed - bool print_above = !simple_loop_stack_.empty(); - // print loops above if needed - if (print_above) { - doc << PrintLoopStack(); - simple_loop_stack_.clear(); - } - if (!simple_loop) { - // print current loop if needed - Doc current_loop; - current_loop << PrintLoop(GetRef(op)); - current_loop << Doc::Indent(4, Doc::NewLine() << PrintBody(op->body)); - doc << (print_above ? Doc::Indent(4, Doc::NewLine() << current_loop) : current_loop); - } else { - doc << Doc::Indent(4, Doc::NewLine() << PrintBody(op->body)); - } - TryDeallocVar(op->loop_var); - loop_var_map_.erase(op->loop_var.get()); - return doc; -} - -Doc TVMScriptPrinter::VisitStmt_(const PrefetchNode* op) { - Doc doc; - doc << tir_prefix_ << ".prefetch(" << Print(op->buffer) << ", " << Print(op->bounds) << ")"; - return doc; -} - -Doc TVMScriptPrinter::VisitStmt_(const WhileNode* op) { - Doc doc; - doc << "while " << Print(op->condition) << ":"; - doc << Doc::Indent(4, Doc::NewLine() << PrintBody(op->body)); - return doc; -} - -Doc TVMScriptPrinter::VisitType_(const PrimTypeNode* node) { - Doc doc; - doc << tir_prefix_ << "."; - if (node->dtype.is_void()) { - doc << "void"; - } else { - doc << runtime::DLDataType2String(node->dtype); - } - return doc; -} - -Doc TVMScriptPrinter::VisitType_(const PointerTypeNode* node) { - Doc doc; - doc << tir_prefix_ << ".Ptr["; - doc << Print(node->element_type); - if (!node->storage_scope.empty()) { - doc << ", " << Doc::StrLiteral(node->storage_scope); - } - doc << "]"; - return doc; -} - -Doc TVMScriptPrinter::VisitType_(const TupleTypeNode* node) { - if (node->fields.empty()) { - return Doc::Text("None"); - } else { - std::vector fields; - for (Type field : node->fields) { - fields.push_back(Print(field)); - } - return Doc::Text(tir_prefix_ + ".Tuple[") << Doc::Concat(fields) << "]"; - } -} - -Doc TVMScriptPrinter::VisitStmt_(const BufferStoreNode* op) { - Doc doc; - if (op->indices.size() == 0) { - doc << Print(op->buffer) << "[()] = " << Print(op->value); - } else { - doc << Print(op->buffer) << PrintBufferIndices(op->indices) << " = " << Print(op->value); - } - return doc; -} - -/*! Helper functions for block printing. */ -Doc TVMScriptPrinter::PrintBlockVar(const IterVar& iter_var, const PrimExpr& value) { - Doc doc; - doc << Print(iter_var->var) << " = " << tir_prefix_ << ".axis."; - switch (iter_var->iter_type) { - case kDataPar: - doc << "spatial"; - break; - case kCommReduce: - doc << "reduce"; - break; - case kOrdered: - doc << "scan"; - break; - case kOpaque: - doc << "opaque"; - break; - default: - LOG(FATAL) << "Unknown block var iter type: " << iter_var->iter_type; - break; - } - doc << "("; - const Range& dom = iter_var->dom; - if (is_zero(dom->min)) { - doc << Print(dom->extent); - } else { - doc << "(" << Print(dom->min) << ", " << Print(dom->min + dom->extent) << ")"; - } - doc << ", " << Print(value) << ")"; - return doc; -} - -Doc TVMScriptPrinter::PrintBlockVarRemaps() { - ICHECK(!block_var_remaps_.empty()); - if (block_var_remaps_.size() == 1) { - const IterVar& iter_var = block_var_remaps_[0].first; - const PrimExpr& value = block_var_remaps_[0].second; - return PrintBlockVar(iter_var, value); - } - Doc doc; - std::vector iter_vars, iter_values; - std::string iter_type; - for (const auto& pair : block_var_remaps_) { - const IterVar& iter_var = pair.first; - const PrimExpr& value = pair.second; - iter_vars.push_back(Print(iter_var->var)); - iter_values.push_back(Print(value)); - if (iter_var->iter_type == kDataPar) { - iter_type += "S"; - } else if (iter_var->iter_type == kCommReduce) { - iter_type += "R"; - } else { - ICHECK(false); - } - } - doc << PrintSep(iter_vars, Doc::Text(", ")) << " = " << tir_prefix_ << ".axis.remap(" - << Doc::StrLiteral(iter_type) << ", [" << PrintSep(iter_values, Doc::Text(", ")) << "])"; - return doc; -} - -Doc TVMScriptPrinter::PrintBlockPredicate(const BlockRealizeNode* op) { - Doc doc; - if (!is_one(op->predicate)) { - doc << Doc::NewLine() << tir_prefix_ << ".where(" << Print(op->predicate) << ")"; - } - return doc; -} - -Doc TVMScriptPrinter::PrintBlockVars(const BlockRealizeNode* op) { - Doc doc; - const auto* block_op = op->block.as(); - ICHECK_EQ(block_op->iter_vars.size(), op->iter_values.size()); - tir::ExprDeepEqual expr_equal; - - auto is_simple_remap = [this, &expr_equal](const IterVar& iter_var, - const PrimExpr& value) -> bool { - if (iter_var->iter_type != kDataPar && iter_var->iter_type != kCommReduce) return false; - if (!value->IsInstance()) return false; - const tir::Var& var = Downcast(value); - auto it = loop_var_map_.find(var.get()); - return it != loop_var_map_.end() && expr_equal(it->second->min, iter_var->dom->min) && - expr_equal(it->second->extent, iter_var->dom->extent); - }; - - for (size_t i = 0; i < block_op->iter_vars.size(); ++i) { - const IterVar& iter_var = block_op->iter_vars[i]; - const PrimExpr& value = op->iter_values[i]; - var_not_in_headers_.insert(iter_var->var.get()); - if (is_simple_remap(iter_var, value)) { - block_var_remaps_.push_back(std::make_pair(iter_var, value)); - } else { - if (!block_var_remaps_.empty()) { - doc << Doc::NewLine() << PrintBlockVarRemaps(); - block_var_remaps_.clear(); - } - doc << Doc::NewLine() << PrintBlockVar(iter_var, value); - } - } - if (!block_var_remaps_.empty()) { - doc << Doc::NewLine() << PrintBlockVarRemaps(); - block_var_remaps_.clear(); - } - return doc; -} - -Doc TVMScriptPrinter::PrintBlockAttr(const BlockRealizeNode* op) { - const auto* block_op = op->block.as(); - Doc block_attr_doc; - // print binding, read/write tensor region, annotations - block_attr_doc << Doc::NewLine() << tir_prefix_ << ".reads(" - << PrintExpandedArray(block_op->reads.as()) << ")"; - block_attr_doc << Doc::NewLine() << tir_prefix_ << ".writes(" - << PrintExpandedArray(block_op->writes.as()) << ")"; - if (!block_op->annotations.empty()) { - block_attr_doc << Doc::NewLine() << tir_prefix_ << ".block_attr({"; - block_attr_doc << PrintAnnotations(block_op->annotations); - block_attr_doc << "})"; - } - return block_attr_doc; -} - -// This function is to make sure arguments of T.reads() and T.writes() is not parsed by printer as a -// List. Therefore the brackets are removed before and after printing arguments out -Doc TVMScriptPrinter::PrintExpandedArray(const ArrayNode* op) { - Doc doc; - for (size_t i = 0; i < op->size(); ++i) { - if (i != 0) { - doc << ", "; - } - doc << Print(op->at(i)); - } - return doc; -} - -Doc TVMScriptPrinter::PrintBlockBody(const BlockNode* op) { - Doc body; - for (const auto& alloc_buf : op->alloc_buffers) { - buf_not_in_headers_.insert(alloc_buf.get()); - body << Print(alloc_buf) << " = " << tir_prefix_ << ".alloc_buffer(" - << memo_buf_decl_[alloc_buf] << ")" << Doc::NewLine(); - } - for (const auto& match_buf : op->match_buffers) { - body << Print(match_buf) << Doc::NewLine(); - } - if (op->init.defined()) { - Doc init_block; - init_block << "with " << tir_prefix_ << ".init():"; - init_block << Doc::Indent(4, Doc::NewLine() << PrintBody(op->init.value())); - body << init_block << Doc::NewLine(); - } - body << PrintBody(op->body); - return body; -} - -/*! - * \brief Print the name of a block - * \param block_op The block node to be printed - */ -Doc TVMScriptPrinter::PrintBlockName(const BlockNode* block_op) { - Doc doc; - doc << "with " << tir_prefix_ << ".block("; - if (!block_op->name_hint.empty()) { - doc << Doc::StrLiteral(block_op->name_hint); - } - doc << "):"; - return doc; -} - -Doc TVMScriptPrinter::VisitStmt_(const BlockRealizeNode* op) { - const auto* block_op = op->block.as(); - Doc doc = PrintOptionalInfo(GetRef(block_op)); - // print block name - doc << PrintBlockName(block_op); - // Print block predicate. - Doc block_predicate = PrintBlockPredicate(op); - // Print the variable bindings, valid to use in block attributes and - // body - Doc block_var = PrintBlockVars(op); - // print read/write tensor region, annotations - Doc block_attr_doc = PrintBlockAttr(op); - // print body - Doc body = PrintBlockBody(block_op); - doc << Doc::Indent(4, block_predicate << block_var << block_attr_doc << Doc::NewLine() << body); - for (const auto& iter_var : block_op->iter_vars) { - TryDeallocVar(iter_var->var); - } - return doc; -} - -Doc TVMScriptPrinter::PrintBody(const Stmt& body) { - int memo_num_child, memo_current_num; - std::swap(memo_num_child, num_child_); - std::swap(memo_current_num, current_num_); - - Doc doc; - if (body->IsInstance()) { - const auto& op = Downcast(body); - num_child_ = op->seq.size(); - current_num_ = 0; - std::vector stmts; - for (Stmt stmt : op->seq) { - stmts.push_back(Print(stmt)); - current_num_++; - } - doc = PrintSep(stmts, Doc::NewLine()); - } else { - num_child_ = 1; - current_num_ = 0; - doc = Print(body); - } - - std::swap(memo_num_child, num_child_); - std::swap(memo_current_num, current_num_); - return doc; -} - -Doc TVMScriptPrinter::PrintIRModule(const IRModule& module) { - auto* op = module.operator->(); - Doc doc; - doc << "@tvm.script.ir_module" << Doc::NewLine(); - doc << "class Module:"; - for (const auto& x : op->functions) { - func2var_[x.second.operator->()] = x.first; - } - Doc body = Doc::NewLine(); - std::vector functions; - for (auto it = op->functions.begin(); it != op->functions.end(); ++it) { - if ((*it).second.as()) { - functions.push_back(Print((*it).second)); - } - } - body << TVMScriptPrinter::PrintSep(functions, Doc::NewLine() << Doc::NewLine()); - body << Doc::NewLine() << DumpMeta(); - doc << Doc::Indent(4, body); - return doc; -} - -Doc TVMScriptPrinter::PrintPrimFunc(const PrimFunc& primFunc) { - auto* op = primFunc.operator->(); - // clear renaming map - memo_var_.clear(); - memo_buf_.clear(); - memo_buf_decl_.clear(); - var_not_in_headers_.clear(); - buf_not_in_headers_.clear(); - // print signature - Doc doc; - doc << "@" << tir_prefix_ << ".prim_func" << Doc::NewLine(); - doc << "def " << (func2var_.find(op) == func2var_.end() ? "func" : func2var_[op]->name_hint) - << "("; - std::vector params; - std::unordered_set simple_buf; - for (const auto& param : op->params) { - var_not_in_headers_.insert(param.get()); - auto it = op->buffer_map.find(param); - // check if this param is a T.handle - if (it != op->buffer_map.end()) { - // check if this match_buffer has only the first two arguments specified - // and whether the match_buffer is a dynamic buffer. - const Buffer& buf = (*it).second; - if (IsSimpleBuffer(buf)) { - simple_buf.insert(buf); - buf_not_in_headers_.insert(buf.get()); - params.push_back(Print(buf) << ": " << PrintInlineBufferBind(buf)); - continue; - } - } - params.push_back(Print(param) << ": " << Print(GetType(param))); - } - doc << PrintSep(params, Doc::Text(", ")) << ")"; - if (primFunc->ret_type.defined()) { - auto as_tuple = primFunc->ret_type.as(); - if (!as_tuple || as_tuple->fields.size()) { - doc << " -> " << Print(primFunc->ret_type); - } - } - doc << ":"; - - Doc body = Doc::NewLine(); - // print buffer_bind - for (const auto& param : op->params) { - auto it = op->buffer_map.find(param); - if (it == op->buffer_map.end()) continue; - const Buffer& buf = (*it).second; - if (simple_buf.count(buf)) continue; - buf_not_in_headers_.insert(buf.get()); - body << Print(buf) << " = " << tir_prefix_ << ".match_buffer("; - ICHECK(memo_buf_decl_.count(buf)); - body << Print((*it).first) << ", " << memo_buf_decl_[buf]; - body << ")" << Doc::NewLine(); - } - // print body - body << "# body" << Doc::NewLine(); - - Optional elided_root_block_body = [&]() -> Optional { - auto block_realize = op->body.as(); - if (!block_realize || block_realize->iter_values.size()) { - return NullOpt; - } - - const auto& block = block_realize->block; - if (block->annotations.size() || ContainsOptionalInfo(block)) { - return NullOpt; - } - - // The autocomplete might recognize the body itself as being a - // root block, and fail to insert it. - bool autocomplete_would_insert_root_block = [&]() -> bool { - if (block->alloc_buffers.size()) { - return true; - } - - auto* block_realize = block->body.as(); - if (block_realize && block_realize->block->iter_vars.size()) { - return true; - } - if (!block_realize && ContainsNode(block->body)) { - return true; - } - return false; - }(); - - if (autocomplete_would_insert_root_block) { - return block; - } else { - return NullOpt; - } - }(); - - if (elided_root_block_body) { - // Skip printing of root block in cases where tvm::tir::ScriptComplete - // would re-insert it. - body << "# with " << tir_prefix_ << ".block(\"root\")" << Doc::NewLine(); - body << PrintBlockBody(elided_root_block_body.value().get()); - } else { - // If this is a non-root block, or is an unskippable root block, - // just print it without skipping. - body << PrintBody(op->body); - } - - // print func attrs - Doc header_attr; - if (primFunc->attrs.defined()) { - header_attr << Doc::NewLine() << "# function attr dict" << Doc::NewLine() << tir_prefix_ - << ".func_attr({"; - std::vector attrs; - for (const auto& it : op->attrs->dict) { - attrs.push_back(Doc::StrLiteral(it.first) << ": " << Print(it.second)); - } - header_attr << PrintSep(attrs, Doc::Text(", ")) << "})"; - } - // print buffer declarations(buffers not defined by buffer_bind or buffer_allocate) - Doc header_buf; - std::vector bufs; - for (const auto& it : memo_buf_) { - if (buf_not_in_headers_.find(it.first.get()) == buf_not_in_headers_.end()) { - bufs.push_back(it.first.get()); - } - } - if (!bufs.empty()) { - header_buf << Doc::NewLine() << "# buffer definition"; - std::sort(bufs.begin(), bufs.end(), [&](const BufferNode* a, const BufferNode* b) { - return memo_buf_[GetRef(a)].str() < memo_buf_[GetRef(b)].str(); - }); - for (const auto& buf : bufs) { - header_buf << Doc::NewLine() << Print(GetRef(buf)) << " = " << tir_prefix_ - << ".buffer_decl("; - header_buf << memo_buf_decl_[GetRef(buf)] << ")"; - } - } - // print var declaration - Doc header_var; - std::vector vars; - for (const auto& it : memo_var_) { - if (var_not_in_headers_.find(it.first.get()) == var_not_in_headers_.end()) { - vars.push_back(it.first.get()); - } - } - if (!var_env_map_.empty()) { - header_var << Doc::NewLine() << "# var definition"; - for (const auto& it : var_env_map_) { - header_var << Doc::NewLine() << Print(it.first) << " = " << tir_prefix_ << ".env_thread(" - << Doc::StrLiteral(it.second) << ")"; - } - } - if (!vars.empty()) { - std::sort(vars.begin(), vars.end(), [&](const tir::VarNode* a, const tir::VarNode* b) { - return memo_var_[GetRef(a)].str() < memo_var_[GetRef(b)].str(); - }); - for (const auto& var : vars) { - auto type = GetRef(var)->type_annotation; - if (auto* ptr_type = type.as()) { - auto* prim_type = ptr_type->element_type.as(); - ICHECK(prim_type); - header_var << Doc::NewLine() << Print(GetRef(var)) << " = " << tir_prefix_ - << ".buffer_var("; - header_var << PrintDType(prim_type->dtype) << ", " - << Doc::StrLiteral(ptr_type->storage_scope) << ")"; - } else { - header_var << Doc::NewLine() << Print(GetRef(var)) << " = " << tir_prefix_ - << ".var("; - header_var << PrintDType(var->dtype) << ")"; - } - } - } - doc << Doc::Indent(4, header_attr << header_var << header_buf << body); - return doc; -} - -Doc TVMScriptPrinter::PrintArray(const ArrayNode* op) { - Doc doc; - doc << '['; - for (size_t i = 0; i < op->size(); ++i) { - if (i != 0) { - doc << ", "; - } - doc << Print(op->at(i)); - } - doc << ']'; - return doc; -} - -Doc TVMScriptPrinter::PrintIterVar(const IterVarNode* op) { - Doc doc; - doc << tir_prefix_ << ".iter_var(" << Print(op->var); - if (op->dom.defined()) { - doc << ", [" << Print(op->dom) << "], "; - } else { - doc << ", None, "; - } - doc << Doc::StrLiteral(IterVarType2String(op->iter_type)) << ", "; - doc << Doc::StrLiteral(op->thread_tag) << ")"; - return doc; -} - -Doc TVMScriptPrinter::PrintRange(const RangeNode* op) { - return Print(op->min) << ":" << Print(op->min + op->extent); -} - -Doc TVMScriptPrinter::PrintBuffer(const BufferNode* op) { - const Buffer& buffer = GetRef(op); - return meta_.InMeta(buffer) ? meta_.GetMetaNode(buffer) : AllocBuf(buffer); -} - -Doc TVMScriptPrinter::PrintBufferIndices(const Array& indices) { - Doc doc; - doc << '['; - for (size_t i = 0; i < indices.size(); ++i) { - if (i != 0) { - doc << ", "; - } - PrimExpr index = indices[i]; - if (const RampNode* ramp = index.as()) { - // specify ramp printing as python index slice - if (auto* stride_imm = ramp->stride.as()) { - doc << Print(ramp->base) << ":" << Print(ramp->base + ramp->lanes * ramp->stride); - if (stride_imm->value != 1) { - doc << ":" << Print(ramp->stride); - } - continue; - } - } - doc << Print(index); - } - doc << ']'; - return doc; -} - -Doc TVMScriptPrinter::PrintNonHeaderBufferDeclarations(const Array& aliasing_buffers) { - Doc decls; - for (const auto& buf_usage : aliasing_buffers) { - decls << Print(buf_usage) << " = " << tir_prefix_ << ".buffer_decl(" - << memo_buf_decl_[buf_usage] << ")" << Doc::NewLine(); - buf_not_in_headers_.insert(buf_usage.get()); - } - return decls; -} - -Doc TVMScriptPrinter::PrintBufferRegion(const BufferRegionNode* op) { - Doc doc; - if (op->region.size() == 0) { - doc << Print(op->buffer) << "[()]"; - } else { - doc << Print(op->buffer) << "["; - for (size_t i = 0; i < op->region.size(); ++i) { - if (i != 0) doc << ", "; - const auto& range = op->region[i]; - if (!is_one(range->extent)) { - doc << Print(range->min) << " : " << Print(ana_.Simplify(range->min + range->extent)); - } else { - doc << Print(range->min); - } - } - doc << "]"; - } - return doc; -} - -Doc TVMScriptPrinter::PrintAnnotations(const Map& annotations) { - Doc res; - std::vector> anno_list; - anno_list.reserve(annotations.size()); - for (const auto& pair : annotations) { - anno_list.emplace_back(pair); - } - sort(anno_list.begin(), anno_list.end()); - for (size_t i = 0; i < anno_list.size(); ++i) { - if (i != 0) { - res << ", "; - } - res << "\"" << anno_list[i].first << "\":" << Print(anno_list[i].second); - } - return res; -} - -Doc TVMScriptPrinter::PrintLoop(const For& loop) { - Doc res; - res << "for " << Print(loop->loop_var) << " in " << tir_prefix_ - << "." + std::string(ForKind2String(loop->kind)) + "("; - if (is_zero(loop->min)) { - res << Print(loop->extent); - } else { - res << Print(loop->min) << ", " << Print(ana_.Simplify(loop->min + loop->extent)); - } - if (loop->thread_binding.defined()) { - res << ", thread="; - res << Print(loop->thread_binding.value()->thread_tag); - } - if (!loop->annotations.empty()) { - res << ", annotations={"; - res << PrintAnnotations(loop->annotations); - res << "}"; - } - res << "):"; - return res; -} - -Doc TVMScriptPrinter::PrintLoopStack() { - Doc res; - if (simple_loop_stack_.size() == 1) { - res << PrintLoop(simple_loop_stack_[0]); - } else if (simple_loop_stack_.size() > 1) { - std::vector vars, extents; - for (const auto& loop : simple_loop_stack_) { - vars.push_back(Print(loop->loop_var)); - extents.push_back(Print(loop->extent)); - } - res << "for " << PrintSep(vars, Doc::Text(", ")) << " in " << tir_prefix_ << ".grid(" - << PrintSep(extents, Doc::Text(", ")) << "):"; - } - return res; -} - -Doc TVMScriptPrinter::PrintTarget(const TargetNode* target) { - Doc res; - res << tir_prefix_ << ".target({"; - Map config = target->Export(); - for (auto it = config.begin(); it != config.end(); ++it) { - if (it != config.begin()) { - res << ", "; - } - res << "\"" << (*it).first << "\":"; - if ((*it).first == "host") { - ICHECK(target->host.defined()); - res << PrintTarget(target->GetHost().value().get()); - } else { - res << Print((*it).second); - } - } - res << "})"; - return res; -} - -/*! - * \brief The printer for TVMScript with diagnostic - * \details The printer obtain the precedence of the top-level operation when printing each - * subexpression to decide whether or not parentheses is needed. - */ -class TVMScriptPrinterWithDiagnostic : public TVMScriptPrinter { - public: - explicit TVMScriptPrinterWithDiagnostic(const String& tir_prefix, bool show_meta, - runtime::TypedPackedFunc annotate) - : TVMScriptPrinter(tir_prefix, show_meta, annotate) {} - - protected: - Doc PrintBlockName(const BlockNode* block_op) override; - Doc PrintUnderline(const Stmt& stmt, int length); - Doc PrintLoop(const For& loop) override; -}; - -Doc TVMScriptPrinterWithDiagnostic::PrintBlockName(const BlockNode* block_op) { - Doc doc = TVMScriptPrinter::PrintBlockName(block_op); - doc << PrintUnderline(GetRef(block_op), doc.str().size()); - return doc; -} - -Doc TVMScriptPrinterWithDiagnostic::PrintUnderline(const Stmt& stmt, int length) { - Doc doc; - // annotation - if (ContainsOptionalInfo(stmt)) { - String underline = std::string(length, '^'); - doc << Doc::NewLine() << underline; - } - return doc; -} - -Doc TVMScriptPrinterWithDiagnostic::PrintLoop(const For& loop) { - Doc res = TVMScriptPrinter::PrintLoop(loop); - res << PrintUnderline(loop, res.str().size()); - return res; -} - -String AsTVMScriptWithDiagnostic(const ObjectRef& mod, const String& tir_prefix, bool show_meta, - runtime::TypedPackedFunc annotate) { - ICHECK(mod->IsInstance() || mod->IsInstance()); - Doc doc; - doc << TVMScriptPrinter::PrintHeader(tir_prefix) - << TVMScriptPrinterWithDiagnostic(tir_prefix, show_meta, annotate).Print(mod); - return doc.str() + "\n"; -} - -TVM_REGISTER_GLOBAL("script.AsTVMScriptWithDiagnostic").set_body_typed(AsTVMScriptWithDiagnostic); - -} // namespace relay -} // namespace tvm diff --git a/src/relay/qnn/op/add.cc b/src/relay/qnn/op/add.cc deleted file mode 100644 index 0e0d3fdbc0dd..000000000000 --- a/src/relay/qnn/op/add.cc +++ /dev/null @@ -1,104 +0,0 @@ -/* - * 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 src/relay/qnn/op/add.cc - * \brief QNN add operator. - */ -#include -#include - -#include "op_common.h" - -namespace tvm { -namespace relay { -namespace qnn { - -/* - * \brief Canonicalizes the QNN add op. - * \param attrs The empty attribute. - * \param new_args The new mutated args to the call node. - * \param arg_types The types of input and output. - * \return The sequence of Relay ops for add op. - */ -Expr QnnAddCanonicalize(const Attrs& attrs, const Array& new_args, - const Array& arg_types) { - // Get the args. - QnnBinaryOpArguments args(new_args); - - // Get the input dtype and shape. - QnnBinaryOpTensorType input_type(arg_types, 0); - - const auto* broadcast_attrs = attrs.as(); - ICHECK(broadcast_attrs != nullptr); - - auto lhs_axis = broadcast_attrs->lhs_axis; - auto rhs_axis = broadcast_attrs->rhs_axis; - - // FIXME (anijain2305) - The lowering can be further optimized. Instead of inserting requantize in - // the start, we can insert requantize at the end if both input tensors have same qnn params. In - // that case, we can first add the tensors, subtract the zero point, and requantize at the end. - // This can be done in future. - - // Since the input qnn params can be different than output qnn params, we first requantize the - // input tensors to the output qnn params. Then we call relay.add on the requantized inputs. This - // addition results in extra addition of the output zero point. We further subtract the zero - // point. The whole process can be represented using following equations - // - // scale_c * (Q_c - zp_c) = scale_a * (Q_a - zp_a) + scale_b * (Q_b - zp_b) - // - // After requantizing Q_a and Q_b, equation becomes, - // scale_c * (Q_c - zp_c) = scale_c * (Q_a' - zp_c) + scale_c * (Q_b' - zp_c) - // scale_c * (Q_c - zp_c) = scale_c * (Q_a' + Q_b' - zp_c - zp_c) - // - // Comparing the LHS and RHS, it results in - // Q_c = Q_a' + Q_b' - zp_c - // The add op is done in int32 precision. - - // Requantize LHS if necessary. Computes Q_a' - auto requantized_lhs = - RequantizeOrUpcast(args.lhs, args.lhs_scale, args.lhs_zero_point, args.output_scale, - args.output_zero_point, input_type.shape, lhs_axis); - // Requantize RHS if necessary. Computes Q_b' - auto requantized_rhs = - RequantizeOrUpcast(args.rhs, args.rhs_scale, args.rhs_zero_point, args.output_scale, - args.output_zero_point, input_type.shape, rhs_axis); - // Computes Q_a' + Q_b' - auto output = Add(requantized_lhs, requantized_rhs); - - // Subtract zero point. Computes (Q_a' + Q_b') - zp_c - auto zero_scalar = MakeConstantScalar(DataType::Int(32), 0); - if (!IsEqualScalar(args.output_zero_point, zero_scalar)) { - output = Subtract(output, args.output_zero_point); - } - - // Go back to lower precision. - return ConvertDtype(output, input_type.dtype); -} - -// QNN Addition operator. -QNN_REGISTER_BINARY_OP("add") - .describe("Elementwise add with broadcasting for quantized tensors.") - .set_support_level(11) - .set_attr("FTVMQnnCanonicalize", QnnAddCanonicalize) - .set_attr("TOpPattern", kBroadcast); - -} // namespace qnn -} // namespace relay -} // namespace tvm diff --git a/src/relay/qnn/op/avg_pool2d.cc b/src/relay/qnn/op/avg_pool2d.cc deleted file mode 100644 index e1a28169ccda..000000000000 --- a/src/relay/qnn/op/avg_pool2d.cc +++ /dev/null @@ -1,225 +0,0 @@ -/* - * 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 src/relay/qnn/op/avg_pool2d.cc - * \brief Quantized avg_pool2d operator - */ - -#include -#include -#include -#include -#include -#include - -#include "../../op/nn/nn.h" -#include "../../op/nn/pooling.h" -#include "../../op/nn/pooling_common.h" -#include "../../op/tensor/transform.h" -#include "../../transforms/infer_layout_utils.h" -#include "../../transforms/pattern_utils.h" -#include "../utils.h" -#include "op_common.h" - -namespace tvm { -namespace relay { -namespace qnn { - -// relay.op.qnn.avg_pool2d -bool QnnAvgPool2DRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - // Expected Types: data, input_zero_point, input_scale, output_zero_point, output_scale - // out_type - - ICHECK_EQ(types.size(), 6); - - const auto* data = types[0].as(); - if (data == nullptr) return false; - ICHECK(data->dtype == DataType::Int(8) || data->dtype == DataType::UInt(8)) - << "Expected quantized avg_pool2d type(int8, uint8) for input but was " << data->dtype; - - const auto* param = attrs.as(); - ICHECK(param != nullptr) << "AvgPool2DAttrs cannot be nullptr."; - - // Check the types of scale and zero points. - for (size_t i = 1; i < 5; ++i) { - if (types[i].as()) { - return false; - } - } - - ICHECK(IsScalarType(types[1], DataType::Float(32))); // input_scale - ICHECK(IsScalarType(types[2], DataType::Int(32))); // input_zero_point - ICHECK(IsScalarType(types[3], DataType::Float(32))); // output_scale - ICHECK(IsScalarType(types[4], DataType::Int(32))); // output_zero_point - - // Find the output shape and data type - const auto dshape = data->shape; - ICHECK_GE(dshape.size(), 2U) - << "Pool2D only support input >= 2-D: input must have height and width"; - - // Check input and output layout - Layout layout(param->layout); - // The Layout is always NHWC - ICHECK(layout.Contains(LayoutAxis::Get('H')) && layout.Contains(LayoutAxis::Get('W')) && - !layout.Contains(LayoutAxis::Get('h')) && !layout.Contains(LayoutAxis::Get('w'))) - << "Invalid input layout " << layout - << ". qnn_avg_pool2d inut layout must have H and W, which cannot be split"; - - // Find the output shape and data type - const auto hidx = layout.IndexOf(LayoutAxis::Get('H')); - const auto widx = layout.IndexOf(LayoutAxis::Get('W')); - - IndexExpr pad_h, pad_w; - if (param->padding.size() == 1) { - pad_h = param->padding[0] * 2; - pad_w = param->padding[0] * 2; - } else if (param->padding.size() == 2) { - // (top, left) - pad_h = param->padding[0] * 2; - pad_w = param->padding[1] * 2; - } else if (param->padding.size() == 4) { - // (top, left, bottom, right) - pad_h = param->padding[0] + param->padding[2]; - pad_w = param->padding[1] + param->padding[3]; - } else { - return false; - } - - std::vector oshape(dshape.begin(), dshape.end()); - if (dshape[hidx].as()) { - oshape[hidx] = dshape[hidx]; - } else { - oshape[hidx] = - calculate_pool_dimension(dshape[hidx], pad_h, param->pool_size[0], param->dilation[0], - param->strides[0], param->ceil_mode); - } - if (dshape[widx].as()) { - oshape[widx] = dshape[widx]; - } else { - oshape[widx] = - calculate_pool_dimension(dshape[widx], pad_w, param->pool_size[1], param->dilation[1], - param->strides[1], param->ceil_mode); - } - - // assign output type - reporter->Assign(types[5], TensorType(oshape, data->dtype)); - return true; -} - -InferCorrectLayoutOutput QnnAvgPoolInferCorrectLayout(const Attrs& attrs, - const Array& new_in_layouts, - const Array& old_in_layouts, - const Array& old_in_types) { - // Use Relay AvgPool2D Infer correct layout. - auto avgpool_new_layouts = - PoolInferCorrectLayout(attrs, new_in_layouts, old_in_layouts, old_in_types); - - // Scales and zero points are scalars, the layouts of these tensors can be treated as channel - // layout. - Layout channel_layout = Layout("C"); - Array input_layouts = {avgpool_new_layouts->input_layouts[0], channel_layout, - channel_layout, channel_layout, channel_layout}; - Array output_layouts = avgpool_new_layouts->output_layouts; - return InferCorrectLayoutOutput(input_layouts, output_layouts, attrs); -} - -/* - * \brief Forward rewrite the qnn avg_pool2d op. - * \param attrs The QNN avg_pool2d attrs. - * \param new_args The new mutated args to the call node. - * \param arg_types The types of input and output. - * \return The sequence of Relay ops for qnn avg_pool2d op. - * \note Lowering of the qnn.avg_pool2d operator - - * Quantized avg_pool2d will take one quantized input tensor and returns another - * quantized tensor. Since the input qnn params can be different from the output - * qnn params, first, we requantize the input tensors with output qnn params and - * cast the results into Int32. Then we call relay.nn.avg_pool2d on that requantized - * inputs. Finally, the results are cast into the quantized output data type. - - * Note: The RequantizeOrUpcast function only perform requantization if the input - * and output qnn params are different, otherwise it only does casting to Int32. - */ - -Expr QnnAvgPoolCanonicalize(const Attrs& attrs, const Array& new_args, - const Array& arg_types) { - ICHECK_EQ(new_args.size(), 5); - Expr input_data = new_args[0]; - Expr input_scale = new_args[1]; - Expr input_zero_point = new_args[2]; - Expr output_scale = new_args[3]; - Expr output_zero_point = new_args[4]; - const auto in_shape = get_shape(arg_types[0]); - const auto* avgpool_attrs = attrs.as(); - auto requantized_input = RequantizeOrUpcast(input_data, input_scale, input_zero_point, - output_scale, output_zero_point, in_shape); - Expr nn_avg = AvgPool2D(requantized_input, avgpool_attrs->pool_size, avgpool_attrs->strides, - avgpool_attrs->dilation, avgpool_attrs->padding, avgpool_attrs->layout, - avgpool_attrs->out_layout, avgpool_attrs->ceil_mode, - avgpool_attrs->count_include_pad); - - const auto* data = arg_types[5].as(); - const int32_t min_val = GetQmin(data->dtype); - const int32_t max_val = GetQmax(data->dtype); - return Cast(Clip(nn_avg, min_val, max_val), data->dtype); -} - -// Positional relay function to create quantized avg_pool2d operator used by frontend FFI. -Expr MakeQuantizedAvgPool2D(Expr data, Expr input_scale, Expr input_zero_point, Expr output_scale, - Expr output_zero_point, Array pool_size, - Array strides, Array padding, - Array dilation, bool ceil_mode, bool count_include_pad, - String layout, String output_layout) { - auto attrs = make_object(); - attrs->pool_size = std::move(pool_size); - attrs->strides = std::move(strides); - attrs->padding = std::move(padding); - attrs->dilation = std::move(dilation); - attrs->layout = std::move(layout); - attrs->out_layout = std::move(output_layout); - attrs->ceil_mode = ceil_mode; - attrs->count_include_pad = count_include_pad; - static const Op& op = Op::Get("qnn.avg_pool2d"); - return Call(op, {data, input_scale, input_zero_point, output_scale, output_zero_point}, - Attrs(attrs), {}); -} - -RELAY_REGISTER_OP("qnn.avg_pool2d") - .describe("Customized? qnn_avg_pool2d for quantized tensors.") - .set_attrs_type() - .set_num_inputs(5) - .add_argument("data", "Quantized Tensor", "The input data.") - .add_argument("input_scale", "Tensor", "The quantization scale of the input tensor.") - .add_argument("input_zero_point", "Tensor", "The quantization zero_point of the input tensor.") - .add_argument("output_scale", "Tensor", "The quantization scale of the output tensor.") - .add_argument("output_zero_point", "Tensor", - "The quantization zero_point of the output tensor.") - .set_support_level(11) - .add_type_rel("QnnAvgPool2D", QnnAvgPool2DRel) - .set_attr("TOpPattern", kOutEWiseFusable) - .set_attr("FInferCorrectLayout", QnnAvgPoolInferCorrectLayout) - .set_attr("FTVMQnnCanonicalize", QnnAvgPoolCanonicalize); - -TVM_REGISTER_GLOBAL("relay.qnn.op._make.avg_pool2d").set_body_typed(MakeQuantizedAvgPool2D); - -} // namespace qnn -} // namespace relay -} // namespace tvm diff --git a/src/relay/qnn/op/batch_matmul.cc b/src/relay/qnn/op/batch_matmul.cc deleted file mode 100644 index a948d9387d6b..000000000000 --- a/src/relay/qnn/op/batch_matmul.cc +++ /dev/null @@ -1,260 +0,0 @@ -/* - * 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 src/relay/qnn/op/batch_matmul.cc - * \brief Property def of qnn batch_matmul operator. - */ - -#include -#include -#include -#include - -#include "../../op/nn/nn.h" -#include "../../transforms/pattern_utils.h" -#include "../utils.h" - -namespace tvm { -namespace relay { -namespace qnn { - -// relay.op.qnn.batch_matmul - -bool QnnBatchMatmulRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - // Expected Types: x, y, x_zero_point, y_zero_point, x_scale, y_scale, - // out_type - ICHECK_EQ(types.size(), 7); - const auto* x = types[0].as(); - const auto* y = types[1].as(); - if (x == nullptr || y == nullptr) return false; - const auto* param = attrs.as(); - ICHECK(param != nullptr) << "BatchMatmulAttrs cannot be nullptr."; - ICHECK(x->dtype == DataType::Int(8) || x->dtype == DataType::UInt(8)) - << "Expected quantized batch_matmul type(int8, uint8) for input but was " << x->dtype; - ICHECK(y->dtype == DataType::Int(8) || y->dtype == DataType::UInt(8)) - << "Expected quantized batch_matmul type(int8, uint8) for weight but was " << y->dtype; - ICHECK(param->out_dtype == DataType::Int(32)) - << "Expected quantized batch_matmul type(int32) for output but was " << param->out_dtype; - - // Check the types of scale and zero points. - for (size_t i = 2; i < 5; ++i) { - if (types[i].as()) { - return false; - } - } - ICHECK(IsScalarType(types[2], DataType::Int(32))); // x_zero_point - ICHECK(IsScalarType(types[3], DataType::Int(32))); // y_zero_point - ICHECK(IsScalarType(types[4], DataType::Float(32))); // x_scale - ICHECK(IsScalarType(types[5], DataType::Float(32))); // y_scale - - ICHECK(param->out_dtype.bits() > 0) << "Output dtype bits should be greater than 0."; - - // Collect the input tensor and output tensor devoid of scale and zero points to reuse Relay - // BatchMatmul infer type function. - Array tensor_types = {types[0], types[1], types[6]}; - return BatchMatmulRel(tensor_types, 3, attrs, reporter); -} - -// Positional relay function to create quantized batch_matmul operator used by frontend FFI. -Expr MakeQuantizedBatchMatmul(Expr x, Expr y, Expr x_zero_point, Expr y_zero_point, Expr x_scale, - Expr y_scale, DataType out_dtype) { - auto attrs = make_object(); - attrs->out_dtype = out_dtype; - // For legacy reason, currently `qnn.batch_matmul` only supports - // (transpose_a=false, transpose_b=true) - // TODO(jcf94): extent to support all tensor format - attrs->transpose_a = false; - attrs->transpose_b = true; - static const Op& op = Op::Get("qnn.batch_matmul"); - return Call(op, {x, y, x_zero_point, y_zero_point, x_scale, y_scale}, Attrs(attrs), {}); -} - -Expr BatchMatmulFirstTerm(const Expr& quantized_x, const Expr& quantized_y, - const BatchMatmulAttrs* attrs) { - ICHECK(attrs->transpose_a == false && attrs->transpose_b == true) - << "Currently qnn.batch_matmul only supports (transpose_a=false, transpose_b=true)."; - return MakeBatchMatmul(quantized_x, quantized_y, attrs->out_dtype, attrs->transpose_a, - attrs->transpose_b); -} - -Expr BatchMatmulSecondTerm(const Expr& x_quantized_data, const Expr& y_zero_point) { - if (IsScalar(y_zero_point)) { - Array axes = {2}; - return Multiply(y_zero_point, - Sum(Cast(x_quantized_data, DataType::Int(32)), axes, true, false)); - } else { - LOG(FATAL) << "Tensor zero point (non-scalar) is not supported"; - return Expr(); - } -} - -Expr BatchMatmulThirdTerm(const Expr& y_quantized_data, const Expr& x_zero_point, - int broadcast_dim_size) { - if (IsScalar(x_zero_point)) { - Array axes = {2}; - auto reducemult = - Multiply(x_zero_point, Sum(Cast(y_quantized_data, DataType::Int(32)), axes, true, false)); - Array newshape; - - // dimension of 0 in reshape copies old dimension size - newshape = {0, 1, broadcast_dim_size}; - return Reshape(reducemult, newshape); - } else { - LOG(FATAL) << "Tensor zero point (non-scalar) is not supported"; - return Expr(); - } -} - -Expr BatchMatmulFourthTerm(Expr x_zero_point, Expr y_zero_point, int reduction_dim_size) { - if (IsScalar(x_zero_point) && IsScalar(y_zero_point)) { - auto zero_point_mul = Multiply(x_zero_point, y_zero_point); - auto const_scale = MakeConstantScalar(DataType::Int(32), reduction_dim_size); - return Multiply(zero_point_mul, const_scale); - } else { - LOG(FATAL) << "Tensor zero point (non-scalar) is not supported"; - return Expr(); - } -} - -Expr BatchMatmulFourthTerm(int x_zero_point_int, int y_zero_point_int, int reduction_dim_size) { - int32_t scalar_term = x_zero_point_int * y_zero_point_int * reduction_dim_size; - return MakeConstantScalar(DataType::Int(32), scalar_term); -} - -Expr BatchMatmulCombineTerms(const Expr& term1, const Expr& term2, const Expr& term3, - const Expr& term4) { - auto data1_term = Subtract(term1, term2); - auto data2_term = Subtract(term4, term3); - return Add(data1_term, data2_term); -} - -/* - * \brief Forward rewrite the qnn batch_matmul op. - * \param attrs The QNN batch_matmul attrs. - * \param new_args The new mutated args to the call node. - * \param arg_types The types of input and output. - * \return The sequence of Relay ops for qnn batch_matmul op. - * \note Lowering of the qnn.batch_matmul operator - * A quantized tensor is represented in following manner - * A = scale_a x (QA - zp_A) - * where QA is quantized tensor, scale_a and zp_A are quantization - * params. - * - * Quantized batch_matmul multiplies two quantized tensors and returns a - * quantized tensor of default dtype of int32, with scale equaling to the - * product of scales of input tensors, and a zero point of zero. - * - * The lowering for asymmetric quantized batch_matmul looks similar to - * quantized conv2d and dense and originally was discussed here: - * https://discuss.tvm.apache.org/t/tf-lite-quantized-conv2d-operator-conversion/2651/7 - * - * The computation gets unrolled into following 4 terms - * C(m, n) = Sigma(k) (X(m, k) * Y(n, k)) - * - * RHS becomes - * Sigma(k) ([QX(m, k) - zp_x] * [QY(n, k) - zp_y]) - * - * Unrolling leads to following sequence - * Sigma(k) QX(m, k) * QX(n, k) // Term1 - * - Sigma(k) zp_y * QX(m, k) // Term2 - * - Sigma(k) zp_x * QY(n, k) // Term3 - * - Sigma(k) * zp_x * zp_y // Term4 - * - * Term4 can be computed at compile time, everything else depending on the - * input type. - */ -Expr QnnBatchMatmulCanonicalize(const Attrs& attrs, const Array& new_args, - const Array& arg_types) { - ICHECK_EQ(new_args.size(), 6); - Expr quantized_x = new_args[0]; - Expr quantized_y = new_args[1]; - Expr x_zero_point = new_args[2]; - Expr y_zero_point = new_args[3]; - - const auto in_shape = get_shape(arg_types[0]); - const int reduction_dim_size = get_const_int(in_shape[2]); - - const auto y_shape = get_shape(arg_types[1]); - const int broadcast_dim_size = get_const_int(y_shape[1]); - - const auto* qnn_batch_matmul_attrs = attrs.as(); - - // Get all the terms as described in the comments. - auto term1 = BatchMatmulFirstTerm(quantized_x, quantized_y, qnn_batch_matmul_attrs); - auto term2 = BatchMatmulSecondTerm(quantized_x, y_zero_point); - auto term3 = BatchMatmulThirdTerm(quantized_y, x_zero_point, broadcast_dim_size); - - if (IsConstScalar(x_zero_point) && IsConstScalar(y_zero_point)) { - // Extract the integer zero points. - auto y_zero_point_int = GetScalarFromConstant(y_zero_point); - auto x_zero_point_int = GetScalarFromConstant(x_zero_point); - auto term4 = BatchMatmulFourthTerm(x_zero_point_int, y_zero_point_int, reduction_dim_size); - // Combine those 4 terms depending on the zero points to get the best lowering. - if (x_zero_point_int == 0 && y_zero_point_int == 0) { - // term 2, 3 and 4 become zero. - return term1; - } else if (x_zero_point_int == 0 && y_zero_point_int != 0) { - // term 3 and term 4 become zero. - return Subtract(term1, term2); - } else if (x_zero_point_int != 0 && y_zero_point_int == 0) { - // term 2 and term 4 become zero. - return Subtract(term1, term3); - } else { - return BatchMatmulCombineTerms(term1, term2, term3, term4); - } - } else { - auto term4 = BatchMatmulFourthTerm(x_zero_point, y_zero_point, reduction_dim_size); - return BatchMatmulCombineTerms(term1, term2, term3, term4); - } -} - -RELAY_REGISTER_OP("qnn.batch_matmul") - .describe(R"code(Compute batch matrix multiplication of `tensor_a` and `tensor_b`. - -Note we expect tensor_b to be transposed to copy the standard nn.batch_matmul conventions. - -.. math:: - - batch\_matmul(A, B)[i, :, :] = matmul(A[i, :, :], B[i, :, :]^T) - -- **data**: quantized(int8, unit8) `(i, m, k)` -- **weight**: quantized(int8, unit8) `(i, n, k)` -- **out**: quantized(int32) `(i, m, n)`. - -)code" TVM_ADD_FILELINE) - .set_attrs_type() - .set_num_inputs(6) - .add_argument("x", "quantized 2D Tensor", "First input data.") - .add_argument("y", "quantized 2D Tensor", "Second input data.") - .add_argument("x_scale", "Tensor", "The quantization scale of the x input tensor.") - .add_argument("x_zero_point", "Tensor", "The quantization zero_point of the x input tensor.") - .add_argument("y_scale", "Tensor", "The quantization scale of the y input tensor.") - .add_argument("y_zero_point", "Tensor", "The quantization zero_point of the y input tensor.") - .set_support_level(11) - .add_type_rel("QBatchMatmul", QnnBatchMatmulRel) - .set_attr("TNonComputational", true) - .set_attr("FTVMQnnCanonicalize", QnnBatchMatmulCanonicalize); - -TVM_REGISTER_GLOBAL("relay.qnn.op._make.batch_matmul").set_body_typed(MakeQuantizedBatchMatmul); - -} // namespace qnn -} // namespace relay -} // namespace tvm diff --git a/src/relay/qnn/op/concatenate.cc b/src/relay/qnn/op/concatenate.cc deleted file mode 100644 index e717d21bc440..000000000000 --- a/src/relay/qnn/op/concatenate.cc +++ /dev/null @@ -1,245 +0,0 @@ -/* - * 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 src/relay/qnn/op/concatenate.cc - * \brief QNN concatenate operator. It concatenates quantized input tensors along a given axis. - */ - -#include -#include -#include -#include - -#include "../../op/tensor/transform.h" -#include "../../transforms/infer_layout_utils.h" -#include "../../transforms/pattern_utils.h" -#include "../utils.h" - -namespace tvm { -namespace relay { -namespace qnn { - -bool QnnConcatenateRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - // Expected Types: data, input_scales, input_zero_points, output_scale, output_zero_point, - // out_type - ICHECK_EQ(types.size(), 6); - - if (types[0].as()) { - return false; - } - // Check the scale and zero point types - const auto* input_scales_tuple = types[1].as(); - if (input_scales_tuple == nullptr) { - if (types[1].as()) { - return false; - } else { - throw CompileError( - ErrorBuilder() - << "qnn concatenate requires a tuple of scales as the second argument, found " - << PrettyPrint(types[1])); - } - } - for (const auto& input_scale : input_scales_tuple->fields) { - if (input_scale.as()) { - return false; - } - ICHECK(IsScalarType(input_scale, DataType::Float(32))); // input_scales[idx] - } - - const auto* input_zero_points_tuple = types[2].as(); - if (input_zero_points_tuple == nullptr) { - if (types[2].as()) { - return false; - } else { - throw CompileError( - ErrorBuilder() - << "qnn concatenate requires a tuple of zero_points as the third argument, found " - << PrettyPrint(types[2])); - } - } - for (const auto& input_zero_point : input_zero_points_tuple->fields) { - if (input_zero_point.as()) { - return false; - } - ICHECK(IsScalarType(input_zero_point, DataType::Int(32))); // input_zero_points[idx] - } - - for (size_t i = 3; i < 5; ++i) { - if (types[i].as()) { - return false; - } - } - ICHECK(IsScalarType(types[3], DataType::Float(32))); // output_scale - ICHECK(IsScalarType(types[4], DataType::Int(32))); // output_zero_point - - // Collect the input tensor and output tensor devoid of scale and zero points to reuse Relay - // Concatenate infer type function. - Array tensor_types = {types[0], types[5]}; - return ConcatenateRel(tensor_types, 2, attrs, reporter); -} - -InferCorrectLayoutOutput QnnConcatenateLayout(const Attrs& attrs, - const Array& new_in_layouts, - const Array& old_in_layouts, - const Array& old_in_types) { - // Collect the layouts and types to reuse Relay Concatenate Infer Correct Layout. - ICHECK_EQ(old_in_types.size(), 5); - auto input_tuple_type = old_in_types[0].as(); - ICHECK(input_tuple_type); - auto num_input_tensors = input_tuple_type->fields.size(); - - Array relay_new_in_layouts(nullptr); - if (new_in_layouts.defined()) { - relay_new_in_layouts = - Array(new_in_layouts.begin(), new_in_layouts.begin() + num_input_tensors); - } - Array relay_old_in_layouts(nullptr); - if (old_in_layouts.defined()) { - relay_old_in_layouts = - Array(old_in_layouts.begin(), old_in_layouts.begin() + num_input_tensors); - } - - // Use Relay Concatenate Infer Correct layout to infer the layouts for data tensors. - auto concat_new_layout = - ConcatenateLayout(attrs, relay_new_in_layouts, relay_old_in_layouts, {old_in_types[0]}); - - // Fill the layouts of remaining input tensors - scales and zero points. The layouts of these - // tensors can be treated as channel layout. Total number of these tensors are 2 * num of data - // tensors (scale and zero point for each input data tensor) + 2 for the output data tensor. - Layout channel_layout = Layout("C"); - Array input_layouts = concat_new_layout->input_layouts; - - for (size_t i = 0; i < 2 * num_input_tensors + 2; i++) { - input_layouts.push_back(channel_layout); - } - Array output_layouts = concat_new_layout->output_layouts; - return InferCorrectLayoutOutput(input_layouts, output_layouts, concat_new_layout->new_attrs); -} - -Expr MakeQnnConcatenate(Expr data, Expr input_scales, Expr input_zero_points, Expr output_scale, - Expr output_zero_point, int axis) { - auto attrs = make_object(); - attrs->axis = axis; - static const Op& op = Op::Get("qnn.concatenate"); - return Call(op, {data, input_scales, input_zero_points, output_scale, output_zero_point}, - Attrs(attrs), {}); -} - -/* - * \brief Canonicalizes the QNN concatenate op. - * \param attrs The QNN concatenate attrs. - * \param new_args The new mutated args to the call node. - * \param arg_types The types of input and output. - * \return The sequence of Relay ops for concatenate op. - */ -Expr ConcatenateQnnCanonicalize(const Attrs& attrs, const Array& new_args, - const Array& arg_types) { - // Get the attrs. - ICHECK_EQ(new_args.size(), 5); - auto& data = new_args[0]; - auto& input_scales = new_args[1]; - auto& input_zero_points = new_args[2]; - auto& output_scale = new_args[3]; - auto& output_zero_point = new_args[4]; - const auto* concatenate_attrs = attrs.as(); - ICHECK(concatenate_attrs != nullptr); - - // Get the input dtype and shape. - ICHECK_GE(arg_types.size(), 1); - auto tuple_type = arg_types[0].as(); - ICHECK(tuple_type != nullptr); - - // FIXME (anijain2305) - The lowering can be further optimized. Instead of inserting requantize in - // the start, we can insert requantize at the end if and only if all the input tensors have same - // qnn params. This can be done in future. - - // If the output qnn params do not match the input qnn params, we can call requantize on the input - // expr first, followed by a concatenate on the requantized input exprs. - - Array tuple_exprs; - if (data->IsInstance()) { - tuple_exprs = data.as()->fields; - } else if (data->IsInstance()) { // if the data is a CallNode, use TupleGetItems - auto call = Downcast(data); - for (size_t i = 0; i < tuple_type->fields.size(); i++) { - tuple_exprs.push_back(TupleGetItem(call, i)); - } - } - ICHECK(!tuple_exprs.empty()); - - auto tuple_input_scales = input_scales.as(); - ICHECK(tuple_input_scales != nullptr); - - auto tuple_input_zero_points = input_zero_points.as(); - ICHECK(tuple_input_zero_points != nullptr); - - int idx = 0; - Array requantized_exprs; - for (auto quantized_expr : tuple_exprs) { - // Get the input scale for the idx quantized input tensor. - auto input_scale = tuple_input_scales->fields[idx]; - - // Get the zero point for the idx quantized input tensor. - auto input_zero_point = tuple_input_zero_points->fields[idx]; - - // Check if output and input qnn params are same. If not, requantize. - if (!IsEqualScalar(input_scale, output_scale) || - !IsEqualScalar(input_zero_point, output_zero_point)) { - // Get the input shape and dtype. - auto tensor_type = tuple_type->fields[idx].as(); - auto input_dtype = tensor_type->dtype; - auto input_shape = tensor_type->shape; - - // Requantize the input. - auto requantized_expr = Requantize(quantized_expr, input_shape, input_scale, input_zero_point, - output_scale, output_zero_point, input_dtype); - requantized_exprs.push_back(requantized_expr); - } else { - requantized_exprs.push_back(quantized_expr); - } - idx++; - } - return MakeConcatenate(Tuple(requantized_exprs), concatenate_attrs->axis); -} - -RELAY_REGISTER_OP("qnn.concatenate") - .describe(R"code(Concatenate the quantized input tensors along the given axis. -)code" TVM_ADD_FILELINE) - .set_attrs_type() - .set_num_inputs(5) - .add_argument("data", "Tensor", "The tensor to concatenate.") - .add_argument("input_scales", "Tensor", "The quantization scales of the input tensors.") - .add_argument("input_zero_points", "Tensor", - "The quantization zero_points of the input tensors.") - .add_argument("output_scale", "Tensor", "The quantization scale of the output tensor.") - .add_argument("output_zero_point", "Tensor", - "The quantization zero_point of the output tensor.") - .set_support_level(11) - .add_type_rel("QnnConcatenate", QnnConcatenateRel) - .set_attr("TNonComputational", true) - .set_attr("FTVMQnnCanonicalize", ConcatenateQnnCanonicalize) - .set_attr("FInferCorrectLayout", QnnConcatenateLayout); - -TVM_REGISTER_GLOBAL("relay.qnn.op._make.concatenate").set_body_typed(MakeQnnConcatenate); - -} // namespace qnn -} // namespace relay -} // namespace tvm diff --git a/src/relay/qnn/op/convolution.cc b/src/relay/qnn/op/convolution.cc deleted file mode 100644 index 2ce5523764f1..000000000000 --- a/src/relay/qnn/op/convolution.cc +++ /dev/null @@ -1,881 +0,0 @@ -/* - * 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 src/relay/qnn/op/convolution.cc - * \brief Property def of qnn convolution operator. - */ -#include "../../op/nn/convolution.h" - -#include -#include -#include -#include -#include -#include -#include - -#include "../../transforms/pattern_utils.h" -#include "../utils.h" - -namespace tvm { -namespace relay { -namespace qnn { - -// relay.op.qnn.conv2d - -bool QnnConv2DRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - // Expected Types: data, weight, input_zero_point, weight_zero_point, input_scale, weight_scale, - // out_type - ICHECK_EQ(types.size(), 7); - const auto* data = types[0].as(); - const auto* weight = types[1].as(); - if (data == nullptr || weight == nullptr) return false; - const auto* param = attrs.as(); - ICHECK(param != nullptr) << "Conv2DAttrs cannot be nullptr."; - ICHECK(data->dtype == DataType::Int(8) || data->dtype == DataType::UInt(8) || - data->dtype == DataType::Int(16)) - << "Expected qnn conv2d type(int8, uint8, int16) for input but was " << data->dtype; - ICHECK(weight->dtype == DataType::Int(8) || weight->dtype == DataType::UInt(8) || - weight->dtype == DataType::Int(16)) - << "Expected qnn conv2d type(int8, uint8, int16) for weight but was " << weight->dtype; - ICHECK(param->out_dtype == DataType::Int(16) || param->out_dtype == DataType::Int(32) || - param->out_dtype == DataType::Int(64)) - << "Expected qnn conv2d type(int16, int32, int64) for output but was " << param->out_dtype; - ICHECK(param->out_dtype.bits() > 0) << "Output dtype bits should be greater than 0."; - - // Check the types of scale and zero points. - for (size_t i = 2; i < 5; ++i) { - if (types[i].as()) { - return false; - } - } - ICHECK(IsScalarType(types[2], DataType::Int(32))); // input_zero_point - ICHECK(IsScalarType(types[4], DataType::Float(32))); // input_scale - // Kernel scale can be a vector of length output_channels or a scalar. - if (param->groups == 1) { - size_t axis = param->kernel_layout.operator std::string().find('O'); - ICHECK(axis != std::string::npos) << "Kernel layout attribute is not defined"; - AssignType(types[5], DataType::Float(32), weight->shape[axis], reporter); // weight_scale - } else { - size_t o_axis = param->kernel_layout.operator std::string().find('O'); - size_t i_axis = param->kernel_layout.operator std::string().find('I'); - size_t d_axis = param->data_layout.operator std::string().find('C'); - ICHECK(o_axis != std::string::npos || i_axis != std::string::npos) - << "Kernel layout attribute is not defined"; - ICHECK(d_axis != std::string::npos) << "Data layout attribute is not defined"; - if (param->channels.defined() && - tvm::tir::ExprDeepEqual()(param->groups, data->shape[d_axis]) && - tvm::tir::ExprDeepEqual()(param->channels, weight->shape[i_axis])) { - // This is depthwise convolution - // Here, total number of output channels depend on depth multiplier. - AssignType(types[5], DataType::Float(32), weight->shape[i_axis] * weight->shape[o_axis], - reporter); // weight_scale - } else { - // This is grouped convolution - AssignType(types[5], DataType::Float(32), param->channels, - reporter); // weight_scale - } - } - - // Collect the input tensor and output tensor devoid of scale and zero points to reuse Relay - // Conv2D infer type function. - Array tensor_types = {types[0], types[1], types[6]}; - return Conv2DRel(tensor_types, 3, attrs, reporter); -} - -InferCorrectLayoutOutput QnnConvInferCorrectLayout(const Attrs& attrs, - const Array& new_in_layouts, - const Array& old_in_layouts, - const Array& old_in_types) { - // Use Relay Conv2D Infer correct layout. - auto conv_new_layouts = - ConvInferCorrectLayout(attrs, new_in_layouts, old_in_layouts, old_in_types); - - // Fill the layouts of remaining input tensors - scales and zero points. The layouts of these - // tensors can be treated as channel layout. - Layout channel_layout = Layout("C"); - Array input_layouts = {conv_new_layouts->input_layouts[0], - conv_new_layouts->input_layouts[1], - channel_layout, - channel_layout, - channel_layout, - channel_layout}; - Array output_layouts = conv_new_layouts->output_layouts; - return InferCorrectLayoutOutput(input_layouts, output_layouts, attrs); -} - -bool is_depthwise(const Conv2DAttrs* param) { - return param->channels.defined() && tvm::tir::ExprDeepEqual()(param->channels, param->groups) && - param->groups != 1; -} - -// Workload - batch_size, in_channels, out_channels, kernel_h, kernel_w, channel_multiplier -using WorkloadType = std::tuple; - -/* - * \brief Get the conv parameters like batch_size, kernel_height etc. - * \param ref_call The original callnode. - * \param param The qnn conv2d attributes. - * \return A tuple of workload. - */ -WorkloadType GetWorkload(const Array& arg_types, const Conv2DAttrs* param) { - // Get conv parameters. - const auto in_shape = get_shape(arg_types[0]); - int batch_size, in_channels; - if (param->data_layout == "NCHW") { - batch_size = get_const_int(in_shape[0]); - in_channels = get_const_int(in_shape[1]); - } else if (param->data_layout == "NHWC") { - batch_size = get_const_int(in_shape[0]); - in_channels = get_const_int(in_shape[3]); - } else { - LOG(FATAL) << "qnn.conv2d does not support " << param->data_layout << " layout"; - } - - const auto kernel_shape = get_shape(arg_types[1]); - int out_channels, kernel_h, kernel_w; - int channel_multiplier = -1; - bool depthwise = is_depthwise(param); - int index_k = 0; - int index_h = 2; - int index_w = 3; - int index_c = 1; - if (param->kernel_layout == "HWIO") { - index_k = 3; - index_h = 0; - index_w = 1; - index_c = 2; - } else if (param->kernel_layout == "HWOI") { - index_k = 2; - index_h = 0; - index_w = 1; - index_c = 3; - } else if (param->kernel_layout == "OHWI") { - index_k = 0; - index_h = 1; - index_w = 2; - index_c = 3; - } else if (param->kernel_layout != "OIHW") { - LOG(FATAL) << "qnn.conv2d does not support " << param->kernel_layout << " layout"; - } - - kernel_h = get_const_int(kernel_shape[index_h]); - kernel_w = get_const_int(kernel_shape[index_w]); - out_channels = get_const_int(kernel_shape[index_k]); - - if (depthwise) { - channel_multiplier = get_const_int(kernel_shape[index_c]); - } - - return std::make_tuple(batch_size, in_channels, out_channels, kernel_h, kernel_w, - channel_multiplier); -} - -/* - * \brief Fallback to simpler lowering for dilation (when non-zero kernel point) or grouped conv. - * \param data The input expr. - * \param weight The weight expr. - * \param input_zero_point The input zero point expr. - * \param kernel_zero_point The kernel zero point expr. - * \param param The qnn conv2d attributes. - * \return The fallback lowered sequence of Relay expr. - * \note In case of dilation with non-zero kernel zero point, normal lowering would require a - * dilated pool. Since, we don't have dilated pool, we fallback to a simpler sequence of Relay - * operations. This will potentially lead to performance degradation as the convolution is called on - * int32 tensors instead of int8 tensors. - */ -Expr Conv2DFallBack(const Expr& data, const Expr& weight, const Expr& input_zero_point, - const Expr& kernel_zero_point, const Conv2DAttrs* param) { - // Upcast the parameters to be at least int32 to avoid overflow - auto upcast_bits = param->out_dtype.bits() < 32 ? 32 : param->out_dtype.bits(); - - auto zp_data = Cast(input_zero_point, DataType::Int(upcast_bits)); - auto zp_kernel = Cast(kernel_zero_point, DataType::Int(upcast_bits)); - - auto shifted_data = Cast(data, DataType::Int(upcast_bits)); - auto zero_scalar = MakeConstantScalar(DataType::Int(upcast_bits), 0); - if (!IsEqualScalar(input_zero_point, zero_scalar)) { - shifted_data = Subtract(Cast(data, DataType::Int(upcast_bits)), zp_data); - } - - auto shifted_kernel = Cast(weight, DataType::Int(upcast_bits)); - if (!IsEqualScalar(kernel_zero_point, zero_scalar)) { - shifted_kernel = Subtract(Cast(weight, DataType::Int(upcast_bits)), zp_kernel); - } - - return Conv2D(shifted_data, shifted_kernel, param->strides, param->padding, param->dilation, - param->groups, param->channels, param->kernel_size, param->data_layout, - param->kernel_layout, param->out_layout, param->out_dtype); -} - -/* - * \brief Pad the input data. - * \param data The input expr. - * \param input_zero_point The input zero point expr. - * \return The padded input expr. - * \note For quantized convolution, the input has to be padded with zero point - * instead of zero. This might lead to performance degradation as pad - * cannot be fused with conv in Relay. In case we see performance - * degradation, we can change the conv2D API to accept a pad_const value. - */ -Expr Conv2DPadInput(const Expr& data, const Expr& input_zero_point, const Conv2DAttrs* param) { - // 1) Pad the input data - auto padded_data = data; - auto pad_top_value = get_const_int(param->padding[0]); - auto pad_left_value = get_const_int(param->padding[1]); - auto pad_bottom_value = get_const_int(param->padding[2]); - auto pad_right_value = get_const_int(param->padding[3]); - bool do_pad = - pad_top_value != 0 || pad_left_value != 0 || pad_bottom_value != 0 || pad_right_value != 0; - if (do_pad) { - Array pad_n({0, 0}); - Array pad_c({0, 0}); - Array pad_h({param->padding[0], param->padding[2]}); - Array pad_w({param->padding[1], param->padding[3]}); - - Array> pad_width; - if (param->data_layout == "NCHW") { - pad_width = {pad_n, pad_c, pad_h, pad_w}; - } else if (param->data_layout == "NHWC") { - pad_width = {pad_n, pad_h, pad_w, pad_c}; - } else { - LOG(FATAL) << "qnn.conv2d does not support " << param->data_layout << " layout"; - } - padded_data = Pad(data, pad_width, input_zero_point, "constant"); - } - return padded_data; -} - -/* - * \brief Calculates the second term in the qnn.conv2d depthwise lowering sequence. - * \param padded_data The padded data expr. - * \param kernel_zero_point The kernel zero point expr. - * \param param The qnn conv2d attributes. - * \param kernel_h The height of kernel. - * \param kernel_w The width of kernel. - * \param channel_multiplier The channel/depth multiplier. - * \return The sequence of Relay operators for term2. - * \note The term2 looks like this - * - * Sigma(r, s) zp_w * Qa(n, oc/cm, oh + r, ow + s) - * - * Second term is not directly representable by one Relay operator. - * However, deeper analysis shows that we can reduce r,s using avg_pool2d, - * followed by repeat on the C axis by cm times. - */ -Expr DepthwiseConv2DSecondTerm(const Expr& padded_data, const Expr& kernel_zero_point, - const Conv2DAttrs* param, int kernel_h, int kernel_w, - int channel_multiplier) { - auto casted_t2 = Cast(padded_data, DataType::Int(32)); - - // We can reduce the H and W axis by using avg_pool2d. However, avg_pool2d averages the sum. - // Since, this is integer division (floor), we can first multiply the data by the pool_size and - // then perform avg_pool2d. Reversing this causes inaccuracy due to floor division. If the - // pool_size and strides are 1x1, we don't need avg_pool2d. - auto reduced_t2 = casted_t2; - if (kernel_h * kernel_w != 1) { - auto scaled_hw_t2 = - Multiply(casted_t2, MakeConstantScalar(DataType::Int(32), kernel_h * kernel_w)); - Array padding({0, 0}); - reduced_t2 = AvgPool2D(scaled_hw_t2, param->kernel_size, param->strides, param->dilation, - padding, param->data_layout, - "", // out_layout - false, // ceil_mode - false); // count_include_pad - } else { - int stride1 = get_const_int(param->strides[0]); - int stride2 = get_const_int(param->strides[1]); - if (stride1 * stride2 != 1) { - Array padding({0, 0}); - reduced_t2 = AvgPool2D(reduced_t2, param->kernel_size, param->strides, param->dilation, - padding, param->data_layout, - "", // out_layout - false, // ceil_mode - false); // count_include_pad - } - } - - auto multiplied_t2 = reduced_t2; - auto one_scalar = MakeConstantScalar(DataType::Int(32), 1); - if (!IsEqualScalar(kernel_zero_point, one_scalar)) { - if (!IsConstScalar(kernel_zero_point)) { - multiplied_t2 = Multiply(MakeRepeat(kernel_zero_point, channel_multiplier, 0), reduced_t2); - } else { - multiplied_t2 = Multiply(kernel_zero_point, reduced_t2); - } - } - - // Reduce the C dimension. Find the dimension. - int axis_t2 = 0; - if (param->data_layout == "NCHW") { - axis_t2 = 1; - } else if (param->data_layout == "NHWC") { - axis_t2 = 3; - } else { - LOG(FATAL) << "qnn.conv2d does not support " << param->data_layout << " layout"; - } - auto repeated_t2 = multiplied_t2; - if (channel_multiplier != 1) { - repeated_t2 = MakeRepeat(multiplied_t2, channel_multiplier, axis_t2); - } - return repeated_t2; -} - -/* - * \brief Calculates the third term in the qnn.conv2d depthwise lowering sequence. - * \param weight The weight expr. - * \param input_zero_point The input zero point expr. - * \param param The qnn conv2d attributes. - * \param out_channels The number of output channels. - * \param channel_multiplier The channel/depth multiplier. - * \return The sequence of Relay operatos for term3. - * \note The term3 looks like this - * - * Sigma(r, s) zp_a * Qw(oc/m, oc%m, r, s) - * - * This can be achieved by calling reduce on r and s axis. The tensor can be then reshaped to - * (1, oc, 1, 1) as (oc/m, oc%m) are just contiguous memory locations. - */ -Expr DepthwiseConv2DThirdTerm(const Expr& weight, const Expr& input_zero_point, - const Conv2DAttrs* param, int out_channels, int channel_multiplier) { - // Find which dimensions are R, S. - Array axes_t3; - if (param->kernel_layout == "OIHW") { - // For OIHW kernel layout, HW are reduce axis - axes_t3 = {2, 3}; - } else if (param->kernel_layout == "HWIO") { - axes_t3 = {0, 1}; - } else if (param->kernel_layout == "HWOI") { - axes_t3 = {0, 1}; - } else { - LOG(FATAL) << "qnn.conv2d does not support " << param->kernel_layout << " layout"; - } - auto reduced_t3 = Sum(Cast(weight, DataType::Int(32)), axes_t3, false, false); - - // Find the newshape depending on NCHW/NHWC layout. - Array newshape; - if (param->data_layout == "NCHW") { - newshape = {1, out_channels * channel_multiplier, 1, 1}; - } else if (param->data_layout == "NHWC") { - newshape = {1, 1, 1, out_channels * channel_multiplier}; - } else { - LOG(FATAL) << "qnn.conv2d does not support " << param->data_layout << " layout"; - } - auto reshaped_t3 = Reshape(reduced_t3, newshape); - - auto one_scalar = MakeConstantScalar(DataType::Int(32), 1); - if (IsEqualScalar(input_zero_point, one_scalar)) { - return reshaped_t3; - } - return Multiply(input_zero_point, reshaped_t3); -} - -/* - * \brief Calculates the fourth term in the qnn.conv2d depthwise lowering sequence. - * \param input_zero_point_int The int value of input zero point. - * \param kernel_zero_point_int The int value of kernel zero point. - * \param kernel_h The height of kernel. - * \param kernel_w The width of kernel. - * \return The sequence of Relay operators for term4. - * \note The term4 looks like this - * - * Sigma(r, s) zp_a * zp_w - */ -Expr DepthwiseConv2DFourthTerm(int input_zero_point_int, int kernel_zero_point_int, int kernel_h, - int kernel_w) { - int scalar_term4 = input_zero_point_int * kernel_zero_point_int * kernel_h * kernel_w; - return MakeConstantScalar(DataType::Int(32), scalar_term4); -} - -/* - * \brief Calculates the fourth term in the qnn.conv2d depthwise lowering sequence - for non-constant zero_points. - * \param input_zero_point The Expr for the input zero point. - * \param kernel_zero_point The Expr for the kernel zero point. - * \param kernel_h The height of kernel. - * \param kernel_w The width of kernel. - * \return The sequence of Relay operators for term4. - * \note The term4 looks like this - * - * Sigma(r, s) zp_a * zp_w - */ -Expr DepthwiseConv2DFourthTerm(const Expr& input_zero_point, const Expr& kernel_zero_point, - int kernel_h, int kernel_w) { - Expr scalar_term4 = MakeConstantScalar(DataType::Int(32), kernel_h * kernel_w); - Expr variable_term4 = Multiply(input_zero_point, kernel_zero_point); - return Multiply(scalar_term4, variable_term4); -} - -/* - * \brief Calculates the first term in the qnn.conv2d lowering sequence. - * \param data The input expr. - * \param weight The weight expr. - * \param param The qnn conv2d attributes. - * \return The sequence of Relay operators for term1. - * \note The term1 is - * Sigma(c,r,s) QW(k, c, r, s) * QA(n, c, h + r, w + s) - * This is just conv2d on int tensors. - */ -Expr Conv2DFirstTerm(const Expr& padded_data, const Expr& weight, const Conv2DAttrs* param) { - // Lowering for Term 1 - Array padding({0, 0, 0, 0}); - return Conv2D(padded_data, weight, param->strides, padding, param->dilation, param->groups, - param->channels, param->kernel_size, param->data_layout, param->kernel_layout, - param->out_layout, param->out_dtype); -} - -/* - * \brief Calculates the second term in the qnn.conv2d lowering sequence. - * \param padded_data The padded data expr. - * \param kernel_zero_point The kernel zero point expr. - * \param param The qnn conv2d attributes. - * \param kernel_h The height of kernel. - * \param kernel_w The width of kernel. - * \return The sequence of Relay operators for term2. - * \note The term2 looks like this - * - * Sigma(c,r,s) zp_w * QA(n, c, h + r, w + s) - * - * Second term is not directly representable by one Relay operator. - * However, deeper analysis shows that we can reduce r,s using avg_pool2d, - * followed by a reduce on the C axis. Using avg_pool2d also gives an - * opportunity to reuse alter_op_layout infrastructure. - */ -Expr Conv2DSecondTerm(const Expr& padded_data, const Expr& kernel_zero_point, - const Conv2DAttrs* param, int kernel_h, int kernel_w, int out_channels) { - auto casted_t2 = Cast(padded_data, DataType::Int(32)); - - // We can reduce the H and W axis by using avg_pool2d. However, avg_pool2d averages the sum. - // Since, this is integer division (floor), we can first multiply the data by the pool_size and - // then perform avg_pool2d. Reversing this causes inaccuracy due to floor division. - Array padding({0, 0}); - - // Reduce the C dimension. Find the dimension. - Array axes_t2; - if (param->data_layout == "NCHW") { - axes_t2 = {1}; - } else if (param->data_layout == "NHWC") { - axes_t2 = {3}; - } else { - LOG(FATAL) << "qnn.conv2d does not support " << param->data_layout << " layout"; - } - // Keep dims true to retain 4D tensor - auto reduced_c_t2 = Sum(casted_t2, axes_t2, true, false); - - // If the pool_size and strides are 1x1, we don't need avg_pool2d. - auto reduced_t2 = reduced_c_t2; - if (kernel_h * kernel_w != 1) { - reduced_c_t2 = - Multiply(reduced_c_t2, MakeConstantScalar(DataType::Int(32), kernel_h * kernel_w)); - reduced_t2 = AvgPool2D(reduced_c_t2, param->kernel_size, param->strides, param->dilation, - padding, param->data_layout, - "", // out_layout - false, // ceil_mode - false); // count_include_pad - } else { - int stride1 = get_const_int(param->strides[0]); - int stride2 = get_const_int(param->strides[1]); - if (stride1 * stride2 != 1) { - reduced_t2 = AvgPool2D(reduced_c_t2, param->kernel_size, param->strides, param->dilation, - padding, param->data_layout, - "", // out_layout - false, // ceil_mode - false); // count_include_pad - } - } - - auto multiplied_t2 = reduced_t2; - auto one_scalar = MakeConstantScalar(DataType::Int(32), 1); - if (!IsEqualScalar(kernel_zero_point, one_scalar)) { - if (!IsConstScalar(kernel_zero_point)) { - Layout layout(param->data_layout); - int channel_axis = layout.IndexOf(LayoutAxis::Get('C')); - reduced_t2 = MakeRepeat(reduced_t2, out_channels, channel_axis); - } - multiplied_t2 = Multiply(kernel_zero_point, reduced_t2); - } - return multiplied_t2; -} - -/* - * \brief Calculates the third term in the qnn.conv2d lowering sequence. - * \param weight The weight expr. - * \param input_zero_point The input zero point expr. - * \param param The qnn conv2d attributes. - * \param out_channels The number of output channels. - * \return The sequence of Relay operators for term3. - * \note The term3 looks like this - * - * Sigma(c,r,s) zp_a * QW(k, c, r, s) - * - * This can be achieved by calling reduce on c, r and s axis, resulting in - * a 1D tensor. The tensor is then reshaped to conform to NHWC/NCHW - * format. - */ -Expr Conv2DThirdTerm(const Expr& weight, const Expr& input_zero_point, const Conv2DAttrs* param, - int out_channels) { - // Find which dimensions are C, R, S. - Array axes_t3; - if (param->kernel_layout == "OIHW") { - // For OIHW kernel layout, IHW are reduce axis - axes_t3 = {1, 2, 3}; - } else if (param->kernel_layout == "HWIO") { - axes_t3 = {0, 1, 2}; - } else if (param->kernel_layout == "HWOI") { - axes_t3 = {0, 1, 3}; - } else if (param->kernel_layout == "OHWI") { - axes_t3 = {1, 2, 3}; - } else { - LOG(FATAL) << "qnn.conv2d does not support " << param->kernel_layout << " layout"; - } - auto reduced_t3 = Sum(Cast(weight, DataType::Int(32)), axes_t3, false, false); - - // Find the newshape depending on NCHW/NHWC layout. - Array newshape; - if (param->data_layout == "NCHW") { - newshape = {1, out_channels, 1, 1}; - } else if (param->data_layout == "NHWC") { - newshape = {1, 1, 1, out_channels}; - } else { - LOG(FATAL) << "qnn.conv2d does not support " << param->data_layout << " layout"; - } - auto reshaped_t3 = Reshape(reduced_t3, newshape); - - auto one_scalar = MakeConstantScalar(DataType::Int(32), 1); - if (IsEqualScalar(input_zero_point, one_scalar)) { - return reshaped_t3; - } - return Multiply(input_zero_point, reshaped_t3); -} - -/* - * \brief Calculates the fourth term in the qnn.conv2d lowering sequence. - * \param input_zero_point_int The int value of input zero point. - * \param kernel_zero_point_int The int value of kernel zero point. - * \param in_channels The number of input channels. - * \param kernel_h The height of kernel. - * \param kernel_w The width of kernel. - * \param param The qnn conv2d attributes. - * \return The sequence of Relay operators for term4. - * \note The term4 looks like this - * - * Sigma(c,r,s) zp_a * zp_w - * - */ -Expr Conv2DFourthTerm(int input_zero_point_int, int kernel_zero_point_int, int in_channels, - int kernel_h, int kernel_w, const Conv2DAttrs* param) { - auto upcast_bits = param->out_dtype.bits() < 32 ? 32 : param->out_dtype.bits(); - int scalar_term4 = - input_zero_point_int * kernel_zero_point_int * in_channels * kernel_h * kernel_w; - return MakeConstantScalar(DataType::Int(upcast_bits), scalar_term4); -} - -/* - * \brief Calculates the fourth term in the qnn.conv2d lowering sequence - for non-constant zero_points. - * \param input_zero_point The Expr for the input zero point. - * \param kernel_zero_point The Expr for the kernel zero point. - * \param in_channels The number of input channels. - * \param kernel_h The height of kernel. - * \param kernel_w The width of kernel. - * \param param The qnn conv2d attributes. - * \return The sequence of Relay operators for term4. - * \note The term4 looks like this - * - * Sigma(c,r,s) zp_a * zp_w - * - */ -Expr Conv2DFourthTerm(const Expr& input_zero_point, const Expr& kernel_zero_point, int in_channels, - int kernel_h, int kernel_w, const Conv2DAttrs* param) { - auto upcast_bits = param->out_dtype.bits() < 32 ? 32 : param->out_dtype.bits(); - Expr scalar_term4 = - MakeConstantScalar(DataType::Int(upcast_bits), in_channels * kernel_h * kernel_w); - Expr variable_term4 = Multiply(input_zero_point, kernel_zero_point); - return Multiply(scalar_term4, variable_term4); -} - -/* - * \brief Combines different terms of qnn conv2d lowering. - * \param term1 The term1 of qnn conv2d lowering. - * \param term2 The term2 of qnn conv2d lowering. - * \param term3 The term3 of qnn conv2d lowering. - * \param term4 The term4 of qnn conv2d lowering. - * \param input_zero_point_int The int value of input zero point. - * \param kernel_zero_point_int The int value of kernel zero point. - * \param param The qnn conv2d attributes. - * \return The combined sequence of relay operations. - * \note The combined operation looks like this - * - * Sigma(c,r,s) QW(k, c, r, s) * QA(n, c, h + r, w + s) // Term1 - * - Sigma(c,r,s) zp_w * QA(n, c, h + r, w + s) // Term2 - * - Sigma(c,r,s) zp_a * QW(k, c, r, s) // Term3 - * + Sigma(c,r,s) zp_a * zp_w // Term4 - * - */ -Expr Conv2DCombineTerms(const Expr& term1, const Expr& term2, const Expr& term3, const Expr& term4, - int input_zero_point_int, int kernel_zero_point_int) { - if (input_zero_point_int == 0 && kernel_zero_point_int == 0) { - // term 2, 3 and 4 become zero. - return term1; - } else if (input_zero_point_int == 0 && kernel_zero_point_int != 0) { - // term 3 and term 4 become zero. - return Subtract(term1, term2); - } else if (input_zero_point_int != 0 && kernel_zero_point_int == 0) { - // term 2 and term 4 become zero. - return Subtract(term1, term3); - } else { - auto data_term = Subtract(term1, term2); - // Putting constant terms together, so that constant folding can fold it. - auto const_term = Subtract(term4, term3); - return Add(data_term, const_term); - } -} - -/* - * \brief Forward rewrite the qnn conv2d op. - * \param attrs The QNN conv2d attrs. - * \param new_args The new mutated args to the call node. - * \param arg_types The types of input and output. - * \return The sequence of Relay ops for qnn cov2d op. - * \node Lowering of the qnn.conv2d operator - * A quantized tensor is represented in following manner - * A = scale_a x (QA - zp_A) - * where QA is quantized tensor, scale_a and zp_A are quantization - * params. - * - * Quantized convolution will convolve two quantized tensors and returns a - * quantized tensor of default dtype of int32, with scale equaling to the - * product of scales of input tensors, and a zero point of zero. - * - * For symmetric quantization, the zp_* for all tensors is 0. So, the - * lowering of qnn.conv2d is - * - * QA(n, ic, oh + r, ow + s) (conv) QW(oc, ic, r, s) - * - * For asymmetric computation, we can perform similar unrolling. We can - * find more details at - * https://discuss.tvm.ai/t/tf-lite-quantized-conv2d-operator-conversion/2651/8?u=janimesh - * The computation gets unrolled into following 4 terms - * - * Sigma(c,r,s) QW(k, c, r, s) * QA(n, c, h + r, w + s) // Term1 - * - Sigma(c,r,s) zp_w * QA(n, c, h + r, w + s) // Term2 - * - Sigma(c,r,s) zp_a * QW(k, c, r, s) // Term3 - * + Sigma(c,r,s) zp_a * zp_w // Term4 - * - * Term3 and Term4 can be computed at compile time. - * - * Key points to notice: - * 1) Padding is done explicitly because the input has to be padded with - * zero point. This might leave some performance opportunity at the - * table. Can be avoided by modifying conv2d API to accept the - * pad_const_value. - * 2) Second term is not directly representable by one Relay operator. - * However, deeper analysis shows that we can reduce r,s using - * avg_pool2d, followed by a reduce on the C axis. Using avg_pool2d also - * gives an opportunity to reuse alter_op_layout infrastructure. - * 3) For dilated conv, in current lowering, we need dilated pool. So as - * a workaround, we fall back to simpler lowering using int32 conv if - * the conv is dilated. We fallback also in case of grouped conv. - * - * For depthwise, we can similarly unroll the computation. The initial compute is as follows - * where cm = channel_multiplier - * - * Qc(n, oc, oh, ow) = Sigma(r, s) (Qw(oc/m, oc%/m, r, s) - zp_w) - * * (Qa(n, oc/cm, oh + r, ow + s) - zp_a) - * - * This can be written as - * - * Sigma(r, s) Qw(oc/m, oc%/m, r, s) * Qa(n, oc/cm, oh + r, ow + s) - * - Sigma(r, s) zp_w * Qa(n, oc/cm, oh + r, ow + s) - * - Sigma(r, s) zp_a * Qw(oc/m, oc%m, r, s) - * - Sigma(r, s) zp_a * zp_w - * - * The whole process can be broken down into following steps - * * Assertion checks for existing support, fallback if necessary - * * Pad the input. - * * Get Term1. - * * Get Term2. - * * Get Term3. - * * Get Term4. - * * Combine the terms. - */ -Expr QnnConv2DCanonicalize(const Attrs& attrs, const Array& new_args, - const Array& arg_types) { - ICHECK_EQ(new_args.size(), 6); - Expr data = new_args[0]; - Expr weight = new_args[1]; - Expr input_zero_point = new_args[2]; - Expr kernel_zero_point = new_args[3]; - const auto* param = attrs.as(); - ICHECK(param != nullptr); - // Assertion checks for existing support. - ICHECK(param->data_layout == "NCHW" || param->data_layout == "NHWC") - << "qnn.conv2d supports only NCHW/NHWC input data layout."; - ICHECK(param->kernel_layout == "OIHW" || param->kernel_layout == "HWIO" || - param->kernel_layout == "HWOI" || param->kernel_layout == "OHWI") - << "qnn.conv2d supports only OIHW/HWIO/HWOI/OHWI kernel data layout."; - ICHECK(param->kernel_size.defined()) << "qnn.conv2d requires kernel size to be specified."; - - auto [batch_size, in_channels, out_channels, kernel_h, kernel_w, channel_multiplier] = - GetWorkload(arg_types, param); - (void)batch_size; // https://gcc.gnu.org/bugzilla/show_bug.cgi?id=81767 - - // zero points are allowed to be non-scalar. Let's check if that's the case. - bool dynamic_zp = false; - // Use -1 zero point as a default for dynamic. - int input_zero_point_int = -1; - int kernel_zero_point_int = -1; - - // Input zero point can either be a constant or a scalar expression. - if (IsConstScalar(input_zero_point) && (IsConstScalar(kernel_zero_point))) { - // Extract the integer zero points. - input_zero_point_int = GetScalarFromConstant(input_zero_point); - kernel_zero_point_int = GetScalarFromConstant(kernel_zero_point); - } else { - // Make kernel_zero_point expression a 1-D tensor for consistent shape. - kernel_zero_point = Reshape(kernel_zero_point, { - -1, - }); - dynamic_zp = true; - } - - // Fallback to int32 conv if there is dilation with non-zero kernel point or grouped conv2d - // For dilated conv, if the kernel zero point is non-zero, the pooling operator also has to - // traverse the elements in dilated manner. Currently, we do not have strided pool. So, in case of - // dilated conv with non-zero kernel point, we fall back to simpler but slow lowering. - - ICHECK_EQ(param->dilation.size(), 2) << "qnn.conv2d only supports 2D dilation"; - auto dilation_h = get_const_int(param->dilation[0]); - auto dilation_w = get_const_int(param->dilation[1]); - // Check if qnn supports the conv2d parameters. If not, fallback to regular conv2d. - bool supported_dilation = (kernel_zero_point_int == 0) || (dilation_h == 1 && dilation_w == 1); - bool supported_groups = (param->groups == 1 || is_depthwise(param)); - bool conv2d_params_supported = supported_dilation && supported_groups; - - // If we need to fall back to default conv2d, kernel zp may need to be broadcast to kernel_layout. - // Otherwise, we broadcast it to data_layout for qnn lowering. - if (dynamic_zp) { - if (!conv2d_params_supported) { - Layout kernel_layout(param->kernel_layout); - int kernel_axis = kernel_layout.IndexOf(LayoutAxis::Get("O")); - kernel_zero_point = ExpandBiasToMatchAxis(kernel_zero_point, 4, {kernel_axis}); - } else { - Layout data_layout(param->data_layout); - int channel_axis = data_layout.IndexOf(LayoutAxis::Get("C")); - kernel_zero_point = ExpandBiasToMatchAxis(kernel_zero_point, 4, {channel_axis}); - } - } - - if (!conv2d_params_supported) { - return Conv2DFallBack(data, weight, input_zero_point, kernel_zero_point, param); - } else if (is_depthwise(param)) { - ICHECK_NE(channel_multiplier, -1); - auto padded_data = Conv2DPadInput(data, input_zero_point, param); - auto term1 = Conv2DFirstTerm(padded_data, weight, param); - auto term2 = DepthwiseConv2DSecondTerm(padded_data, kernel_zero_point, param, kernel_h, - kernel_w, channel_multiplier); - auto term3 = - DepthwiseConv2DThirdTerm(weight, input_zero_point, param, out_channels, channel_multiplier); - Expr term4; - if (dynamic_zp) { - term4 = DepthwiseConv2DFourthTerm(input_zero_point, kernel_zero_point, kernel_h, kernel_w); - } else { - term4 = DepthwiseConv2DFourthTerm(input_zero_point_int, kernel_zero_point_int, kernel_h, - kernel_w); - } - return Conv2DCombineTerms(term1, term2, term3, term4, input_zero_point_int, - kernel_zero_point_int); - } - - auto padded_data = Conv2DPadInput(data, input_zero_point, param); - auto term1 = Conv2DFirstTerm(padded_data, weight, param); - auto term2 = - Conv2DSecondTerm(padded_data, kernel_zero_point, param, kernel_h, kernel_w, out_channels); - auto term3 = Conv2DThirdTerm(weight, input_zero_point, param, out_channels); - Expr term4; - if (dynamic_zp) { - term4 = Conv2DFourthTerm(input_zero_point, kernel_zero_point, in_channels, kernel_h, kernel_w, - param); - } else { - term4 = Conv2DFourthTerm(input_zero_point_int, kernel_zero_point_int, in_channels, kernel_h, - kernel_w, param); - } - return Conv2DCombineTerms(term1, term2, term3, term4, input_zero_point_int, - kernel_zero_point_int); -} - -// Positional relay function to create quantized conv2d operator -// used by frontend FFI. -Expr MakeQnnConv2D(Expr data, Expr weight, Expr input_zero_point, Expr kernel_zero_point, - Expr input_scale, Expr kernel_scale, Array strides, - Array padding, Array dilation, int groups, - IndexExpr channels, Array kernel_size, String data_layout, - String kernel_layout, String out_layout, DataType out_dtype) { - auto attrs = make_object(); - attrs->strides = std::move(strides); - attrs->padding = std::move(padding); - attrs->dilation = std::move(dilation); - attrs->groups = groups; - attrs->channels = std::move(channels); - attrs->kernel_size = std::move(kernel_size); - attrs->data_layout = std::move(data_layout); - attrs->kernel_layout = std::move(kernel_layout); - attrs->out_layout = std::move(out_layout); - attrs->out_dtype = std::move(out_dtype); - static const Op& op = Op::Get("qnn.conv2d"); - return Call(op, {data, weight, input_zero_point, kernel_zero_point, input_scale, kernel_scale}, - Attrs(attrs), {}); -} - -RELAY_REGISTER_OP("qnn.conv2d") - .describe(R"code(2D quantized convolution layer. -This operator convolves quantized weight with quantized data. The scale of the -output quantized tensor is the product of the weight_scale and input_scale of -the input quantized tensors. The zero point of the output quantized tensor is -0. By default, the dtype of output is int32. Please also refer to Requantize -operator to understand how to scale back the int32 output to (u)int8 or (u)int16. -- **data**: This depends on the `layout` parameter. Input is 4D array of shape - (batch_size, in_channels, height, width) if `layout` is `NCHW`. -- **weight**: (channels, in_channels, kernel_size[0], kernel_size[1]) -- **out**: This depends on the `layout` parameter. Output is 4D array of shape - (batch_size, channels, out_height, out_width) if `layout` is `NCHW`. -)code" TVM_ADD_FILELINE) - .set_attrs_type() - .set_num_inputs(6) - .add_argument("data", "Tensor", "The quantized input data tensor.") - .add_argument("weight", "Tensor", "The quantized weight tensor.") - .add_argument("input_scale", "Tensor", "The quantization scale of the input tensor.") - .add_argument("input_zero_point", "Tensor", "The quantization zero_point of the input tensor.") - .add_argument("weight_scale", "Tensor", "The quantization scale of the weight tensor.") - .add_argument("weight_zero_point", "Tensor", - "The quantization zero_point of the weight tensor.") - .set_support_level(11) - .add_type_rel("QnnConv2D", QnnConv2DRel) - .set_attr("TNonComputational", true) - .set_attr("FTVMQnnCanonicalize", QnnConv2DCanonicalize) - .set_attr("FInferCorrectLayout", QnnConvInferCorrectLayout) - .set_attr("TOpPattern", kOutEWiseFusable); - -TVM_REGISTER_GLOBAL("relay.qnn.op._make.conv2d").set_body_typed(MakeQnnConv2D); - -} // namespace qnn -} // namespace relay -} // namespace tvm diff --git a/src/relay/qnn/op/convolution_transpose.cc b/src/relay/qnn/op/convolution_transpose.cc deleted file mode 100644 index 0b24ae71ca8c..000000000000 --- a/src/relay/qnn/op/convolution_transpose.cc +++ /dev/null @@ -1,180 +0,0 @@ -/* - * 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 src/relay/qnn/op/convolution_transpose.cc - * \brief Property def of qnn transpose convolution operator. - */ -#include -#include -#include -#include -#include -#include -#include - -#include "../../op/nn/convolution.h" -#include "../../transforms/pattern_utils.h" -#include "../utils.h" - -namespace tvm { -namespace relay { -namespace qnn { - -// relay.op.qnn.conv2d_transpose - -inline Expr MakeQnnConv2DTranspose(Expr data, Expr weight, Expr input_zero_point, - Expr kernel_zero_point, Expr input_scale, Expr kernel_scale, - Array strides, Array padding, - Array dilation, int groups, IndexExpr channels, - Array kernel_size, std::string data_layout, - std::string kernel_layout, std::string out_layout, - Array output_padding, DataType out_dtype) { - auto attrs = make_object(); - attrs->strides = std::move(strides); - attrs->padding = std::move(padding); - attrs->dilation = std::move(dilation); - attrs->groups = groups; - attrs->channels = std::move(channels); - attrs->kernel_size = std::move(kernel_size); - attrs->data_layout = std::move(data_layout); - attrs->kernel_layout = std::move(kernel_layout); - attrs->out_layout = std::move(out_layout); - attrs->output_padding = std::move(output_padding); - attrs->out_dtype = std::move(out_dtype); - const Op& op = Op::Get("qnn.conv2d_transpose"); - return Call(op, {data, weight, input_zero_point, kernel_zero_point, input_scale, kernel_scale}, - Attrs(attrs), {}); -} - -InferCorrectLayoutOutput QnnConvTransposeInferCorrectLayout( - const Attrs& attrs, const Array& new_in_layouts, const Array& old_in_layouts, - const Array& old_in_types) { - // Use Relay Conv2D transpose Infer correct layout. - auto conv_transpose_new_layouts = ConvInferCorrectLayout( - attrs, new_in_layouts, old_in_layouts, old_in_types); - - // Fill the layouts of remaining input tensors - scales and zero points. The layouts of these - // tensors can be treated as channel layout. - Layout channel_layout = Layout("C"); - Array input_layouts = {conv_transpose_new_layouts->input_layouts[0], - conv_transpose_new_layouts->input_layouts[1], - channel_layout, - channel_layout, - channel_layout, - channel_layout}; - Array output_layouts = conv_transpose_new_layouts->output_layouts; - return InferCorrectLayoutOutput(input_layouts, output_layouts, attrs); -} - -bool QnnConv2DTransposeRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - // Expected Types: data, weight, input_zero_point, weight_zero_point, input_scale, weight_scale, - // out_type - ICHECK_EQ(types.size(), 7); - const auto* data = types[0].as(); - const auto* weight = types[1].as(); - if (data == nullptr || weight == nullptr) return false; - const auto* param = attrs.as(); - ICHECK(param != nullptr) << "Conv2DTransposeAttrs cannot be nullptr."; - ICHECK(data->dtype == DataType::Int(8) || data->dtype == DataType::UInt(8) || - data->dtype == DataType::Int(16) || data->dtype == DataType::UInt(16)) - << "Expected qnn conv2d type(int8, uint8, int16) for input but was " << data->dtype; - ICHECK(weight->dtype == DataType::Int(8) || weight->dtype == DataType::UInt(8)) - << "Expected qnn conv2d type(int8, uint8) for weight but was " << weight->dtype; - ICHECK(param->out_dtype == DataType::Int(16) || param->out_dtype == DataType::Int(32) || - param->out_dtype == DataType::Int(64)) - << "Expected qnn conv2d type(int16, int32, int64) for output but was " << param->out_dtype; - ICHECK(param->out_dtype.bits() > 0) << "Output dtype bits should be greater than 0."; - - // Check the types of scale and zero points. - for (size_t i = 2; i < 5; ++i) { - if (types[i].as()) { - return false; - } - } - - const auto* weight_zp_type = types[3].as(); - ICHECK(weight_zp_type->dtype == DataType::Int(32)); // weight_zero_point - - bool input_zp_is_scalar = (types[2].as())->shape.size() == 0 || - get_const_int((types[2].as())->Size()) == 1; - bool input_scale_is_scalar = (types[4].as())->shape.size() == 0 || - get_const_int((types[4].as())->Size()) == 1; - - ICHECK(input_scale_is_scalar && input_zp_is_scalar) - << "Zero point or scale should be scalar or a vector with one element."; - - // Assign types for input scale and zero point. - AssignType(types[2], DataType::Int(32), Integer(1), reporter); // input_zero_point - AssignType(types[4], DataType::Float(32), Integer(1), reporter); // input_scale - - // Kernel scale can be a vector of length output_channels or a scalar. - if (param->groups == 1) { - size_t axis = param->kernel_layout.find('O'); - ICHECK(axis != std::string::npos) << "Kernel layout attribute is not defined"; - AssignType(types[5], DataType::Float(32), weight->shape[axis], reporter); // weight_scale - } else { - // Here, total number of output channels depend on depth multiplier. - size_t o_axis = param->kernel_layout.find('O'); - size_t i_axis = param->kernel_layout.find('I'); - ICHECK(o_axis != std::string::npos || i_axis != std::string::npos) - << "Kernel layout attribute is not defined"; - AssignType(types[5], DataType::Float(32), weight->shape[i_axis] * weight->shape[o_axis], - reporter); // kernel scale - } - - // Collect the input tensor and output tensor devoid of scale and zero points to reuse Relay - // Conv2D infer type function. - Array tensor_types = {types[0], types[1], types[6]}; - return Conv2DTransposeRel(tensor_types, 3, attrs, reporter); -} - -RELAY_REGISTER_OP("qnn.conv2d_transpose") - .describe(R"code(Quantized transposed 2D convolution layer (sometimes called Deconvolution). -This operator deconvolves quantized weight with quantized data. The scale of the -output quantized tensor is the product of the weight_scale and input_scale of -the input quantized tensors. The zero point of the output quantized tensor is -0. By default, the dtype of output is int32. Please also refer to Requantize -operator to understand how to scale back the int32 output to (u)int8. -- **data**: This depends on the `layout` parameter. Input is 4D array of shape - (batch_size, in_channels, height, width) if `layout` is `NCHW`. -- **weight**: (channels, in_channels, kernel_size[0], kernel_size[1]) -- **out**: This depends on the `layout` parameter. Output is 4D array of shape - (batch_size, channels, out_height, out_width) if `layout` is `NCHW`. -)code" TVM_ADD_FILELINE) - .set_attrs_type() - .set_num_inputs(6) - .add_argument("data", "Tensor", "The quantized input data tensor.") - .add_argument("weight", "Tensor", "The quantized weight tensor.") - .add_argument("input_scale", "Tensor", "The quantization scale of the input tensor.") - .add_argument("input_zero_point", "Tensor", "The quantization zero_point of the input tensor.") - .add_argument("weight_scale", "Tensor", "The quantization scale of the weight tensor.") - .add_argument("weight_zero_point", "Tensor", - "The quantization zero_point of the weight tensor.") - .set_support_level(11) - .add_type_rel("QnnConv2DTranspose", QnnConv2DTransposeRel) - .set_attr("TNonComputational", true) - .set_attr("FInferCorrectLayout", QnnConvTransposeInferCorrectLayout); - -TVM_REGISTER_GLOBAL("relay.qnn.op._make.conv2d_transpose").set_body_typed(MakeQnnConv2DTranspose); - -} // namespace qnn -} // namespace relay -} // namespace tvm diff --git a/src/relay/qnn/op/dense.cc b/src/relay/qnn/op/dense.cc deleted file mode 100644 index 48f2a813d0e7..000000000000 --- a/src/relay/qnn/op/dense.cc +++ /dev/null @@ -1,332 +0,0 @@ -/* - * 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 src/relay/qnn/op/dense.cc - * \brief Property def of qnn dense operator. - */ - -#include -#include - -#include "../../op/nn/nn.h" -#include "../../transforms/pattern_utils.h" -#include "../utils.h" - -namespace tvm { -namespace relay { -namespace qnn { - -// relay.op.qnn.dense - -bool QnnDenseRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - // Expected Types: data, weight, input_zero_point, weight_zero_point, input_scale, weight_scale, - // out_type - ICHECK_EQ(types.size(), 7); - const auto* data = types[0].as(); - const auto* weight = types[1].as(); - if (data == nullptr || weight == nullptr) return false; - const auto* param = attrs.as(); - ICHECK(param != nullptr) << "DenseAttrs cannot be nullptr."; - ICHECK(data->dtype == DataType::Int(8) || data->dtype == DataType::UInt(8) || - data->dtype == DataType::Int(16) || data->dtype == DataType::UInt(16)) - << "Expected quantized dense type(int8, uint8, int16, uint16) for input but was " - << data->dtype; - ICHECK(weight->dtype == DataType::Int(8) || weight->dtype == DataType::UInt(8)) - << "Expected quantized dense type(int8, uint8) for weight but was " << weight->dtype; - ICHECK(param->out_dtype == DataType::Int(32) || param->out_dtype == DataType::Int(64)) - << "Expected quantized dense type(int32, int64) for output but was " << param->out_dtype; - - // Check the types of scale and zero points. - for (size_t i = 2; i < 5; ++i) { - if (types[i].as()) { - return false; - } - } - ICHECK(IsScalarType(types[2], DataType::Int(32))); // input_zero_point - ICHECK(IsScalarType(types[4], DataType::Float(32))); // input_scale - // weight_zero_point can be a scalar or a vector of the same shape as the weight_scale - AssignType(types[5], DataType::Float(32), param->units, reporter); // weight_scale - - ICHECK(param->out_dtype.bits() > 0) << "Output dtype bits should be greater than 0."; - - // Collect the input tensor and output tensor devoid of scale and zero points to reuse Relay - // Dense infer type function. - Array tensor_types = {types[0], types[1], types[6]}; - return MatmulRel(tensor_types, 3, attrs, reporter); -} - -InferCorrectLayoutOutput QnnDenseInferCorrectLayout(const Attrs& attrs, - const Array& new_in_layouts, - const Array& old_in_layouts, - const Array& old_in_types) { - // Use Relay Dense Infer correct layout. - auto dense_new_layouts = - DenseInferCorrectLayout(attrs, new_in_layouts, old_in_layouts, old_in_types); - - // Fill the layouts of remaining input tensors - scales and zero points. The layouts of these - // tensors can be treated as channel layout. - Layout channel_layout = Layout("N"); - Array input_layouts = {dense_new_layouts->input_layouts[0], - dense_new_layouts->input_layouts[1], - channel_layout, - channel_layout, - channel_layout, - channel_layout}; - Array output_layouts = dense_new_layouts->output_layouts; - return InferCorrectLayoutOutput(input_layouts, output_layouts, attrs); -} - -// Positional relay function to create quantized dense operator used by frontend FFI. -Expr MakeQuantizedDense(Expr data, Expr weight, Expr input_zero_point, Expr kernel_zero_point, - Expr input_scale, Expr kernel_scale, IndexExpr units, DataType out_dtype) { - auto attrs = make_object(); - attrs->units = std::move(units); - attrs->out_dtype = out_dtype; - static const Op& op = Op::Get("qnn.dense"); - return Call(op, {data, weight, input_zero_point, kernel_zero_point, input_scale, kernel_scale}, - Attrs(attrs), {}); -} - -Expr DenseFirstTerm(const Expr& quantized_data, const Expr& quantized_kernel, - const DenseAttrs* attrs) { - return Dense(quantized_data, quantized_kernel, attrs->units, attrs->out_dtype); -} - -Expr DenseSecondTerm(const Expr& quantized_data, const Expr& kernel_zero_point, - const int out_dim_size) { - Array axes = {1}; - Expr reduced_t2 = Sum(Cast(quantized_data, DataType::Int(32)), axes, true, false); - Expr multiplied_t2; - if (!IsConstScalar(kernel_zero_point)) { - multiplied_t2 = Multiply(kernel_zero_point, MakeRepeat(reduced_t2, out_dim_size, 1)); - } else { - multiplied_t2 = Multiply(kernel_zero_point, reduced_t2); - } - return multiplied_t2; -} - -Expr DenseThirdTerm(const Expr& quantized_kernel, const Expr& input_zero_point) { - Array axes = {1}; - return Multiply(input_zero_point, - Sum(Cast(quantized_kernel, DataType::Int(32)), axes, false, false)); -} - -Expr DenseFourthTerm(int input_zero_point_int, int kernel_zero_point_int, int reduction_dim_size) { - int32_t scalar_term = input_zero_point_int * kernel_zero_point_int * reduction_dim_size; - return MakeConstantScalar(DataType::Int(32), scalar_term); -} - -Expr DenseFourthTerm(const Expr& input_zero_point, const Expr& kernel_zero_point, - int reduction_dim_size) { - auto reduction_dim = MakeConstantScalar(DataType::Int(32), reduction_dim_size); - return Multiply(Multiply(input_zero_point, kernel_zero_point), reduction_dim); -} - -Expr DenseCombineTerms(const Expr& term1, const Expr& term2, const Expr& term3, const Expr& term4) { - auto data_term = Subtract(term1, term2); - // Putting constant terms together, so that constant folding can fold it. - auto const_term = Subtract(term4, term3); - return Add(data_term, const_term); -} -/* - * \brief Forward rewrite the qnn dense op. - * \param attrs The QNN dense attrs. - * \param new_args The new mutated args to the call node. - * \param arg_types The types of input and output. - * \return The sequence of Relay ops for qnn cov2d op. - * \note Lowering of the qnn.dense operator - * A quantized tensor is represented in following manner - * A = scale_a x (QA - zp_A) - * where QA is quantized tensor, scale_a and zp_A are quantization - * params. - * - * Quantized dense multiplies two quantized tensors and returns a - * quantized tensor of default dtype of int32, with scale equaling to the - * product of scales of input tensors, and a zero point of zero. - * - * The lowering for asymmetric quantized dense looks as follows. More details at - * https://discuss.tvm.ai/t/tf-lite-quantized-conv2d-operator-conversion/2651/8 - * The computation gets unrolled into following 4 terms - * C(m, n) = Sigma(k) (A(m, k) * W(n, k)) - * - * RHS becomes - * Sigma(k) ([QA(m, k) - zp_a] * [QW(n, k) - zp_w]) - * - * Unrolling leads to following sequence - * Sigma(k) QA(m, k) * QW(n, k) // Term1 - * - Sigma(k) zp_w * QA(m, k) // Term2 - * - Sigma(k) zp_a * QW(n, k) // Term3 - * - Sigma(k) * zp_a * zp_w // Term4 - * - * Term3 and Term4 can be computed at compile time. - */ -Expr QnnDenseCanonicalize(const Attrs& attrs, const Array& new_args, - const Array& arg_types) { - ICHECK_EQ(new_args.size(), 6); - Expr quantized_data = new_args[0]; - Expr quantized_kernel = new_args[1]; - Expr input_zero_point = new_args[2]; - Expr kernel_zero_point = new_args[3]; - - const auto in_shape = get_shape(arg_types[0]); - const auto w_shape = get_shape(arg_types[1]); - const int reduction_dim_size = get_const_int(in_shape[1]); - const int out_dim_size = get_const_int(w_shape[0]); - - const auto* qnn_dense_attrs = attrs.as(); - - auto term1 = DenseFirstTerm(quantized_data, quantized_kernel, qnn_dense_attrs); - auto term2 = DenseSecondTerm(quantized_data, kernel_zero_point, out_dim_size); - auto term3 = DenseThirdTerm(quantized_kernel, input_zero_point); - - // Extract the integer zero points. - - if (!IsConstScalar(input_zero_point) || !IsConstScalar(kernel_zero_point)) { - auto term4 = DenseFourthTerm(input_zero_point, kernel_zero_point, reduction_dim_size); - return DenseCombineTerms(term1, term2, term3, term4); - } - - auto kernel_zero_point_int = GetScalarFromConstant(kernel_zero_point); - auto input_zero_point_int = GetScalarFromConstant(input_zero_point); - - // Get all the terms as described in the comments. - auto term4 = DenseFourthTerm(input_zero_point_int, kernel_zero_point_int, reduction_dim_size); - - // Combine those 4 terms depending on the zero points to get the best lowering. - if (input_zero_point_int == 0 && kernel_zero_point_int == 0) { - // term 2, 3 and 4 become zero. - return term1; - } else if (input_zero_point_int == 0 && kernel_zero_point_int != 0) { - // term 3 and term 4 become zero. - return Subtract(term1, term2); - } else if (input_zero_point_int != 0 && kernel_zero_point_int == 0) { - // term 2 and term 4 become zero. - return Subtract(term1, term3); - } else { - return DenseCombineTerms(term1, term2, term3, term4); - } -} - -RELAY_REGISTER_OP("qnn.dense") - .describe(R"code(Applies a linear transformation: :math:`Y = XW^T`. -- **data**: quantized(int8, unit8) `(x1, x2, ..., xn, input_dim)` -- **weight**: quantized(int8, unit8) `(units, input_dim)` -- **out**: quantized(int32) `(x1, x2, ..., xn, units)`. -)code" TVM_ADD_FILELINE) - .set_attrs_type() - .set_num_inputs(6) - .add_argument("data", "quantized nD Tensor", "Input data.") - .add_argument("weight", "quantized 2D Tensor", "Weight matrix.") - .add_argument("input_scale", "Tensor", "The quantization scale of the input tensor.") - .add_argument("input_zero_point", "Tensor", "The quantization zero_point of the input tensor.") - .add_argument("weight_scale", "Tensor", "The quantization scale of the weight tensor.") - .add_argument("weight_zero_point", "Tensor", - "The quantization zero_point of the weight tensor.") - .set_support_level(11) - .add_type_rel("QDense", QnnDenseRel) - .set_attr("FInferCorrectLayout", QnnDenseInferCorrectLayout) - .set_attr("TNonComputational", true) - .set_attr("FTVMQnnCanonicalize", QnnDenseCanonicalize) - .set_attr("TOpPattern", kOutEWiseFusable); - -TVM_REGISTER_GLOBAL("relay.qnn.op._make.dense").set_body_typed(MakeQuantizedDense); - -// ------------------- relay.qnn.op.contrib_dense_pack - -bool QnnDensePackRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - // Expected types: data, weight, input_zero_point, weight_zero_point, input_scale, weight_scale, - // out_type - ICHECK_EQ(types.size(), 7); - const auto* data = types[0].as(); - const auto* weight = types[1].as(); - if (data == nullptr || weight == nullptr) return false; - - const DensePackAttrs* param = attrs.as(); - ICHECK(param != nullptr); - - ICHECK_EQ(data->shape.size(), 2) << "Only 2D data is supported"; - ICHECK(weight->shape.size() == 4) << "Expect weight to be 4D tensor"; - - Array oshape = data->shape; - oshape.Set(1, weight->shape[0] * weight->shape[2]); - - ICHECK(param->out_dtype.bits() > 0) << "Output dtype bits should be greater than 0."; - // assign output type - reporter->Assign(types[6], TensorType(oshape, param->out_dtype)); - return true; -} - -InferCorrectLayoutOutput QnnDensePackInferCorrectLayout( - const Attrs& attrs, const Array& new_in_layouts, const Array& old_in_layouts, - const Array& old_in_types) { - auto params = attrs.as(); - ICHECK(params); - return InferCorrectLayoutOutput({"NC", params->weight_layout, "N", "N", "N", "N"}, {"NC"}, attrs); -} - -Expr QnnDensePackCanonicalize(const Attrs& attrs, const Array& new_args, - const Array& arg_types) { - LOG(FATAL) << "Canonicalization function for qnn.contrib_dense_pack is not implemented"; - return Expr(); -} - -Expr MakeQuantizedDensePack(Expr data, Expr weight, Expr input_zero_point, Expr kernel_zero_point, - Expr input_scale, Expr kernel_scale, tvm::String weight_layout, - IndexExpr units, DataType out_dtype) { - auto attrs = make_object(); - attrs->units = std::move(units); - attrs->out_dtype = out_dtype; - attrs->weight_layout = weight_layout; - static const Op& op = Op::Get("qnn.contrib_dense_pack"); - return Call(op, {data, weight, input_zero_point, kernel_zero_point, input_scale, kernel_scale}, - Attrs(attrs), {}); -} - -RELAY_REGISTER_OP("qnn.contrib_dense_pack") - .describe(R"code(Applies a linear transformation: :math:`Y = XW^T`. -- **data**: quantized(int8, uint8) `(x1, x2, ..., xn, input_dim)` -- **weight**: quantized(int8, uint8) `(units, input_dim)` -- **out**: quantized(int32) `(x1, x2, ..., xn, units)`. -)code" TVM_ADD_FILELINE) - .set_attrs_type() - .set_num_inputs(6) - .add_argument("data", "quantized nD Tensor", "Input data.") - .add_argument("weight", "quantized 2D Tensor", "Weight matrix.") - .add_argument("input_scale", "Tensor", "The quantization scale of the input tensor.") - .add_argument("input_zero_point", "Tensor", "The quantization zero_point of the input tensor.") - .add_argument("weight_scale", "Tensor", "The quantization scale of the weight tensor.") - .add_argument("weight_zero_point", "Tensor", - "The quantization zero_point of the weight tensor.") - .set_support_level(11) - .add_type_rel("QnnDensePack", QnnDensePackRel) - .set_attr("FInferCorrectLayout", QnnDensePackInferCorrectLayout) - .set_attr("TNonComputational", true) - .set_attr("FTVMQnnCanonicalize", QnnDensePackCanonicalize) - .set_attr("TOpPattern", kOutEWiseFusable); - -TVM_REGISTER_GLOBAL("relay.qnn.op._make.contrib_dense_pack").set_body_typed(MakeQuantizedDensePack); - -// ------------------- relay.qnn.op.contrib_dense_pack - -} // namespace qnn -} // namespace relay -} // namespace tvm diff --git a/src/relay/qnn/op/dequantize.cc b/src/relay/qnn/op/dequantize.cc deleted file mode 100644 index 5e2ef39edacb..000000000000 --- a/src/relay/qnn/op/dequantize.cc +++ /dev/null @@ -1,179 +0,0 @@ -/* - * 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 src/relay/qnn/op/dequantize.cc - * \brief QNN dequantize operator. Dequantize operator converts from quantized - * domain to unquantized domain. - */ - -#include -#include -#include - -#include "../../transforms/pattern_utils.h" -#include "../utils.h" - -namespace tvm { -namespace relay { -namespace qnn { - -TVM_REGISTER_NODE_TYPE(DequantizeAttrs); - -bool DequantizeRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - ICHECK_EQ(types.size(), 4); - const auto* data = types[0].as(); - - if (data == nullptr) { - return false; - } - - const auto input_dtype = data->dtype; - ICHECK(input_dtype == DataType::Int(8) || input_dtype == DataType::UInt(8) || - input_dtype == DataType::Int(16) || input_dtype == DataType::UInt(16) || - input_dtype == DataType::Int(32)) - << "Input type should be one of the quantized types [int8, unit8, int16, uint16, int32] but " - << "was " << input_dtype; - - const auto* dequantize_attrs = attrs.as(); - int axis = dequantize_attrs->axis; - auto rank = static_cast(data->shape.size()); - axis = (axis < 0) ? ((rank > 0) ? data->shape.size() + axis : 0) : axis; - - // If zero point and scale are scalar or have arbitrary rank with one element, - // then axis doesn't matter. - bool scale_is_scalar = (types[1].as())->shape.size() == 0 || - get_const_int((types[1].as())->Size()) == 1; - bool zp_is_scalar = (types[2].as())->shape.size() == 0 || - get_const_int((types[2].as())->Size()) == 1; - - if (!scale_is_scalar || !zp_is_scalar) { - ICHECK_LT(axis, rank > 0 ? rank : 1) << "axis " << dequantize_attrs->axis << " is out of range"; - ICHECK_GE(axis, 0) << "axis " << dequantize_attrs->axis << " is out of range"; - } - - PrimExpr axis_shape; - if (!scale_is_scalar || !zp_is_scalar) { - axis_shape = data->shape[axis]; - } else { - axis_shape = Integer(1); - } - // Check and assign types for scale and zero points. - AssignType(types[1], DataType::Float(32), axis_shape, reporter); // scale - AssignType(types[2], DataType::Int(32), axis_shape, reporter); // zero point - - const Array oshape = data->shape; - const DataType out_dtype = dequantize_attrs->out_dtype; - ICHECK(out_dtype == DataType::Float(16) || out_dtype == DataType::Float(32)) - << "Output type should be one of [float16, float32] but was " << out_dtype; - // assign output type. - reporter->Assign(types[3], TensorType(oshape, out_dtype)); - return true; -} - -Expr MakeDequantize(Expr data, Expr input_scale, Expr input_zero_point, int axis, - DataType out_dtype) { - // real_value = scale * (quantized_value - zero_point) - // A more detailed explanation can be found here - - // https://github.com/google/gemmlowp/blob/master/doc/quantization.md - auto attrs = make_object(); - attrs->axis = axis; - attrs->out_dtype = out_dtype; - static const Op& op = Op::Get("qnn.dequantize"); - return Call(op, {data, input_scale, input_zero_point}, Attrs(attrs), {}); -} - -Expr DequantizeLower(const Expr& input_tensor, const Expr& input_scale, - const Expr& input_zero_point, const Array& types, - const DequantizeAttrs* attrs) { - auto axis = attrs->axis; - - ICHECK_EQ(types.size(), 4); - auto in_type = types[0]; - auto in_tensor_type = in_type.as(); - ICHECK(in_tensor_type != nullptr) << "Type information missing" - << " Please run infer_type pass."; - Array input_shape = in_tensor_type->shape; - - size_t n_dim = input_shape.size(); - - // Wrap axis from negative to positive if needed. - if (axis < 0) { - axis = static_cast(n_dim) + axis; - } - - // Expand scale and zero point if the input tensor is channel quantized - auto expanded_input_scale = input_scale; - if (!IsConstScalar(input_scale) && !IsScalarType(types[1])) { - expanded_input_scale = ExpandBiasToMatchAxis(input_scale, n_dim, {axis}); - } - - auto expanded_input_zero_point = input_zero_point; - if (!IsConstScalar(input_zero_point) && !IsScalarType(types[2])) { - expanded_input_zero_point = ExpandBiasToMatchAxis(input_zero_point, n_dim, {axis}); - } - - auto shift = Subtract(Cast(input_tensor, DataType::Int(32)), expanded_input_zero_point); - auto scaled_output = Multiply(Cast(shift, DataType::Float(32)), expanded_input_scale); - - const DataType out_dtype = attrs->out_dtype; - if (out_dtype.is_float() && out_dtype.bits() == 32) return scaled_output; - - double min_val = tvm::min_value(out_dtype).as()->value; - double max_val = tvm::max_value(out_dtype).as()->value; - auto clamped_output = Clip(scaled_output, min_val, max_val); - return Cast(clamped_output, out_dtype); -} - -Expr DequantizeQnnCanonicalize(const Attrs& attrs, const Array& new_args, - const Array& types) { - ICHECK_EQ(new_args.size(), 3); - auto& data = new_args[0]; - auto& input_scale = new_args[1]; - auto& input_zero_point = new_args[2]; - ICHECK_EQ(types.size(), 4); - - // Get attrs. - const auto* dequantize_attrs = attrs.as(); - ICHECK(dequantize_attrs != nullptr); - - return DequantizeLower(data, input_scale, input_zero_point, types, dequantize_attrs); -} - -RELAY_REGISTER_OP("qnn.dequantize") - .describe(R"code(Dequantizes the input and produces float32 output. -The input is always quantized (int8, uint8) and will be converted to float32 given input scale and zero_point. -- **data**: Quantized tensor of any shape to dequantize. The input data can be of floating point -)code" TVM_ADD_FILELINE) - .set_attrs_type() - .set_num_inputs(3) - .add_argument("data", "Tensor", "The tensor to dequantize.") - .add_argument("input_scale", "Tensor", "The quantization scale of the input tensor.") - .add_argument("input_zero_point", "Tensor", "The quantization zero_point of the input tensor.") - .set_support_level(11) - .add_type_rel("Dequantize", DequantizeRel) - .set_attr("TNonComputational", true) - .set_attr("FTVMQnnCanonicalize", DequantizeQnnCanonicalize); - -TVM_REGISTER_GLOBAL("relay.qnn.op._make.dequantize").set_body_typed(MakeDequantize); - -} // namespace qnn -} // namespace relay -} // namespace tvm diff --git a/src/relay/qnn/op/leaky_relu.cc b/src/relay/qnn/op/leaky_relu.cc deleted file mode 100644 index 458fde0d8a08..000000000000 --- a/src/relay/qnn/op/leaky_relu.cc +++ /dev/null @@ -1,161 +0,0 @@ -/* - * 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 src/relay/qnn/op/leaky_relu.cc - * \brief QNN leaky relu operator. - */ -#include -#include - -#include "op_common.h" - -namespace tvm { -namespace relay { -namespace qnn { - -bool QnnLeakyReluRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - // Expected Types: data, input_scale, input_zero_point, output_scale, output_zero_point, out_type - ICHECK_EQ(types.size(), 6); - const auto* x = types[0].as(); - if (x == nullptr) return false; - ICHECK(x->dtype == DataType::Int(8) || x->dtype == DataType::UInt(8)) - << "Expected quantized leaky_relu type(int8, uint8) for input but was " << x->dtype; - const auto* param = attrs.as(); - ICHECK(param != nullptr) << "LeakyReluAttrs cannot be nullptr."; - - // Check the types of scale and zero points. - for (size_t i = 1; i < 5; ++i) { - if (types[i].as()) { - return false; - } - } - - ICHECK(IsScalarType(types[1], DataType::Float(32))); // input_scale - ICHECK(IsScalarType(types[2], DataType::Int(32))); // input_zero_point - ICHECK(IsScalarType(types[3], DataType::Float(32))); // output_scale - ICHECK(IsScalarType(types[4], DataType::Int(32))); // output_zero_point - - // Assign types for scale and zero points. - reporter->Assign(types[1], TensorType({}, DataType::Float(32))); // input_scale - reporter->Assign(types[2], TensorType({}, DataType::Int(32))); // input_zero_point - reporter->Assign(types[3], TensorType({}, DataType::Float(32))); // output_scale - reporter->Assign(types[4], TensorType({}, DataType::Int(32))); // output_zero_point - - // Collect the input tensor and output tensor devoid of scale and zero points to reuse Relay - // IdentityRel infer type function. - Array tensor_types = {types[0], types[5]}; - return IdentityRel(tensor_types, 2, attrs, reporter); -} - -// Positional relay function to create quantized leaky relu operator used by frontend FFI. -Expr MakeQuantizedLeakyRelu(Expr x, double alpha, Expr input_scale, Expr input_zero_point, - Expr output_scale, Expr output_zero_point) { - auto attrs = make_object(); - attrs->alpha = alpha; - static const Op& op = Op::Get("qnn.leaky_relu"); - return Call(op, {x, input_scale, input_zero_point, output_scale, output_zero_point}, Attrs(attrs), - {}); -} - -/* - * \brief Canonicalizes the QNN leaky relu op. - * \param attrs The empty attribute. - * \param new_args The new mutated args to the call node. - * \param arg_types The types of input and output. - * \return The sequence of Relay ops for leaky relu op. - */ -Expr QnnLeakyReluCanonicalize(const Attrs& attrs, const Array& new_args, - const Array& arg_types) { - // We rely on fixed point arithmetic to preserve the precision of multiplication - // by a small alpha value < 1. - // - // We assume the same scale and zero point for alpha and the input tensor. - // LeakyReLU can be written in terms of respective quantized tensors, scales and - // zero points as - // - // scale_o * (Q_o - zp_o) = alpha * scale_i * (Q_i - zp_i) when Q_i < zp_i (1) - // scale_o * (Q_o - zp_o) = scale_i * (Q_i - zp_i) when Q_i >= zp_i (2) - // - // Since the input qnn params can be different than output qnn params, we first requantize the - // input tensor to the output qnn params. After requantizing Q_i, equation (1) becames equation - // (3) where Q_i' is the requantized data from Q_i. - // - // scale_o * (Q_o - zp_o) = alpha * scale_o * (Q_i' - zp_o) when Q_i < zp_i (3) - // Q_o = alpha * Q_i' + (1 - alpha) * zp_o when Q_i < zp_i (4) - // - // It is equal to requantize Q_i to Q_o using scale_o and zp_o in equation (2). - // So equation (2) becomes - // - // Q_o = requantize(Q_i) when Q_i >= zp_i (5) - // - // Finnally, Q_o could be calculated by equation (4) and equation (5). - ICHECK_EQ(new_args.size(), 5); - Expr data = Cast(new_args[0], DataType::Int(32)); - Expr input_scale = new_args[1]; - Expr input_zero_point = Cast(new_args[2], DataType::Int(32)); - Expr output_scale = new_args[3]; - Expr output_zero_point = Cast(new_args[4], DataType::Int(32)); - - const auto* q_attrs = attrs.as(); - auto alpha = q_attrs->alpha; - - const auto input_shape = get_shape(arg_types[0]); - const auto input_dtype = arg_types[0].as()->dtype; - - // requantize the input to Q_i' - auto requantized_expr = RequantizeOrUpcast(data, input_scale, input_zero_point, output_scale, - output_zero_point, input_shape); - - // alpha * Q_i' - auto [fixed_point_multiplier, shift] = GetFixedPointMultiplierShift(alpha); - auto prod = FixedPointMultiply(requantized_expr, fixed_point_multiplier, shift); - - // (1 - alpha) * zp_o - auto [fixed_point_multiplier_z, shift_z] = GetFixedPointMultiplierShift(1 - alpha); - auto scaled_z = FixedPointMultiply(output_zero_point, fixed_point_multiplier_z, shift_z); - - // alpha * Q_i' + (1 - alpha) * zp_o - auto add = Add(prod, scaled_z); - auto output = Where(Less(data, input_zero_point), add, requantized_expr); - - return ConvertDtype(output, input_dtype); -} - -RELAY_REGISTER_OP("qnn.leaky_relu") - .describe("Leaky relu for quantized tensors.") - .set_attrs_type() - .set_num_inputs(5) - .add_argument("data", "Quantized Tensor", "The input data.") - .add_argument("input_scale", "Tensor", "The quantization scale of the input tensor.") - .add_argument("input_zero_point", "Tensor", "The quantization zero_point of the input tensor.") - .add_argument("output_scale", "Tensor", "The quantization scale of the output tensor.") - .add_argument("output_zero_point", "Tensor", - "The quantization zero_point of the output tensor.") - .set_support_level(11) - .add_type_rel("QLeakyRelu", QnnLeakyReluRel) - .set_attr("TNonComputational", true) - .set_attr("FTVMQnnCanonicalize", QnnLeakyReluCanonicalize); - -TVM_REGISTER_GLOBAL("relay.qnn.op._make.leaky_relu").set_body_typed(MakeQuantizedLeakyRelu); - -} // namespace qnn -} // namespace relay -} // namespace tvm diff --git a/src/relay/qnn/op/mul.cc b/src/relay/qnn/op/mul.cc deleted file mode 100644 index 73c6eed44889..000000000000 --- a/src/relay/qnn/op/mul.cc +++ /dev/null @@ -1,170 +0,0 @@ -/* - * 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 src/relay/qnn/op/mul.cc - * \brief QNN mul operator. - */ -#include -#include -#include - -#include "../../transforms/pattern_utils.h" -#include "../utils.h" -#include "op_common.h" - -namespace tvm { -namespace relay { -namespace qnn { - -/* - * \brief Canonicalizes the QNN mul op. - * \param attrs The QNN concatenate attrs. - * \param new_args The new mutated args to the call node. - * \param arg_types The types of input and output. - * \return The sequence of Relay ops for mul op. - */ -Expr QnnMulCanonicalize(const Attrs& attrs, const Array& new_args, - const Array& arg_types) { - Expr output; - - // Get the attrs. - QnnBinaryOpArguments args(new_args); - - // Get the input dtype and shape. - QnnBinaryOpTensorType input_type(arg_types, 0); - // data types - const auto int32_dtype = DataType::Int(32); - const auto float32_dtype = DataType::Float(32); - - const auto* broadcast_attrs = attrs.as(); - ICHECK(broadcast_attrs != nullptr); - - auto lhs_axis = broadcast_attrs->lhs_axis; - auto rhs_axis = broadcast_attrs->rhs_axis; - - if (IsConstScalar(args.lhs_scale) && IsConstScalar(args.rhs_scale)) { - /* - This is per-tensor quantized multiply. - - A tensor multiplication c = a * b can be written in terms of respective - quantized tensors, scales and zero points as - S_c * (Q_c - zp_c) = S_a * (Q_a - zp_a) * S_b * (Q_b - zp_b). - - We can consider the product (Q_a - zp_a) * (Q_b - zp_b) as a different - quantized tensor of c, Q', with corresponding scale S' = S_a * S_b and zp' = - 0. The quantized multiplication then becomes - Q_c = S'/S_c Q' + z_c, - which is essentially a requantization of tensor Q' into tensor Q_c. - */ - - auto lhs_shifted = Cast(args.lhs, int32_dtype); - auto rhs_shifted = Cast(args.rhs, int32_dtype); - - auto zero_scalar = MakeConstantScalar(int32_dtype, 0); - if (!IsEqualScalar(args.lhs_zero_point, zero_scalar)) { - lhs_shifted = Subtract(lhs_shifted, args.lhs_zero_point); - } - - if (!IsEqualScalar(args.rhs_zero_point, zero_scalar)) { - rhs_shifted = Subtract(rhs_shifted, args.rhs_zero_point); - } - - // Create a new tensor Q' - output = Multiply(lhs_shifted, rhs_shifted); - - // Get the adjusted new scale and zero points. - float lhs_scale_float = GetScalarFromConstant(args.lhs_scale); - float rhs_scale_float = GetScalarFromConstant(args.rhs_scale); - float new_scale_float = lhs_scale_float * rhs_scale_float; - auto new_input_scale = MakeConstantScalar(float32_dtype, new_scale_float); - auto new_input_zero_point = zero_scalar; - - // Requantize to get Q_c - output = Requantize(output, input_type.shape, new_input_scale, new_input_zero_point, - args.output_scale, args.output_zero_point, input_type.dtype); - } else if (lhs_axis == rhs_axis) { - /* - This is per-channel quantized multiply, assumming lhs_axis and rhs_axis are the same. - The subtract is done on the specified axis via broadcast. Then, we multiply lhs and rhs. - The output is requantized using new scale and axis. TODO: support different axes. - */ - - auto lhs_data = Cast(args.lhs, int32_dtype); - auto rhs_data = Cast(args.rhs, int32_dtype); - - auto zero_scalar = MakeConstantScalar(int32_dtype, 0); - if (!IsEqualScalar(args.lhs_zero_point, zero_scalar)) { - // Broadcast lhs zero point if needed - int rank = static_cast(input_type.shape.size()); - int axis = (lhs_axis < 0) ? ((rank > 0) ? rank + lhs_axis : 0) : lhs_axis; - Expr lhs_zero_broadcast = ExpandBiasToMatchAxis(Reshape(args.lhs_zero_point, - { - -1, - }), - rank, {axis}); - lhs_data = Subtract(lhs_data, Cast(lhs_zero_broadcast, DataType::Int(32))); - } - - if (!IsEqualScalar(args.rhs_zero_point, zero_scalar)) { - // Broadcast rhs zero point if needed - int rank = static_cast(input_type.shape.size()); - int axis = (rhs_axis < 0) ? ((rank > 0) ? rank + rhs_axis : 0) : rhs_axis; - Expr rhs_zero_broadcast = ExpandBiasToMatchAxis(Reshape(args.rhs_zero_point, - { - -1, - }), - rank, {axis}); - rhs_data = Subtract(rhs_data, Cast(rhs_zero_broadcast, DataType::Int(32))); - } - - // Create a new tensor Q' - output = Multiply(lhs_data, rhs_data); - - // Requantize to get Q_c - auto lhs_scales = GetFloatVectorFromConstant(args.lhs_scale); - auto rhs_scales = GetFloatVectorFromConstant(args.rhs_scale); - std::vector output_multipliers; - for (size_t i = 0; i < lhs_scales.size(); i++) { - double multiplier = static_cast(lhs_scales[i]) * static_cast(rhs_scales[i]); - output_multipliers.push_back(multiplier); - } - auto new_input_scale = MakeConstantTensor( - DataType::Float(32), {(int64_t)output_multipliers.size()}, output_multipliers); - - output = Requantize(output, input_type.shape, new_input_scale, zero_scalar, args.output_scale, - args.output_zero_point, input_type.dtype, lhs_axis); - - } else { - LOG(FATAL) << "Not supported: lhs_axis and rhs_axis are not the same."; - } - - return output; -} - -// QNN Multiplication operator. -QNN_REGISTER_BINARY_OP("mul") - .describe("Elementwise mul with broadcasting for quantized tensors.") - .set_support_level(11) - .set_attr("FTVMQnnCanonicalize", QnnMulCanonicalize) - .set_attr("TOpPattern", kBroadcast); - -} // namespace qnn -} // namespace relay -} // namespace tvm diff --git a/src/relay/qnn/op/op_common.h b/src/relay/qnn/op/op_common.h deleted file mode 100644 index 7ace12a26cfa..000000000000 --- a/src/relay/qnn/op/op_common.h +++ /dev/null @@ -1,456 +0,0 @@ -/* - * 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 src/relay/qnn/op/op_common.h - * \brief A set of utilities and common functionality for QNN ops. - */ -#ifndef TVM_RELAY_QNN_OP_OP_COMMON_H_ -#define TVM_RELAY_QNN_OP_OP_COMMON_H_ - -#include -#include -#include -#include -#include - -#include - -#include "../../op/type_relations.h" -#include "../../transforms/infer_layout_utils.h" -#include "../utils.h" - -namespace tvm { -namespace relay { -namespace qnn { - -TVM_REGISTER_NODE_TYPE(BroadcastAttrs); - -/* - * Number of inputs for the Qnn binary operators. - * Refer the QNN_REGISTER_BINARY_OP macro to see - * what the operators are. - */ -static constexpr int kNumQnnBinaryOpInputs = 8; - -/* - * Number of expected arg types. - */ -static constexpr int kNumQnnBinaryOpArgTypes = 9; - -/* - * \brief Simple struct to organize the inputs to the Qnn - * binary operators. The main reason to have a struct - * is to be able to perform the common checks needed at a - * central location. - */ -struct QnnBinaryOpArguments { - Expr lhs; - Expr rhs; - Expr lhs_scale; - Expr lhs_zero_point; - Expr rhs_scale; - Expr rhs_zero_point; - Expr output_scale; - Expr output_zero_point; - - explicit QnnBinaryOpArguments(const Array& new_args) { - ICHECK_EQ(new_args.size(), kNumQnnBinaryOpInputs); - int idx = 0; - lhs = new_args[idx++]; - rhs = new_args[idx++]; - lhs_scale = new_args[idx++]; - lhs_zero_point = new_args[idx++]; - rhs_scale = new_args[idx++]; - rhs_zero_point = new_args[idx++]; - output_scale = new_args[idx++]; - output_zero_point = new_args[idx++]; - ICHECK_EQ(idx, kNumQnnBinaryOpInputs); - } -}; - -/* - * Number of inputs for the Qnn unary operators. - */ -static constexpr int kNumQnnUnaryOpInputs = 5; - -/* - * Number of expected arg types. - */ -static constexpr int kNumQnnUnaryOpArgTypes = 6; - -/* - * \brief Simple struct to organize the inputs to the Qnn - * unary operators. The main reason to have a struct - * is to be able to perform the common checks needed at a - * central location. - */ -struct QnnUnaryOpArguments { - Expr x; - Expr scale; - Expr zero_point; - Expr output_scale; - Expr output_zero_point; - - explicit QnnUnaryOpArguments(const Array& new_args) { - ICHECK_EQ(new_args.size(), kNumQnnUnaryOpInputs); - int idx = 0; - x = new_args[idx++]; - scale = new_args[idx++]; - zero_point = new_args[idx++]; - output_scale = new_args[idx++]; - output_zero_point = new_args[idx++]; - ICHECK_EQ(idx, kNumQnnUnaryOpInputs); - } -}; - -/* - * \brief Simple structure to hold the input tensor's dtype - * and shape. This structure allows a common point to do - * all the validation checks for Qnn unary operators. - */ -struct QnnUnaryOpTensorType { - DataType dtype; - Array shape; - - explicit QnnUnaryOpTensorType(const Array& arg_types, const int32_t arg_idx) { - ICHECK_EQ(arg_types.size(), kNumQnnUnaryOpArgTypes); - auto tensor_type = arg_types[arg_idx].as(); - ICHECK(tensor_type != nullptr); - dtype = tensor_type->dtype; - shape = tensor_type->shape; - } -}; - -/* - * \brief Simple structure to hold the input tensor's dtype - * and shape. This structure allows a common point to do - * all the validation checks for Qnn binary operators. - */ -struct QnnBinaryOpTensorType { - DataType dtype; - Array shape; - - explicit QnnBinaryOpTensorType(const Array& arg_types, const int32_t arg_idx) { - ICHECK_EQ(arg_types.size(), kNumQnnBinaryOpArgTypes); - auto tensor_type = arg_types[arg_idx].as(); - ICHECK(tensor_type != nullptr); - dtype = tensor_type->dtype; - shape = tensor_type->shape; - } -}; - -/* - * \brief Converts the expression from expression's dtype - * to target dtype. This is mainly used for converting - * computations done in Int32 to lower precision Int8 or - * UInt8. - * \param expr The expression to whose dtype needs conversion. - * \param target_dtype The dtype of the target expression - * \return New expression with target dtype and possibly lower - * precision. - */ -inline Expr ConvertDtype(const Expr& expr, const DataType& target_dtype) { - auto q_min = GetQmin(target_dtype); - auto q_max = GetQmax(target_dtype); - auto output = Clip(expr, q_min, q_max); - return Cast(output, target_dtype); -} - -/* - * \brief Requantizes the given expression if expression's - * scale and zero point both do not match target scale and - * zero point. This is mainly needed for requantizing the - * input tensors with output tensor's scale and zero point - * to ease the computation of final quantized tensor. - * \param expr The expression on which the check needs to be performed. - * \param expr_scale The scale of the expression. - * \param expr_zero_point The zero point of the expression. - * \param target_scale The scale of the output tensor. - * \param target_zero_point The zero point of the output tensor. - * \param expr_shape The shape of the input expression. - * \return New expression that is requantized to target scale and zero - * point if the expression scale and zero points are different otherwise - * it simply casts the given expression to Int32 as no requantization is - * needed in this case. - */ -inline Expr RequantizeOrUpcast(const Expr& expr, const Expr& expr_scale, - const Expr& expr_zero_point, const Expr& target_scale, - const Expr& target_zero_point, const Array& expr_shape, - const int& axis = -1, - const DataType& target_dtype = DataType::Int(32)) { - auto result = expr; - if (!IsEqualScalar(expr_scale, target_scale) || - !IsEqualScalar(expr_zero_point, target_zero_point)) { - result = Requantize(expr, expr_shape, expr_scale, expr_zero_point, target_scale, - target_zero_point, target_dtype, axis); - } else { - result = Cast(result, target_dtype); - } - return result; -} - -/*! \brief Infer layout for QNN binary broadcast operators */ -inline InferCorrectLayoutOutput QnnBinaryBroadcastLayout( - const Attrs& attrs, const Array& new_in_layouts, const Array& old_in_layouts, - const Array& old_in_types) { - // Use Relay Binary Broadcast Infer correct layout. - auto layouts = BinaryBroadcastLayout(attrs, new_in_layouts, old_in_layouts, old_in_types); - - // Fill the layouts of remaining input tensors - scales and zero points. The layouts of these - // tensors can be treated as C. - Layout channel_layout = Layout("C"); - Array input_layouts = {layouts->input_layouts[0], - layouts->input_layouts[1], - channel_layout, - channel_layout, - channel_layout, - channel_layout, - channel_layout, - channel_layout}; - Array output_layouts = layouts->output_layouts; - return InferCorrectLayoutOutput(input_layouts, output_layouts, attrs); -} - -static inline bool QnnBroadcastRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - // Expected Types: lhs, rhs, lhs_scale, lhs_zero_point, rhs_scale, rhs_zero_point, output_scale, - // output_zero_point, out_type - ICHECK_EQ(types.size(), kNumQnnBinaryOpArgTypes); - - // Check the lhs and rhs types - for (size_t i = 0; i < 2; ++i) { - if (types[i].as()) { - return false; - } - } - // Check the scale and zero point types - for (size_t i = 2; i < 8; ++i) { - if (types[i].as()) { - return false; - } - } - - const auto* lhs_data = types[0].as(); - const auto* rhs_data = types[1].as(); - - if (lhs_data == nullptr || rhs_data == nullptr) { - return false; - } - - const BroadcastAttrs* broadcast_attrs = attrs.as(); - ICHECK(broadcast_attrs); - - auto lhs_rank = static_cast(lhs_data->shape.size()); - auto rhs_rank = static_cast(rhs_data->shape.size()); - - auto get_channel_axis = [](int rank, int axis_from_attr) { - if (rank <= 1) return 0; - if (axis_from_attr < 0) return rank + axis_from_attr; - return axis_from_attr; - }; - - const int lhs_axis = get_channel_axis(lhs_rank, broadcast_attrs->lhs_axis); - const int rhs_axis = get_channel_axis(rhs_rank, broadcast_attrs->rhs_axis); - - // If zero point and scale are scalar then axis doesn't matter. - bool lhs_scale_is_scalar = (types[2].as())->shape.size() == 0; - bool lhs_zp_is_scalar = (types[3].as())->shape.size() == 0; - bool rhs_scale_is_scalar = (types[4].as())->shape.size() == 0; - bool rhs_zp_is_scalar = (types[5].as())->shape.size() == 0; - - if (!(lhs_scale_is_scalar && lhs_zp_is_scalar)) { - ICHECK_LT(lhs_axis, lhs_rank > 0 ? lhs_rank : 1) - << "lhs_axis " << broadcast_attrs->lhs_axis << " is out of range"; - ICHECK_GE(lhs_axis, 0) << "lhs_axis " << broadcast_attrs->lhs_axis << " is out of range"; - } - - if (!(rhs_scale_is_scalar && rhs_zp_is_scalar)) { - ICHECK_LT(rhs_axis, rhs_rank > 0 ? rhs_rank : 1) - << "rhs_axis " << broadcast_attrs->rhs_axis << " is out of range"; - ICHECK_GE(rhs_axis, 0) << "rhs_axis " << broadcast_attrs->rhs_axis << " is out of range"; - } - - PrimExpr lhs_axis_shape; - if (lhs_rank > 0) { - lhs_axis_shape = lhs_data->shape[lhs_axis]; - } else { - lhs_axis_shape = Integer(1); - } - - PrimExpr rhs_axis_shape; - if (rhs_rank > 0) { - rhs_axis_shape = rhs_data->shape[rhs_axis]; - } else { - rhs_axis_shape = Integer(1); - } - - // Check and assign types for scale and zero points. - AssignType(types[2], DataType::Float(32), lhs_axis_shape, reporter); // lhs_scale - AssignType(types[3], DataType::Int(32), lhs_axis_shape, reporter); // lhs_zero_point - AssignType(types[4], DataType::Float(32), rhs_axis_shape, reporter); // rhs_scale - AssignType(types[5], DataType::Int(32), rhs_axis_shape, reporter); // rhs_zero_point - - ICHECK(IsScalarType(types[6], DataType::Float(32))); // output_scale - ICHECK(IsScalarType(types[7], DataType::Int(32))); // output_zero_point - - // Collect the input tensor and output tensor devoid of scale and zero points to reuse Relay - // BroadcastRel infer type function. - Array tensor_types = {types[0], types[1], types[8]}; - return BroadcastRel(tensor_types, 3, attrs, reporter); -} - -/*! Quick helper macro - * - Expose a positional make function to construct the node. - * - Register op to the registry. - * - * We make the decision to always only expose positional argument. - * We will do rewrapping in the frontend to support language - * sugars such as keyword arguments and default value. - * - * \param OpName the name of registry. - */ -#define QNN_REGISTER_BINARY_OP(OpName) \ - TVM_REGISTER_GLOBAL("relay.qnn.op._make." OpName) \ - .set_body_typed([](Expr lhs, Expr rhs, Expr lhs_scale, Expr lhs_zero_point, Expr rhs_scale, \ - Expr rhs_zero_point, Expr output_scale, Expr output_zero_point, \ - int lhs_axis, int rhs_axis) { \ - static const Op& op = Op::Get("qnn." OpName); \ - auto attrs = make_object(); \ - attrs->lhs_axis = lhs_axis; \ - attrs->rhs_axis = rhs_axis; \ - return Call(op, \ - {lhs, rhs, lhs_scale, lhs_zero_point, rhs_scale, rhs_zero_point, output_scale, \ - output_zero_point}, \ - Attrs(attrs), {}); \ - }); \ - RELAY_REGISTER_OP("qnn." OpName) \ - .set_attrs_type() \ - .set_num_inputs(kNumQnnBinaryOpInputs) \ - .add_argument("lhs", "Tensor", "The left hand side quantized tensor.") \ - .add_argument("rhs", "Tensor", "The right hand side quantized tensor.") \ - .add_argument("lhs_scale", "Tensor", "The scale of the lhs tensor.") \ - .add_argument("lhs_zero_point", "Tensor", "The zero_point of the lhs tensor.") \ - .add_argument("rhs_scale", "Tensor", "The scale of the rhs tensor.") \ - .add_argument("rhs_zero_point", "Tensor", "The zero_point of the rhs tensor.") \ - .add_argument("output_scale", "Tensor", "The scale of the output tensor.") \ - .add_argument("output_zero_point", "Tensor", "The zero_point of the output tensor.") \ - .add_argument("lhs_axis", "Tensor", "The channel quantization of the lhs tensor.") \ - .add_argument("rhs_axis", "Tensor", "The channel quantization of the rhs tensor.") \ - .add_type_rel("QnnBroadcast", QnnBroadcastRel) \ - .set_attr("TNonComputational", true) \ - .set_attr("FInferCorrectLayout", QnnBinaryBroadcastLayout) - -static inline bool QnnElementwiseUnaryFuncRel(const Array& types, int num_inputs, - const Attrs& attrs, const TypeReporter& reporter) { - // Expected Types: data, scale, zero_point, output_scale, output_zero_point - ICHECK_EQ(types.size(), 6); - const auto* x = types[0].as(); - if (x == nullptr) return false; - ICHECK(x->dtype == DataType::Int(8) || x->dtype == DataType::UInt(8)) - << "Expected quantized type(int8, uint8) for input but was " << x->dtype; - - // Check the types of scale and zero points. - for (size_t i = 1; i < 5; ++i) { - if (types[i].as()) { - return false; - } - } - ICHECK(IsScalarType(types[1], DataType::Float(32))); // scale - ICHECK(IsScalarType(types[2], DataType::Int(32))); // zero_point - ICHECK(IsScalarType(types[3], DataType::Float(32))); // output_scale - ICHECK(IsScalarType(types[4], DataType::Int(32))); // output_zero_point - - // Assign types for scale and zero points. - reporter->Assign(types[1], TensorType({}, DataType::Float(32))); // scale - reporter->Assign(types[2], TensorType({}, DataType::Int(32))); // zero_point - reporter->Assign(types[3], TensorType({}, DataType::Float(32))); // output_scale - reporter->Assign(types[4], TensorType({}, DataType::Int(32))); // output_zero_point - - // Collect the input tensor and output tensor devoid of scale and zero points to reuse Relay - // IdentityRel infer type function. - Array tensor_types = {types[0], types[5]}; - return IdentityRel(tensor_types, 2, attrs, reporter); -} - -static inline Expr LegalizeExpr(const Expr& expr) { - // Canonicalizations should not contain qnn ops, so use this - // to lower expressions automatically after using things like qnn.dequantize - // in the lowering process. - auto mod = IRModule::FromExpr(expr); - mod = transform::Legalize()(mod); - if (expr.as()) { - return mod->Lookup("main"); - } else { - return mod->Lookup("main").as()->body; - } -} - -/*! Quick helper macro - * - Expose a positional make function to construct the node. - * - Register op to the registry. - * - * For Unary Operators which also take in QParams. - * - * \param OpName the name of registry. - */ -#define QNN_CREATE_UNARY_ELEMENTWISE_OP(OpName) \ - TVM_REGISTER_GLOBAL("relay.qnn.op._make." OpName) \ - .set_body_typed( \ - [](Expr x, Expr scale, Expr zero_point, Expr output_scale, Expr output_zero_point) { \ - return Call(Op::Get("qnn." OpName), \ - {x, scale, zero_point, output_scale, output_zero_point}, Attrs(), {}); \ - }); \ - \ - RELAY_REGISTER_OP("qnn." OpName) \ - .describe("Elementwise " OpName " for quantized tensors.") \ - .set_num_inputs(5) \ - .add_argument("data", "Quantized Tensor", "The input data.") \ - .add_argument("scale", "Tensor", "The quantization scale of the input tensor.") \ - .add_argument("zero_point", "Tensor", "The quantization zero_point of the input tensor.") \ - .add_argument("output_scale", "Tensor", "The quantization scale of the output tensor.") \ - .add_argument("output_zero_point", "Tensor", \ - "The quantization zero_point of the output tensor.") \ - .set_support_level(11) \ - .add_type_rel("qnn." OpName, QnnElementwiseUnaryFuncRel) \ - .set_attr("TNonComputational", true) - -/*! Quick helper macro - * Create a default canonicalization for a QNN operator, which dequantizes the operator - * runs the calculation using the provided Call func, and then requantizes. - * - * FloatingPointFunc is usually a handle from "src/relay/transforms/pattern_utils.h" - * - * \param FloatingPointFunc the floating point function with function signature `Expr Erf(Expr e)` - */ -#define QNN_UNARY_OP_DEFAULT_CANONICALIZATION(FloatingPointFunc) \ - [](const Attrs& attrs, const Array& new_args, const Array& arg_types) { \ - QnnUnaryOpArguments args(new_args); \ - QnnUnaryOpTensorType input_type(arg_types, 0); \ - Expr dequantized_arg = MakeDequantize(args.x, args.scale, args.zero_point, -1); \ - Expr output = FloatingPointFunc(dequantized_arg); \ - Expr result = \ - MakeQuantize(output, args.output_scale, args.output_zero_point, -1, input_type.dtype); \ - return LegalizeExpr(result); \ - } -} // namespace qnn -} // namespace relay -} // namespace tvm - -#endif // TVM_RELAY_QNN_OP_OP_COMMON_H_ diff --git a/src/relay/qnn/op/quantize.cc b/src/relay/qnn/op/quantize.cc deleted file mode 100644 index 8ed1f9ef4c4f..000000000000 --- a/src/relay/qnn/op/quantize.cc +++ /dev/null @@ -1,190 +0,0 @@ -/* - * 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 src/relay/qnn/op/quantize.cc - * \brief QNN quantize operator. Quantize operator converts from unquantized - * domain to quantized domain. - */ - -#include -#include -#include - -#include "../../transforms/pattern_utils.h" -#include "../utils.h" - -namespace tvm { -namespace relay { -namespace qnn { - -TVM_REGISTER_NODE_TYPE(QuantizeAttrs); - -bool QuantizeRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - ICHECK_EQ(types.size(), 4); - const auto* data = types[0].as(); - - if (data == nullptr) { - return false; - } - - const auto input_dtype = data->dtype; - ICHECK(input_dtype == DataType::Float(32)) - << "Input type should be one of float32 but was " << input_dtype; - - const auto* quantize_attrs = attrs.as(); - int axis = quantize_attrs->axis; - auto rank = static_cast(data->shape.size()); - axis = (axis < 0) ? ((rank > 0) ? data->shape.size() + axis : 0) : axis; - - // If zero point and scale are scalar then axis doesnt matter. - bool scale_is_scalar, zp_is_scalar; - - if (auto ttype = types[1].as()) { - scale_is_scalar = ttype->shape.size() == 0; - } else { - ICHECK(types[1].as()) - << "Quantize: expect to be TensorType but get " << types[1]; - return false; - } - - if (auto ttype = types[2].as()) { - zp_is_scalar = ttype->shape.size() == 0; - } else { - ICHECK(types[2].as()) - << "Quantize: expect to be TensorType but get " << types[2]; - return false; - } - - if (!(scale_is_scalar && zp_is_scalar)) { - ICHECK_LT(axis, rank > 0 ? rank : 1) << "axis " << quantize_attrs->axis << " is out of range"; - ICHECK_GE(axis, 0) << "axis " << quantize_attrs->axis << " is out of range"; - } - - PrimExpr axis_shape; - if (rank > 0) { - axis_shape = data->shape[axis]; - } else { - axis_shape = Integer(1); - } - // Check and assign types for scale and zero points. - AssignType(types[1], DataType::Float(32), axis_shape, reporter); // scale - AssignType(types[2], DataType::Int(32), axis_shape, reporter); // zero point - - const Array oshape = data->shape; - const DataType out_dtype = quantize_attrs->out_dtype; - ICHECK(out_dtype == DataType::Int(8) || out_dtype == DataType::UInt(8) || - out_dtype == DataType::Int(16) || out_dtype == DataType::UInt(16) || - out_dtype == DataType::Int(32)) - << "Output type should be one of [int8, unit8, int16, uint16, int32] but was " << out_dtype; - // assign output type - reporter->Assign(types[3], TensorType(oshape, out_dtype)); - return true; -} - -Expr MakeQuantize(Expr data, Expr output_scale, Expr output_zero_point, int axis, - DataType out_dtype) { - auto attrs = make_object(); - attrs->axis = axis; - attrs->out_dtype = std::move(out_dtype); - // result_quantized_value = result_zero_point + result_real_value / result_scale. - // A more detailed explanation can be found here - - // https://github.com/google/gemmlowp/blob/master/doc/quantization.md - static const Op& op = Op::Get("qnn.quantize"); - return Call(op, {data, output_scale, output_zero_point}, Attrs(attrs), {}); -} - -Expr QuantizeLower(const Expr& input_tensor, const Expr& output_scale, - const Expr& output_zero_point, const Array& types, - const QuantizeAttrs* attrs) { - ICHECK_EQ(types.size(), 4); - auto in_type = types[0]; - auto in_tensor_type = in_type.as(); - ICHECK(in_tensor_type != nullptr) << "Type information missing." - << " Please run infer_type pass."; - Array input_shape = in_tensor_type->shape; - - const auto out_dtype = attrs->out_dtype; - auto axis = attrs->axis; - - size_t n_dim = input_shape.size(); - - // Wrap axis from negative to positive if needed. - if (axis < 0) { - axis = static_cast(n_dim) + axis; - } - - auto expanded_output_scale = output_scale; - if (!IsConstScalar(output_scale) && !IsScalarType(types[1])) { - expanded_output_scale = ExpandBiasToMatchAxis(output_scale, n_dim, {axis}); - } - - auto expanded_output_zero_point = output_zero_point; - if (!IsConstScalar(output_zero_point) && !IsScalarType(types[2])) { - expanded_output_zero_point = ExpandBiasToMatchAxis(output_zero_point, n_dim, {axis}); - } - - const int32_t min_val = GetQmin(out_dtype); - const int32_t max_val = GetQmax(out_dtype); - auto scale_data = Round(Divide(input_tensor, expanded_output_scale)); - auto add_zero_point = Add(scale_data, Cast(expanded_output_zero_point, DataType::Float(32))); - auto clamped_output = Clip(add_zero_point, min_val, max_val); - return Cast(clamped_output, out_dtype); -} - -Expr QuantizeQnnCanonicalize(const Attrs& attrs, const Array& new_args, - const Array& types) { - ICHECK_EQ(new_args.size(), 3); - auto& data = new_args[0]; - auto& output_scale = new_args[1]; - auto& output_zero_point = new_args[2]; - const auto* quantize_attrs = attrs.as(); - ICHECK(quantize_attrs != nullptr); - - return QuantizeLower(data, output_scale, output_zero_point, types, quantize_attrs); -} - -RELAY_REGISTER_OP("qnn.quantize") - .describe(R"code(Quantizes the input and produces quantized output. -The input can be either float or quantized(int8, unit8). If the input is float, -this op takes scale and zero point and quantize the float value to -quantized output, in int8 or uint8 format. If the input is quantized value, -the op requantize the input (of a certain type, with a given scale and zero -point) to the output of the same or different type with a same or different -scale and zero point. -- **data**: Tensor of any shape to quantize. The input data can be of floating point - or quantized. -)code" TVM_ADD_FILELINE) - .set_attrs_type() - .set_num_inputs(3) - .add_argument("data", "Tensor", "The tensor to quantize.") - .add_argument("output_scale", "Tensor", "The quantization scale of the output tensor.") - .add_argument("output_zero_point", "Tensor", - "The quantization zero_point of the output tensor.") - .set_support_level(11) - .add_type_rel("Quantize", QuantizeRel) - .set_attr("TNonComputational", true) - .set_attr("FTVMQnnCanonicalize", QuantizeQnnCanonicalize); - -TVM_REGISTER_GLOBAL("relay.qnn.op._make.quantize").set_body_typed(MakeQuantize); - -} // namespace qnn -} // namespace relay -} // namespace tvm diff --git a/src/relay/qnn/op/requantize.cc b/src/relay/qnn/op/requantize.cc deleted file mode 100644 index 2dd74e1321bf..000000000000 --- a/src/relay/qnn/op/requantize.cc +++ /dev/null @@ -1,566 +0,0 @@ -/* - * 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 src/relay/qnn/op/requantize.cc - * \brief QNN requantize operator. - */ - -#include -#include -#include - -#include "../../op/op_common.h" -#include "../../transforms/infer_layout_utils.h" -#include "../../transforms/pattern_utils.h" -#include "../utils.h" -#include "./requantize_config.h" - -namespace tvm { -namespace relay { -namespace qnn { - -TVM_REGISTER_NODE_TYPE(RequantizeAttrs); - -InferCorrectLayoutOutput RequantizeInferCorrectLayout(const Attrs& attrs, - const Array& new_in_layouts, - const Array& old_in_layouts, - const Array& old_in_types) { - const auto* attrs_ptr = attrs.as(); - ICHECK(attrs_ptr); - ObjectPtr param = make_object(*attrs_ptr); - - Array> old_in_shapes; - for (auto old_in_t : old_in_types) { - ICHECK(old_in_t.as()); - old_in_shapes.push_back(old_in_t.as()->shape); - } - - Array input_layouts, output_layouts; - if (new_in_layouts.defined()) { - // Adapt to new layout. The axis has to change. - // Record original reduce axis. Convert to the modified layout axis. - ICHECK_EQ(new_in_layouts.size(), 5); - ICHECK_EQ(old_in_layouts.size(), 5); - - // 1) Get the axis. - int axis = param->axis; - axis = (axis == -1) ? old_in_shapes[0].size() - 1 : axis; - - // 2) Collect the original axis - std::string old_dim = old_in_layouts[0][axis].name(); - - // 3) Collect the new axes by walking new_layout. - tvm::Integer new_axis; - std::string new_layout_string = ""; - int axis_index = 0; - for (auto iter_var : new_in_layouts[0]->axes) { - const auto& layout_axis = LayoutAxis::Get(iter_var); - const std::string& layout_dim = layout_axis.name(); - if (old_dim == layout_dim) { - new_axis = tvm::Integer(axis_index); - } - - if (layout_axis.IsPrimal()) { - new_layout_string += layout_dim; - axis_index++; - } else { - // Propogate layout if input_zero_point and input_scale are scalar values. - ICHECK_GE(old_in_types.size(), 3); - if (IsScalarType(old_in_types[1]) && IsScalarType(old_in_types[2])) { - new_layout_string += std::to_string(new_in_layouts[0].FactorOf(layout_axis)) + layout_dim; - axis_index++; - } - } - } - - // 4) Set the new axis and layout. - Layout new_layout = Layout(new_layout_string); - - // Fill the layouts of remaining input tensors - scales and zero points. The layouts of these - // tensors can be treated as channel layout. - Layout channel_layout = Layout("C"); - input_layouts = {new_layout, channel_layout, channel_layout, channel_layout, channel_layout}; - output_layouts = {new_layout}; - param->axis = new_axis.IntValue(); - } else if (old_in_layouts.defined()) { - // If the new layout is undefined, set the old layout as the inferred layout. - ICHECK_EQ(old_in_layouts.size(), 5); - - Layout old_layout = old_in_layouts[0]; - - // Fill the layouts of remaining input tensors - scales and zero points. The layouts of these - // tensors can be treated as channel layout. - Layout channel_layout = Layout("C"); - input_layouts = {old_layout, channel_layout, channel_layout, channel_layout, channel_layout}; - output_layouts = {old_layout}; - } else { - // Set the layouts to undef. - Layout undef = Layout::Undef(); - input_layouts = Array(5, undef); - output_layouts = {undef}; - } - - return InferCorrectLayoutOutput(input_layouts, output_layouts, Attrs(param)); -} - -bool has_current_target_sse41_support() { - auto target_has_feature_fn_ptr = tvm::runtime::Registry::Get("target.target_has_feature"); - ICHECK(target_has_feature_fn_ptr) << "Function target.target_has_feature not found"; - return (*target_has_feature_fn_ptr)("sse4.1", Target::Current(true)); -} - -/* - * \brief TONEAREST is the standard rounding where the value is rounded away - * from zero at midpoints (for example, -1.5 rounds to -2). - * \param input_tensor The input tensor to rounding op. - * \return The sequence of existing Relay ops. - */ -template -Expr Tonearest(const Expr& input_tensor) { - if (has_current_target_sse41_support()) return Round(input_tensor); - - auto half = MakeConstantScalar(DataType::Float(Bits), 0.5f); - auto zero = MakeConstantScalar(DataType::Float(Bits), 0.f); - auto pos_one = MakeConstantScalar(DataType::Float(Bits), +1.f); - auto neg_one = MakeConstantScalar(DataType::Float(Bits), -1.f); - auto multiplier = Where(Less(input_tensor, zero), neg_one, pos_one); - auto half_multiplied = Multiply(half, multiplier); - auto input_tensor_biased = Add(input_tensor, half_multiplied); - auto input_tensor_biased_multiplied = Multiply(input_tensor_biased, multiplier); - auto input_tensor_biased_multiplied_int = - Cast(input_tensor_biased_multiplied, DataType::Int(Bits)); - auto input_tensor_biased_multiplied_float = - Cast(input_tensor_biased_multiplied_int, DataType::Float(Bits)); - auto input_tensor_rounded = Multiply(input_tensor_biased_multiplied_float, multiplier); - return Where(IsFinite(input_tensor), input_tensor_rounded, input_tensor); -} - -/* - * \brief UPWARD is the standard rounding except at midpoints where the value - * is rounded to positive infinity (for example, -1.5 rounds to -1). - * \param input_tensor The input tensor to rounding op. - * \return The sequence of existing Relay ops. - */ -template -Expr Upward(const Expr& input_tensor) { - auto half = MakeConstantScalar(DataType::Float(Bits), 0.5f); - auto input_tensor_biased = Add(input_tensor, half); - if (has_current_target_sse41_support()) return Floor(input_tensor_biased); - - auto zero = MakeConstantScalar(DataType::Float(Bits), 0.f); - auto one = MakeConstantScalar(DataType::Float(Bits), +1.f); - auto input_tensor_biased_int = Cast(input_tensor_biased, DataType::Int(Bits)); - auto input_tensor_biased_float = Cast(input_tensor_biased_int, DataType::Float(Bits)); - auto is_subtraction_not_necessary = - LogicalOr(Equal(input_tensor_biased, input_tensor_biased_float), - GreaterEqual(input_tensor_biased, zero)); - auto input_tensor_rounded = Where(is_subtraction_not_necessary, input_tensor_biased_float, - Subtract(input_tensor_biased_float, one)); - return Where(IsFinite(input_tensor), input_tensor_rounded, input_tensor); -} - -// Lowering of qnn.requantize op - -/* - * \brief Lower requantize to a sequence of ops. - * \param input_tensor The input tensor to requantize op. - * \param param The requantize op attrs. - * \param input_shape The input tensor shape of the requantize op. - * \return The sequence of existing Relay ops. - * \note RequantizationInt using only integer computation. Here, the computation is - * converted to a fixed point computation by computing output multiplier - * and shift. This is useful, if the target device does not support/have - * very expensive floating point computations. - * - * The whole computation this can be broken down into following steps - * 1) Calculate the integer multiplier and integer shift. - * 2) Subtract the input integer zero point. - * 3) Perform fixed point multiplication. - * 4) Add the output zero point. - * 5) Cast to the out_dtype. - */ -Expr RequantizeLowerInt(const Expr& input_tensor, const Expr& input_scale, - const Expr& input_zero_point, const Expr& output_scale, - const Expr& output_zero_point, const RequantizeAttrs* param, - const Array& input_shape, const DataType& out_dtype) { - auto tensor = Cast(input_tensor, DataType::Int(32)); - auto zero_scalar = MakeConstantScalar(DataType::Int(32), 0); - if (!IsEqualScalar(input_zero_point, zero_scalar)) { - // Broadcast input zero point if needed. - int rank = static_cast(input_shape.size()); - int axis = (param->axis < 0) ? ((rank > 0) ? rank + param->axis : 0) : param->axis; - Expr input_zero_broadcast = ExpandBiasToMatchAxis(Reshape(input_zero_point, - { - -1, - }), - rank, {axis}); - tensor = Subtract(tensor, Cast(input_zero_broadcast, DataType::Int(32))); - } - - // 2) If the input and output scales are same, we can skip the fixed point multiplication. Check - // if the input scale is per-tensor or per-channel. If it is per-tensor, there is single scale for - // the whole tensor. For per-channel (aka per-axis), there is a vector of scales for the input - // tensor. Depending on the quantization type, the fixed point multiplication routing is called. - const bool is_upward_rounding = (param->rounding == "UPWARD"); - auto scaled_int32_t = tensor; - float output_scale_float = GetScalarFromConstant(output_scale); - if (IsConstScalar(input_scale)) { - // This is per-tensor quantization. Single scale. - float input_scale_float = GetScalarFromConstant(input_scale); - double double_multiplier = - static_cast(input_scale_float) / static_cast(output_scale_float); - // Skip if input and output scales are same. - if (!IsEqualScalar(input_scale, output_scale)) { - auto [fixed_point_multiplier, shift] = GetFixedPointMultiplierShift(double_multiplier); - - // When using upward rounding (i.e., x.5 rounded to x+1), leverage - // the FixedPointMultiply operator - scaled_int32_t = - (is_upward_rounding - ? FixedPointMultiply(scaled_int32_t, fixed_point_multiplier, shift) - : FixedPointMultiplyToNearest(scaled_int32_t, double_multiplier, input_shape)); - } - - } else { - // This is per-channel (per=axis) quantization. - std::vector double_multipliers; - auto input_axis_scales = GetFloatVectorFromConstant(input_scale); - for (auto input_axis_scale : input_axis_scales) { - double multiplier = - static_cast(input_axis_scale) / static_cast(output_scale_float); - double_multipliers.push_back(multiplier); - } - int axis = param->axis; - axis = (axis == -1) ? input_shape.size() - 1 : axis; - - // When using "upward" rounding, leverage the FixedPointMultiplyPerAxis operator, - // for "tonearest" rounding - lower to multiply, add, shift operators sequence. - scaled_int32_t = is_upward_rounding - ? FixedPointMultiplyPerChannel(scaled_int32_t, double_multipliers, axis) - : FixedPointMultiplyPerChannelToNearest(scaled_int32_t, double_multipliers, - input_shape, axis); - } - - // 3) Add the output zero point. - auto shifted_int32_t = scaled_int32_t; - if (!IsEqualScalar(output_zero_point, zero_scalar)) { - shifted_int32_t = Add(Cast(output_zero_point, DataType::Int(32)), scaled_int32_t); - } - - // 4) Clip to the out_dtype min/max. Skip clipping if out_dtype is Int32. The fixed point - // multiplication keeps the value in int32 range. - if (out_dtype == DataType::Int(32)) { - return shifted_int32_t; - } - - auto q_min = GetQmin(out_dtype); - auto q_max = GetQmax(out_dtype); - auto clipped_t = Clip(shifted_int32_t, q_min, q_max); - return Cast(clipped_t, out_dtype); -} - -// Lowering of qnn.requantize op - -/* - * \brief Lower requantize to a sequence of ops. - * \param input_tensor The input tensor to requantize op. - * \param param The requantize op attrs. - * \param input_shape The input tensor shape of the requantize op. - * \return The sequence of existing Relay ops. - * \note RequantizationFP using floating computation. All multiplication/sub/sum - * occurs in floating point data type and only at the end is converted to - * int32 data type and clamped for output data type. - * - * The whole computation this can be broken down into following steps - * 1) Subtract the input zero point. - * 2) Perform multiplication. - * 3) Add the output zero point. - * 4) Cast to the out_dtype. - */ -template -Expr RequantizeLowerFP(const Expr& input_tensor, const Expr& input_scale, - const Expr& input_zero_point, const Expr& output_scale, - const Expr& output_zero_point, const RequantizeAttrs* param, - const Array& input_shape, const DataType& out_dtype) { - auto tensor = Cast(input_tensor, DataType::Float(Bits)); - auto zero_scalar = MakeConstantScalar(DataType::Int(32), 0); - if (!IsEqualScalar(input_zero_point, zero_scalar)) { - // Broadcast input zero point if needed. - int rank = static_cast(input_shape.size()); - int axis = (param->axis < 0) ? ((rank > 0) ? rank + param->axis : 0) : param->axis; - Expr input_zero_broadcast = ExpandBiasToMatchAxis(Reshape(input_zero_point, - { - -1, - }), - rank, {axis}); - tensor = Subtract(tensor, Cast(input_zero_broadcast, DataType::Float(Bits))); - } - - // 2) If the input and output scales are same, we can skip the multiplication. Check - // if the input scale is per-tensor or per-channel. If it is per-tensor, there is single scale for - // the whole tensor. For per-channel (aka per-axis), there is a vector of scales for the input - // tensor. Depending on the quantization type, the fixed point multiplication routing is called. - auto scaled_fp_t = tensor; - double output_scale_float = GetScalarFromConstant(output_scale); - if (IsConstScalar(input_scale)) { - // This is per-tensor quantization. Single scale. - double input_scale_float = GetScalarFromConstant(input_scale); - double double_multiplier = static_cast(input_scale_float) / output_scale_float; - // Skip if input and output scales are same. - if (!IsEqualScalar(input_scale, output_scale)) { - double multiplier = double_multiplier; - auto m_scalar = MakeConstantScalar(DataType::Float(Bits), multiplier); - scaled_fp_t = Multiply(m_scalar, scaled_fp_t); - } - - } else { - // This is per-channel (per=axis) quantization. - std::vector double_multipliers; - auto input_axis_scales = GetFloatVectorFromConstant(input_scale); - double output_scale_float = GetScalarFromConstant(output_scale); - for (auto input_axis_scale : input_axis_scales) { - double multiplier = static_cast(input_axis_scale) / output_scale_float; - double_multipliers.push_back(multiplier); - } - int axis = param->axis; - axis = (axis == -1) ? input_shape.size() - 1 : axis; - - auto fixed_pt_multiplier_expr = MakeConstantTensor( - DataType::Float(Bits), {(int64_t)double_multipliers.size()}, double_multipliers); - size_t n_dim = input_shape.size(); - auto exp_fixed_pt_multiplier_expr = - ExpandBiasToMatchAxis(fixed_pt_multiplier_expr, n_dim, {axis}); - - scaled_fp_t = Multiply(scaled_fp_t, exp_fixed_pt_multiplier_expr); - } - - // 3) Add the output zero point. - auto shifted_fp_t = scaled_fp_t; - if (!IsEqualScalar(output_zero_point, zero_scalar)) { - shifted_fp_t = Add(shifted_fp_t, Cast(output_zero_point, DataType::Float(Bits))); - } - - if (param->rounding == "UPWARD") { - shifted_fp_t = Upward(shifted_fp_t); - } else /*if (param->rounding == "TONEAREST")*/ { - shifted_fp_t = Tonearest(shifted_fp_t); - } - - shifted_fp_t = Cast(shifted_fp_t, DataType::Int(32)); - // 4) Clip to the out_dtype min/max. Skip clipping if out_dtype is Int32. The fixed point - // multiplication keeps the value in int32 range. - if (out_dtype == DataType::Int(32)) { - return shifted_fp_t; - } - - auto q_min = GetQmin(out_dtype); - auto q_max = GetQmax(out_dtype); - auto clipped_t = Clip(shifted_fp_t, q_min, q_max); - return Cast(clipped_t, out_dtype); -} - -// Lowering of qnn.requantize op -/* - * \brief Lower requantize to a sequence of ops. - * \param input_tensor The input tensor to requantize op. - * \param param The requantize op attrs. - * \param input_shape The input tensor shape of the requantize op. - * \return The sequence of existing Relay ops. - */ -Expr RequantizeLower(const Expr& input_tensor, const Expr& input_scale, - const Expr& input_zero_point, const Expr& output_scale, - const Expr& output_zero_point, const RequantizeAttrs* param, - const Array& input_shape, const DataType& out_dtype) { - // Check output scale validity. - ICHECK_NE(GetScalarFromConstant(output_scale), 0.0) - << "QNN requantize output scale can not be equal to 0.0"; - // Check rounding validity. - ICHECK(param->rounding == "UPWARD" || param->rounding == "TONEAREST") - << "QNN requantize supports two rounding modes - UPWARD and " - << "TONEAREST"; - // Check compute_dtype validity. - ICHECK(param->compute_dtype == "int64" || param->compute_dtype == "float32" || - param->compute_dtype == "float64") - << "QNN requantize supports three compute_dtype variants - \"int64\", \"float32\" and " - "\"float64\""; - if (param->compute_dtype == "float32") { - return RequantizeLowerFP<32>(input_tensor, input_scale, input_zero_point, output_scale, - output_zero_point, param, input_shape, out_dtype); - } else if (param->compute_dtype == "float64") { - return RequantizeLowerFP<64>(input_tensor, input_scale, input_zero_point, output_scale, - output_zero_point, param, input_shape, out_dtype); - } else /*if (param->compute_dtype == "int64") */ { - return RequantizeLowerInt(input_tensor, input_scale, input_zero_point, output_scale, - output_zero_point, param, input_shape, out_dtype); - } -} - -/* - * \brief Forward rewrite the requantize op. - * \param ref_call The original call that will be lowered. - * \param new_args The new mutated args to the call node. - * \param ctx The node context. - * \return The sequence of Relay ops for requantize op. - * \note Lowering of the requantize operation. The requantize operator converts - * one quantized tensor to another quantized tensor. For the output - * tensor, we are provided with output scale and zero point. The - * computation looks like this - * - * Q_output = zp_output + (scale_input)/(scale_ouptut) * (Q_input - zp_input) - */ -Expr RequantizeQnnCanonicalize(const Attrs& attrs, const Array& new_args, - const Array& types) { - ICHECK_EQ(new_args.size(), 5); - auto& quantized_data = new_args[0]; - auto& input_scale = new_args[1]; - auto& input_zero_point = new_args[2]; - auto& output_scale = new_args[3]; - auto& output_zero_point = new_args[4]; - const auto* param = attrs.as(); - const RequantizeConfig& cfg = RequantizeConfig::Current(); - - ICHECK(param != nullptr); - - const_cast(param)->rounding = - SelectRequntizeParameter(param->rounding, cfg->get_rounding(), cfg->is_default, "rounding"); - const_cast(param)->compute_dtype = SelectRequntizeParameter( - param->compute_dtype, cfg->get_compute_dtype(), cfg->is_default, "compute_dtype"); - - // Find input shape. - ICHECK_EQ(types.size(), 6); - auto in_type = types[0]; - auto in_tensor_type = in_type.as(); - ICHECK(in_tensor_type != nullptr) << "Type information missing." - << " Please run infer_type pass."; - Array input_shape = in_tensor_type->shape; - - // Find the output dtype. - auto out_type = types[5]; - auto out_tensor_type = out_type.as(); - ICHECK(out_tensor_type != nullptr) << "Type information missing." - << " Please run infer_type pass."; - auto out_dtype = out_tensor_type->dtype; - return RequantizeLower(quantized_data, input_scale, input_zero_point, output_scale, - output_zero_point, param, input_shape, out_dtype); -} - -/* - * \brief Infer shape function of Requantize op. - * \param types The types of input args. - * \param num_inputs The number of inputs. - * \param attrs The op attributes. - * \param reporter The type reporter that sets the dtype and shapes. - * \return True if the infer shape succeeded. - */ -bool RequantizeRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - // Expected Types: data, input_scale, input_zero_point, output_scale, output_zero_point, output - ICHECK_EQ(types.size(), 6); - const auto* data = types[0].as(); - - if (data == nullptr) { - return false; - } - - // Check the scale and zero point types - for (size_t i = 3; i < 5; ++i) { - if (types[i].as()) { - return false; - } - } - const auto in_dtype = data->dtype; - ICHECK(in_dtype == DataType::Int(8) || in_dtype == DataType::UInt(8) || - in_dtype == DataType::Int(16) || in_dtype == DataType::Int(32) || - in_dtype == DataType::Int(64)) - << "Input type should be one of [int8, uint8, int16, int32, int64] but was " << in_dtype; - - const RequantizeAttrs* requantize_attrs = attrs.as(); - int axis = requantize_attrs->axis; - auto rank = static_cast(data->shape.size()); - axis = (axis < 0) ? ((rank > 0) ? data->shape.size() + axis : 0) : axis; - ICHECK_LT(axis, rank > 0 ? rank : 1) << "axis " << requantize_attrs->axis << " is out of range"; - ICHECK_GE(axis, 0) << "axis " << requantize_attrs->axis << " is out of range"; - - PrimExpr axis_shape; - if (rank > 0) { - axis_shape = data->shape[axis]; - } else { - axis_shape = Integer(1); - } - // Check and assign types for scale and zero points. - AssignType(types[1], DataType::Float(32), axis_shape, reporter); // input_scale - AssignType(types[2], DataType::Int(32), axis_shape, reporter); // input_zero_pt - // For now, requantize output tensor is limited to full tensor uniform quantization. - ICHECK(IsScalarType(types[3], DataType::Float(32))); // output_scale - ICHECK(IsScalarType(types[4], DataType::Int(32))); // output_zero_point - - const Array oshape = data->shape; - // assign output type - auto out_dtype = requantize_attrs->out_dtype; - ICHECK(out_dtype == DataType::Int(8) || out_dtype == DataType::UInt(8) || - out_dtype == DataType::Int(16) || out_dtype == DataType::Int(32)) - << "Output type should be one of [int8, uint8, int16, int32] but was " << out_dtype; - reporter->Assign(types[5], TensorType(oshape, out_dtype)); - return true; -} - -// Positional relay function to create qnn requantize operator -// used by frontend FFI. -Expr MakeRequantize(Expr data, Expr input_scale, Expr input_zero_point, Expr output_scale, - Expr output_zero_point, int axis, String rounding, String compute_dtype, - DataType out_dtype) { - auto attrs = make_object(); - attrs->axis = axis; - attrs->rounding = std::move(rounding); - attrs->out_dtype = std::move(out_dtype); - attrs->compute_dtype = std::move(compute_dtype); - static const Op& op = Op::Get("qnn.requantize"); - return Call(op, {data, input_scale, input_zero_point, output_scale, output_zero_point}, - Attrs(attrs), {}); -} - -RELAY_REGISTER_OP("qnn.requantize") - .describe(R"code(Requantize operator. -The requantize operator converts one quantized tensor to another quantized -tensor. For the output tensor, we are provided with output scale and zero -point. The computation looks like this - -Q_output = zp_output + (scale_input)/(scale_output) * (Q_input - zp_input) - -)code" TVM_ADD_FILELINE) - .set_attrs_type() - .set_num_inputs(5) - .add_argument("data", "Tensor", "The quantized input tensor.") - .add_argument("input_scale", "Tensor", "The quantization scale of the input tensor.") - .add_argument("input_zero_point", "Tensor", "The quantization zero_point of the input tensor.") - .add_argument("output_scale", "Tensor", "The quantization scale of the output tensor.") - .add_argument("output_zero_point", "Tensor", - "The quantization zero_point of the output tensor.") - .set_support_level(11) - .add_type_rel("Requantize", RequantizeRel) - .set_attr("TNonComputational", true) - .set_attr("FTVMQnnCanonicalize", RequantizeQnnCanonicalize) - .set_attr("FInferCorrectLayout", RequantizeInferCorrectLayout); - -TVM_REGISTER_GLOBAL("relay.qnn.op._make.requantize").set_body_typed(MakeRequantize); - -} // namespace qnn -} // namespace relay -} // namespace tvm diff --git a/src/relay/qnn/op/requantize_config.cc b/src/relay/qnn/op/requantize_config.cc deleted file mode 100644 index 4a52f56400c9..000000000000 --- a/src/relay/qnn/op/requantize_config.cc +++ /dev/null @@ -1,93 +0,0 @@ -/* - * 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 src/relay/qnn/op/requantize_config.cc - * \brief QNN requantize config. - */ - -#include "./requantize_config.h" - -#include -#include -#include - -#include - -namespace tvm { -namespace relay { -namespace qnn { - -/*! \brief Entry to hold the BuildConfig context stack. */ -struct TVMRequantizeConfigThreadLocalEntry { - /*! \brief The default build config if the stack is empty */ - RequantizeConfig default_config; - - /*! \brief The current build config context */ - std::stack context_stack; - - TVMRequantizeConfigThreadLocalEntry() : default_config(make_object(true)) {} -}; - -/*! \brief Thread local store to hold the BuildConfig context stack. */ -typedef dmlc::ThreadLocalStore - TVMRequantizeConfigThreadLocalStore; - -void RequantizeConfig::EnterRequantizeConfigScope(const RequantizeConfig& build_config) { - TVMRequantizeConfigThreadLocalEntry* entry = TVMRequantizeConfigThreadLocalStore::Get(); - entry->context_stack.push(build_config); -} - -void RequantizeConfig::ExitRequantizeConfigScope() { - TVMRequantizeConfigThreadLocalEntry* entry = TVMRequantizeConfigThreadLocalStore::Get(); - entry->context_stack.pop(); -} - -RequantizeConfig& RequantizeConfig::Current() { - TVMRequantizeConfigThreadLocalEntry* entry = TVMRequantizeConfigThreadLocalStore::Get(); - if (entry->context_stack.size() > 0) { - return entry->context_stack.top(); - } - - return entry->default_config; -} - -TVM_REGISTER_NODE_TYPE(RequantizeConfigNode); - -TVM_STATIC_IR_FUNCTOR(ReprPrinter, vtable) - .set_dispatch([](const ObjectRef& ref, ReprPrinter* p) { - auto* op = static_cast(ref.get()); - p->stream << "requantize_config("; - p->stream << "rounding==" << op->get_rounding() << ", "; - p->stream << "compute_dtype==" << op->get_compute_dtype(); - p->stream << ")"; - }); - -TVM_REGISTER_GLOBAL("relay._requantize._GetCurrentRequantizeConfig") - .set_body_typed([]() -> RequantizeConfig { return RequantizeConfig::Current(); }); - -TVM_REGISTER_GLOBAL("relay._requantize._EnterRequantizeConfigScope") - .set_body_typed(RequantizeConfig::EnterRequantizeConfigScope); - -TVM_REGISTER_GLOBAL("relay._requantize._ExitRequantizeConfigScope") - .set_body_typed(RequantizeConfig::ExitRequantizeConfigScope); - -} // namespace qnn -} // namespace relay -} // namespace tvm diff --git a/src/relay/qnn/op/requantize_config.h b/src/relay/qnn/op/requantize_config.h deleted file mode 100644 index a4238fa498c6..000000000000 --- a/src/relay/qnn/op/requantize_config.h +++ /dev/null @@ -1,125 +0,0 @@ -/* - * 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 src/relay/qnn/op/requantize_config.h - * \brief QNN requantize config. - */ - -#ifndef TVM_RELAY_QNN_OP_REQUANTIZE_CONFIG_H_ -#define TVM_RELAY_QNN_OP_REQUANTIZE_CONFIG_H_ - -#include -#include -#include -#include -#include - -#include - -#include "../../op/op_common.h" - -namespace tvm { -namespace relay { -namespace qnn { - -class RequantizeConfig; -/*! - * \brief Container for build configuration options - */ -class RequantizeConfigNode : public Object { - std::string rounding; - std::string compute_dtype; - - public: - explicit RequantizeConfigNode(bool is_default = false) : is_default(is_default) {} - - std::string get_rounding() const { - if (!rounding.empty()) return rounding; - return "UPWARD"; - } - - std::string get_compute_dtype() const { - if (!compute_dtype.empty()) return compute_dtype; - - // For the x86 architecture, the float32 computation is expected to give significant speedup, - // with little loss in the accuracy of the requantize operation. - auto target = Target::Current(true); - auto target_has_feature_fn_ptr = tvm::runtime::Registry::Get("target.target_has_feature"); - ICHECK(target_has_feature_fn_ptr) << "Function target.target_has_feature not found"; - if (target.defined() && target->kind->name == "llvm") { - if ((*target_has_feature_fn_ptr)("sse4.1", target)) { - return "float32"; - } - } - return "int64"; - } - - const bool is_default = false; - - void VisitAttrs(AttrVisitor* v) { - v->Visit("rounding", &rounding); - v->Visit("compute_dtype", &compute_dtype); - } - - static constexpr const char* _type_key = "relay.qnn.op.RequantizeConfig"; - TVM_DECLARE_FINAL_OBJECT_INFO(RequantizeConfigNode, Object); -}; - -/*! - * \brief Container for build configuration options - */ -class RequantizeConfig : public ObjectRef { - public: - RequantizeConfig() {} - explicit RequantizeConfig(ObjectPtr n) : ObjectRef(n) {} - - const RequantizeConfigNode* operator->() const { - return static_cast(get()); - } - - RequantizeConfigNode* operator->() { return static_cast(get_mutable()); } - - /*! - * \brief Push a new BuildConfig context onto the thread local stack. - * \param build_config The configuration to set as the current context. - */ - static void EnterRequantizeConfigScope(const RequantizeConfig& requantize_config); - - /*! - * \brief Pop a build config off the thread local context stack, restoring the previous - * configuration as the current context. - */ - static void ExitRequantizeConfigScope(); - - /*! - * \brief Get the current BuildConfig context from thread local storage, or a default - * configuration if a BuildConfig scope has not been entered. - * \return The configuration that is the current context. - */ - static RequantizeConfig& Current(); - - using ContainerType = RequantizeConfigNode; -}; - -} // namespace qnn -} // namespace relay -} // namespace tvm - -#endif // TVM_RELAY_QNN_OP_REQUANTIZE_CONFIG_H_ diff --git a/src/relay/qnn/op/simulated_dequantize.cc b/src/relay/qnn/op/simulated_dequantize.cc deleted file mode 100644 index e1fc47d700c9..000000000000 --- a/src/relay/qnn/op/simulated_dequantize.cc +++ /dev/null @@ -1,80 +0,0 @@ -/* - * 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 src/relay/qnn/op/simulated_dequantize.cc - * \brief QNN simulated dequantize operator. Mimics the behavior - * of QNN dequantize in floating point with added flexibility. - */ - -#include -#include -#include - -#include "../../transforms/pattern_utils.h" -#include "../utils.h" - -namespace tvm { -namespace relay { -namespace qnn { - -bool SimulatedDequantizeRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - // types = [data_type, datatype_type, scale_type, zp_type, ret_type] - ICHECK_EQ(types.size(), 5); - const auto* data = types[0].as(); - const auto* dtype = types[1].as(); - - if ((data == nullptr) || (dtype == nullptr)) { - return false; - } - - // assign output type - reporter->Assign(types[4], TensorType(data->shape, data->dtype)); - return true; -} - -Expr MakeSimulatedDequantize(Expr data, Expr in_dtype, Expr input_scale, Expr input_zero_point, - int axis) { - auto attrs = make_object(); - attrs->axis = axis; - static const Op& op = Op::Get("qnn.simulated_dequantize"); - return Call(op, {data, in_dtype, input_scale, input_zero_point}, Attrs(attrs), {}); -} - -RELAY_REGISTER_OP("qnn.simulated_dequantize") - .describe(R"code(Simulates the functionality of qnn.dequantize but allows more flexible - dynamic input type conversion and always operates on float values. -)code" TVM_ADD_FILELINE) - .set_attrs_type() - .set_num_inputs(4) - .add_argument("data", "Tensor", "The tensor to dequantize.") - .add_argument("in_dtype", "Tensor", - "A code corresponding to the type of quantization to convert from.") - .add_argument("input_scale", "Tensor", "The quantization scale of the input tensor.") - .add_argument("input_zero_point", "Tensor", "The quantization zero_point of the input tensor.") - .set_support_level(11) - .add_type_rel("QNNSimulatedDequantize", SimulatedDequantizeRel); - -TVM_REGISTER_GLOBAL("relay.qnn.op._make.simulated_dequantize") - .set_body_typed(MakeSimulatedDequantize); - -} // namespace qnn -} // namespace relay -} // namespace tvm diff --git a/src/relay/qnn/op/simulated_quantize.cc b/src/relay/qnn/op/simulated_quantize.cc deleted file mode 100644 index 089762a6ade0..000000000000 --- a/src/relay/qnn/op/simulated_quantize.cc +++ /dev/null @@ -1,82 +0,0 @@ -/* - * 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 src/relay/qnn/op/simulated_quantize.cc - * \brief QNN simulated quantize operator. Mimics the behavior - * of QNN quantize in floating point with added flexibility. - */ - -#include -#include -#include - -#include "../../transforms/pattern_utils.h" -#include "../utils.h" - -namespace tvm { -namespace relay { -namespace qnn { - -TVM_REGISTER_NODE_TYPE(SimulatedQuantizeAttrs); - -bool SimulatedQuantizeRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - // types = [data_type, datatype_type, scale_type, zp_type, ret_type] - ICHECK_EQ(types.size(), 5); - const auto* data = types[0].as(); - const auto* dtype = types[1].as(); - - if ((data == nullptr) || (dtype == nullptr)) { - return false; - } - - // assign output type - reporter->Assign(types[4], TensorType(data->shape, data->dtype)); - return true; -} - -Expr MakeSimulatedQuantize(Expr data, Expr out_dtype, Expr output_scale, Expr output_zero_point, - int axis) { - auto attrs = make_object(); - attrs->axis = axis; - static const Op& op = Op::Get("qnn.simulated_quantize"); - return Call(op, {data, out_dtype, output_scale, output_zero_point}, Attrs(attrs), {}); -} - -RELAY_REGISTER_OP("qnn.simulated_quantize") - .describe(R"code(Simulates the functionality of qnn.quantize but allows more flexible - dynamic input type conversion and always outputs float values. -)code" TVM_ADD_FILELINE) - .set_attrs_type() - .set_num_inputs(4) - .add_argument("data", "Tensor", "The tensor to quantize.") - .add_argument("out_dtype", "Tensor", - "A code corresponding to the type of quantization to apply.") - .add_argument("output_scale", "Tensor", "The quantization scale of the output tensor.") - .add_argument("output_zero_point", "Tensor", - "The quantization zero_point of the output tensor.") - .set_support_level(11) - .add_type_rel("QNNSimulatedQuantize", SimulatedQuantizeRel); - -TVM_REGISTER_GLOBAL("relay.qnn.op._make.simulated_quantize").set_body_typed(MakeSimulatedQuantize); - -} // namespace qnn -} // namespace relay -} // namespace tvm diff --git a/src/relay/qnn/op/softmax.cc b/src/relay/qnn/op/softmax.cc deleted file mode 100644 index f848ba9384e3..000000000000 --- a/src/relay/qnn/op/softmax.cc +++ /dev/null @@ -1,154 +0,0 @@ -/* - * 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 src/relay/qnn/op/softmax.cc - * \brief QNN softmax operator. - */ -#include -#include - -#include "op_common.h" -#include "tvm/ir/expr.h" -#include "tvm/relay/attrs/nn.h" -#include "tvm/relay/type.h" -#include "tvm/runtime/data_type.h" -#include "tvm/runtime/logging.h" -#include "tvm/topi/reduction.h" - -namespace tvm { -namespace relay { -namespace qnn { - -bool QnnSoftmaxRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - // Expected Types: input, scale, zero_point, output_scale, output_zero_point, output - ICHECK_EQ(types.size(), 6); - const auto* x = types[0].as(); - if (x == nullptr) return false; - ICHECK(x->dtype == DataType::Int(8)) - << "Expected quantized softmax type(int8) for input but was " << x->dtype; - - // Check the types of scale and zero points. - for (size_t i = 1; i < 5; ++i) { - if (types[i].as()) { - return false; - } - } - - ICHECK(IsScalarType(types[1], DataType::Float(32))); // scale - ICHECK(IsScalarType(types[2], DataType::Int(32))); // zero_point - ICHECK(IsScalarType(types[3], DataType::Float(32))); // scale - ICHECK(IsScalarType(types[4], DataType::Int(32))); // zero_point - - // Assign types for scale and zero points. - reporter->Assign(types[1], TensorType({}, DataType::Float(32))); // scale - reporter->Assign(types[2], TensorType({}, DataType::Int(32))); // zero_point - reporter->Assign(types[3], TensorType({}, DataType::Float(32))); // scale - reporter->Assign(types[4], TensorType({}, DataType::Int(32))); // zero_point - - // Collect the input tensor and output tensor devoid of scale and zero points to reuse Relay - // IdentityRel infer type function. - Array tensor_types = {types[0], types[5]}; - return IdentityRel(tensor_types, 2, attrs, reporter); -} - -// Positional relay function to create quantized softmax operator used by frontend FFI. -Expr MakeQuantizedSoftmax(Expr x, int axis, Expr scale, Expr zero_point, Expr output_scale, - Expr output_zero_point) { - auto attrs = make_object(); - attrs->axis = axis; - static const Op& op = Op::Get("qnn.softmax"); - return Call(op, {x, scale, zero_point, output_scale, output_zero_point}, Attrs(attrs), {}); -} - -/* - * \brief Canonicalizes the QNN softmax op. - * \param attrs The Softmax attrs. - * \param new_args The new mutated args to the call node. - * \param arg_types The types of input and output. - * \return The sequence of Relay ops for softmax op. - * \note This op is highly experimental and sometimes lacks accuracy. - * Be aware that the input scale must be in the range of 0 to 1. - */ -Expr QnnSoftmaxCanonicalize(const Attrs& attrs, const Array& new_args, - const Array& arg_types) { - // Expected: input, scale, zero_point, output_scale, output_zero_point - ICHECK_EQ(new_args.size(), 5); - - const auto const_i32 = [&](int32_t val) { return MakeConstantScalar(DataType::Int(32), val); }; - const auto const_f32 = [&](float val) { return MakeConstantScalar(DataType::Float(32), val); }; - - const auto const_input_scale = new_args[1].as(); - ICHECK(const_input_scale) << "Input scale should be constant."; - ICHECK(const_input_scale->is_scalar()) << "Input scale should be scalar."; - const float input_scale = static_cast(const_input_scale->data->data)[0]; - ICHECK(input_scale <= 1.f) << "Input scale should be less than or equal to 1."; - - const Expr input_zero_point = new_args[2]; - const Expr output_scale = new_args[3]; - const Expr output_zero_point = new_args[4]; - const int axis = attrs.as()->axis; - - // Refer to the Algorithm 1 in https://arxiv.org/pdf/2207.01405.pdf - - const Expr quantized_data = Subtract(Cast(new_args[0], DataType::Int(32)), input_zero_point); - - const Expr x_0 = ConvertDtype(const_f32(std::round(1.f / input_scale)), DataType::Int(32)); - const Expr max = Max(quantized_data, {axis}, true, false); - const Expr x = Subtract(quantized_data, max); - - const int m = 30; - const int bits = 8; - const Expr x_p = Subtract(Add(x, RightShift(x, const_i32(1))), RightShift(x, const_i32(4))); - const Expr q = Clip(Divide(x_p, Negative(x_0)), 0, 20); - const Expr max_q = Max(q, {axis}, true, false); - const Expr r = Subtract(x_p, Multiply(q, Negative(x_0))); - const Expr x_b = Add(RightShift(r, const_i32(1)), x_0); - const Expr exps = LeftShift(x_b, Subtract(max_q, q)); - const Expr sums = Sum(exps, {axis}, true, false); - const Expr output = - RightShift(Multiply(Divide(const_i32(1 << m), sums), exps), const_i32(m - (bits - 1))); - const Expr requantized = Requantize(output, arg_types[0].as()->shape, - const_f32(1.f / (1 << (bits - 1))), const_i32(0), - output_scale, output_zero_point, DataType::Int(bits), 0); - - return requantized; -} - -RELAY_REGISTER_OP("qnn.softmax") - .describe("Softmax for quantized tensors.") - .set_attrs_type() - .set_num_inputs(5) - .add_argument("data", "Quantized Tensor", "The input data.") - .add_argument("scale", "Tensor", "The quantization scale of the input tensor.") - .add_argument("zero_point", "Tensor", "The quantization zero_point of the input tensor.") - .add_argument("output_scale", "Tensor", "The quantization scale of the output tensor.") - .add_argument("output_zero_point", "Tensor", - "The quantization zero_point of the output tensor.") - .set_support_level(11) - .add_type_rel("QSoftmax", QnnSoftmaxRel) - .set_attr("TNonComputational", true) - .set_attr("FTVMQnnCanonicalize", QnnSoftmaxCanonicalize); - -TVM_REGISTER_GLOBAL("relay.qnn.op._make.softmax").set_body_typed(MakeQuantizedSoftmax); - -} // namespace qnn -} // namespace relay -} // namespace tvm diff --git a/src/relay/qnn/op/subtract.cc b/src/relay/qnn/op/subtract.cc deleted file mode 100644 index 962a3434cb72..000000000000 --- a/src/relay/qnn/op/subtract.cc +++ /dev/null @@ -1,105 +0,0 @@ -/* - * 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 src/relay/qnn/op/subtract.cc - * \brief QNN subtract operator. - */ -#include -#include - -#include "op_common.h" - -namespace tvm { -namespace relay { -namespace qnn { - -/* - * \brief Canonicalizes the QNN subtract op. - * \param attrs The empty attribute. - * \param new_args The new mutated args to the call node. - * \param arg_types The types of input and output. - * \return The sequence of Relay ops for add op. - */ -Expr QnnSubtractCanonicalize(const Attrs& attrs, const Array& new_args, - const Array& arg_types) { - // Get the args. - QnnBinaryOpArguments args(new_args); - - // Get the input dtype and shape. - QnnBinaryOpTensorType input_type(arg_types, 0); - - const auto* broadcast_attrs = attrs.as(); - ICHECK(broadcast_attrs != nullptr); - - auto lhs_axis = broadcast_attrs->lhs_axis; - auto rhs_axis = broadcast_attrs->rhs_axis; - - // TODO(shoubhik) - The lowering can be further optimized. Instead of inserting requantize in - // the start, we can insert requantize at the end if both input tensors have same qnn params. In - // that case, we can first subtract the tensors, add the zero point, and requantize at the end. - // This can be done in future. - - // Since the input qnn params can be different than output qnn params, we first requantize the - // input tensors to the output qnn params. Then we call relay.subtract on the requantized inputs. - // This subtraction results in extra subtraction of the output zero point. We further add - // the zero point. The whole process can be represented using following equations - // - // scale_c * (Q_c - zp_c) = scale_a * (Q_a - zp_a) - scale_b * (Q_b - zp_b) - // - // After requantizing Q_a and Q_b, equation becomes, - // scale_c * (Q_c - zp_c) = scale_c * (Q_a' - zp_c) - scale_c * (Q_b' - zp_c) - // scale_c * (Q_c - zp_c) = scale_c * (Q_a' - Q_b') - // - // Comparing the LHS and RHS, it results in - // Q_c = Q_a' - Q_b' + zp_c - // The subtract op is done in int32 precision. - - // Requantize LHS if necessary. Computes Q_a' - auto requantized_lhs = - RequantizeOrUpcast(args.lhs, args.lhs_scale, args.lhs_zero_point, args.output_scale, - args.output_zero_point, input_type.shape, lhs_axis); - // Requantize RHS if necessary. Computes Q_b' - auto requantized_rhs = - RequantizeOrUpcast(args.rhs, args.rhs_scale, args.rhs_zero_point, args.output_scale, - args.output_zero_point, input_type.shape, rhs_axis); - - // Computes Q_a' - Q_b' - auto output = Subtract(requantized_lhs, requantized_rhs); - - // Add zero point. Computes (Q_a' - Q_b') + zp_c - auto zero_scalar = MakeConstantScalar(DataType::Int(32), 0); - if (!IsEqualScalar(args.output_zero_point, zero_scalar)) { - output = Add(output, args.output_zero_point); - } - - // Go back to lower precision. - return ConvertDtype(output, input_type.dtype); -} - -// QNN Subtraction operator. -QNN_REGISTER_BINARY_OP("subtract") - .describe("Elementwise subtract with broadcasting for quantized tensors.") - .set_support_level(11) - .set_attr("FTVMQnnCanonicalize", QnnSubtractCanonicalize) - .set_attr("TOpPattern", kBroadcast); - -} // namespace qnn -} // namespace relay -} // namespace tvm diff --git a/src/relay/qnn/op/unary_elementwise_op.cc b/src/relay/qnn/op/unary_elementwise_op.cc deleted file mode 100644 index cdd3ea63b0e7..000000000000 --- a/src/relay/qnn/op/unary_elementwise_op.cc +++ /dev/null @@ -1,61 +0,0 @@ -/* - * 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 src/relay/qnn/op/unary_elementwise_op.cc - * \brief QNN unary elementwise operators. - */ - -#include "op_common.h" - -namespace tvm { -namespace relay { -namespace qnn { - -QNN_CREATE_UNARY_ELEMENTWISE_OP("tanh").set_attr( - "FTVMQnnCanonicalize", QNN_UNARY_OP_DEFAULT_CANONICALIZATION(Tanh)); - -QNN_CREATE_UNARY_ELEMENTWISE_OP("exp").set_attr( - "FTVMQnnCanonicalize", QNN_UNARY_OP_DEFAULT_CANONICALIZATION(Exp)); - -QNN_CREATE_UNARY_ELEMENTWISE_OP("sqrt").set_attr( - "FTVMQnnCanonicalize", QNN_UNARY_OP_DEFAULT_CANONICALIZATION(Sqrt)); - -QNN_CREATE_UNARY_ELEMENTWISE_OP("rsqrt").set_attr( - "FTVMQnnCanonicalize", QNN_UNARY_OP_DEFAULT_CANONICALIZATION(Rsqrt)); - -QNN_CREATE_UNARY_ELEMENTWISE_OP("erf").set_attr( - "FTVMQnnCanonicalize", QNN_UNARY_OP_DEFAULT_CANONICALIZATION(Erf)); - -QNN_CREATE_UNARY_ELEMENTWISE_OP("sigmoid").set_attr( - "FTVMQnnCanonicalize", QNN_UNARY_OP_DEFAULT_CANONICALIZATION(Sigmoid)); - -QNN_CREATE_UNARY_ELEMENTWISE_OP("hardswish") - .set_attr("FTVMQnnCanonicalize", - QNN_UNARY_OP_DEFAULT_CANONICALIZATION(Hardswish)); - -QNN_CREATE_UNARY_ELEMENTWISE_OP("log").set_attr( - "FTVMQnnCanonicalize", QNN_UNARY_OP_DEFAULT_CANONICALIZATION(Log)); - -QNN_CREATE_UNARY_ELEMENTWISE_OP("abs").set_attr( - "FTVMQnnCanonicalize", QNN_UNARY_OP_DEFAULT_CANONICALIZATION(Abs)); - -} // namespace qnn -} // namespace relay -} // namespace tvm diff --git a/src/relay/qnn/pass/legalize.cc b/src/relay/qnn/pass/legalize.cc deleted file mode 100644 index fd88c4df8c06..000000000000 --- a/src/relay/qnn/pass/legalize.cc +++ /dev/null @@ -1,65 +0,0 @@ -/* - * 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 relay/qnn/pass/legalize.cc - * \brief The Legalize wrapper for QNN. - */ - -#include - -namespace tvm { -namespace relay { -namespace qnn { - -namespace transform { - -// QnnLegalize pass is a wrapper for relay::legalize::Legalize pass. -Pass QnnLegalize() { - runtime::TypedPackedFunc pass_func = - [=](Function f, IRModule m, relay::transform::PassContext pc) { - return Downcast(relay::legalize::Legalize(f, "FTVMQnnLegalize")); - }; - return relay::transform::CreateFunctionPass(pass_func, 1, "QnnLegalize", {"InferType"}); -} - -// QnnCanonicalize pass is a wrapper for relay::legalize::Legalize pass. -Pass QnnCanonicalize() { - runtime::TypedPackedFunc pass_func = - [=](Function f, IRModule m, relay::transform::PassContext pc) { - return Downcast(relay::legalize::Legalize(f, "FTVMQnnCanonicalize")); - }; - return relay::transform::CreateFunctionPass(pass_func, 1, "QnnCanonicalize", {"InferType"}); -} - -Pass Legalize() { - Array pass_seqs; - pass_seqs.push_back(QnnLegalize()); - pass_seqs.push_back(QnnCanonicalize()); - relay::transform::Pass seq = relay::transform::Sequential(pass_seqs, "qnn.Legalize"); - return seq; -} - -TVM_REGISTER_GLOBAL("relay.qnn._transform.Legalize").set_body_typed(Legalize); - -} // namespace transform - -} // namespace qnn -} // namespace relay -} // namespace tvm diff --git a/src/relay/qnn/utils.cc b/src/relay/qnn/utils.cc deleted file mode 100644 index ab72bd957080..000000000000 --- a/src/relay/qnn/utils.cc +++ /dev/null @@ -1,249 +0,0 @@ -/* - * 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 src/relay/qnn/utils.cc - * \brief Utility functions for QNN. - */ - -#include "utils.h" - -#include "../transforms/pattern_utils.h" - -namespace tvm { -namespace relay { -namespace qnn { - -std::pair GetFixedPointMultiplierShift(double double_multiplier) { - int32_t significand, exponent; - if (double_multiplier == 0.) { - significand = 0; - exponent = 0; - return std::make_pair(significand, exponent); - } - - // Get the significand and exponent. - double significand_d = std::frexp(double_multiplier, &exponent); - - // Convert the double significand to int significand, i.e., convert into a - // integer where the decimal point is between bit 31 and 30. This is done by - // multiplying the double value with 2^31 and then casting to int. - significand_d = std::round(significand_d * (1ll << 31)); - auto significand_int64 = static_cast(significand_d); - ICHECK_LE(significand_int64, (1ll << 31)); - if (significand_int64 == (1ll << 31)) { - significand_int64 /= 2; - ++exponent; - } - ICHECK_LE(significand_int64, std::numeric_limits::max()); - significand = static_cast(significand_int64); - return std::make_pair(significand, exponent); -} - -Expr FixedPointMultiplyToNearest(Expr tensor, double multiplier, - const Array& input_shape) { - // Choose high precision datatype to be int64. This is for avoiding overflow - // in multiplication of two int32 values. - DataType hp_dtype = DataType::Int(64); - tensor = Cast(tensor, hp_dtype); - - // 1) Calculating the integer multiplier and integer shift - auto [fixed_point_multiplier, shift] = GetFixedPointMultiplierShift(multiplier); - int left_shift = shift > 0 ? shift : 0; - int right_shift = shift > 0 ? 0 : -shift; - - // 2) Multiply the integer multiplier - if (left_shift != 0) { - tensor = LeftShift(tensor, MakeConstantScalar(hp_dtype, left_shift)); - } - - // 3) Perform the multiplication in higher precision. - // The scalar is a fixed point value of int32 where the decimal point is - // between bits 31 and 30. After multiplying with input_tensor, the result - // is in int64 where the decimal point is sitting between bits 31 and 30 - // (from the right, rightmost bit is bit 0). The computation is performed in - // higher precision to avoid overflow in multiplying two int32 values. - Expr scalar = MakeConstantScalar(hp_dtype, fixed_point_multiplier); - tensor = Multiply(tensor, scalar); - - // 4) Find the rounding scalar. This depends on where the final decimal - // point sits. As we will be right shifting the multiplied_t, we need to - // first calculate the total_right_shift. - int total_right_shift = right_shift + 31; - int64_t pos_rounding_value = (1ll << (total_right_shift - 1)); - - Expr round_scalar; - - auto pos_rounder = MakeConstantScalar(hp_dtype, pos_rounding_value); - auto neg_rounder = MakeConstantScalar(hp_dtype, pos_rounding_value - 1); - auto pos_rounder_t = Full(pos_rounder, input_shape, hp_dtype); - auto neg_rounder_t = Full(neg_rounder, input_shape, hp_dtype); - - auto zero_t = Zeros(input_shape, hp_dtype); - round_scalar = Where(GreaterEqual(tensor, zero_t), pos_rounder_t, neg_rounder_t); - - // Add the rounding scalar. - tensor = Add(tensor, round_scalar); - - // 5) Simply right shift the result to get the final output. - tensor = RightShift(tensor, MakeConstantScalar(hp_dtype, total_right_shift)); - - // 6) The fixed point multiplication keeps the value in int32 range. Casting back to int32. - return Cast(tensor, DataType::Int(32)); -} - -Expr FixedPointMultiplyPerChannel(Expr tensor, const std::vector& multipliers, int axis) { - DataType dtype = DataType::Int(32); - int64_t n_channels = static_cast(multipliers.size()); - - std::vector fixed_pt_multipliers, lshifts, rshifts; - bool is_lshift_required = false, is_rshift_required = false; - for (auto multiplier : multipliers) { - auto [fixed_pt_multiplier, shift] = GetFixedPointMultiplierShift(multiplier); - int lshift = shift > 0 ? shift : 0; - int rshift = shift > 0 ? 0 : -shift; - fixed_pt_multipliers.push_back(fixed_pt_multiplier); - lshifts.push_back(lshift); - rshifts.push_back(rshift); - is_lshift_required = is_lshift_required | (lshift != 0); - is_rshift_required = is_rshift_required | (rshift != 0); - } - - auto left_shift_expr = MakeConstantTensor(dtype, {n_channels}, lshifts); - auto right_shift_expr = MakeConstantTensor(dtype, {n_channels}, rshifts); - auto fixed_pt_multiplier_expr = MakeConstantTensor(dtype, {n_channels}, fixed_pt_multipliers); - - return FixedPointMultiplyPerAxis(tensor, fixed_pt_multiplier_expr, left_shift_expr, - right_shift_expr, is_lshift_required, is_rshift_required, - {axis}); -} - -Expr FixedPointMultiplyPerChannel(Expr tensor, std::vector multipliers, - const Array& input_shape, int channel_axis, - const std::string& rounding) { - // Get the n dim. This will be used to expand the multiplier to match the axis. - size_t n_dim = input_shape.size(); - - // Get the num of channels/axis along which the tensor was quantized. - int64_t n_channels = (int64_t)multipliers.size(); - - // Choose high precision datatype to be int64. This is for avoiding overflow - // in multiplication of two int32 values. - DataType hp_dtype = DataType::Int(64); - tensor = Cast(tensor, hp_dtype); - - // 1) Calculating the integer multiplier and integer shift. These are calculated per axis/per - // channel. - std::vector fixed_pt_multipliers, lshifts, rshifts; - bool is_lshift_required = false; - for (auto multiplier : multipliers) { - auto [fixed_pt_multiplier, shift] = GetFixedPointMultiplierShift(multiplier); - int lshift = shift > 0 ? shift : 0; - int rshift = shift > 0 ? 0 : -shift; - fixed_pt_multipliers.push_back(fixed_pt_multiplier); - lshifts.push_back(lshift); - rshifts.push_back(rshift); - is_lshift_required = is_lshift_required | (lshift != 0); - } - - // 2) Multiply the integer multiplier. Convert lefts shifts into expr and multiply. - if (is_lshift_required) { - auto lshift_expr = MakeConstantTensor(hp_dtype, {n_channels}, lshifts); - auto exp_lshift_expr = ExpandBiasToMatchAxis(lshift_expr, n_dim, {channel_axis}); - tensor = LeftShift(tensor, exp_lshift_expr); - } - - // 3) Perform the multiplication in higher precision. - // The scalar is a fixed point value of int32 where the decimal point is - // between bits 31 and 30. After multiplying with input_tensor, the result - // is in int64 where the decimal point is sitting between bits 31 and 30 - // (from the right, rightmost bit is bit 0). The computation is performed in - // higher precision to avoid overflow in multiplying two int32 values. - auto fixed_pt_multiplier_expr = MakeConstantTensor(hp_dtype, {n_channels}, fixed_pt_multipliers); - auto exp_fixed_pt_multiplier_expr = - ExpandBiasToMatchAxis(fixed_pt_multiplier_expr, n_dim, {channel_axis}); - tensor = Multiply(tensor, exp_fixed_pt_multiplier_expr); - - // 4) Find the rounding scalar. This depends on where the final decimal point sits. As we will be - // right shifting the multiplied_t, we need to first calculate the total_rshift. Further, we can - // calculate the pos and neg rounding offset. - std::vector pos_rounding_values, neg_rounding_values, total_rshifts; - for (auto rshift : rshifts) { - int total_rshift = rshift + 31; - total_rshifts.push_back(total_rshift); - pos_rounding_values.push_back((1ll << (total_rshift - 1))); - neg_rounding_values.push_back((1ll << (total_rshift - 1)) - 1); - } - // Make a Relay expr from positive and negative rounding offset values. - auto pos_rounding_value_expr = MakeConstantTensor(hp_dtype, {n_channels}, pos_rounding_values); - auto exp_pos_rounding_value_expr = - ExpandBiasToMatchAxis(pos_rounding_value_expr, n_dim, {channel_axis}); - auto neg_rounding_value_expr = MakeConstantTensor(hp_dtype, {n_channels}, neg_rounding_values); - auto exp_neg_rounding_value_expr = - ExpandBiasToMatchAxis(neg_rounding_value_expr, n_dim, {channel_axis}); - - Expr round_scalar; - if (rounding == "UPWARD") { - round_scalar = exp_pos_rounding_value_expr; - } else if (rounding == "TONEAREST") { - // To satisfy where op shape requirements, the rounding values are broadcasted. - auto pos_rounder = BroadCastTo(exp_pos_rounding_value_expr, input_shape); - auto neg_rounder = BroadCastTo(exp_neg_rounding_value_expr, input_shape); - - auto zero_t = Zeros(input_shape, hp_dtype); - round_scalar = Where(GreaterEqual(tensor, zero_t), pos_rounder, neg_rounder); - } else { - LOG(FATAL) << "Rounding mode " << rounding << " not supported."; - } - // Add the rounding scalar. - tensor = Add(tensor, round_scalar); - - // 5) Simply right shift the result to get the final output. - auto total_rshift_expr = MakeConstantTensor(hp_dtype, {n_channels}, total_rshifts); - auto exp_total_rshift_expr = ExpandBiasToMatchAxis(total_rshift_expr, n_dim, {channel_axis}); - tensor = RightShift(tensor, exp_total_rshift_expr); - - // 6) The fixed point multiplication keeps the value in int32 range. Casting back to int32. - return Cast(tensor, DataType::Int(32)); -} - -Expr FixedPointMultiplyPerChannelToNearest(Expr tensor, std::vector multipliers, - const Array& input_shape, int channel_axis) { - return FixedPointMultiplyPerChannel(tensor, multipliers, input_shape, channel_axis, "TONEAREST"); -} - -std::string SelectRequntizeParameter(const std::string& arg_value, const std::string& cfg_value, - const bool is_cfg_default, const std::string& name) { - if (arg_value == "None") { - return cfg_value; - } else { - if (!is_cfg_default && arg_value != cfg_value) { - DLOG(INFO) << "The value of parameter \"" << name - << "\" from the non-default requantize config will not be used. The value " - "provided from " - "requantize function argument will be used instead. The value used is \"" - << arg_value << "\"."; - } - return arg_value; - } -} - -} // namespace qnn -} // namespace relay -} // namespace tvm diff --git a/src/relay/qnn/utils.h b/src/relay/qnn/utils.h deleted file mode 100644 index 4102fb29a6fe..000000000000 --- a/src/relay/qnn/utils.h +++ /dev/null @@ -1,302 +0,0 @@ -/* - * 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 src/relay/qnn/utils.h - * \brief Utility methods needs for quantized ops that can be shared - */ - -#ifndef TVM_RELAY_QNN_UTILS_H_ -#define TVM_RELAY_QNN_UTILS_H_ - -#include -#include -#include -#include - -#include -#include -#include -#include - -#include "./op/requantize_config.h" - -namespace tvm { -namespace relay { -namespace qnn { - -static inline Array get_shape(const Type& type) { - auto input_tt = type.as(); - ICHECK(input_tt != nullptr) << "Type information missing." - << " Please run infer_type pass."; - return input_tt->shape; -} - -static inline int32_t GetQmin(const DataType& dtype) { - ICHECK_LE(dtype.bits(), 32) << "QNN ops support int32 or lower precision"; - if (dtype.is_int() || dtype.is_uint()) { - auto min_value_expr = tvm::min_value(dtype); - auto* min_value = tir::as_const_int(min_value_expr); - ICHECK(min_value != nullptr); - return static_cast(min_value[0]); - } else { - LOG(FATAL) << "Type not supported " << dtype; - } -} - -static inline int32_t GetQmax(const DataType& dtype) { - ICHECK_LE(dtype.bits(), 32) << "QNN ops support int32 or lower precision"; - if (dtype.is_int() || dtype.is_uint()) { - auto max_value_expr = tvm::max_value(dtype); - auto* max_value = tir::as_const_int(max_value_expr); - ICHECK(max_value != nullptr); - return static_cast(max_value[0]); - } else { - LOG(FATAL) << "Type not supported " << dtype; - } -} - -/* - * \brief Convert FP32 representation into fixed point representation. - * \param double_multplier The input FP32 number. - * \return The pair of multiplier and shift for fixed point representation. - * \note Converts a floating point number so that it can be represented by - * integers. The representation is - * float_number = (significand) * 2^(exponent) - * - * The significand is a number between 0.5 and 1. This is represented by - * an integer number. For example, if it is int32, then the decimal point - * exists between bit 31 and 30 from LSB (or between first and second bit - * from the left). - * - * Some examples are - * 0.25 = (0.5) * 2^(-1) - * 0.125 = (0.5) * 2^(-2) - * - * Credit to TFLite reference implementation. - */ -std::pair GetFixedPointMultiplierShift(double double_multiplier); - -Expr RequantizeLower(const Expr& input_tensor, const Expr& input_scale, - const Expr& input_zero_point, const Expr& output_scale, - const Expr& output_zero_point, const RequantizeAttrs* param, - const Array& input_shape, const DataType& out_dtype); - -std::string SelectRequntizeParameter(const std::string& arg_value, const std::string& cfg_value, - const bool is_cfg_default, const std::string& name); - -static inline Expr Requantize(const Expr& data, const Array& input_shape, - const Expr& input_scale, const Expr& input_zero_point, - const Expr& output_scale, const Expr& output_zero_point, - const DataType& out_dtype, const int& axis = -1, - const std::string& rounding = "None", - const std::string& compute_dtype = "None") { - auto attrs = make_object(); - attrs->axis = axis; - attrs->out_dtype = std::move(out_dtype); - const RequantizeConfig& cfg = RequantizeConfig::Current(); - attrs->rounding = - SelectRequntizeParameter(rounding, cfg->get_rounding(), cfg->is_default, "rounding"); - attrs->compute_dtype = SelectRequntizeParameter(compute_dtype, cfg->get_compute_dtype(), - cfg->is_default, "compute_dtype"); - return RequantizeLower(data, input_scale, input_zero_point, output_scale, output_zero_point, - attrs.operator->(), input_shape, attrs->out_dtype); -} - -Expr MakeRequantize(Expr data, Expr input_scale, Expr input_zero_point, Expr output_scale, - Expr output_zero_point, int axis, String rounding, String compute_dtype, - DataType out_dtype); - -Expr DequantizeLower(const Expr& input_tensor, const Expr& input_scale, - const Expr& input_zero_point, const Array& types, - const DequantizeAttrs* attrs); - -static inline Expr Dequantize(const Expr& data, const Expr& input_scale, - const Expr& input_zero_point, const Array& types, - const int& axis = -1) { - auto attrs = make_object(); - attrs->axis = std::move(axis); - - return DequantizeLower(data, input_scale, input_zero_point, types, attrs.operator->()); -} -Expr MakeDequantize(Expr data, Expr input_scale, Expr input_zero_point, int axis, - DataType out_dtype = DataType::Float(32)); - -Expr QuantizeLower(const Expr& input_tensor, const Expr& output_scale, - const Expr& output_zero_point, const Array& types, - const QuantizeAttrs* attrs); - -static inline Expr Quantize(const Expr& data, const Expr& output_scale, - const Expr& output_zero_point, const DataType& out_dtype, - const Array& types, const int& axis = -1) { - auto attrs = make_object(); - attrs->axis = std::move(axis); - attrs->out_dtype = std::move(out_dtype); - - return QuantizeLower(data, output_scale, output_zero_point, types, attrs.operator->()); -} -Expr MakeQuantize(Expr data, Expr output_scale, Expr output_zero_point, int axis, - DataType out_dtype); - -static inline int64_t get_const_int(const tvm::PrimExpr& x) { - auto* value_ptr = tir::as_const_int(x); - ICHECK(value_ptr) << "Expr is not a constant int"; - return value_ptr[0]; -} - -/* - * \brief Fixed point multiplication between integer tensor with floating point - * scalar. This implementation rounds to the nearest value when it is midway - * between two representable values. - * \param tensor The quantized input tensor of dtype int64. - * \param multiplier The scalar multiplier. - * \param input_shape Shape of the input tensor. - * \return The sequence of Relay ops for fixed point multiplication with TONEARES rounding. - - * \note Original compuation is scale_fp32 * quantized_tensor. To convert into - * integer computation, the multiplication with fp32 scalar can be - * replaced by multiplication with an int value and then right shifting - * the result. This approximates the floating point computation with a - * fixed point computation. - * - * Computation of fixed point multiplication is consist of following - steps: - * 1) Multiply the fixed point multiplier with quantized tensor. - * 2) Round the result. - * 3) Right shift the result - */ -Expr FixedPointMultiplyToNearest(Expr tensor, double multiplier, - const Array& input_shape); - -/* - * \brief Fixed point multiplication between integer tensor with floating point - scalar where the input tensor is per-axis/per-channel quantized.. - * \param tensor The quantized input tensor of dtype int64. - * \param multiplier The scalar multiplier. - * \param input_shape Shape of the input tensor. - * \param channel_axis The channel_axis along which the input tensor is quantized. Default value is - -1 which corresponds to the last channel_axis. - * \param rounding "UPWARD" or "TONEAREST". The rounding direction when the value - is midway between" "two representable values. - * \return The sequence of Relay ops for fixed point multiplication. - - * \note Original compuation is scale_fp32 * quantized_tensor. To convert into - * integer computation, the multiplication with fp32 vector can be - * replaced by multiplication with an int vector and then right shifting - * the result. This approximates the floating point computation with a - * fixed point computation. - * - * Computation of fixed point multiplication is consist of following - steps: - * 1) Multiply the fixed point multiplier with quantized tensor. - * 2) Round the result. - * 3) Right shift the result - */ -Expr FixedPointMultiplyPerChannel(Expr tensor, std::vector multiplier, - const Array& input_shape, int channel_axis, - const std::string& rounding); - -/* - * Wrapper for 'FixedPointMultiplyPerChannel' with rounding parameter == "TONEAREST". - */ -Expr FixedPointMultiplyPerChannelToNearest(Expr tensor, std::vector multiplier, - const Array& input_shape, int channel_axis); - -/* - * \brief Creates FixedPointMultiply operation where the input tensor is - per-axis/per-channel quantized.. - * \param tensor The quantized input tensor. - * \param multipliers List of scalar multipliers. - * \param channel_axis The channel_axis along which the input tensor is quantized. - * \return The Relay op. - */ -Expr FixedPointMultiplyPerChannel(Expr tensor, const std::vector& multipliers, int axis); - -/* - * \brief Checks whether an expr type is scalar of a given data type. - * \param expr_type The type of expr to be checked. - * \param dtype The expected dtype. - * \return True if the type is a scalar of given dtype - */ -static inline bool IsScalarType(const Type& expr_type, const DataType& dtype) { - const auto* tensor_type = expr_type.as(); - ICHECK(tensor_type) << "Only tensor type can be checked for scalar values. But got" - << AsText(expr_type, false); - ICHECK_EQ(tensor_type->shape.size(), 0); - ICHECK(tensor_type->dtype == dtype) << "Expected " << dtype << " but got " << tensor_type->dtype; - return true; -} - -/* - * \brief Checks whether an expr type is scalar. - * \param expr_type The type of expr to be checked. - * \return True if the type is a scalar - */ -static inline bool IsScalarType(const Type& expr_type) { - const auto* tensor_type = expr_type.as(); - CHECK(tensor_type) << "Only tensor type can be checked for scalar values. But got" - << AsText(expr_type, false); - return tensor_type->shape.size() == 0; -} - -/* - * \brief Checks and assigns types to scale and zero points. - * \param expr_type The type of expr to be checked. - * \param dtype The expected dtype. - * \param shape The shape at C dim of original tensor. - * \param reporter The type reported of original InferType call. - */ -static inline void AssignType(const Type& expr_type, const DataType& dtype, const IndexExpr& shape, - const TypeReporter& reporter) { - // Scale/Zero_points can be either const scalar or a vector with C axis num elems. - const auto* tensor_type = expr_type.as(); - ICHECK(tensor_type) << "Can assign type to Tensor type only. But got " - << AsText(expr_type, false); - const auto tensor_dtype = tensor_type->dtype; - ICHECK(tensor_dtype == dtype) << "Expected type is " << dtype << " but received " << tensor_dtype; - if (tensor_type->shape.size() != 0) { - reporter->Assign(expr_type, TensorType({shape}, tensor_type->dtype)); - } -} - -static inline std::vector GetFloatVectorFromConstant(const Expr& expr) { - const auto* n = expr.as(); - std::vector vals; - ICHECK(n) << "Expr must be a constant expr - " << AsText(expr, false); - int64_t num_elems = 1; - auto shape = n->data.Shape(); - for (size_t i = 0; i < shape.size(); i++) { - num_elems *= shape[i]; - } - for (int64_t i = 0; i < num_elems; i++) { - vals.push_back(static_cast(n->data->data)[i]); - } - return vals; -} - -Expr MakeQnnConv2D(Expr data, Expr weight, Expr input_zero_point, Expr kernel_zero_point, - Expr input_scale, Expr kernel_scale, Array strides, - Array padding, Array dilation, int groups, - IndexExpr channels, Array kernel_size, String data_layout, - String kernel_layout, String out_layout, DataType out_dtype); - -} // namespace qnn -} // namespace relay -} // namespace tvm -#endif // TVM_RELAY_QNN_UTILS_H_ diff --git a/src/relay/quantize/annotate.cc b/src/relay/quantize/annotate.cc deleted file mode 100644 index c704bcbc466b..000000000000 --- a/src/relay/quantize/annotate.cc +++ /dev/null @@ -1,112 +0,0 @@ -/* - * 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 annotate.cc - * - * \brief Annotating the graph with simulated quantize operators. - */ - -#include -#include - -#include "./quantize.h" - -namespace tvm { -namespace relay { -namespace quantize { - -using namespace relay::transform; - -class QAnnotateExpr; -class QAnnotateExprNode : public TempExprNode { - public: - Expr expr; - QAnnotateKind kind; - - void VisitAttrs(tvm::AttrVisitor* v) { - v->Visit("expr", &expr); - v->Visit("kind", &kind); - } - - Expr Realize() const final; - - static constexpr const char* _type_key = "relay.QAnnotateExpr"; - TVM_DECLARE_FINAL_OBJECT_INFO(QAnnotateExprNode, TempExprNode); -}; - -class QAnnotateExpr : public TempExpr { - public: - /*! - * \brief The constructor - * \param expr The original relay expression. - * \param kind The annotation kind. - */ - TVM_DLL QAnnotateExpr(Expr expr, QAnnotateKind kind); - - TVM_DEFINE_OBJECT_REF_METHODS(QAnnotateExpr, TempExpr, QAnnotateExprNode); -}; - -Expr QAnnotateExprNode::Realize() const { return expr; } - -QAnnotateExpr::QAnnotateExpr(Expr expr, QAnnotateKind kind) { - auto rnode = make_object(); - rnode->expr = std::move(expr); - rnode->kind = kind; - data_ = std::move(rnode); -} - -TVM_REGISTER_GLOBAL("relay._quantize.make_annotate_expr").set_body_typed([](Expr expr, int kind) { - return QAnnotateExpr(expr, static_cast(kind)); -}); - -Pass QuantizeAnnotate() { - // TODO(tvm-teams): since partition has added cast_hint in different - // branches, try to remove this in the future. - std::function fmulti_ref = [](const Expr& e) { - if (e->IsInstance()) { - const auto* n = e.as(); - ICHECK(n); - const PackedFunc* f = runtime::Registry::Get("relay.quantize.attach_simulated_quantize"); - Expr ret = (*f)(n->expr, static_cast(kQInput)); - return static_cast(QAnnotateExpr(ret, kQInput)); - } - return e; - }; - - runtime::TypedPackedFunc pass_func = - [=](Function f, IRModule m, PassContext pc) { - auto func = Downcast(ForwardRewrite(f, "FQAnnotateRewrite", nullptr, fmulti_ref)); - auto new_params = func->params; - for (const auto& x : FreeVars(func)) { - new_params.push_back(x); - } - return WithFields(func, new_params); - }; - return CreateFunctionPass(pass_func, 1, "QuantizeAnnotate", {}); -} - -TVM_REGISTER_GLOBAL("relay._quantize.QuantizeAnnotate").set_body_typed(QuantizeAnnotate); - -TVM_REGISTER_NODE_TYPE(QAnnotateExprNode); - -} // namespace quantize -} // namespace relay -} // namespace tvm diff --git a/src/relay/quantize/calibrate.cc b/src/relay/quantize/calibrate.cc deleted file mode 100644 index 2b831ee1403f..000000000000 --- a/src/relay/quantize/calibrate.cc +++ /dev/null @@ -1,224 +0,0 @@ -/* - * 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 calibrate.cc - * - * \brief Create profile graph and calibrate on dataset - */ -#include -#include -#include - -#include - -#include "./quantize.h" - -namespace tvm { -namespace relay { -namespace quantize { - -// KL divergence minimization code is adapted from MXNet. -// The original one is in incubator-mxnet/src/operator/quantization/calibrate.cc -static std::vector SmoothDistribution(const std::vector& p, - const float eps = 0.0001) { - std::vector is_zeros(p.size()); - std::vector is_nonzeros(p.size()); - { - auto it = p.begin(); - std::generate(is_zeros.begin(), is_zeros.end(), - [&it]() { return static_cast(*(it++) == 0.f); }); - } - { - auto it = p.begin(); - std::generate(is_nonzeros.begin(), is_nonzeros.end(), - [&it]() { return static_cast(*(it++) != 0.f); }); - } - size_t n_zeros = std::accumulate(is_zeros.begin(), is_zeros.end(), 0); - size_t n_nonzeros = p.size() - n_zeros; - if (!n_nonzeros) { - // The discrete probability distribution is malformed. All entries are 0. - return std::vector(); - } - float eps1 = eps * static_cast(n_zeros) / static_cast(n_nonzeros); - if (eps1 >= 1.0) return std::vector(); - auto ret = p; - for (size_t i = 0; i < p.size(); i++) { - ret[i] += eps * is_zeros[i] - eps1 * is_nonzeros[i]; - } - return ret; -} - -static float ComputeEntropy(float* p, float* q, size_t size) { - float p_sum = std::accumulate(p, p + size, 0.f); - float q_sum = std::accumulate(q, q + size, 0.f); - float ret = 0; - for (size_t i = 0; i < size; i++) { - ICHECK(p[i] > 0 && q[i] > 0); - p[i] /= p_sum; - q[i] /= q_sum; - if (p[i] && q[i]) ret += p[i] * std::log(p[i] / q[i]); - } - return ret; -} - -float MinimizeKL(const std::vector& hist, const std::vector& hist_edges, int num_bins, - int num_quantized_bins) { - const int zero_bin_idx = num_bins / 2; - const int num_half_quantized_bins = num_quantized_bins / 2; - std::vector thresholds(num_bins / 2 + 1 - num_quantized_bins / 2, 0.f); - std::vector divergence(thresholds.size(), 0.f); - std::vector quantized_bins(num_quantized_bins, 0); - for (int i = num_quantized_bins / 2; i < zero_bin_idx + 1; ++i) { - const int p_bin_idx_start = zero_bin_idx - i; - const int p_bin_idx_stop = zero_bin_idx + i + 1; - thresholds[i - num_half_quantized_bins] = hist_edges[p_bin_idx_stop]; - - std::vector sliced_nd_hist(p_bin_idx_stop - p_bin_idx_start); - std::vector p(sliced_nd_hist.size()); - p[0] = 0; - p.back() = 0; - for (int j = 0; j < num_bins; j++) { - if (j <= p_bin_idx_start) { - p[0] += hist[j]; - } else if (j >= p_bin_idx_stop) { - p.back() += hist[j]; - } else { - sliced_nd_hist[j - p_bin_idx_start] = hist[j]; - p[j - p_bin_idx_start] = hist[j]; - } - } - // calculate how many bins should be merged to generate quantized distribution q - const auto num_merged_bins = sliced_nd_hist.size() / num_quantized_bins; - for (int j = 0; j < num_quantized_bins; j++) { - const int start = j * num_merged_bins; - const int stop = (j + 1) * num_merged_bins; - quantized_bins[j] = - std::accumulate(sliced_nd_hist.begin() + start, sliced_nd_hist.begin() + stop, 0); - } - quantized_bins.back() += std::accumulate( - sliced_nd_hist.begin() + static_cast(num_quantized_bins * num_merged_bins), - sliced_nd_hist.end(), 0); - // expand quantized_bins into p.size bins - std::vector q(sliced_nd_hist.size(), 0); - for (int j = 0; j < num_quantized_bins; j++) { - const int start = j * num_merged_bins; - const int stop = (j == num_quantized_bins - 1) ? q.size() : ((j + 1) * num_merged_bins); - int norm = std::count_if(sliced_nd_hist.begin() + start, sliced_nd_hist.begin() + stop, - [](size_t i) { return i != 0; }); - if (norm) { - for (int k = start; k < stop; k++) { - if (p[k]) q[k] = quantized_bins[j] / norm; - } - } - } - p = SmoothDistribution(p); - q = SmoothDistribution(q); - - if (!q.size()) { - divergence[i - num_half_quantized_bins] = std::numeric_limits::infinity(); - } else { - divergence[i - num_half_quantized_bins] = ComputeEntropy(p.data(), q.data(), p.size()); - } - } - auto min_divergence_idx = - std::distance(divergence.begin(), std::min_element(divergence.begin(), divergence.end())); - return thresholds[min_divergence_idx]; -} - -class StatsCollector : private ExprMutator { - public: - StatsCollector() : simulated_quantize_op_(Op::Get("relay.op.annotation.simulated_quantize")) {} - - Expr Collect(const Expr& expr) { - auto new_e = this->Mutate(expr); - const FunctionNode* func = new_e.as(); - ICHECK(func) << "Input shoule be Function"; - Expr new_body = Tuple(std::move(profile_data_)); - Function ret_func = WithFields(GetRef(func), FreeVars(new_body), new_body); - - // We are changing the function's ret_type to an empty type. Unfortunately, Optional() is - // indistinguishable from NullValue(), so we can't express "update to nullptr" in - // WithFields. - ret_func.CopyOnWrite()->ret_type = NullValue(); - return std::move(ret_func); - } - - private: - Array profile_data_; - const Op& simulated_quantize_op_; - - Expr VisitExpr_(const CallNode* call) { - Expr new_e = ExprMutator::VisitExpr_(call); - const CallNode* new_call = new_e.as(); - ICHECK(new_call); - if (new_call->op == simulated_quantize_op_) { - auto attrs = new_call->attrs.as(); - // rewrite the annotation - auto new_attrs = make_object(); - const Expr& quantize_input = new_call->args[0]; // expression being quantized - auto placeholder = MakeConstantScalar(DataType::Float(32), 0.); // unused argument - Array new_args{quantize_input, placeholder, placeholder, placeholder}; - new_attrs->kind = QAnnotateKind::kQIdentity; - new_attrs->sign = attrs->sign; - new_attrs->rounding = attrs->rounding; - Expr identity_quantize = Call(new_call->op, new_args, Attrs{new_attrs}, {}); - - // add non-const expressions to profile data - if (attrs->kind != QAnnotateKind::kQWeight) { - ICHECK(!quantize_input.as()); - profile_data_.push_back(identity_quantize); - } - return identity_quantize; - } else { - return new_e; - } - } -}; - -/* - * \brief Given an annotated graph, create a profile graph to collect profile data from the - * calibration dataset. - * - * This pass collects simulated_quantize op into a tuple. Simulated_quantize ops are rewritten to - * identity mode. The tuple is the output of the profile graph. Both input and output of this pass - * are relay::Function. - * - * \param expr The simulation graph after annotation. - * \return The profile graph. - */ -Expr CreateStatsCollector(const Expr& expr) { return StatsCollector().Collect(expr); } - -TVM_REGISTER_GLOBAL("relay._quantize.CreateStatsCollector").set_body_typed(CreateStatsCollector); - -TVM_REGISTER_GLOBAL("relay._quantize.FindScaleByKLMinimization") - .set_body([](TVMArgs args, TVMRetValue* ret) { - int* hist_ptr = static_cast(static_cast(args[0])); - float* hist_edges_ptr = static_cast(static_cast(args[1])); - int num_bins = args[2]; - int num_quantized_bins = args[3]; - std::vector hist(hist_ptr, hist_ptr + num_bins); - std::vector hist_edges(hist_edges_ptr, hist_edges_ptr + num_bins + 1); - ret[0] = MinimizeKL(hist, hist_edges, num_bins, num_quantized_bins); - }); - -} // namespace quantize -} // namespace relay -} // namespace tvm diff --git a/src/relay/quantize/partition.cc b/src/relay/quantize/partition.cc deleted file mode 100644 index 6cd596a814ac..000000000000 --- a/src/relay/quantize/partition.cc +++ /dev/null @@ -1,95 +0,0 @@ -/* - * 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 partition.cc - * - * \brief Partition a graph into sections for quantization. - */ - -#include - -#include "../op/annotation/annotation.h" -#include "./quantize.h" - -namespace tvm { -namespace relay { -namespace quantize { - -using namespace relay::transform; - -class QPartitionExpr; -class QPartitionExprNode : public TempExprNode { - public: - /*! \brief The original expression */ - Expr expr; - - void VisitAttrs(tvm::AttrVisitor* v) { v->Visit("expr", &expr); } - - Expr Realize() const final; - - static constexpr const char* _type_key = "relay.QPartitionExpr"; - TVM_DECLARE_FINAL_OBJECT_INFO(QPartitionExprNode, TempExprNode); -}; - -class QPartitionExpr : public TempExpr { - public: - /*! - * \brief The constructor - * \param expr The original relay expression. - */ - TVM_DLL explicit QPartitionExpr(Expr expr); - - TVM_DEFINE_OBJECT_REF_METHODS(QPartitionExpr, TempExpr, QPartitionExprNode); -}; - -Expr QPartitionExprNode::Realize() const { - // insert cast hint and stop fusion - const QConfig& cfg = QConfig::Current(); - Expr ret = CastHint(this->expr, cfg->dtype_input); - return StopFusion(ret); -} - -QPartitionExpr::QPartitionExpr(Expr expr) { - auto rnode = make_object(); - rnode->expr = std::move(expr); - data_ = std::move(rnode); -} - -TVM_REGISTER_GLOBAL("relay._quantize.make_partition_expr").set_body_typed([](Expr expr) { - return QPartitionExpr(expr); -}); - -Pass QuantizePartition() { - runtime::TypedPackedFunc pass_func = - [=](Function f, IRModule m, PassContext pc) { - auto ret = Downcast(ForwardRewrite(f, "FQPartitionRewrite", nullptr, nullptr)); - return ret; - }; - return CreateFunctionPass(pass_func, 1, "QuantizePartition", {}); -} - -TVM_REGISTER_GLOBAL("relay._quantize.QuantizePartition").set_body_typed(QuantizePartition); - -TVM_REGISTER_NODE_TYPE(QPartitionExprNode); - -} // namespace quantize -} // namespace relay -} // namespace tvm diff --git a/src/relay/quantize/quantize.cc b/src/relay/quantize/quantize.cc deleted file mode 100644 index afd5e522657d..000000000000 --- a/src/relay/quantize/quantize.cc +++ /dev/null @@ -1,149 +0,0 @@ -/* - * 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 quantize.cc - * - * \brief transform a graph to a low-bit graph - * for compression and acceleration. - */ -#include "./quantize.h" - -#include -#include -#include - -#include - -namespace tvm { -namespace relay { -namespace quantize { - -TVM_REGISTER_NODE_TYPE(SimulatedQuantizeAttrs); - -bool SimulatedQuantizeRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - ICHECK_EQ(types.size(), 5); - const auto param = attrs.as(); - ICHECK(param != nullptr); - - const auto* data = types[0].as(); - - if (data == nullptr) { - return false; - } - - ICHECK_NE(data->shape.size(), 0) << "Input shape cannot be empty"; - - reporter->Assign(types[1], TensorType({}, DataType::Float(32))); // dom_scale - reporter->Assign(types[2], TensorType({}, DataType::Float(32))); // clip_min - reporter->Assign(types[3], TensorType({}, DataType::Float(32))); // clip_max - reporter->Assign(types[4], types[0]); // output - return true; -} - -RELAY_REGISTER_OP("relay.op.annotation.simulated_quantize") - .describe(R"code(simulated quantize op)code" TVM_ADD_FILELINE) - .set_num_inputs(4) - .add_argument("data", "Tensor", "The input data.") - .add_argument("dom_scale", "Tensor", "The domain scale of input data. It should be a scalar") - .add_argument("clip_min", "Tensor", "lower bound. It should be a scalar") - .add_argument("clip_max", "Tensor", "upper bound. It should be a scalar") - .set_attrs_type() - .set_support_level(11) - .add_type_rel("SimulatedQuantize", SimulatedQuantizeRel); - -TVM_REGISTER_GLOBAL("relay._quantize.simulated_quantize") - .set_body_typed([](Expr data, Expr dom_scale, Expr clip_min, Expr clip_max, int kind, bool sign, - String rounding) { - auto attrs = make_object(); - attrs->kind = kind; - attrs->sign = sign; - attrs->rounding = rounding; - static const Op& op = Op::Get("relay.op.annotation.simulated_quantize"); - return Call(op, {data, dom_scale, clip_min, clip_max}, Attrs(attrs), {}); - }); - -/*! \brief Entry to hold the BuildConfig context stack. */ -struct TVMQConfigThreadLocalEntry { - /*! \brief The default build config if the stack is empty */ - QConfig default_config; - - /*! \brief The current build config context */ - std::stack context_stack; - - TVMQConfigThreadLocalEntry() : default_config(make_object()) {} -}; - -/*! \brief Thread local store to hold the BuildConfig context stack. */ -typedef dmlc::ThreadLocalStore TVMQConfigThreadLocalStore; - -void QConfig::EnterQConfigScope(const QConfig& build_config) { - TVMQConfigThreadLocalEntry* entry = TVMQConfigThreadLocalStore::Get(); - entry->context_stack.push(build_config); -} - -void QConfig::ExitQConfigScope() { - TVMQConfigThreadLocalEntry* entry = TVMQConfigThreadLocalStore::Get(); - entry->context_stack.pop(); -} - -QConfig& QConfig::Current() { - TVMQConfigThreadLocalEntry* entry = TVMQConfigThreadLocalStore::Get(); - if (entry->context_stack.size() > 0) { - return entry->context_stack.top(); - } - - return entry->default_config; -} - -TVM_REGISTER_NODE_TYPE(QConfigNode); - -TVM_STATIC_IR_FUNCTOR(ReprPrinter, vtable) - .set_dispatch([](const ObjectRef& ref, ReprPrinter* p) { - auto* op = static_cast(ref.get()); - p->stream << "qconfig("; - p->stream << "nbit_input=" << op->nbit_input << ", "; - p->stream << "nbit_weight=" << op->nbit_weight << ", "; - p->stream << "nbit_activation=" << op->nbit_activation << ", "; - p->stream << "calibrate_mode=" << op->calibrate_mode << ", "; - p->stream << "global_scale=" << op->global_scale << ", "; - p->stream << "weight_scale=" << op->weight_scale << ", "; - p->stream << "skip_conv_layers==" << op->skip_conv_layers << ", "; - p->stream << "skip_dense_layer==" << op->skip_dense_layer << ", "; - p->stream << "do_simulation==" << op->do_simulation << ", "; - p->stream << "round_for_shift==" << op->round_for_shift << ", "; - p->stream << "debug_enabled_ops==" << op->debug_enabled_ops << ", "; - p->stream << "rounding==" << op->rounding << ", "; - p->stream << "partition_conversions==" << op->partition_conversions; - p->stream << ")"; - }); - -TVM_REGISTER_GLOBAL("relay._quantize._GetCurrentQConfig").set_body_typed([]() -> QConfig { - return QConfig::Current(); -}); - -TVM_REGISTER_GLOBAL("relay._quantize._EnterQConfigScope") - .set_body_typed(QConfig::EnterQConfigScope); - -TVM_REGISTER_GLOBAL("relay._quantize._ExitQConfigScope").set_body_typed(QConfig::ExitQConfigScope); - -} // namespace quantize -} // namespace relay -} // namespace tvm diff --git a/src/relay/quantize/quantize.h b/src/relay/quantize/quantize.h deleted file mode 100644 index 7c2acbcb06d4..000000000000 --- a/src/relay/quantize/quantize.h +++ /dev/null @@ -1,156 +0,0 @@ -/* - * 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 tvm/relay/quantize.h - * \brief Header of definitions for quantization - */ -#ifndef TVM_RELAY_QUANTIZE_QUANTIZE_H_ -#define TVM_RELAY_QUANTIZE_QUANTIZE_H_ - -#include -#include - -#include - -#include "../transforms/pattern_utils.h" - -namespace tvm { -namespace relay { -namespace quantize { - -/*! \brief Kind of annotate field */ -enum QAnnotateKind : int { kQIdentity = 0, kQInput = 1, kQWeight = 2, kQActivation = 3 }; - -/*! \brief Attribute for simulated quantize operator */ -struct SimulatedQuantizeAttrs : public tvm::AttrsNode { - int kind; - bool sign; - std::string rounding; - - TVM_DECLARE_ATTRS(SimulatedQuantizeAttrs, "relay.attrs.SimulatedQuantizeAttrs") { - TVM_ATTR_FIELD(kind).describe("kind of field, hint for nbit/dtype configuration."); - TVM_ATTR_FIELD(sign).set_default(true).describe("whether to use signed data type."); - TVM_ATTR_FIELD(rounding).set_default("round").describe( - "rounding mode. Can be 'floor', 'ceil', 'round'"); - } -}; - -class QConfig; -/*! - * \brief Container for build configuration options - */ -class QConfigNode : public Object { - public: - int nbit_input = 8; - int nbit_weight = 8; - int nbit_activation = 32; - DataType dtype_input = DataType::Int(8); - DataType dtype_weight = DataType::Int(8); - DataType dtype_activation = DataType::Int(32); - std::string calibrate_mode = "global_scale"; - double global_scale = 8.0; - std::string weight_scale = "power2"; - bool skip_dense_layer = true; - Array skip_conv_layers = Array(ObjectPtr(nullptr)); - bool do_simulation = false; - bool round_for_shift = true; - Array debug_enabled_ops = Array(ObjectPtr(nullptr)); - std::string rounding = "UPWARD"; - int calibrate_chunk_by = -1; - std::string partition_conversions = "disabled"; - - void VisitAttrs(AttrVisitor* v) { - v->Visit("nbit_input", &nbit_input); - v->Visit("nbit_weight", &nbit_weight); - v->Visit("nbit_activation", &nbit_activation); - v->Visit("dtype_input", &dtype_input); - v->Visit("dtype_weight", &dtype_weight); - v->Visit("dtype_activation", &dtype_activation); - v->Visit("calibrate_mode", &calibrate_mode); - v->Visit("global_scale", &global_scale); - v->Visit("weight_scale", &weight_scale); - v->Visit("skip_dense_layer", &skip_dense_layer); - v->Visit("skip_conv_layers", &skip_conv_layers); - v->Visit("do_simulation", &do_simulation); - v->Visit("round_for_shift", &round_for_shift); - v->Visit("debug_enabled_ops", &debug_enabled_ops); - v->Visit("rounding", &rounding); - v->Visit("calibrate_chunk_by", &calibrate_chunk_by); - v->Visit("partition_conversions", &partition_conversions); - } - - static constexpr const char* _type_key = "relay.quantize.QConfig"; - TVM_DECLARE_FINAL_OBJECT_INFO(QConfigNode, Object); -}; - -/*! - * \brief Container for build configuration options - */ -class QConfig : public ObjectRef { - public: - QConfig() {} - explicit QConfig(ObjectPtr n) : ObjectRef(n) {} - - const QConfigNode* operator->() const { return static_cast(get()); } - - QConfigNode* operator->() { return static_cast(get_mutable()); } - - /*! - * \brief Push a new BuildConfig context onto the thread local stack. - * \param build_config The configuration to set as the current context. - */ - static void EnterQConfigScope(const QConfig& qconfig); - - /*! - * \brief Pop a build config off the thread local context stack, restoring the previous - * configuration as the current context. - */ - static void ExitQConfigScope(); - - /*! - * \brief Get the current BuildConfig context from thread local storage, or a default - * configuration if a BuildConfig scope has not been entered. - * \return The configuration that is the current context. - */ - static QConfig& Current(); - - using ContainerType = QConfigNode; -}; - -/*! - * \brief RAII container to provide a scoped BuildConfig context. Pushes a configuration onto the - * context stack when constructed, and pops it when destructed. - */ -struct QConfigContext { - /*! - * \brief Enter a new BuildConfig context. The given BuildConfig becomes the new current - * context. When the BuildConfigContext is destructed, the previous context is restored. - * \param build_config The BuildConfig to set as the new current context. - */ - explicit QConfigContext(const QConfig& qconfig) { QConfig::EnterQConfigScope(qconfig); } - - /*! \brief Destructor. Pops the context off the thread local stack. */ - ~QConfigContext() { QConfig::ExitQConfigScope(); } -}; - -} // namespace quantize -} // namespace relay -} // namespace tvm -#endif // TVM_RELAY_QUANTIZE_QUANTIZE_H_ diff --git a/src/relay/quantize/realize.cc b/src/relay/quantize/realize.cc deleted file mode 100644 index 514be1f0a71c..000000000000 --- a/src/relay/quantize/realize.cc +++ /dev/null @@ -1,562 +0,0 @@ -/* - * 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 realize.cc - * - * \brief Realizing the simulated graph into real low-precision - * graph. - */ - -#include "./realize.h" - -#include -#include -#include - -#include "../op/annotation/annotation.h" -#include "../qnn/utils.h" -#include "../transforms/fold_constant.h" -#include "../transforms/infer_layout_utils.h" -#include "./quantize.h" - -namespace tvm { -namespace relay { -namespace quantize { - -using namespace relay::transform; - -Expr QRealizeIntExprNode::Realize() const { - Expr data = this->data; - // dequantize - data = Cast(data, DataType::Float(32)); - data = Multiply(data, this->dom_scale); - return data; -} - -QRealizeIntExpr::QRealizeIntExpr(Expr data, Expr dom_scale, DataType dtype) { - ObjectPtr n = make_object(); - n->data = std::move(data); - n->dom_scale = std::move(dom_scale); - n->dtype = std::move(dtype); - data_ = std::move(n); -} - -inline Expr ForwardOp(const Call& ref_call, const Array& args) { - return Call(ref_call->op, args, ref_call->attrs, ref_call->type_args); -} - -/* calculate `data * s1 / s2`, use shift if possible */ -inline Expr MulAndDiv(Expr data, float s1, float s2, DataType dtype, - const Array& data_shape) { - const QConfig& cfg = QConfig::Current(); - // here we assume the dtype of data is dtype activation - if (s1 == s2) return data; - - float factor = s1 / s2; - float shift_factor = std::log2(factor); - ICHECK_GT(shift_factor, 0); - if (static_cast(shift_factor) == shift_factor) { - return LeftShift(data, MakeConstantScalar(dtype, static_cast(shift_factor))); - } else if (static_cast(factor) == factor) { - return Multiply(data, MakeConstantScalar(dtype, factor)); - } else { - if (cfg->rounding == "UPWARD") { - auto [fixed_point_multiplier, shift] = qnn::GetFixedPointMultiplierShift(factor); - data = relay::FixedPointMultiply(data, fixed_point_multiplier, shift); - } else { - data = qnn::FixedPointMultiplyToNearest(data, factor, data_shape); - } - - return Cast(data, dtype); - } -} - -Expr QuantizeRealize(const Call& ref_call, const Array& new_args, const ObjectRef& ctx) { - const QConfig& cfg = QConfig::Current(); - // do not handle data type cast - const auto param = ref_call->attrs.as(); - ICHECK_EQ(param->rounding, "round"); - - Expr dom_scale = new_args[1]; - Expr clip_min = new_args[2]; - Expr clip_max = new_args[3]; - - float dom_scale_imm = GetScalarFromConstant(dom_scale); - float clip_min_imm = GetScalarFromConstant(clip_min); - float clip_max_imm = GetScalarFromConstant(clip_max); - - // x * idom_scale = y * odom_scale - // => y = x * idom_scale / odom_scale - if (const auto* n = new_args[0].as()) { - // int32->int8 - Expr data = n->data; - float idom_scale_imm = GetScalarFromConstant(n->dom_scale); - float odom_scale_imm = GetScalarFromConstant(dom_scale); - if (idom_scale_imm == odom_scale_imm) { - // same domain scale, only clip - data = Clip(data, clip_min_imm, clip_max_imm); - return QRealizeIntExpr(data, dom_scale, n->dtype); - } - - float shift_nbit = std::log2(odom_scale_imm / idom_scale_imm); - ICHECK_NE(shift_nbit, 0); - if (static_cast(shift_nbit) == shift_nbit) { - if (shift_nbit > 0) { - // use right shift - if (cfg->round_for_shift) { - float round_bias = std::pow(2.0, shift_nbit - 1); - data = Add(data, MakeConstantScalar(cfg->dtype_activation, static_cast(round_bias))); - } - data = RightShift(data, - MakeConstantScalar(cfg->dtype_activation, static_cast(shift_nbit))); - } else { - data = LeftShift(data, - MakeConstantScalar(cfg->dtype_activation, static_cast(-shift_nbit))); - } - data = Clip(data, clip_min_imm, clip_max_imm); - return QRealizeIntExpr(data, dom_scale, n->dtype); - } else { - data = Cast(data, DataType::Int(64)); - if (cfg->rounding == "UPWARD") { - auto [fixed_point_multiplier, shift] = - qnn::GetFixedPointMultiplierShift(idom_scale_imm / odom_scale_imm); - data = relay::FixedPointMultiply(data, fixed_point_multiplier, shift); - } else { - data = qnn::FixedPointMultiplyToNearest(data, idom_scale_imm / odom_scale_imm, - ref_call->type_as()->shape); - } - data = Cast(Clip(data, clip_min_imm, clip_max_imm), n->dtype); - return QRealizeIntExpr(data, dom_scale, n->dtype); - } - } - - // quantize from real - ICHECK(!new_args[0]->IsInstance()); - Expr data = new_args[0]; - Expr scaled_data = Multiply(data, MakeConstantScalar(DataType::Float(32), 1 / dom_scale_imm)); - Expr round_data = Clip(Round(scaled_data), clip_min_imm, clip_max_imm); - return QRealizeIntExpr(round_data, dom_scale, DataType::Float(32)); -} - -InferCorrectLayoutOutput SimQuantizeLayout(const Attrs& attrs, const Array& new_in_layouts, - const Array& old_in_layouts, - const Array& old_in_types) { - Layout ret; - - if (new_in_layouts.defined()) { - ICHECK_GE(new_in_layouts.size(), 1); - ret = new_in_layouts[0]; - } else { - ICHECK_GE(old_in_layouts.size(), 1); - ret = old_in_layouts[0]; - } - Layout channel_layout = Layout("C"); - Array input_layouts = {ret, channel_layout, channel_layout, channel_layout}; - return InferCorrectLayoutOutput(input_layouts, {ret}, attrs); -} - -RELAY_REGISTER_OP("relay.op.annotation.simulated_quantize") - .set_attr("FQRealizeRewrite", QuantizeRealize) - .set_attr("FInferCorrectLayout", SimQuantizeLayout); - -Expr Conv2dRealize(const Call& ref_call, const Array& new_args, const ObjectRef& ctx) { - const QConfig& cfg = QConfig::Current(); - ICHECK_EQ(new_args.size(), 2); - if (new_args[0].as() && new_args[1].as()) { - const auto* lhs = new_args[0].as(); - const auto* rhs = new_args[1].as(); - Expr ldata = lhs->data; - if (lhs->dtype != cfg->dtype_input) { - ldata = Cast(ldata, cfg->dtype_input); - } - Expr rdata = Cast(rhs->data, cfg->dtype_weight); - - const auto ref_attrs = ref_call->attrs.as(); - auto attrs = make_object(); - *attrs = *ref_attrs; - DataType out_dtype = cfg->dtype_activation; - attrs->out_dtype = out_dtype; - - Expr ret = Call(ref_call->op, {ldata, rdata}, Attrs(attrs), ref_call->type_args); - Expr mul = Multiply(lhs->dom_scale, rhs->dom_scale); - Expr dom_scale = FoldConstantExpr(mul); - return QRealizeIntExpr(ret, dom_scale, out_dtype); - } - ICHECK(!new_args[0]->IsInstance() || !new_args[1]->IsInstance()); - return Expr(nullptr); -} - -RELAY_REGISTER_OP("nn.conv2d").set_attr("FQRealizeRewrite", Conv2dRealize); - -Expr Conv1dRealize(const Call& ref_call, const Array& new_args, const ObjectRef& ctx) { - const QConfig& cfg = QConfig::Current(); - CHECK_EQ(new_args.size(), 2); - if (!new_args[0]->IsInstance() && !new_args[1]->IsInstance()) { - return Expr(nullptr); - } - const auto* lhs = new_args[0].as(); - CHECK(lhs); - const auto* rhs = new_args[1].as(); - CHECK(rhs); - - Expr ldata = lhs->data; - if (lhs->dtype != cfg->dtype_input) { - ldata = Cast(ldata, cfg->dtype_input); - } - Expr rdata = Cast(rhs->data, cfg->dtype_weight); - - const auto ref_attrs = ref_call->attrs.as(); - auto attrs = make_object(); - *attrs = *ref_attrs; - DataType out_dtype = cfg->dtype_activation; - attrs->out_dtype = out_dtype; - - Expr ret = Call(ref_call->op, {ldata, rdata}, Attrs(attrs), ref_call->type_args); - Expr mul = Multiply(lhs->dom_scale, rhs->dom_scale); - Expr dom_scale = FoldConstantExpr(mul); - return QRealizeIntExpr(ret, dom_scale, out_dtype); -} - -RELAY_REGISTER_OP("nn.conv1d").set_attr("FQRealizeRewrite", Conv1dRealize); - -Expr DenseRealize(const Call& ref_call, const Array& new_args, const ObjectRef& ctx) { - const QConfig& cfg = QConfig::Current(); - ICHECK_EQ(new_args.size(), 2); - if (!new_args[0]->IsInstance() || !new_args[1]->IsInstance()) { - return Expr(nullptr); - } - const auto* lhs = new_args[0].as(); - const auto* rhs = new_args[1].as(); - - Expr ldata = lhs->data; - if (lhs->dtype != cfg->dtype_input) { - ldata = Cast(ldata, cfg->dtype_input); - } - Expr rdata = Cast(rhs->data, cfg->dtype_weight); - - const auto ref_attrs = ref_call->attrs.as(); - auto attrs = make_object(); - *attrs = *ref_attrs; - DataType out_dtype = cfg->dtype_activation; - attrs->out_dtype = out_dtype; - - Expr ret = Call(ref_call->op, {ldata, rdata}, Attrs(attrs), ref_call->type_args); - Expr mul = Multiply(lhs->dom_scale, rhs->dom_scale); - Expr dom_scale = FoldConstantExpr(mul); - return QRealizeIntExpr(ret, dom_scale, out_dtype); -} - -RELAY_REGISTER_OP("nn.dense").set_attr("FQRealizeRewrite", DenseRealize); - -Expr MulRealize(const Call& ref_call, const Array& new_args, const ObjectRef& ctx) { - const QConfig& cfg = QConfig::Current(); - ICHECK_EQ(new_args.size(), 2); - if (new_args[0].as() && new_args[1].as()) { - // execute the operation with activation data type. - const auto* lhs = new_args[0].as(); - const auto* rhs = new_args[1].as(); - Expr ldata = lhs->data; - Expr rdata = rhs->data; - - DataType dtype = cfg->dtype_activation; - if (lhs->dtype != dtype) { - ldata = Cast(ldata, dtype); - } - if (rhs->dtype != dtype) { - rdata = Cast(rdata, dtype); - } - - Expr ret = ForwardOp(ref_call, {ldata, rdata}); - Expr mul = Multiply(lhs->dom_scale, rhs->dom_scale); - Expr dom_scale = FoldConstantExpr(mul); - return QRealizeIntExpr(ret, dom_scale, dtype); - } - ICHECK(!new_args[0]->IsInstance() || !new_args[1]->IsInstance()); - return Expr(nullptr); -} - -RELAY_REGISTER_OP("multiply").set_attr("FQRealizeRewrite", MulRealize); - -float ChooseDomScale(const std::vector& nptrs) { - if (nptrs.size() == 2) { - // x = a * s1, y = b * s2 - // x + y = (a * s1 / s2 + b) * s2, if s1 > s2 - // = (a + b * s2 / s1) * s1, if s2 > s1 - float s1 = GetScalarFromConstant(nptrs[0]->dom_scale); - float s2 = GetScalarFromConstant(nptrs[1]->dom_scale); - return s1 > s2 ? s2 : s1; - } else { - const QConfig& cfg = QConfig::Current(); - float scale = cfg->global_scale; - return scale / std::pow(2.0, cfg->nbit_activation - 1); - } -} - -/* \brief Unify the dom scale of arguments */ -Array UnifyDTypeScale(const Array& ref_args, const Array& args, - DataType* dtype_ptr, Expr* scale_ptr, - DataType dtype = DataType::Void()) { - static const Op& simulated_quantize = Op::Get("relay.op.annotation.simulated_quantize"); - const QConfig& cfg = QConfig::Current(); - - std::vector nptrs; - Array ret; - for (auto arg : args) { - const auto* nptr = arg.as(); - ICHECK(nptr); - nptrs.push_back(nptr); - ret.push_back(nptr->data); - } - - // unify the data type - ICHECK_EQ(ref_args.size(), args.size()); - - if (dtype.is_void()) { - if (ret.size() == 2 && nptrs[1]->dtype == cfg->dtype_input) { - dtype = cfg->dtype_input; - } else { - dtype = cfg->dtype_activation; - } - } - - for (size_t i = 0; i < ret.size(); ++i) { - auto ref_arg = ref_args[i].as(); - if (nptrs[i]->dtype != dtype) { - ret.Set(i, Cast(ret[i], dtype)); - } else if (ref_arg && ref_arg->op.same_as(simulated_quantize) && - ref_arg->attrs.as()->kind == kQInput) { - auto new_arg = Cast(ret[i], cfg->dtype_input); - new_arg = StopFusion(new_arg); - ret.Set(i, Cast(new_arg, dtype)); - } - } - - // unify the dom_scale - float s = ChooseDomScale(nptrs); - Expr dom_scale = MakeConstantScalar(DataType::Float(32), s); - for (size_t i = 0; i < ret.size(); ++i) { - float cur_s = GetScalarFromConstant(nptrs[i]->dom_scale); - ret.Set(i, MulAndDiv(ret[i], cur_s, s, dtype, ref_args[i]->type_as()->shape)); - } - - *dtype_ptr = dtype; - *scale_ptr = dom_scale; - return ret; -} - -Expr AddRealize(const Call& ref_call, const Array& new_args, const ObjectRef& ctx) { - ICHECK_EQ(new_args.size(), 2); - if (new_args[0].as() && new_args[1].as()) { - DataType dtype; - Expr dom_scale; - // execute the operation with activation data type. - const QConfig& cfg = QConfig::Current(); - Array ret_args = - UnifyDTypeScale(ref_call->args, new_args, &dtype, &dom_scale, cfg->dtype_activation); - for (size_t i = 0; i < ret_args.size(); ++i) { - // do not fuse float32 arg - if (new_args[i].as()->dtype == DataType::Float(32)) { - ret_args.Set(i, StopFusion(ret_args[i])); - } - } - Expr ret = ForwardOp(ref_call, ret_args); - return QRealizeIntExpr(ret, dom_scale, dtype); - } - - ICHECK(!new_args[0]->IsInstance() && !new_args[1]->IsInstance()); - return Expr(nullptr); -} - -RELAY_REGISTER_OP("add").set_attr("FQRealizeRewrite", AddRealize); - -Expr ClipRealize(const Call& ref_call, const Array& new_args, const ObjectRef& ctx) { - ICHECK_EQ(new_args.size(), 1); - if (const auto* n = new_args[0].as()) { - const auto ref_attrs = ref_call->attrs.as(); - auto attrs = make_object(); - double dom_scale = GetScalarFromConstant(n->dom_scale); - attrs->a_min = ref_attrs->a_min / dom_scale; - attrs->a_max = ref_attrs->a_max / dom_scale; - - Expr ret = Call(ref_call->op, {n->data}, Attrs(attrs), ref_call->type_args); - return QRealizeIntExpr(ret, n->dom_scale, n->dtype); - } - ICHECK(!new_args[0]->IsInstance()); - return Expr(nullptr); -} - -RELAY_REGISTER_OP("clip").set_attr("FQRealizeRewrite", ClipRealize); - -Expr ConcatenateRealize(const Call& ref_call, const Array& new_args, const ObjectRef& ctx) { - ICHECK_EQ(new_args.size(), 1); - ICHECK_EQ(ref_call->args.size(), 1); - - const auto* tuple = new_args[0].as(); - const auto* ref_tuple = ref_call->args[0].as(); - ICHECK(tuple); - ICHECK(ref_tuple); - const Array& arr = tuple->fields; - const Array& ref_arr = ref_tuple->fields; - - if (arr[0].as()) { - DataType dtype; - Expr dom_scale; - Array ret_args = UnifyDTypeScale(ref_arr, arr, &dtype, &dom_scale); - Expr ret = ForwardOp(ref_call, {Tuple(ret_args)}); - return QRealizeIntExpr(ret, dom_scale, dtype); - } else { - for (auto arg : new_args) { - ICHECK(!arg->IsInstance()); - } - return Expr(nullptr); - } -} - -RELAY_REGISTER_OP("concatenate").set_attr("FQRealizeRewrite", ConcatenateRealize); - -/* \brief forward the original operator */ -Expr IdentityRealize(const Call& ref_call, const Array& new_args, const ObjectRef& ctx) { - ICHECK_EQ(new_args.size(), 1); - if (const auto* n = new_args[0].as()) { - Expr ret = ForwardOp(ref_call, {n->data}); - return QRealizeIntExpr(ret, n->dom_scale, n->dtype); - } - ICHECK(!new_args[0]->IsInstance()); - return Expr(nullptr); -} - -RELAY_REGISTER_OP("nn.relu").set_attr("FQRealizeRewrite", IdentityRealize); - -RELAY_REGISTER_OP("reshape").set_attr("FQRealizeRewrite", IdentityRealize); - -RELAY_REGISTER_OP("strided_slice").set_attr("FQRealizeRewrite", IdentityRealize); - -RELAY_REGISTER_OP("nn.batch_flatten") - .set_attr("FQRealizeRewrite", IdentityRealize); - -RELAY_REGISTER_OP("transpose").set_attr("FQRealizeRewrite", IdentityRealize); - -RELAY_REGISTER_OP("annotation.stop_fusion") - .set_attr("FQRealizeRewrite", IdentityRealize); - -/* \brief for unary operators which requantize its input to dtype_nbit */ -Expr CastDtypeInputRealize(const Call& ref_call, const Array& new_args, - const ObjectRef& ctx) { - const QConfig& cfg = QConfig::Current(); - ICHECK_EQ(new_args.size(), 1); - if (const auto* n = new_args[0].as()) { - Expr data = Cast(n->data, cfg->dtype_input); - Expr ret = ForwardOp(ref_call, {data}); - return QRealizeIntExpr(ret, n->dom_scale, cfg->dtype_input); - } - ICHECK(!new_args[0]->IsInstance()); - return Expr(nullptr); -} - -RELAY_REGISTER_OP("nn.max_pool2d") - .set_attr("FQRealizeRewrite", CastDtypeInputRealize); - -RELAY_REGISTER_OP("nn.max_pool1d") - .set_attr("FQRealizeRewrite", CastDtypeInputRealize); - -Expr AvgPoolRealize(const Call& ref_call, const Array& new_args, const ObjectRef& ctx) { - const QConfig& cfg = QConfig::Current(); - ICHECK_EQ(new_args.size(), 1); - if (const auto* n = new_args[0].as()) { - Expr data = n->data; - if (n->dtype != cfg->dtype_activation) { - data = Cast(n->data, cfg->dtype_activation); - } - Expr ret = ForwardOp(ref_call, {data}); - return QRealizeIntExpr(ret, n->dom_scale, cfg->dtype_activation); - } - ICHECK(!new_args[0]->IsInstance()); - return Expr(nullptr); -} - -RELAY_REGISTER_OP("nn.avg_pool2d").set_attr("FQRealizeRewrite", AvgPoolRealize); - -RELAY_REGISTER_OP("nn.global_avg_pool2d") - .set_attr("FQRealizeRewrite", AvgPoolRealize); - -Expr CastHintRealize(const Call& ref_call, const Array& new_args, const ObjectRef& ctx) { - const auto param = ref_call->attrs.as(); - ICHECK_EQ(new_args.size(), 1); - if (const auto* n = new_args[0].as()) { - Expr ret = Cast(n->data, param->dtype); - return QRealizeIntExpr(ret, n->dom_scale, param->dtype); - } - ICHECK(!new_args[0]->IsInstance()); - return Expr(nullptr); -} - -RELAY_REGISTER_OP("annotation.cast_hint") - .set_attr("FQRealizeRewrite", CastHintRealize); - -Expr BatchMatmulRealize(const Call& ref_call, const Array& new_args, const ObjectRef& ctx) { - const QConfig& cfg = QConfig::Current(); - ICHECK_EQ(new_args.size(), 2); - if (!new_args[0]->IsInstance() || !new_args[1]->IsInstance()) { - return Expr(nullptr); - } - const auto* lhs = new_args[0].as(); - const auto* rhs = new_args[1].as(); - - Expr ldata = lhs->data; - Expr rdata = rhs->data; - DataType dtype_input = cfg->dtype_input; - DataType dtype_weight = cfg->dtype_weight; - - if (lhs->dtype != dtype_input) { - ldata = Cast(ldata, dtype_input); - } - if (rhs->dtype != dtype_weight) { - rdata = Cast(rdata, dtype_weight); - } - - const auto ref_attrs = ref_call->attrs.as(); - auto attrs = make_object(); - *attrs = *ref_attrs; - DataType out_dtype = cfg->dtype_activation; - attrs->out_dtype = out_dtype; - - Expr ret = Call(ref_call->op, {ldata, rdata}, Attrs(attrs), ref_call->type_args); - Expr mul = Multiply(lhs->dom_scale, rhs->dom_scale); - Expr dom_scale = FoldConstantExpr(mul); - return QRealizeIntExpr(ret, dom_scale, out_dtype); -} - -RELAY_REGISTER_OP("nn.batch_matmul") - .set_attr("FQRealizeRewrite", BatchMatmulRealize); - -Pass QuantizeRealizePass() { - runtime::TypedPackedFunc pass_func = - [=](Function f, IRModule m, PassContext pc) { - return Downcast(ForwardRewrite(f, "FQRealizeRewrite", nullptr, nullptr)); - }; - return CreateFunctionPass(pass_func, 1, "QuantizeRealize", {}); -} - -TVM_REGISTER_GLOBAL("relay._quantize.QuantizeRealize").set_body_typed(QuantizeRealizePass); - -} // namespace quantize -} // namespace relay -} // namespace tvm diff --git a/src/relay/quantize/realize.h b/src/relay/quantize/realize.h deleted file mode 100644 index 6eba69e9c9b1..000000000000 --- a/src/relay/quantize/realize.h +++ /dev/null @@ -1,75 +0,0 @@ -/* - * 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 realize.h - * - * \brief Header of definitions for op realizations - * - */ -#ifndef TVM_RELAY_QUANTIZE_REALIZE_H_ -#define TVM_RELAY_QUANTIZE_REALIZE_H_ - -#include - -namespace tvm { -namespace relay { -namespace quantize { - -class QRealizeExprNode : public TempExprNode { - public: - Expr data; - static constexpr const char* _type_key = "relay.quantize.QRealizeExpr"; - TVM_DECLARE_BASE_OBJECT_INFO(QRealizeExprNode, TempExprNode); -}; - -class QRealizeExpr : public TempExpr { - public: - TVM_DEFINE_OBJECT_REF_METHODS(QRealizeExpr, TempExpr, QRealizeExprNode); -}; - -class QRealizeIntExprNode : public QRealizeExprNode { - public: - Expr dom_scale; - DataType dtype; - - void VisitAttrs(tvm::AttrVisitor* v) { - v->Visit("data", &data); - v->Visit("dom_scale", &dom_scale); - v->Visit("dtype", &dtype); - } - - Expr Realize() const final; - - static constexpr const char* _type_key = "relay.quantize.QRealizeIntExpr"; - TVM_DECLARE_FINAL_OBJECT_INFO(QRealizeIntExprNode, QRealizeExprNode); -}; - -class QRealizeIntExpr : public QRealizeExpr { - public: - TVM_DLL QRealizeIntExpr(Expr data, Expr dom_scale, DataType dtype); - - TVM_DEFINE_OBJECT_REF_METHODS(QRealizeIntExpr, QRealizeExpr, QRealizeIntExprNode); -}; - -} // namespace quantize -} // namespace relay -} // namespace tvm -#endif // TVM_RELAY_QUANTIZE_REALIZE_H_ diff --git a/src/relay/transforms/alter_op_layout.cc b/src/relay/transforms/alter_op_layout.cc deleted file mode 100644 index f347eddae760..000000000000 --- a/src/relay/transforms/alter_op_layout.cc +++ /dev/null @@ -1,144 +0,0 @@ -/* - * 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 alter_op_layout.cc - * \brief Alternate the layouts of operators or replace primitive operators with - other expressions. This pass can be used for computing convolution in - custom layouts or other general weight pre-transformation. - */ -#include -#include -#include -#include -#include - -#include -#include -#include -#include -#include -#include - -#include "pattern_utils.h" -#include "transform_layout.h" - -namespace tvm { -namespace relay { - -namespace alter_op_layout { - -/*! - * \brief Container to instantiate a Node for alter op layouts. - */ -class AlterTransformMemorizerNode : public TransformMemorizerNode { - public: - static constexpr const char* _type_key = "relay.alter_op_layout.AlterTransformMemorizerNode"; - - /*! - * \brief Defines the call transformation for AlterOpLayout pass. The new layouts are defined by - * used for different targets using a packed func. - * \param ref_call The original call. - * \param new_attrs Updated attributes consistent with new layouts. - * \param new_args The traversed/recursed args to the call. - * \return The new Call after calling the packed func. - */ - Call CallWithNewLayouts(const Call& ref_call, Attrs new_attrs, - const std::vector& new_args) override { - static auto falter_layout = Op::GetAttrMap("FTVMAlterOpLayout"); - Op op = Downcast(ref_call->op); - - Expr new_e; - bool modified = false; - if (falter_layout.count(op)) { - tvm::Array tinfos; - for (auto expr : ref_call->args) { - auto ttype = expr->type_as(); - tinfos.push_back(tvm::te::placeholder(ttype->shape, ttype->dtype)); - } - // TODO(@kevinthesun, @icemelon9): This won't work if inputs/outputs are dynamic shapes. - // Probably we need to disable the AlterOpLayout when compiling dynamic models. - Expr altered_value = falter_layout[op](new_attrs, new_args, tinfos, ref_call->checked_type()); - if (altered_value.defined()) { - new_e = altered_value; - modified = true; - } - } - if (!modified) { - new_e = Call(ref_call->op, new_args, new_attrs); - } - - const CallNode* new_call = new_e.as(); - ICHECK(new_call) << "Can only replace the original operator with another call node"; - return GetRef(new_call); - } - - Call CallWithNewLayouts(const Call& ref_call, const std::vector& new_args) override { - return CallWithNewLayouts(ref_call, ref_call->attrs, new_args); - } -}; - -/*! - * \brief Container that provides the transformation function for alter layout.. - */ -class AlterTransformMemorizer : public TransformMemorizer { - public: - AlterTransformMemorizer() = default; - explicit AlterTransformMemorizer(ObjectPtr n) : TransformMemorizer(n) {} - - AlterTransformMemorizerNode* operator->() { - return static_cast(get_mutable()); - } - - using ContainerType = AlterTransformMemorizerNode; -}; - -/*! - * Limitations: - * 1. The altered op should have the same number of arguments as the previous one. - * 2. Do not support nested tuple arguments. - */ -Expr AlterOpLayout(const Expr& expr) { - // TODO(@icemelon9): need to rerun type inference after applying an alter op. - AlterTransformMemorizer alter_memorizer(make_object()); - std::function fcontext = [=](const Call& call) -> ObjectRef { - return alter_memorizer; - }; - FForwardRewrite rewrite_func = LayoutRewriter; - return ForwardRewrite(expr, rewrite_func, fcontext); -} - -} // namespace alter_op_layout - -namespace transform { - -Pass AlterOpLayout() { - runtime::TypedPackedFunc pass_func = - [=](Function f, IRModule m, PassContext pc) { - return Downcast(relay::alter_op_layout::AlterOpLayout(f)); - }; - return CreateFunctionPass(pass_func, 3, "AlterOpLayout", {"InferType"}); -} - -TVM_REGISTER_GLOBAL("relay._transform.AlterOpLayout").set_body_typed(AlterOpLayout); - -} // namespace transform - -} // namespace relay -} // namespace tvm diff --git a/src/relay/transforms/annotate_target.cc b/src/relay/transforms/annotate_target.cc deleted file mode 100644 index eb6f9ec00432..000000000000 --- a/src/relay/transforms/annotate_target.cc +++ /dev/null @@ -1,446 +0,0 @@ -/* - * 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 src/relay/transforms/annotate_target.cc - * \brief Wraps an expr with compiler_begin and compiler_end to indicate that - * this expr should be handled by the external compiler. - */ - -#include -#include -#include -#include - -#include "pass_utils.h" - -namespace tvm { -namespace relay { -namespace annotate_target { - -static const PackedFunc* make_begin_op = - runtime::Registry::Get("relay.op.annotation._make.compiler_begin"); -static const PackedFunc* make_end_op = - runtime::Registry::Get("relay.op.annotation._make.compiler_end"); -static const char default_target[] = "default"; -// A helper class to insert annotation boundaries for all the ops of a program -// region that will be handled by a specific compiler. -class AnnotateTargetRewriter : public ExprRewriter { - public: - explicit AnnotateTargetRewriter(Array targets) : targets_(std::move(targets)) {} - - protected: - /*! \brief The target backends for annotation. */ - Array targets_; - /*! \brief Maintain the decision of the target for each op expr. */ - std::unordered_map op_expr_to_target_; - - /*! - * \brief This function annotates a compiler end and a compiler begin to all arguments. - * - * The compiler end is based on the arg target while the compiler begin is based on the given - * target. If target is not given and all arguments are going to the same target, then we will - * use that target; otherwise we use default for this op. Note that all arg exprs must be - * available in op_expr_to_target before calling this function. - * - * \param args An array of arguments of the given node. - * \param target The target of the current node. - * \return A pair of target and annotated argument expressions. - */ - std::pair> AnnotateArgs(const Array& args, - const std::string& target = "") { - std::string ref_target = ""; - Array compiler_begins; - Array compiler_ends; - for (auto arg : args) { - std::string arg_target = default_target; - const CallNode* call = arg.as(); - - if (call && call->op == CompilerBeginOp()) { - // Argument is already compiler begin node meaning that this is not the first time - // running this pass, so we simply remove it and will add a new one later. - ICHECK_EQ(call->args.size(), 1U); - // Do not alter existing annotation if not default - if (default_target != call->attrs.as()->compiler) { - compiler_begins.push_back(arg); - } else { - // Remove default - compiler_ends.push_back(call->args[0]); - } - const CallNode* end = call->args[0].as(); - if (end && end->op == CompilerEndOp()) { - arg_target = end->attrs.as()->compiler; - } - } else if (op_expr_to_target_.find(arg) != op_expr_to_target_.end()) { - arg_target = op_expr_to_target_[arg]; - // If an argument is a call node and has no argument, then it should be tensor ops such as - // zeros, so we treat it as input vars. - if (call && call->args.size() == 0) { - compiler_ends.push_back(arg); - } else { - compiler_ends.push_back(InsertAnnotation(arg, arg_target, make_end_op)); - } - } else { - // Input vars. - compiler_ends.push_back(arg); - } - - // Maintain reference target in case the target of the current node is unassigned. - if (ref_target == "") { - ref_target = arg_target; - } else if (ref_target != arg_target) { - ref_target = default_target; - } - } - - // Determine compiler begin target. - std::string op_target = (target == "") ? ref_target : target; - - if (ref_target != "") { - for (const auto& end : compiler_ends) { - compiler_begins.push_back(InsertAnnotation(end, op_target, make_begin_op)); - } - } else { - return {op_target, args}; - } - return {op_target, compiler_begins}; - } - - Expr InsertAnnotation(const Expr& expr, const std::string& target, const PackedFunc* ann_op) { - Expr new_op = (*ann_op)(expr, target); - new_op->checked_type_ = expr->checked_type_; - return new_op; - } - - Expr InsertCompilerEndAndPropogateTarget(const Expr& expr) { - /*! - * \brief This function inserts compiler end to expr and maps the corresponding target to the - * new expression. - * - * This function checks for expr existence within the map and inserts the annotation. - * If the expression has a free variable (e.g: relay.zeros, relay.ones) we do not insert - * compiler end, since there are no compiler begins for it. - * Further, it propagates the target to the new expression and returns it - * - * \param expr A relay expression - * \return An annotated and target-propagated relay expression. - */ - Expr new_expr = expr; - const CallNode* call = expr.as(); - const TupleNode* tup = expr.as(); - if (op_expr_to_target_.find(expr) != op_expr_to_target_.end()) { - // Check whether expr has args, if not - do not insert compiler_end. - if (expr->IsInstance() || expr->IsInstance() || - expr->IsInstance() || expr->IsInstance() || - (call && !call->args.empty()) || (tup && !tup->fields.empty())) { - std::string target = op_expr_to_target_[new_expr]; - new_expr = InsertAnnotation(new_expr, target, make_end_op); - op_expr_to_target_[new_expr] = target; - } - } else if (call && call->op == CompilerEndOp()) { - if (default_target == call->attrs.as()->compiler) { - ICHECK_EQ(call->args.size(), 1U); - new_expr = call->args[0]; - std::string target = op_expr_to_target_[new_expr]; - new_expr = InsertAnnotation(new_expr, target, make_end_op); - op_expr_to_target_[new_expr] = target; - } - } - - return std::move(new_expr); - } - - public: - Expr Rewrite_(const CallNode* pre, const Expr& post) override { - // Supported targets for this node. The order implies the priority. - std::vector supported_targets; - - auto op_node = pre->op.as(); - - // This graph has annotations, meaning that this is not the first time running this pass. - if (op_node && pre->op == CompilerBeginOp()) { - // Bypass compiler begin due to lack of target information. It will be processed - // when the following op handling arguments. - ICHECK_EQ(pre->args.size(), 1U); - // Preserve annotations - return post; - } else if (op_node && pre->op == CompilerEndOp()) { - // Override compiler end with the new target. - ICHECK_EQ(pre->args.size(), 1U); - auto input_expr = post.as()->args[0]; - // Already annotated. Recover target - if (op_expr_to_target_.find(input_expr) == op_expr_to_target_.end()) { - op_expr_to_target_[input_expr] = post.as()->attrs.as()->compiler; - } - ICHECK(op_expr_to_target_.find(input_expr) != op_expr_to_target_.end()); - // Preserve annotated nodes - return post; - } - // Check prior to peeking first argument - if (pre->args.size()) { - // Peek the first argument. If it is compiler begin then this node had annotated by - // another target before, so we also consider that target as a supported target. - const CallNode* first_arg_call = pre->args[0].as(); - if (first_arg_call && first_arg_call->op == CompilerBeginOp()) { - std::string arg_target = first_arg_call->attrs.as()->compiler; - if (arg_target != default_target) { - // annotated already - return post; - } - } - } - - // Check which targets this op can be offloaded. - if (op_node) { - // TVM operators: Check target specific op checking function and add to supported_targets - // if it is supported. - Op op = Downcast(pre->op); - ICHECK(op.defined()); - for (const auto& target : this->targets_) { - if (!Op::HasAttrMap("target." + std::string(target))) { - continue; - } - auto fannotate = Op::GetAttrMap("target." + std::string(target)); - const Expr& ex = GetRef(pre); - if (fannotate.count(op) && fannotate[op](ex)) { - supported_targets.push_back(target); - } - } - } else if (pre->op->IsInstance()) { - // Composite function: Add the target of a composite function to supported_targets - // if it is in the target list. - Function func = Downcast(pre->op); - ICHECK(func.defined()); - if (auto comp_name = func->GetAttr(attr::kComposite)) { - std::string comp_name_str = comp_name.value(); - size_t i = comp_name_str.find('.'); - if (i != std::string::npos) { - std::string comp_target = comp_name_str.substr(0, i); - for (const auto& target : this->targets_) { - if (std::string(target) == comp_target) { - supported_targets.push_back(comp_target); - break; - } - } - } - } - } - supported_targets.push_back(default_target); // Make default as the last option. - // Visit and mutate arguments after the target of this op has been determined. - Call post_call = Downcast(post); - if (pre->op->IsInstance()) { - auto new_call = RewriteVarCall(post_call); - if (nullptr != new_call) return GetRef(new_call->get()); - } - // TODO(@comaniac, @zhiics): Now we simply assign this node to the target with - // the highest priority, but we should preserve all supported targets so that - // we can make a better decision. - std::string target = supported_targets[0]; - - // Add annotations to each arg. - auto target_n_args = AnnotateArgs(post_call->args, target); - Array compiler_begins = std::get<1>(target_n_args); - Call new_call = Call(post_call->op, compiler_begins, post_call->attrs); - new_call->checked_type_ = pre->checked_type_; - new_call->span = pre->span; - - // Update the target map. - op_expr_to_target_[new_call] = target; - return std::move(new_call); - } - - virtual std::unique_ptr RewriteVarCall(const Call& post_call) { return nullptr; } - - Expr Rewrite_(const TupleNode* tuple_node, const Expr& post) override { - auto tuple = Downcast(post); - - auto target_n_args = AnnotateArgs(tuple->fields); - auto new_expr = WithFields(tuple, std::get<1>(target_n_args)); - op_expr_to_target_[new_expr] = std::get<0>(target_n_args); - return std::move(new_expr); - } - - Expr Rewrite_(const TupleGetItemNode* op, const Expr& post) override { - auto expr = Downcast(post); - - auto target_n_args = AnnotateArgs(Array({expr->tuple})); - auto new_expr = TupleGetItem(std::get<1>(target_n_args)[0], expr->index); - op_expr_to_target_[new_expr] = std::get<0>(target_n_args); - return std::move(new_expr); - } - - Expr Rewrite_(const FunctionNode* fn, const Expr& post) override { - Function func; - Expr new_body; - // don't step into composite functions - if (fn->GetAttr(attr::kComposite).defined()) { - func = GetRef(fn); - new_body = func->body; - } else { - func = Downcast(post); - new_body = InsertCompilerEndAndPropogateTarget(func->body); - } - return WithFields(func, func->params, new_body); - } - - Expr Rewrite_(const LetNode* op, const Expr& post) override { - auto let = Downcast(post); - - Expr new_expr; - std::pair> target_n_args; - Expr new_body = InsertCompilerEndAndPropogateTarget(let->body); - // Do not annotate function literal with let binding. - if (let->value->IsInstance()) { - new_expr = Let(let->var, let->value, new_body); - } else { - target_n_args = AnnotateArgs({let->value}); - new_expr = Let(let->var, std::get<1>(target_n_args)[0], new_body); - } - - return std::move(new_expr); - } - - Expr Rewrite_(const IfNode* op, const Expr& post) override { - auto expr = Downcast(post); - Expr new_cond = InsertCompilerEndAndPropogateTarget(expr->cond); - Expr new_true_branch = InsertCompilerEndAndPropogateTarget(expr->true_branch); - Expr new_false_branch = InsertCompilerEndAndPropogateTarget(expr->false_branch); - - auto new_expr = If(new_cond, new_true_branch, new_false_branch); - return std::move(new_expr); - } - - Expr Rewrite_(const RefCreateNode* op, const Expr& post) override { - auto expr = Downcast(post); - - auto target_n_args = AnnotateArgs(Array({expr->value})); - auto new_expr = RefCreate(std::get<1>(target_n_args)[0]); - op_expr_to_target_[new_expr] = std::get<0>(target_n_args); - return std::move(new_expr); - } - - Expr Rewrite_(const RefReadNode* op, const Expr& post) override { - auto expr = Downcast(post); - - auto target_n_args = AnnotateArgs(Array({expr->ref})); - auto new_expr = RefRead(std::get<1>(target_n_args)[0]); - op_expr_to_target_[new_expr] = std::get<0>(target_n_args); - return std::move(new_expr); - } - - Expr Rewrite_(const RefWriteNode* op, const Expr& post) override { - auto expr = Downcast(post); - - auto target_n_args = AnnotateArgs(Array({expr->ref, expr->value})); - auto new_expr = RefWrite(std::get<1>(target_n_args)[0], std::get<1>(target_n_args)[1]); - op_expr_to_target_[new_expr] = std::get<0>(target_n_args); - return std::move(new_expr); - } -}; - -// A helper class to insert annotation boundaries for call ops and function nodes -// in a program region that will be handled by a specific compiler. -class CallOpsTargetRewriter : public AnnotateTargetRewriter { - public: - explicit CallOpsTargetRewriter(Array targets) - : AnnotateTargetRewriter(std::move(targets)) {} - - std::unique_ptr RewriteVarCall(const Call& post_call) override { - Array ends; - for (auto arg : post_call->args) { - ends.push_back(InsertCompilerEndAndPropogateTarget(arg)); - } - auto new_call = std::make_unique(post_call->op, ends, post_call->attrs); - (*new_call)->checked_type_ = post_call->checked_type_; - return new_call; - } - - Expr Rewrite_(const TupleNode* tuple_node, const Expr& post) override { - auto tuple = Downcast(post); - Array new_fields; - new_fields.reserve(tuple->fields.size()); - - for (auto f : tuple->fields) { - new_fields.push_back(InsertCompilerEndAndPropogateTarget(f)); - } - return WithFields(tuple, new_fields); - } - - Expr Rewrite_(const TupleGetItemNode* op, const Expr& post) override { - auto expr = Downcast(post); - return std::move(TupleGetItem(InsertCompilerEndAndPropogateTarget(expr->tuple), expr->index)); - } - - Expr Rewrite_(const IfNode* op, const Expr& post) override { - auto expr = Downcast(post); - Expr new_cond = InsertCompilerEndAndPropogateTarget(expr->cond); - Expr new_true_branch = InsertCompilerEndAndPropogateTarget(expr->true_branch); - Expr new_false_branch = InsertCompilerEndAndPropogateTarget(expr->false_branch); - - auto new_expr = If(new_cond, new_true_branch, new_false_branch); - return std::move(new_expr); - } - - Expr Rewrite_(const RefCreateNode* op, const Expr& post) override { - auto expr = Downcast(post); - auto new_expr = RefCreate(InsertCompilerEndAndPropogateTarget(expr->value)); - return std::move(new_expr); - } - - Expr Rewrite_(const RefReadNode* op, const Expr& post) override { - auto expr = Downcast(post); - auto new_expr = RefRead(InsertCompilerEndAndPropogateTarget(expr->ref)); - return std::move(new_expr); - } - - Expr Rewrite_(const RefWriteNode* op, const Expr& post) override { - auto expr = Downcast(post); - auto new_expr = RefWrite(InsertCompilerEndAndPropogateTarget(expr->ref), - InsertCompilerEndAndPropogateTarget(expr->value)); - return std::move(new_expr); - } -}; - -Expr AnnotateTarget(const Expr& expr, const Array& targets, - bool include_non_call_ops) { - auto r = include_non_call_ops ? std::make_unique(targets) - : std::make_unique(targets); - return PostOrderRewrite(expr, r.get()); -} - -} // namespace annotate_target - -namespace transform { - -Pass AnnotateTarget(const Array& targets, bool include_non_call_ops) { - runtime::TypedPackedFunc pass_func = - [=](Function f, IRModule m, PassContext pc) { - return Downcast( - relay::annotate_target::AnnotateTarget(f, targets, include_non_call_ops)); - }; - auto func_pass = CreateFunctionPass(pass_func, 0, "AnnotateTargetFunc", {"InferType"}); - return transform::Sequential({func_pass, InferType()}, "AnnotateTarget"); -} - -TVM_REGISTER_GLOBAL("relay._transform.AnnotateTarget").set_body_typed(AnnotateTarget); - -} // namespace transform - -} // namespace relay -} // namespace tvm diff --git a/src/relay/transforms/annotate_texture_storage.cc b/src/relay/transforms/annotate_texture_storage.cc deleted file mode 100644 index 9ccb2171d8e9..000000000000 --- a/src/relay/transforms/annotate_texture_storage.cc +++ /dev/null @@ -1,686 +0,0 @@ -/* - * 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 annotate_texture_storage.cc - * \brief Collection of target specific relay passes which - * storage scope related information. - * - * - CollectStorageInfo returns a mapping from relay expr - * to a map of storage scopes for each call argument. - * These scopes are used during memory planning as well - * as downstream when doing codegen and in the graph runtime when doing runtime dataspace - * allocations. - * - * - AnnotateMemoryScope calls *target.CollectStorageInfo for all target been represented - * in the graph and rewrites graph modifying or inserting of VirtualDevice with required - * memory_scope collected from the CollectStorageInfo - */ - -#include -#include -#include -#include -#include - -#include -#include -#include - -#include "../op/memory/device_copy.h" -#include "../op/memory/memory.h" -#include "../transforms/device_aware_visitors.h" - -namespace tvm { -namespace relay { -namespace { - -/** - * @brief Analyzes the graph and returns mapping of expressions vs desired memory scope - */ -class StorageInfo : private transform::DeviceAwareExprVisitor { - public: - StorageInfo() : transform::DeviceAwareExprVisitor(Optional()) {} - - static Map>> GetStorageMap(const Expr& expr) { - StorageInfo storage_info; - storage_info.VisitExpr(expr); - storage_info.LegalizeProducerStorage(); - Map>> storage_map = storage_info.accept_textures_; - for (auto& kv : storage_info.storage_scope_) { - std::vector storage_scopes; - std::copy(kv.second.begin(), kv.second.end(), std::back_inserter(storage_scopes)); - Map> ent; - ent.Set(Expr(), Array{storage_scopes}); - storage_map.Set(GetRef(kv.first), ent); - } - - // Filling the input arguments by "global" scope to handle PlanDevice algo which propagates - // virtual devices from outputs to inputs. At the same time outputs must be unconstrained - // to avoid useless device_copy - for (const auto& cs : storage_info.consumer_storage_scopes_) { - // we have record in consumers that mean that potentially consumer - // dealt with textures anyhow, it's safe to mark this expr as global scope - // even without verification of the consumer's outputs scope - if (storage_info.CanConsumeTextures(cs.second) && - storage_map.find(GetRef(cs.first)) == storage_map.end()) { - Map> ent; - ent.Set(Expr(), Array{"global"}); - storage_map.Set(GetRef(cs.first), ent); - } - } - - // initial algo assumes mapping of outputs of the expr that is not enough, need to update - // VirtualDevice for function variables to get proper codegen. Adding vars to storage_map - for (const auto& a : storage_info.args_to_vars_) { - if (storage_map.count(a.first)) { - for (const auto& v : a.second) { - if (storage_info.buffers_params.find(v) != storage_info.buffers_params.end()) { - Map> ent; - ent.Set(Expr(), Array{"global"}); - storage_map.Set(v, ent); - } else { - storage_map.Set(v, storage_map[a.first]); - if (storage_map[a.first][Expr()][0] == "global" && - storage_info.accept_textures_.count(v)) { - Map> ent; - ent.Set(Expr(), storage_info.accept_textures_[v][Expr()]); - storage_map.Set(v, ent); - for (const auto& calls : storage_info.accept_textures_[v]) { - if (calls.first != Expr()) { - if (storage_map.count(a.first)) { - Map> ent_call = storage_map[a.first]; - ent_call.Set(calls.first, calls.second); - storage_map.Set(a.first, ent_call); - } else { - Map> ent_call; - ent_call.Set(calls.first, calls.second); - storage_map.Set(a.first, ent_call); - } - } - } - } - } - } - } - } - return storage_map; - } - - private: - using transform::DeviceAwareExprVisitor::VisitExpr_; - - void Visit(const Expr& expr) { - // Pre-order traversal to enable upward propagation - // of consumer storage scopes to producers when desirable. - if (const auto* fn = expr.as()) { - this->VisitExpr(fn->body); - for (const auto& param : fn->params) { - this->VisitExpr(param); - } - } else { - this->VisitExpr(expr); - } - } - - void VisitExpr_(const VarNode* vn) final { ApplyConsumerScopeToInputs(vn); } - - void VisitExpr_(const ConstantNode* cn) final { ApplyConsumerScopeToInputs(cn); } - - void DeviceAwareVisitExpr_(const CallNode* call) final { - // Check the contents of this primitive function - if (const auto* fn = call->op.as()) { - if (fn->HasNonzeroAttr(attr::kPrimitive)) { - primitive_supports_texture_ = false; - Visit(call->op); - if (primitive_supports_texture_) { - if (call->checked_type().as()) { - std::string scope = "global.texture"; - if (const auto* ttype = call->checked_type().as()) { - scope = Scope(ttype->shape, GetVirtualDevice(GetRef(call))); - } - storage_scope_[call].push_back(scope); - } else { - const auto* tuple_type = call->type_as(); - ICHECK(tuple_type); - // TODO(csullivan): Add support for mixed output storage scope. - // In current adreno storage planner all outputs of a - // primitive function are assumed to be of the same storage - // type. This should be easy to extend in the future. - for (size_t i = 0; i < tuple_type->fields.size(); i++) { - storage_scope_[call].push_back("global.texture"); - } - } - const int weights_pos = 1; - for (size_t i = 0; i < fn->params.size(); i++) { - args_to_vars_[call->args[i]].push_back(fn->params[i]); - // adding info about arguments if they can be converted to texture - for (const auto& ttype : FlattenTupleType(fn->params[i]->checked_type())) { - std::string scope = Scope(ttype->shape, GetVirtualDevice(GetRef(call))); - if (expr_attrib.as() || expr_attrib.as()) { - String kernel_layout = expr_attrib.as() - ? expr_attrib.as()->kernel_layout - : expr_attrib.as()->kernel_layout; - if ((i == weights_pos) && !ttype->dtype.is_float16() && - CanUseBuffers(call->args[i], ttype->shape, kernel_layout)) { - buffers_params.insert(fn->params[i]); - buffers_args.insert(call->args[i]); - scope = "global"; - } - } - if (scope.find("global.texture") != std::string::npos) { - if (accept_textures_.count(fn->params[i])) { - Map> ent = accept_textures_[fn->params[i]]; - ent.Set(GetRef(call), Array{scope}); - ent.Set(Expr(), Array{scope}); - accept_textures_.Set(fn->params[i], ent); - } else { - Map> ent; - ent.Set(GetRef(call), Array{scope}); - ent.Set(Expr(), Array{scope}); - accept_textures_.Set(fn->params[i], ent); - } - } - } - } - } - // Add consumer storage scope information for call arguments - for (auto& arg : call->args) { - if (storage_scope_.count(call)) { - ICHECK(!HasMixedStorageOutputs(call)) - << "Mixed output storage scopes are not currently supported"; - consumer_storage_scopes_[arg.operator->()].push_back("global.texture"); - } else { - consumer_storage_scopes_[arg.operator->()].push_back("global"); - } - } - } - } - if (!primitive_supports_texture_) { - expr_attrib = call->attrs; - primitive_supports_texture_ = SupportsTextureStorage(call); - } - - for (auto& arg : call->args) { - if (buffers_args.find(arg) == buffers_args.end()) { - Visit(arg); - } - } - // We have all callees filled into storage_scope_ if they support textures - // We need to verify if this call expects texture and if it does not, remove from - // storage_scope_ since initially storage_scope_ is filled only based on knowledge - // that function able to work with textures, but not necessary that this texture is - // expected by function callee - for (auto& arg : call->args) { - if (consumer_storage_scopes_.count(arg.operator->()) && - GetConsumerScope(consumer_storage_scopes_[arg.operator->()]) != "global.texture") { - storage_scope_.erase(arg.operator->()); - } - } - } - - /** - * Defines the name of the memory scope which can fit the tensor of required shape - * - * The scope stands for "global" if tensor does not satisfy current flattening rules for textures - * (texture currently has to be 5d tensors with value eq 4 in the last dimension) - * - * The packing layout inside the texture scope (the part after the dash) is defined - * during the shape itself. Hardware can have limitations on the texture spatial dimensions - * we must not exceed these sizes. In addition to the fitting of h/w limitation we want to - * get balanced packing where final spatial sizes of textures will not be too different - * @param shape shape to be analyzed - * @param vd VirtualDevice for the tensors determined of memory scope - * @return string representing memory scope either "global" or "global.texture-layout" - */ - std::string Scope(Array shape, const VirtualDevice& vd) { - // currently we support only textures been made from 5d tensors - // 5d requirement is not limitation of textures in general, it is limitation how - // we are representing memory scopes/layout and flattening of textures in tir - if (vd != VirtualDevice::FullyUnconstrained() && shape.size() == 5 && - shape[4].as()->value == 4) { - std::map diffs; - int limit = - vd->target->GetAttr("texture_spatial_limit").value_or(Integer(16384))->value; - int a0 = shape[0].as()->value; - int a1 = shape[1].as()->value; - int a2 = shape[2].as()->value; - int a3 = shape[3].as()->value; - - int d3l = a0 * a1 * a2; - int d3r = a3; - int diff3 = d3l > d3r ? d3l - d3r : d3r - d3l; - if (d3l < limit && d3r < limit) diffs[diff3] = ""; - - int d2l = a0 * a1; - int d2r = a2 * a3; - int diff2 = d2l > d2r ? d2l - d2r : d2r - d2l; - if (d2l < limit && d2r < limit) diffs[diff2] = "nhwc"; - - int d1l = a0; - int d1r = a1 * a2 * a3; - int diff1 = d1l > d1r ? d1l - d1r : d1r - d1l; - if (d1l < limit && d1r < limit) diffs[diff1] = "weight"; - if (!diffs.empty()) { - std::string scope = "global.texture"; - if (!diffs.begin()->second.empty()) { - scope += ("-" + diffs.begin()->second); - } - return scope; - } - } - return "global"; - } - - void ApplyConsumerScopeToInputs(const ExprNode* expr) { - std::string scope; - auto consumer_scopes_it = consumer_storage_scopes_.find(expr); - if (consumer_scopes_it != consumer_storage_scopes_.end()) { - std::string consumer_scope = GetConsumerScope(consumer_scopes_it->second); - ICHECK(!storage_scope_.count(expr)) - << "Already propagated consumer scopes to input: " << GetRef(expr); - - bool expr_is_rgba_vectorizable = false; - if (const auto* ttype = expr->checked_type().as()) { - scope = Scope(ttype->shape, GetVirtualDevice(GetRef(expr))); - if (scope != "global") { - auto inner_dim = ttype->shape.back().as(); - if (inner_dim && inner_dim->value == 4) { - expr_is_rgba_vectorizable = true; - } - } - } - - // Only propagate texture scope from consumers to input expr if - // the input shape of the input expr is rgba vectorizable. - if (consumer_scope.find("global.texture") != std::string::npos) { - if (expr_is_rgba_vectorizable) { - storage_scope_[expr].push_back(scope); - } - } else { - storage_scope_[expr].push_back(consumer_scope); - } - } - } - - void LegalizeProducerStorage() { - for (auto& kv : consumer_storage_scopes_) { - const ExprNode* producer = kv.first; - std::string legal_scope = GetConsumerScope(kv.second); - if (storage_scope_.count(producer)) { - ICHECK(!HasMixedStorageOutputs(producer)) - << "Mixed output storage scopes are not currently supported"; - if (storage_scope_[producer][0].find(legal_scope) == std::string::npos) { - for (size_t i = 0; i < storage_scope_[producer].size(); i++) { - // Only support uniform storage scope across all outputs for now - storage_scope_[producer][i] = legal_scope; - } - } - } - } - } - - std::string GetConsumerScope(const std::vector& consumer_scopes) const { - if (!consumer_scopes.size()) { - return "global"; - } - std::string texture_tag = "global.texture"; - for (auto& consumer_scope : consumer_scopes) { - if (consumer_scope.find(texture_tag) == std::string::npos) { - return "global"; - } - } - return texture_tag; - } - - bool CanConsumeTextures(const std::vector& consumer_scopes) const { - std::string texture_tag = "global.texture"; - for (auto& consumer_scope : consumer_scopes) { - if (consumer_scope.find(texture_tag) == 0) { - return true; - } - } - return false; - } - - bool HasMixedStorageOutputs(const ExprNode* expr) { - if (storage_scope_.count(expr)) { - std::string ref_scope = storage_scope_[expr][0]; - for (std::string& scope : storage_scope_[expr]) { - if (scope != ref_scope) { - return true; - } - } - } - return false; - } - - bool SupportsTextureStorage(const CallNode* call) const { - bool supports_texture_storage = false; - // we need to verify only entry functions since one of entry op defines main schedule - for (const auto& arg : call->args) { - if (!arg.as()) { - return false; - } - } - if (auto attrs = call->attrs.as()) { - if (attrs->data_layout == "NCHW4c" && attrs->kernel_layout == "OIHW4o") { - supports_texture_storage = true; - } else if (attrs->data_layout == "NHWC4c" && - (attrs->kernel_layout == "HWOI4o" || attrs->kernel_layout == "HWIO4o" || - attrs->kernel_layout == "OIHW4o")) { - supports_texture_storage = true; - } - } else if (auto attrs = call->attrs.as()) { - if ((attrs->data_layout == "NCHW4c" || attrs->data_layout == "NHWC4c") && - (attrs->kernel_layout == "OIHW4o" || attrs->kernel_layout == "HWIO4o")) { - supports_texture_storage = true; - } - } else if (auto attrs = call->attrs.as()) { - if (attrs->data_layout == "NCHW4c" && attrs->kernel_layout == "IOHW4o") { - supports_texture_storage = true; - } - } else if (auto attrs = call->attrs.as()) { - if (attrs->layout == "NCHW4c") { - supports_texture_storage = true; - } - } else if (auto attrs = call->attrs.as()) { - if (attrs->layout == "NCHW4c") { - supports_texture_storage = true; - } - } else if (auto attrs = call->attrs.as()) { - if (attrs->layout == "NCHW4c") { - supports_texture_storage = true; - } - } else if (const OpNode* opnode = call->op.as()) { - auto fpattern = Op::GetAttrMap("TOpPattern"); - auto pattern = fpattern[GetRef(opnode)]; - if (pattern <= kCommReduce) { - if (const auto* ttype = call->checked_type().as()) { - if (ttype->shape.size() == 5) { - auto node0 = ttype->shape[0].as(); - auto node1 = ttype->shape[1].as(); - auto node2 = ttype->shape[2].as(); - auto node3 = ttype->shape[3].as(); - auto node4 = ttype->shape[4].as(); - // if tensor has any dimension then textures are not supported - if (!node0 || !node1 || !node2 || !node3 || !node4) { - return false; - } - supports_texture_storage = true; - } - } - } - } - - return supports_texture_storage; - } - - bool CanUseBuffers(const Expr param, const Array shape, - const String kernel_layout) const { - bool use_buffer = false; - if (param.as() && shape.size() == 5) { - if (kernel_layout == "HWOI4o" || kernel_layout == "HWIO4o") { - int a0 = shape[0].as()->value; - int a1 = shape[1].as()->value; - if (a0 != 1 && a1 != 1) { - use_buffer = true; - } - } else if (kernel_layout == "OIHW4o") { - int a2 = shape[2].as()->value; - int a3 = shape[3].as()->value; - if (a2 != 1 && a3 != 1) { - use_buffer = true; - } - } - } - return use_buffer; - } - - /*! \brief Temporary state for marking whether a visited function - * primitive supports texture storage scope */ - bool primitive_supports_texture_ = false; - /*! \brief expr storage scope mapping for each output */ - std::unordered_map> storage_scope_; - /*! \brief output storage scopes used by consumers of expr key */ - std::unordered_map> consumer_storage_scopes_; - /*! \brief mapping of arguments to call to function variables*/ - std::unordered_map, ObjectPtrHash, ObjectPtrEqual> args_to_vars_; - /*! \brief mapping of arguments that can be converted to texture*/ - Map>> accept_textures_; - /*! \brief main attribute for expression*/ - tvm::Attrs expr_attrib; - /*! \brief parameters that filter out from storage_map to use buffers*/ - std::unordered_set buffers_params; - /*! \brief arguments in expression that will use buffers*/ - std::unordered_set buffers_args; -}; - -} // namespace - -/** - * @brief rewrite of virtual devices, memory_scope part for expressions defined - * by the StorageInfo analysis pass - * - * Currently this workflow supports analysis and rewriting of VirtualDevice for - * Constants and function Variables - */ -class RewriteVDStorageScopes : public transform::DeviceAwareExprMutator { - using VarMap = std::unordered_map; - - public: - using transform::DeviceAwareExprMutator::VisitExpr_; - - explicit RewriteVDStorageScopes(const Map>>& storage_scope) - : transform::DeviceAwareExprMutator(Optional()), storage_scope_(storage_scope) {} - - Function Rewrite(const Expr& expr) { return Downcast(Mutate(expr)); } - - Expr VisitExpr_(const VarNode* vn) final { - if (storage_scope_.find(GetRef(vn)) != storage_scope_.end() && - storage_scope_[GetRef(vn)].find(Expr()) != storage_scope_[GetRef(vn)].end() && - storage_scope_[GetRef(vn)][Expr()][0] != "global") { - Var c = Var(vn->vid, vn->type_annotation, vn->span); - auto virtual_device = GetVirtualDevice(GetRef(vn)); - c->virtual_device_ = - VirtualDevice(virtual_device->device_type(), virtual_device->virtual_device_id, - virtual_device->target, storage_scope_[GetRef(vn)][Expr()][0]); - return std::move(c); - } - return GetRef(vn); - } - - Expr VisitExpr_(const ConstantNode* vn) final { - if (storage_scope_.find(GetRef(vn)) != storage_scope_.end() && - storage_scope_[GetRef(vn)].find(Expr()) != storage_scope_[GetRef(vn)].end()) { - Expr c = Constant(vn->data, vn->span); - auto virtual_device = GetVirtualDevice(GetRef(vn)); - c = OnDevice( - c, - VirtualDevice(virtual_device->device_type(), virtual_device->virtual_device_id, - virtual_device->target, storage_scope_[GetRef(vn)][Expr()][0]), - true); - return c; - } - return GetRef(vn); - } - - Expr DeviceAwareVisitExpr_(const CallNode* call_node) final { - // we need to duplicate ExprMutator::VisitExpr_ to correct argument scopes and - // put device_copy - auto new_op = this->Mutate(call_node->op); - - tvm::Array ty_args; - ty_args.reserve(call_node->type_args.size()); - - for (auto ty_arg : call_node->type_args) { - auto new_ty_arg = this->VisitType(ty_arg); - ty_args.push_back(new_ty_arg); - } - - tvm::Array call_args; - call_args.reserve(call_node->args.size()); - for (auto arg : call_node->args) { - auto new_arg = this->Mutate(arg); - // verification if we need to put device_copy - if (storage_scope_.count(arg) && storage_scope_[arg].count(GetRef(call_node))) { - auto virtual_device = GetVirtualDevice(GetRef(call_node)); - VirtualDevice virtual_device_from = - VirtualDevice(virtual_device->device_type(), virtual_device->virtual_device_id, - virtual_device->target, virtual_device->memory_scope); - VirtualDevice virtual_device_to = - VirtualDevice(virtual_device->device_type(), virtual_device->virtual_device_id, - virtual_device->target, storage_scope_[arg][GetRef(call_node)][0]); - new_arg = DeviceCopy(new_arg, virtual_device_from, virtual_device_to); - new_arg = OnDevice( - new_arg, - VirtualDevice(virtual_device->device_type(), virtual_device->virtual_device_id, - virtual_device->target, storage_scope_[arg][GetRef(call_node)][0]), - true); - } - call_args.push_back(new_arg); - } - - auto new_call = WithFields(GetRef(call_node), new_op, call_args, {}, ty_args); - - auto virtual_device = GetVirtualDevice(GetRef(call_node)); - std::string memory_scope = ""; - if (storage_scope_.find(GetRef(call_node)) != storage_scope_.end() && - storage_scope_[GetRef(call_node)].find(Expr()) != - storage_scope_[GetRef(call_node)].end()) { - memory_scope = storage_scope_[GetRef(call_node)][Expr()][0]; - } else if (virtual_device->memory_scope != "") { - memory_scope = virtual_device->memory_scope; - } else if (!call_node->op.as()) { - memory_scope = ""; - } - if (!memory_scope.empty()) { - new_call = - OnDevice(new_call, - VirtualDevice(virtual_device->device_type(), virtual_device->virtual_device_id, - virtual_device->target, memory_scope), - true); - } - return std::move(new_call); - } - - private: - Map>> storage_scope_; - VarMap new_vars_; - Array current_function_scope_; -}; - -Map>> CollectTextureStorage(const Expr& expr) { - return StorageInfo::GetStorageMap(expr); -} - -/** - * @brief Collects all target devices participated in graph - */ -class CollectVirtualDevices : public transform::DeviceAwareExprVisitor { - public: - CollectVirtualDevices() : transform::DeviceAwareExprVisitor(Optional()) {} - /** - * @brief Get all unique device elements from target of each VirtualDevice - * - * @param expr - IR - * @return set of devices - */ - std::set GetDevices(const Expr& expr) { - this->Run(expr); - return std::move(devices_); - } - - void Visit(const Expr& expr) { - // Pre-order traversal to enable upward propagation - // of consumer storage scopes to producers when desirable. - if (const auto* fn = expr.as()) { - this->VisitExpr(fn->body); - for (const auto& param : fn->params) { - this->VisitExpr(param); - } - } else { - this->VisitExpr(expr); - } - } - - void DeviceAwareVisitExpr_(const CallNode* call) final { - auto vd = GetVirtualDevice(GetRef(call)); - if (vd != VirtualDevice::FullyUnconstrained()) { - if (Optional t_device = vd->target->GetAttr("device")) { - devices_.insert(vd->target->kind->name + "." + t_device.value()); - } - } - for (auto& arg : call->args) { - Visit(arg); - } - } - - void Run(const Expr& expr) { VisitExpr(expr); } - using transform::DeviceAwareExprVisitor::VisitExpr_; - std::set devices_; -}; - -/*! - * \brief Collect the target specific tensor storage info for each expression's output. - * \param expr The expression. - * \return The device based storage mapping. - */ -Map>> CollectStorageInfo(const Expr& expr) { - std::set device_types = CollectVirtualDevices().GetDevices(expr); - // TODO(amalyshe): current approach collects all targets withing graph and call the only - // function corresponding to all these targets in alphabetic order - // this will work reliable only for case of only one device and should be redesigned - // to handle common case - std::string ftarget_prefix = "relay.backend"; - for (auto& dev_id : device_types) { - ftarget_prefix += (std::string(".") + dev_id); - } - - Map>> storage_info = {}; - if (const auto* f = runtime::Registry::Get(ftarget_prefix + "._CollectStorageInfo")) { - storage_info = (*f)(expr); - } - return storage_info; -} - -Expr AnnotateMemoryScopeExpr(const Expr& expr, const IRModule& mod) { - auto storage_scope = CollectStorageInfo(expr); - if (storage_scope.size()) { - return RewriteVDStorageScopes(storage_scope).Rewrite(expr); - } else { - return expr; - } -} - -namespace transform { -tvm::transform::Pass AnnotateMemoryScope() { - runtime::TypedPackedFunc pass_func = - [](Function f, IRModule m, PassContext pc) { - return Downcast(AnnotateMemoryScopeExpr(f, m)); - }; - return CreateFunctionPass(pass_func, 2, "AnnotateMemoryScope", {}); -} -} // namespace transform - -TVM_REGISTER_GLOBAL("relay.backend.opencl.adreno._CollectStorageInfo") - .set_body_typed(CollectTextureStorage); - -} // namespace relay -} // namespace tvm diff --git a/src/relay/transforms/auto_scheduler_layout_rewrite.cc b/src/relay/transforms/auto_scheduler_layout_rewrite.cc deleted file mode 100644 index 532b25769f87..000000000000 --- a/src/relay/transforms/auto_scheduler_layout_rewrite.cc +++ /dev/null @@ -1,186 +0,0 @@ -/* - * 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 auto_scheduler_layout_rewrite.h - * \brief Rewrite the layout of "layout free" tensors (e.g., the weight tensors in - * conv2d and dense layers) according to the tile structure generated by the auto-scheduler. - */ - -#include "auto_scheduler_layout_rewrite.h" - -#include -#include -#include -#include - -#include -#include -#include - -#include "../backend/te_compiler.h" -#include "pattern_utils.h" - -namespace tvm { -namespace relay { - -// Two global variables for receiving layout information from python -std::deque AutoSchedulerLayoutRewriter::global_ori_layouts_queue; -std::deque AutoSchedulerLayoutRewriter::global_new_layouts_queue; - -// Copy an Attrs but with a new auto_scheduler_rewritten_layout filed. -template -Attrs CopyAttrsWithNewLayout(const T* ptr, const std::string& layout) { - auto n = make_object(*ptr); - n->auto_scheduler_rewritten_layout = layout; - return Attrs(n); -} - -// Mutate ops in a function -class FuncMutator : public ExprMutator { - public: - FuncMutator(const std::deque& ori_layouts_queue, - const std::deque& new_layouts_queue) - : ExprMutator(), - ori_layouts_queue_(ori_layouts_queue), - new_layouts_queue_(new_layouts_queue) {} - - Expr VisitExpr_(const CallNode* n) { - auto new_n = ExprMutator::VisitExpr_(n); - - const auto* call = new_n.as(); - if (call && call->op.as() && - (std::find(target_ops_.begin(), target_ops_.end(), n->op.as()->name) != - target_ops_.end()) && - !ori_layouts_queue_.empty() && !new_layouts_queue_.empty()) { - // Pop a new layout from the queue - const std::string ori_layout = ori_layouts_queue_.front(); - const std::string new_layout = new_layouts_queue_.front(); - ori_layouts_queue_.pop_front(); - new_layouts_queue_.pop_front(); - - // Insert a new op to do layout transform. (This will be simplified by FoldConstant later). - Expr updated_kernel = MakeAutoSchedulerLayoutTransform(call->args[1], ori_layout, new_layout); - Array updated_args = {call->args[0], updated_kernel}; - - // Update the attrs - Attrs updated_attrs; - if (auto pattr = call->attrs.as()) { - updated_attrs = CopyAttrsWithNewLayout(pattr, new_layout); - } else if (auto pattr = call->attrs.as()) { - updated_attrs = CopyAttrsWithNewLayout(pattr, new_layout); - } else if (auto pattr = call->attrs.as()) { - updated_attrs = CopyAttrsWithNewLayout(pattr, new_layout); - } else if (auto pattr = call->attrs.as()) { - updated_attrs = CopyAttrsWithNewLayout(pattr, new_layout); - } else if (auto pattr = call->attrs.as()) { - updated_attrs = CopyAttrsWithNewLayout(pattr, new_layout); - } else if (auto pattr = call->attrs.as()) { - updated_attrs = CopyAttrsWithNewLayout(pattr, new_layout); - } else { - LOG(FATAL) << "Unhandled attribute: " << call->attrs; - } - new_n = Call(call->op, updated_args, updated_attrs); - } - return new_n; - } - - private: - std::deque ori_layouts_queue_; - std::deque new_layouts_queue_; - - std::vector target_ops_{ - "nn.conv2d", "nn.conv3d", "nn.contrib_conv2d_winograd_without_weight_transform", - "nn.matmul", "nn.dense", "nn.batch_matmul"}; -}; - -Expr AutoSchedulerLayoutRewriter::VisitExpr_(const CallNode* n) { - auto new_n = ExprMutator::VisitExpr_(n); - - if (const auto* call = new_n.as()) { - if (const auto* func = call->op.as()) { - global_ori_layouts_queue.clear(); - global_new_layouts_queue.clear(); - - // Use ScheduleGetter to call python lower functions. - // This is used to get the layout transform information. - // The layout transformation will be recorded to global_ori_layout_queue - // and global_new_layouts_queue in ComputeDAG::RewriteLayout. - auto f = runtime::Registry::Get("auto_scheduler.enter_layout_rewrite"); - CHECK(f) << "Could not find auto_scheduler.enter_layout_rewrite function."; - (*f)(); - - tec::PrimFuncFor(GetRef(func), Target::Current()); - - f = runtime::Registry::Get("auto_scheduler.exit_layout_rewrite"); - CHECK(f) << "Could not find ansor.exit_layout_rewrite function."; - (*f)(); - - // Mutate the called function - if (!global_ori_layouts_queue.empty() && !global_new_layouts_queue.empty()) { - auto ret = FuncMutator(global_ori_layouts_queue, global_new_layouts_queue).VisitExpr(new_n); - return ret; - } - } - } - - return new_n; -} - -Expr AutoSchedulerLayoutRewrite(const Expr& expr) { - return AutoSchedulerLayoutRewriter().Mutate(expr); -} - -namespace transform { - -Pass AutoSchedulerLayoutRewrite() { - runtime::TypedPackedFunc pass_func = - [=](Function f, IRModule m, PassContext pc) { - return Downcast(relay::AutoSchedulerLayoutRewrite(f)); - }; - return CreateFunctionPass(pass_func, 3, "AutoSchedulerLayoutRewrite", {"InferType"}); -} - -TVM_REGISTER_GLOBAL("relay._transform.AutoSchedulerLayoutRewrite") - .set_body_typed(AutoSchedulerLayoutRewrite); - -TVM_REGISTER_GLOBAL("relay.attrs.get_auto_scheduler_rewritten_layout") - .set_body_typed([](const Attrs& attrs) { - if (attrs->IsInstance()) { - return attrs.as()->auto_scheduler_rewritten_layout; - } else if (attrs->IsInstance()) { - return attrs.as()->auto_scheduler_rewritten_layout; - } else if (attrs->IsInstance()) { - return attrs.as()->auto_scheduler_rewritten_layout; - } else if (attrs->IsInstance()) { - return attrs.as()->auto_scheduler_rewritten_layout; - } else if (attrs->IsInstance()) { - return attrs.as()->auto_scheduler_rewritten_layout; - } else if (attrs->IsInstance()) { - return attrs.as()->auto_scheduler_rewritten_layout; - } else { - LOG(FATAL) << "Unhandled attribute: " << attrs; - } - return tvm::String(); - }); - -} // namespace transform - -} // namespace relay -} // namespace tvm diff --git a/src/relay/transforms/auto_scheduler_layout_rewrite.h b/src/relay/transforms/auto_scheduler_layout_rewrite.h deleted file mode 100644 index d0d89db42e68..000000000000 --- a/src/relay/transforms/auto_scheduler_layout_rewrite.h +++ /dev/null @@ -1,49 +0,0 @@ -/* - * 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 auto_scheduler_layout_rewrite.h - * \brief Rewrite the layout of "layout free" tensors (e.g., the weight tensors in - * conv2d and dense layers) according to the tile structure generated by the auto-scheduler. - */ - -#ifndef TVM_RELAY_TRANSFORMS_AUTO_SCHEDULER_LAYOUT_REWRITE_H_ -#define TVM_RELAY_TRANSFORMS_AUTO_SCHEDULER_LAYOUT_REWRITE_H_ - -#include - -#include -#include - -namespace tvm { -namespace relay { - -class AutoSchedulerLayoutRewriter : public ExprMutator { - public: - Expr VisitExpr_(const CallNode* n) final; - - // Two global variables for receiving layout information from python - static std::deque global_ori_layouts_queue; - static std::deque global_new_layouts_queue; -}; - -} // namespace relay -} // namespace tvm - -#endif // TVM_RELAY_TRANSFORMS_AUTO_SCHEDULER_LAYOUT_REWRITE_H_ diff --git a/src/relay/transforms/canonicalize_cast.cc b/src/relay/transforms/canonicalize_cast.cc deleted file mode 100644 index f268665ce212..000000000000 --- a/src/relay/transforms/canonicalize_cast.cc +++ /dev/null @@ -1,142 +0,0 @@ -/* - * 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 canonicalize_cast.cc - * \brief Canonicalize cast expressions to make operator fusion more efficient. - */ -#include -#include -#include -#include - -#include "pass_utils.h" -#include "pattern_utils.h" - -namespace tvm { -namespace relay { - -// This pass finds upcast that is referred by multiple elemwise/broadcast operators, and creates a -// copy of it in each branch such that after fusion the previous function have output with fewer -// bits. -// -// Consider the following example: -// \code -// def @main(x: int8) { -// %1 = cast(%x, f32) -// %2 = exp(%1) -// %3 = log(%1) -// (%3, 4) -// } -// \endcode -// -// We would like to prevent sharing of the cast expression such that operator fusion can produce -// more efficient result as below. -// \code -// def @main(x: int8) { -// %1 = fn (%p1: i8) { -// exp(cast(%p1, f32) -// } -// %3 = %1(%x) -// %2 = fn (%p1: i8) { -// log(cast(%p1, f32) -// } -// %4 = %2(%x) -// (%3, 4) -// } -// \endcode -class CastCanonicalizer : public ExprMutator { - public: - CastCanonicalizer() : cast_op_(Op::Get("cast")) {} - - Expr VisitExpr_(const CallNode* call) { - static auto fpattern = Op::GetAttrMap("TOpPattern"); - - if (auto call_op = call->op.as()) { - auto pattern = fpattern[call_op.value()]; - if (pattern <= kBroadcast) { - Array call_args = call->args; - bool unchanged = true; - for (size_t i = 0; i < call_args.size(); ++i) { - Expr arg = call_args[i]; - Expr new_arg = GetNewCallArg(arg); - if (!arg.same_as(new_arg)) { - call_args.Set(i, new_arg); - unchanged = false; - } - } - if (unchanged) { - return GetRef(call); - } - return Call(call->op, call_args, call->attrs, call->type_args); - } - } - - Expr new_expr = ExprMutator::VisitExpr_(call); - return new_expr; - } - - private: - std::unordered_map ref_counter_; - // cast op is frequently checked for equivalence. Therefore, we cache it to - // reduce lookup overhead. - const Op& cast_op_; - - Expr GetNewCallArg(const Expr& e) { - // if e is a upcast and ref count > 1, create an copy; otherwise call the default visitor - Expr new_expr = this->VisitExpr(e); - - if (const CallNode* call = e.as()) { - if (call->op == cast_op_) { - auto attrs = call->attrs.as(); - const auto* from_type = call->args[0]->type_as(); - ICHECK(from_type); - - if (from_type->dtype.bits() < attrs->dtype.bits()) { - if (++ref_counter_[call] > 1) { - const CallNode* new_call = new_expr.as(); - ICHECK(new_call); - ICHECK(new_call->op == cast_op_); - return Call(new_call->op, new_call->args, new_call->attrs, new_call->type_args); - } - } - } - } - return new_expr; - } -}; - -Expr CanonicalizeCast(const Expr& e) { return CastCanonicalizer().Mutate(e); } - -namespace transform { - -Pass CanonicalizeCast() { - runtime::TypedPackedFunc pass_func = - [=](Function f, IRModule m, PassContext pc) { - return Downcast(CanonicalizeCast(f)); - }; - return CreateFunctionPass(pass_func, 3, "CanonicalizeCast", {"InferType"}); -} - -TVM_REGISTER_GLOBAL("relay._transform.CanonicalizeCast").set_body_typed(CanonicalizeCast); - -} // namespace transform - -} // namespace relay -} // namespace tvm diff --git a/src/relay/transforms/canonicalize_ops.cc b/src/relay/transforms/canonicalize_ops.cc deleted file mode 100644 index cf14ddcb7c5b..000000000000 --- a/src/relay/transforms/canonicalize_ops.cc +++ /dev/null @@ -1,86 +0,0 @@ -/* - * 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 canonicalize_ops.cc - * \brief Canonicalize special operators to basic operators. - This can simplify latter analysis. (e.g. Expand bias_add to expand_dims and broadcast_add.) - */ -#include -#include -#include -#include -#include - -#include "pattern_utils.h" - -namespace tvm { -namespace relay { - -class BiasAddSimplifier : public ExprRewriter { - public: - BiasAddSimplifier() : bias_add_op_(Op::Get("nn.bias_add")) {} - - Expr Rewrite_(const CallNode* n, const Expr& post) override { - auto new_n = post; - if (n->op == bias_add_op_) { - Call call = Downcast(new_n); - ICHECK_EQ(call->args.size(), 2); - const BiasAddAttrs* param = call->attrs.as(); - - auto ttype = n->args[0]->type_as(); - size_t n_dim = ttype->shape.size(); - int axis = param->axis; - if (axis < 0) { - axis += n_dim; - } - Expr expanded_bias = ExpandBiasToMatchAxis(call->args[1], n_dim, {axis}); - Expr ret = Add(call->args[0], expanded_bias); - ret->checked_type_ = n->checked_type_; - return ret; - } - return new_n; - } - - private: - // Cache the bias_add for equivalence checking. - const Op& bias_add_op_; -}; - -Expr CanonicalizeOps(const Expr& e) { - auto rewriter = BiasAddSimplifier(); - return PostOrderRewrite(e, &rewriter); -} - -namespace transform { - -Pass CanonicalizeOps() { - runtime::TypedPackedFunc pass_func = - [=](Function f, IRModule m, PassContext pc) { - return Downcast(CanonicalizeOps(f)); - }; - return CreateFunctionPass(pass_func, 3, "CanonicalizeOps", {"InferType"}); -} - -TVM_REGISTER_GLOBAL("relay._transform.CanonicalizeOps").set_body_typed(CanonicalizeOps); - -} // namespace transform - -} // namespace relay -} // namespace tvm diff --git a/src/relay/transforms/capture_postdfsindex_in_spans.cc b/src/relay/transforms/capture_postdfsindex_in_spans.cc deleted file mode 100644 index 17c7e59c7f60..000000000000 --- a/src/relay/transforms/capture_postdfsindex_in_spans.cc +++ /dev/null @@ -1,134 +0,0 @@ -/* - * 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 tvm/relay/transform/capture_index_in_spans.cc - * \brief A pass to set spans to capture the post-dfs index of every node. - */ - -#include -#include - -#include "../ir/indexed_graph.h" - -namespace tvm { -namespace relay { -namespace transform { - -namespace { - -/*! \brief Update all the spans to capture their post-dfs index. */ -class SpansRewriter : public ExprRewriter { - public: - explicit SpansRewriter(const IndexedGraph* indexed_graph) - : source_name_(SourceName::Get("index")), indexed_graph_(indexed_graph) {} - - private: - Expr Rewrite_(const VarNode* var_node, const Expr& post) final { - return WithFields(Downcast(post), {}, {}, {}, MakeSpan(GetRef(var_node))); - } - - Expr Rewrite_(const GlobalVarNode* global_var_node, const Expr& post) final { - return WithFields(Downcast(post), {}, {}, {}, - MakeSpan(GetRef(global_var_node))); - } - - Expr Rewrite_(const ConstantNode* constant_node, const Expr& post) final { - return WithFields(Downcast(post), {}, {}, MakeSpan(GetRef(constant_node))); - } - - Expr Rewrite_(const TupleNode* tuple_node, const Expr& post) final { - return WithFields(Downcast(post), {}, {}, MakeSpan(GetRef(tuple_node))); - } - - Expr Rewrite_(const FunctionNode* function_node, const Expr& post) final { - return WithFields(Downcast(post), {}, {}, {}, {}, {}, {}, - MakeSpan(GetRef(function_node))); - } - - Expr Rewrite_(const CallNode* call_node, const Expr& post) final { - return WithFields(Downcast(post), {}, {}, {}, {}, {}, MakeSpan(GetRef(call_node))); - } - - Expr Rewrite_(const LetNode* let_node, const Expr& post) final { - return WithFields(Downcast(post), {}, {}, {}, {}, MakeSpan(GetRef(let_node))); - } - - Expr Rewrite_(const IfNode* if_node, const Expr& post) final { - return WithFields(Downcast(post), {}, {}, {}, {}, MakeSpan(GetRef(if_node))); - } - - // OpNodes are not rewritten. - - Expr Rewrite_(const TupleGetItemNode* tuple_get_item_node, const Expr& post) final { - return WithFields(Downcast(post), {}, {}, {}, - MakeSpan(GetRef(tuple_get_item_node))); - } - - Expr Rewrite_(const RefCreateNode* ref_create_node, const Expr& post) final { - return WithFields(Downcast(post), {}, {}, - MakeSpan(GetRef(ref_create_node))); - } - - Expr Rewrite_(const RefReadNode* ref_read_node, const Expr& post) final { - return WithFields(Downcast(post), {}, {}, MakeSpan(GetRef(ref_read_node))); - } - - Expr Rewrite_(const RefWriteNode* ref_write_node, const Expr& post) final { - return WithFields(Downcast(post), {}, {}, {}, - MakeSpan(GetRef(ref_write_node))); - } - - // ConstructorNodes are not rewritten. - - Expr Rewrite_(const MatchNode* match_node, const Expr& post) final { - return WithFields(Downcast(post), {}, {}, {}, MakeSpan(GetRef(match_node))); - } - - Span MakeSpan(const Expr& expr) { - auto node = indexed_graph_->item_to_node(expr); - int node_index = static_cast(node->index_); - int dominator_index = - node->dominator_parent_ ? static_cast(node->dominator_parent_->index_) : -1; - Span span(source_name_, /*line=*/node_index, /*end_line=*/node_index, - /*column=*/dominator_index, /*end_column=*/dominator_index); - return span; - } - - SourceName source_name_; - const IndexedGraph* indexed_graph_; -}; - -} // namespace - -tvm::transform::Pass CapturePostDfsIndexInSpans() { - auto pass_func = [](Function f, IRModule m, transform::PassContext ctxt) { - std::unique_ptr> indexed_graph = CreateIndexedGraph(f); - SpansRewriter rewriter(indexed_graph.get()); - return Downcast(PostOrderRewrite(f, &rewriter)); - }; - return CreateFunctionPass(pass_func, 0, "CapturePostDfsIndexInSpans", {}); -} - -TVM_REGISTER_GLOBAL("relay._transform.CapturePostDfsIndexInSpans") - .set_body_typed(CapturePostDfsIndexInSpans); - -} // namespace transform -} // namespace relay -} // namespace tvm diff --git a/src/relay/transforms/combine_parallel_batch_matmul.cc b/src/relay/transforms/combine_parallel_batch_matmul.cc deleted file mode 100644 index ddab87a4893e..000000000000 --- a/src/relay/transforms/combine_parallel_batch_matmul.cc +++ /dev/null @@ -1,174 +0,0 @@ -/* - * 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 combine_parallel_batch_matmul.cc - * \brief Combine parallel batch matmuls into a single one. - * - * This pass replaces batch_matmul that share the same lhs node with a - * single batch matmul.Elemwise and broadcast ops following batch_matmul are also - * combined if possible. - * - * This prevents launching multiple kernels in networks with multiple - * convolution branches, such as Inception block. - */ - -#include -#include -#include -#include -#include -#include - -#include -#include - -#include "./combine_parallel_op.h" -#include "./expr_subst.h" -#include "pattern_utils.h" - -namespace tvm { -namespace relay { - -class ParallelBatchMatmulCombiner : public ParallelOpCombiner { - public: - explicit ParallelBatchMatmulCombiner(uint64_t min_num_branches) - : ParallelOpCombiner("nn.batch_matmul", min_num_branches) {} - - protected: - bool IsSupportedOp(const CallNode* n) { return true; } - - bool CanOpsBeCombined(const CallNode* a, const CallNode* b) { - StructuralEqual eq; - const auto* attrs_a = a->attrs.as(); - const auto* attrs_b = b->attrs.as(); - ICHECK(attrs_a); - ICHECK(attrs_b); - const auto* rhs_a = a->args[1]->type_as(); - const auto* rhs_b = b->args[1]->type_as(); - const auto* restype_a = a->type_as(); - const auto* restype_b = b->type_as(); - // shape[2] is the contraction axis and automatically consistent - // if it were valid batch_matmul ops - - // TODO(jcf94): Add full support of layout format - if (!(attrs_a->transpose_a == false && attrs_a->transpose_b == true && - attrs_b->transpose_a == false && attrs_b->transpose_b == true)) { - LOG(WARNING) << "For legacy reason, this pass only supports" - << " (transpose_a=false, transpose_b=true) now, skip combining these two with:" - << " batch_matmul_a: " << attrs_a->transpose_a << ", " << attrs_a->transpose_b - << " batch_matmul_b: " << attrs_b->transpose_a << ", " << attrs_b->transpose_b; - return false; - } - - auto res = eq(rhs_a->dtype, rhs_b->dtype) && eq(restype_a->dtype, restype_b->dtype) && - (rhs_a->shape.size() == 3) && (rhs_b->shape.size() == 3) && - eq(rhs_a->shape[0], rhs_b->shape[0]) && eq(attrs_a->out_dtype, attrs_b->out_dtype); - return res; - } - - Call MakeCombinedOp(const Group& branches) { - Expr data = branches[0][0]->args[0]; - - Array weights; - for (const auto& branch : branches) { - auto call = branch[0]; - weights.push_back(call->args[1]); - } - Expr new_weight = MakeConcatenate(Tuple(weights), 1); - - const auto* origin_attrs = branches[0][0]->attrs.as(); - ICHECK(origin_attrs); - return Downcast(MakeBatchMatmul(data, new_weight, origin_attrs->out_dtype, - origin_attrs->transpose_a, origin_attrs->transpose_b)); - } - - bool IsArgCompatible(const CallNode* a, const CallNode* b, size_t index) { return true; } - - Call MakeCombinedCallFromFollowingOps(const Expr& data, const Group& branches, size_t depth, - size_t parent_index) { - Array new_args; - const CallNode* call = branches[0][depth]; - - for (size_t i = 0; i < call->args.size(); i++) { - if (i == parent_index) { - new_args.push_back(data); - continue; - } - - Array tuple; - for (const auto& branch : branches) { - tuple.push_back(branch[depth]->args[i]); - } - - auto concat = MakeConcatenate(Tuple(tuple), -1); - new_args.push_back(std::move(concat)); - } - - return Call(call->op, new_args, call->attrs, {}); - } - - void UpdateGroupOutput(const Expr& data, const Group& branches, size_t depth, - ExprSubstMap* subst_map) { - int64_t index = 0; - - for (const auto& branch : branches) { - const CallNode* batch_matmul = branch[0]; - auto feature_dim = batch_matmul->args[1]->type_as()->shape[1]; - auto fpp = tir::as_const_int(feature_dim); - int64_t features = *fpp; - Array begin; - Array end; - for (size_t i = 0; i < 2; i++) { - begin.push_back(0); - end.push_back(-1); - } - begin.push_back(index); - index += features; - end.push_back(features); - Array strides(begin.size(), 1); - auto slice = MakeStridedSlice(data, begin, end, strides, "size"); - subst_map->insert({GetRef(branch[depth]), slice}); - } - } -}; - -/*! \brief Combine parallel batch_matmul if number of branches >= min_num_branches */ -Expr CombineParallelBatchMatmul(const Expr& expr, uint64_t min_num_branches) { - return ParallelBatchMatmulCombiner(min_num_branches).Combine(expr); -} - -namespace transform { - -Pass CombineParallelBatchMatmul(uint64_t min_num_branches) { - runtime::TypedPackedFunc pass_func = - [=](Function f, IRModule m, PassContext pc) { - return Downcast(CombineParallelBatchMatmul(f, min_num_branches)); - }; - return CreateFunctionPass(pass_func, 4, "CombineParallelBatchMatmul", {"InferType"}); -} - -TVM_REGISTER_GLOBAL("relay._transform.CombineParallelBatchMatmul") - .set_body_typed(CombineParallelBatchMatmul); - -} // namespace transform - -} // namespace relay -} // namespace tvm diff --git a/src/relay/transforms/combine_parallel_conv2d.cc b/src/relay/transforms/combine_parallel_conv2d.cc deleted file mode 100644 index 9c7bcc27ec82..000000000000 --- a/src/relay/transforms/combine_parallel_conv2d.cc +++ /dev/null @@ -1,225 +0,0 @@ -/* - * 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 combine_parallel_conv2d.cc - * \brief Combine parallel 2d convolutions into a single convolution. - * - * This pass replaces convolutions that share the same input node and the same - * arguments (except that the number of output channels can be different) with a - * single convolution. The weight of the new 2d convolution is the concatenation - * of the original weights. Elemwise and broadcast ops following conv2d are also - * combined if possible. - * - * This prevents launching multiple kernels in networks with multiple - * convolution branches, such as Inception block. - */ - -#include -#include -#include -#include -#include -#include - -#include -#include - -#include "./combine_parallel_op.h" -#include "./expr_subst.h" -#include "pattern_utils.h" - -namespace tvm { -namespace relay { - -class ParallelConv2DCombiner : public ParallelOpCombiner { - public: - explicit ParallelConv2DCombiner(uint64_t min_num_branches) - : ParallelOpCombiner("nn.conv2d", min_num_branches) {} - - protected: - bool IsSupportedOp(const CallNode* n) { return n->attrs.as()->groups == 1; } - - bool CanOpsBeCombined(const CallNode* a, const CallNode* b) { - StructuralEqual eq; - const Layout kOIHW("OIHW"); - const auto* attrs_a = a->attrs.as(); - const auto* attrs_b = b->attrs.as(); - ICHECK(attrs_a); - ICHECK(attrs_b); - const auto* tweight_a = a->args[1]->type_as(); - const auto* tweight_b = b->args[1]->type_as(); - const auto shape_a = - tir::BijectiveLayout(Layout(attrs_a->kernel_layout), kOIHW).ForwardShape(tweight_a->shape); - const auto shape_b = - tir::BijectiveLayout(Layout(attrs_b->kernel_layout), kOIHW).ForwardShape(tweight_b->shape); - - return eq(attrs_a->strides, attrs_b->strides) && eq(attrs_a->padding, attrs_b->padding) && - eq(attrs_a->dilation, attrs_b->dilation) && eq(attrs_a->groups, attrs_b->groups) && - eq(attrs_a->data_layout, attrs_b->data_layout) && - eq(attrs_a->kernel_layout, attrs_b->kernel_layout) && - eq(attrs_a->out_dtype, attrs_b->out_dtype) && - eq(attrs_a->out_layout, attrs_b->out_layout) && eq(shape_a[2], shape_b[2]) && - eq(shape_a[3], shape_b[3]); - } - - Call MakeCombinedOp(const Group& branches) { - const Op& conv2d = Op::Get("nn.conv2d"); - Expr data = branches[0][0]->args[0]; - auto [new_weight, new_channels] = TransformWeight(branches); - - const CallNode* group_root = branches[0][0]; - const auto* attrs = group_root->attrs.as(); - ICHECK(attrs); - const auto new_attrs = make_object(); - new_attrs->strides = attrs->strides; - new_attrs->padding = attrs->padding; - new_attrs->dilation = attrs->dilation; - new_attrs->groups = attrs->groups; - new_attrs->kernel_size = attrs->kernel_size; - new_attrs->data_layout = attrs->data_layout; - new_attrs->kernel_layout = attrs->kernel_layout; - new_attrs->out_layout = attrs->out_layout; - new_attrs->out_dtype = attrs->out_dtype; - new_attrs->channels = new_channels; - - const std::string& layout = - new_attrs->out_layout == "" ? new_attrs->data_layout : new_attrs->out_layout; - channel_pos_ = layout.find('C'); - ICHECK_NE(channel_pos_, std::string::npos); - - return Call(conv2d, {data, new_weight}, Attrs{new_attrs}, {}); - } - - bool IsArgCompatible(const CallNode* a, const CallNode* b, size_t index) { - StructuralEqual eq; - auto ta = a->args[index]->type_as(); - auto tb = b->args[index]->type_as(); - auto toutput_a = a->type_as(); - auto toutput_b = b->type_as(); - - if (!eq(ta->dtype, tb->dtype) || ta->shape.size() != tb->shape.size()) return false; - - // Position of the 'C' dimension in the argument - size_t arg_channel_pos = channel_pos_ - toutput_a->shape.size() + ta->shape.size(); - - // Channel super-dimension shoule be present and not broadcasted - if ((arg_channel_pos > channel_pos_) || // size_t overflow - !eq(ta->shape[arg_channel_pos], toutput_a->shape[channel_pos_]) || - !eq(tb->shape[arg_channel_pos], toutput_b->shape[channel_pos_])) - return false; - - for (size_t i = 0; i < ta->shape.size(); i++) { - if (i == arg_channel_pos) continue; - if (!eq(ta->shape[i], tb->shape[i])) return false; - } - return true; - } - - Call MakeCombinedCallFromFollowingOps(const Expr& data, const Group& branches, size_t depth, - size_t parent_index) { - Array new_args; - const CallNode* call = branches[0][depth]; - size_t ndim = call->type_as()->shape.size(); - - for (size_t i = 0; i < call->args.size(); i++) { - if (i == parent_index) { - new_args.push_back(data); - continue; - } - - size_t arg_ndim = call->args[i]->type_as()->shape.size(); - size_t arg_channel_pos = channel_pos_ - ndim + arg_ndim; - Array tuple; - for (const auto& branch : branches) { - tuple.push_back(branch[depth]->args[i]); - } - - auto concat = MakeConcatenate(Tuple(tuple), arg_channel_pos); - new_args.push_back(std::move(concat)); - } - - return Call(call->op, new_args, call->attrs, {}); - } - - void UpdateGroupOutput(const Expr& data, const Group& branches, size_t depth, - ExprSubstMap* subst_map) { - int64_t index = 0; - - for (const auto& branch : branches) { - const CallNode* conv2d = branch[0]; - int64_t channels = GetConv2DSuperChannelsDim(conv2d); - Array begin; - Array end; - for (size_t i = 0; i < channel_pos_; i++) { - begin.push_back(0); - end.push_back(-1); - } - begin.push_back(index); - index += channels; - end.push_back(channels); - Array strides(begin.size(), 1); - auto slice = MakeStridedSlice(data, begin, end, strides, "size"); - subst_map->insert({GetRef(branch[depth]), slice}); - } - } - - private: - /* \brief index of channel dimension */ - size_t channel_pos_; - - std::tuple TransformWeight(const Group& branches) { - int64_t num_filters = 0; // number of filters of the transformed weight - Array weights; - for (const auto& branch : branches) { - auto conv2d = branch[0]; - weights.push_back(conv2d->args[1]); - auto channels = GetConv2DSuperChannelsDim(conv2d); - num_filters += channels; - } - auto index = - branches[0][0]->attrs.as()->kernel_layout.operator std::string().find('O'); - ICHECK_NE(index, std::string::npos); - return std::make_tuple(MakeConcatenate(Tuple(weights), index), - tir::make_const(DataType::Int(32), num_filters)); - } -}; - -/*! \brief Combine parallel conv2d if number of branches >= min_num_branches */ -Expr CombineParallelConv2D(const Expr& expr, uint64_t min_num_branches) { - return ParallelConv2DCombiner(min_num_branches).Combine(expr); -} - -namespace transform { - -Pass CombineParallelConv2D(uint64_t min_num_branches) { - runtime::TypedPackedFunc pass_func = - [=](Function f, IRModule m, PassContext pc) { - return Downcast(CombineParallelConv2D(f, min_num_branches)); - }; - return CreateFunctionPass(pass_func, 4, "CombineParallelConv2d", {"InferType"}); -} - -TVM_REGISTER_GLOBAL("relay._transform.CombineParallelConv2D").set_body_typed(CombineParallelConv2D); - -} // namespace transform - -} // namespace relay -} // namespace tvm diff --git a/src/relay/transforms/combine_parallel_dense.cc b/src/relay/transforms/combine_parallel_dense.cc deleted file mode 100644 index e5f7e0b975f4..000000000000 --- a/src/relay/transforms/combine_parallel_dense.cc +++ /dev/null @@ -1,275 +0,0 @@ -/* - * 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 combine_parallel_dense.cc - * \brief Combine parallel dense ops into a single dense. - * - * This pass replaces dense ops that share the same input node, same shape, - * and don't have "units" defined with a single batch matrix multiplication. - * The inputs of the new batch_matmul is the stack of the original inputs. - * Elemwise and broadcast ops following dense are also combined if possible. - * - * This prevents launching multiple kernels in networks with multiple - * dense branches, such as BERT. - */ - -#include -#include -#include -#include -#include -#include - -#include -#include - -#include "./combine_parallel_op_batch.h" -#include "./expr_subst.h" -#include "pattern_utils.h" - -namespace tvm { -namespace relay { - -/* - * Class that find and combine parallel dense ops into batch_matmul. - */ -class ParallelDenseToBatchCombiner : public ParallelOpBatchCombiner { - public: - explicit ParallelDenseToBatchCombiner(uint64_t min_num_branches) - : ParallelOpBatchCombiner("nn.dense", "nn.batch_matmul", min_num_branches) {} - - protected: - Call MakeCombinedOp(const Group& branches) { - Array new_args; - size_t num_args = branches[0][0]->args.size(); - for (size_t i = 0; i < num_args; i++) { - Array arg_from_all_branches; - for (const auto& branch : branches) { - arg_from_all_branches.push_back(branch[0]->args[i]); - } - - new_args.push_back(MakeStack(Tuple(arg_from_all_branches), 0)); - } - - CHECK_EQ(num_args, 2); - const auto* origin_attrs = branches[0][0]->attrs.as(); - ICHECK(origin_attrs); - return Downcast( - MakeBatchMatmul(new_args[0], new_args[1], origin_attrs->out_dtype, false, true)); - } - - virtual bool CanOpsBeCombined(const CallNode* a, const CallNode* b) { - StructuralEqual eq; - const auto* attrs_a = a->attrs.as(); - const auto* attrs_b = b->attrs.as(); - ICHECK(attrs_a); - ICHECK(attrs_b); - const auto* weight_a = a->args[1]->type_as(); - const auto* weight_b = b->args[1]->type_as(); - - return eq(attrs_a->out_dtype, attrs_b->out_dtype) && - eq(weight_a->shape[0], weight_b->shape[0]) && eq(weight_a->shape[1], weight_b->shape[1]); - } -}; - -/* - * Class that find and combine parallel dense ops into one dense op - * whose num of output units equals to sum of each sub-ops. - */ -class ParallelDenseToDenseCombiner : public ParallelOpCombiner { - public: - explicit ParallelDenseToDenseCombiner(uint64_t min_num_branches) - : ParallelOpCombiner("nn.dense", min_num_branches) {} - - protected: - bool IsSupportedOp(const CallNode* n) { return true; } - - bool CanOpsBeCombined(const CallNode* a, const CallNode* b) { - StructuralEqual eq; - const auto* attrs_a = a->attrs.as(); - const auto* attrs_b = b->attrs.as(); - const auto* weight_a = a->args[1]->type_as(); - const auto* weight_b = b->args[1]->type_as(); - ICHECK(attrs_a != nullptr && attrs_b != nullptr && weight_a != nullptr && weight_b != nullptr); - // output dims (weight->shape[0]) can be different - return eq(attrs_a->out_dtype, attrs_b->out_dtype) && eq(weight_a->shape[1], weight_b->shape[1]); - } - - Call MakeCombinedOp(const Group& branches) { - const Op& dense_op = Op::Get("nn.dense"); - Expr input = branches[0][0]->args[0]; - // concat all weights into one - auto [new_weight, new_output_dims] = TransformWeight(branches); - const auto* origin_attrs = branches[0][0]->attrs.as(); - ICHECK(origin_attrs); - const auto dense_attrs = make_object(); - dense_attrs->units = new_output_dims; - dense_attrs->out_dtype = origin_attrs->out_dtype; - return Call(dense_op, {input, new_weight}, Attrs{dense_attrs}, {}); - } - - bool IsArgCompatible(const CallNode* a, const CallNode* b, size_t index) { - StructuralEqual eq; - auto ta = a->args[index]->type_as(); - auto tb = b->args[index]->type_as(); - auto toutput_a = a->type_as(); - auto toutput_b = b->type_as(); - ICHECK(ta != nullptr && tb != nullptr && toutput_a != nullptr && toutput_b != nullptr); - - if (!eq(ta->dtype, tb->dtype) || ta->shape.size() != tb->shape.size()) { - return false; - } - if (toutput_a->shape.size() < ta->shape.size() || toutput_b->shape.size() < tb->shape.size()) { - return false; // not broadcast/elemwise - } - if (ta->shape.size() > 0) { - for (size_t i = 0; i < ta->shape.size() - 1; i++) { - // shape dims must match except last dim - if (!eq(ta->shape[i], tb->shape[i])) return false; - } - } - return true; - } - - Call MakeCombinedCallFromFollowingOps(const Expr& data, const Group& branches, size_t depth, - size_t parent_index) { - Array new_args; - const CallNode* call = branches[0][depth]; - for (size_t i = 0; i < call->args.size(); i++) { - if (i == parent_index) { - new_args.push_back(data); - continue; - } - size_t arg_ndim = call->args[i]->type_as()->shape.size(); - size_t concat_axis = arg_ndim == 0 ? 0 : arg_ndim - 1; - Array tuple; - for (const auto& branch : branches) { - auto parent = branch[depth]->args[parent_index]; - auto& parent_shape = parent->type_as()->shape; - auto out_dim = tir::as_const_int(parent_shape[parent_shape.size() - 1]); - ICHECK(out_dim != nullptr); - - auto arg = branch[depth]->args[i]; - auto& arg_shape = arg->type_as()->shape; - bool repeat_last_dim = false; - if (arg_ndim == 0) { - repeat_last_dim = true; - arg = MakeExpandDims(arg, -1, 1); - } else { - auto arg_last_dim = tir::as_const_int(arg_shape[arg_shape.size() - 1]); - ICHECK(arg_last_dim != nullptr); - if (*out_dim > 1 && *arg_last_dim == 1) { - repeat_last_dim = true; - } - } - if (repeat_last_dim) { - // ensure broadcast is valid after concat args - arg = MakeRepeat(arg, *out_dim, concat_axis); - } - tuple.push_back(arg); - } - auto concat = MakeConcatenate(Tuple(tuple), concat_axis); - new_args.push_back(std::move(concat)); - } - return Call(call->op, new_args, call->attrs, {}); - } - - void UpdateGroupOutput(const Expr& data, const Group& branches, size_t depth, - ExprSubstMap* subst_map) { - int index = 0; - const auto dense_op = Op::Get("nn.dense"); - for (const auto& branch : branches) { - const CallNode* call = branch[depth]; - auto& out_shape = call->type_as()->shape; - - const CallNode* dense = branch[0]; - ICHECK(dense->op.same_as(dense_op)); - auto& dense_shape = dense->type_as()->shape; - auto dense_out_dims = tir::as_const_int(dense_shape[1]); - ICHECK(dense_out_dims != nullptr); - - // dense can be followed by shape-changing operations, so the slicing axis is - // not necessarily the last one. - // TODO(masahi): The following logic is incorrect if (1) there is no axis in - // out_shape[i] that directly corresponds to the output channel of dense or (2) there - // is another axis that happens to have the same size as the output channel of dense. - // Such cases might arise due to reshape / transpose / split etc. Revisit this logic - // when we encounter them in practice. - auto slice_axis = -1; - for (size_t i = out_shape.size() - 1; i >= 0; --i) { - ICHECK(tir::as_const_int(out_shape[i])); - if (*tir::as_const_int(out_shape[i]) == *dense_out_dims) { - slice_axis = i; - break; - } - } - ICHECK(slice_axis != -1); - - Array begin(out_shape.size(), 0); - Array end(out_shape.size(), -1); - Array strides(out_shape.size(), 1); - begin.Set(slice_axis, index); - end.Set(slice_axis, *dense_out_dims); - index += *dense_out_dims; - auto slice = MakeStridedSlice(data, begin, end, strides, "size"); - subst_map->insert({GetRef(branch[depth]), slice}); - } - } - - private: - std::tuple TransformWeight(const Group& branches) { - int64_t out_dims = 0; - Array weights; - for (const auto& branch : branches) { - auto weight = branch[0]->args[1]; - weights.push_back(weight); - out_dims += *tir::as_const_int(weight->type_as()->shape[0]); - } - return std::make_tuple(MakeConcatenate(Tuple(weights), 0), - tir::make_const(DataType::Int(32), out_dims)); - } -}; - -/*! \brief Combine parallel dense if number of branches >= min_num_branches */ -Expr CombineParallelDense(const Expr& expr, uint64_t min_num_branches, bool to_batch) { - if (to_batch) { - return ParallelDenseToBatchCombiner(min_num_branches).Combine(expr); - } else { - return ParallelDenseToDenseCombiner(min_num_branches).Combine(expr); - } -} - -namespace transform { - -Pass CombineParallelDense(uint64_t min_num_branches, bool to_batch_matmul) { - runtime::TypedPackedFunc pass_func = - [=](Function f, IRModule m, PassContext pc) { - return Downcast(CombineParallelDense(f, min_num_branches, to_batch_matmul)); - }; - return CreateFunctionPass(pass_func, 4, "CombineParallelDense", {"InferType"}); -} - -TVM_REGISTER_GLOBAL("relay._transform.CombineParallelDense").set_body_typed(CombineParallelDense); - -} // namespace transform - -} // namespace relay -} // namespace tvm diff --git a/src/relay/transforms/combine_parallel_op.cc b/src/relay/transforms/combine_parallel_op.cc deleted file mode 100644 index 1c9a58f49824..000000000000 --- a/src/relay/transforms/combine_parallel_op.cc +++ /dev/null @@ -1,178 +0,0 @@ -/* - * 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 combine_parallel_op.cc - * \brief Abstract class to combine parallel ops and their successive element-wise ops. - */ - -#include "combine_parallel_op.h" - -#include -#include -#include -#include -#include -#include -#include -#include - -#include -#include -#include -#include - -#include "expr_subst.h" -#include "pattern_utils.h" - -namespace tvm { -namespace relay { - -BranchGroupFinder::BranchGroupFinder(const Op& op, FIsSupportedOp fis_supported_op, - FAreCompatibleOps fare_compatible_ops) - : cached_op_(op), - fis_supported_op_(fis_supported_op), - fare_compatible_ops_(fare_compatible_ops) {} - -std::vector BranchGroupFinder::Find(const Expr& expr) { - this->VisitExpr(expr); - - std::vector groups; - for (const auto& root : op_roots_) { - const auto& children = children_map_.at(root); - size_t ngroups = groups.size(); - for (const CallNode* child : children) { - if (child->op != cached_op_) continue; - - auto&& branch = CreateBranch(child); - // add the branch to a group, or create a new group - auto it = std::find_if(groups.begin() + ngroups, groups.end(), [&](const Group& group) { - ICHECK(!group.empty() && !group[0].empty()); - return fare_compatible_ops_(child, group[0][0]); - }); - if (it != groups.end()) { - it->push_back(branch); - } else { - groups.emplace_back(); - // each group has at least one branch - groups.back().push_back(branch); - } - } - } - return groups; -} - -// Create a branch starting from op. -Branch BranchGroupFinder::CreateBranch(const CallNode* op) { - auto fpattern = Op::GetAttrMap("TOpPattern"); - // each branch has at least one element, the first element is always op - Branch branch{op}; - auto it = children_map_.find(GetRef(branch.back())); - while (it != children_map_.end() && it->second.size() == 1) { - const CallNode* call = it->second[0]; - auto pattern = fpattern[Downcast(call->op)]; - if (pattern <= kBroadcast) { - branch.push_back(call); - it = children_map_.find(GetRef(branch.back())); - } else { - break; - } - } - return branch; -} - -void BranchGroupFinder::VisitExpr_(const CallNode* n) { - ExprVisitor::VisitExpr_(n); - if (n->op == cached_op_ && fis_supported_op_(n)) { - op_roots_.insert(n->args[0]); - children_map_[n->args[0]].push_back(n); - } else { - for (size_t i = 0; i < n->args.size(); i++) { - children_map_[n->args[i]].push_back(n); - } - } -} - -ParallelOpCombiner::ParallelOpCombiner(const std::string& op_name, uint64_t min_num_branches) - : cached_op_(Op::Get(op_name)), min_num_branches_(min_num_branches) {} - -Expr ParallelOpCombiner::Combine(const Expr& expr) { - auto groups = BranchGroupFinder( - cached_op_, [&](const CallNode* n) { return IsSupportedOp(n); }, - [&](const CallNode* a, const CallNode* b) { return CanOpsBeCombined(a, b); }) - .Find(expr); - for (const Group& group : groups) { - if (group.size() < min_num_branches_) { - continue; - } - CombineBranches(group); - } - return ExprSubst(expr, std::move(subst_map_)); -} - -void ParallelOpCombiner::CombineBranches(const Group& branches) { - Call combined = MakeCombinedOp(branches); - auto it = std::min_element(branches.begin(), branches.end(), - [](const Branch& branch_a, const Branch& branch_b) { - return branch_a.size() < branch_b.size(); - }); - size_t depth = it->size(); - size_t i; - // starting from 1 to skip the op - for (i = 1; i < depth; i++) { - size_t parent_index; - for (parent_index = 0; parent_index < branches[0][i]->args.size(); parent_index++) { - if (branches[0][i]->args[parent_index].get() == branches[0][i - 1]) break; - } - ICHECK_NE(parent_index, branches[0][i]->args.size()); - if (!CheckLevel(branches, i, parent_index)) break; - combined = MakeCombinedCallFromFollowingOps(combined, branches, i, parent_index); - } - UpdateGroupOutput(combined, branches, i - 1, &subst_map_); -} - -bool ParallelOpCombiner::CheckLevel(const Group& branches, size_t depth, size_t parent_index) { - const CallNode* call = branches[0][depth]; - tvm::StructuralEqual attrs_equal; - // check if all branches in current depth can be combined - for (auto it = branches.begin() + 1; it != branches.end(); it++) { - const Branch& branch = *it; - if (!branch[depth]->op.same_as(call->op) || !attrs_equal(branch[depth]->attrs, call->attrs) || - branch[depth]->args.size() != call->args.size()) { - return false; - } - - if (branch[depth]->args[parent_index].get() != branch[depth - 1]) return false; - - // Check args - for (size_t i = 0; i < call->args.size(); i++) { - if (i == parent_index) continue; - - if (!IsArgCompatible(call, branch[depth], i) || - !attrs_equal(call->attrs, branch[depth]->attrs)) { - return false; - } - } - } - return true; -} - -} // namespace relay -} // namespace tvm diff --git a/src/relay/transforms/combine_parallel_op.h b/src/relay/transforms/combine_parallel_op.h deleted file mode 100644 index 9785a366299b..000000000000 --- a/src/relay/transforms/combine_parallel_op.h +++ /dev/null @@ -1,237 +0,0 @@ -/* - * 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 combine_parallel_op.h - * \brief Abstract class to combine parallel ops and their successive element-wise ops. - */ -#ifndef TVM_RELAY_TRANSFORMS_COMBINE_PARALLEL_OP_H_ -#define TVM_RELAY_TRANSFORMS_COMBINE_PARALLEL_OP_H_ - -#include -#include -#include -#include -#include -#include - -#include -#include -#include -#include - -#include "./expr_subst.h" -#include "pattern_utils.h" - -namespace tvm { -namespace relay { - -using Branch = std::vector; -using Group = std::vector; -using FIsSupportedOp = std::function; -using FAreCompatibleOps = std::function; -using ExprSubstMap = std::unordered_map; - -/* - * Class to find parallel branches starting with op that are - * grouped if they are able to be combined. They are eligible to - * be combined if they have the same input data. - * Op can be followed by zero or more elemwise or broadcast ops, - * which are included in the group. - * Intermediate nodes have exactly one successor. It is possible that branches meet at a point, - * which should be handled in ParallelOpCombiner. - * - * data - * / \ - * op op - * | | - * elem-wise elem-wise - * | | - */ -class BranchGroupFinder : private ExprVisitor { - public: - /* - * \brief Constructor - * \param op The op that indicates the start of each group - * \param fis_supported_op function that returns true if op - * is supported for combining - * \param fare_compatible_ops function that returns true if - * two ops are compatible for combining - */ - BranchGroupFinder(const Op& op, FIsSupportedOp fis_supported_op, - FAreCompatibleOps fare_compatible_ops); - - /* - * \brief Finds all groups that can be combined. - * \param expr Relay expression that represents function - * to look at for groups to be combined - * \return Vector of groups which can be combined. - */ - std::vector Find(const Expr& expr); - - private: - /* \brief Cache the op for finding parallel branches */ - const Op& cached_op_; - - /* \brief function to return true if op is eligible to be combined, - * false otherwise - */ - FIsSupportedOp fis_supported_op_; - - /* \brief function to return true if two parallel ops are eligible - * to be combined, false otherwise - */ - FAreCompatibleOps fare_compatible_ops_; - - /* \brief ops that are on the first (logically, leftmost) branch - * of parallel ops and are eligible to be combined - */ - std::unordered_set op_roots_; - - /* \brief map of Expr to CallNodes that follow it */ - std::unordered_map, ObjectPtrHash, ObjectPtrEqual> - children_map_; - - /* - * \brief Creates new branch from op and its children that have - * elementwise or broadcast patterns - * \return New branch - */ - Branch CreateBranch(const CallNode* op); - - /* - * \brief Expression visitor function - */ - void VisitExpr_(const CallNode* n) final; -}; - -/* - * Abstract class to find and combine parallel ops and the elementwise ops that follow. - */ -class ParallelOpCombiner { - public: - /*! \brief virtual destructor */ - virtual ~ParallelOpCombiner() {} - /* - * \brief Constructor. - * \param op_name name of op to combine - * \param min_num_branches min number of parallel branches beginning with op - * to start combining - */ - explicit ParallelOpCombiner(const std::string& op_name, uint64_t min_num_branches); - - /* - * \brief Combines ops and following elementwise or broadcast ops - * \param expr function to modify - * \return new function with combined ops - */ - Expr Combine(const Expr& expr); - - protected: - /* - * \brief Checks if node is supported to be combined - * \param n node in question - * \return True if the op represented by n is supported to be the root of a branch - * to be combined. False otherwise. - */ - virtual bool IsSupportedOp(const CallNode* n) = 0; - - /* - * \brief Checks if two ops can be combined - * \param a node a - * \param b node b - * \return True if a and b can be combined. False otherwise. - */ - virtual bool CanOpsBeCombined(const CallNode* a, const CallNode* b) = 0; - - /* - * \brief Makes combined op from parallel ops in branches. This usually involves - * concatenating or stacking inputs, then creating a new call. - * \param branches branches that are to be combined - * \return new call with branches combined. - */ - virtual Call MakeCombinedOp(const Group& branches) = 0; - - /* - * \brief Checks if argument of op following combined ops are able to be combined - * \param a node a - * \param b node b - * \param index index of argument in question - * \return True if argument of a and b and index can be combined - */ - virtual bool IsArgCompatible(const CallNode* a, const CallNode* b, size_t index) = 0; - - /* - * \brief Create combined call from ops that follow the initial combined op at the depth-th level. - * This usually involves concatenating or stacking inputs, then creating a new call. - * Only called if IsArgCompatbile returns true for each arg. - * \param data combined op - * \param branches branches of parallel ops to be combined - * \param depth depth at which to combine ops - * \param parent_index index of arg that corresponds to original input that was shared among - * all combined ops - * \return new combined call - */ - virtual Call MakeCombinedCallFromFollowingOps(const Expr& data, const Group& branches, - size_t depth, size_t parent_index) = 0; - - /* - * \brief Updates map of expr to substitute with combined expr. This usually involves - * slicing or splitting data. - * \param data combined op - * \param branches branches of parallel ops to be combined - * \param depth depth at which to substitute - * \param subst_map map of Expr to replace with Expr to replace it with - */ - virtual void UpdateGroupOutput(const Expr& data, const Group& branches, size_t depth, - ExprSubstMap* subst_map) = 0; - - private: - /* \brief Cache the op to be combined */ - const Op& cached_op_; - - /* \brief minimum number of parallel branches to combine */ - uint64_t min_num_branches_; - - /* \brief map of Expr to Expr to substitute it with after running pass */ - ExprSubstMap subst_map_; - - /* - * \brief Combine parallel branches and updates subst_map_ with Exprs - * to be substituted - * \param branches branches to be combined - */ - void CombineBranches(const Group& branches); - - /* - * \brief Combine parallel branches and updates subst_map_ with Exprs - * to be substituted - * \param branches parallel branches to potentially be combined - * \param depth depth at which to look at op - * \param parent_index index of arg that corresponds to original input that was shared among - * all combined ops - * \return true if parallel ops at depth can be combined, false otherwise - */ - bool CheckLevel(const Group& branches, size_t depth, size_t parent_index); -}; - -} // namespace relay -} // namespace tvm -#endif // TVM_RELAY_TRANSFORMS_COMBINE_PARALLEL_OP_H_ diff --git a/src/relay/transforms/combine_parallel_op_batch.cc b/src/relay/transforms/combine_parallel_op_batch.cc deleted file mode 100644 index 74827f166b51..000000000000 --- a/src/relay/transforms/combine_parallel_op_batch.cc +++ /dev/null @@ -1,194 +0,0 @@ -/* - * 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 combine_parallel_op_batch.cc - * \brief Combine parallel ops into a single batch op. - * - * This pass replaces ops that share the same input node and same shape - * with a single op that takes in batched input. The inputs of the new - * batched op are the stack of the original inputs. Elementwise and - * broadcast ops following the original op are also stacked - * and fused if possible. For example: - * - * data - * / \ - * add (2,2) add (2,2) - * | | - * elemwise (2,2) elemwise (2,2) - * | | - * - * Would become: - * - * data - * | - * add+elemwise (2,2,2) - * / \ - * - */ - -#include "./combine_parallel_op_batch.h" - -#include -#include -#include -#include -#include -#include - -#include -#include - -#include "./combine_parallel_op.h" -#include "./expr_subst.h" -#include "pattern_utils.h" - -namespace tvm { -namespace relay { - -ParallelOpBatchCombiner::ParallelOpBatchCombiner(const std::string& op_name, - const std::string& batch_op_name, - uint64_t min_num_branches) - : ParallelOpCombiner(op_name, min_num_branches), batch_op_name_(batch_op_name) {} - -bool ParallelOpBatchCombiner::IsSupportedOp(const CallNode* n) { return true; } - -bool ParallelOpBatchCombiner::CanOpsBeCombined(const CallNode* a, const CallNode* b) { - if (a->args.size() != b->args.size()) { - return false; - } - - StructuralEqual eq; - for (size_t i = 0; i < a->args.size(); i++) { - auto ta = a->args[i]->type_as(); - auto tb = b->args[i]->type_as(); - if (ta->shape.size() != tb->shape.size() || !eq(ta->dtype, tb->dtype)) { - return false; - } - - for (size_t j = 0; j < ta->shape.size(); j++) { - if (!eq(ta->shape[j], tb->shape[j])) { - return false; - } - } - } - - return true; -} - -Call ParallelOpBatchCombiner::MakeCombinedOp(const Group& branches) { - const Op& batch_op = Op::Get(batch_op_name_); - - Array new_args; - size_t num_args = branches[0][0]->args.size(); - for (size_t i = 0; i < num_args; i++) { - Array arg_from_all_branches; - for (const auto& branch : branches) { - arg_from_all_branches.push_back(branch[0]->args[i]); - } - - new_args.push_back(MakeStack(Tuple(arg_from_all_branches), 0)); - } - - return Call(batch_op, new_args, Attrs(), {}); -} - -bool ParallelOpBatchCombiner::IsArgCompatible(const CallNode* a, const CallNode* b, size_t index) { - StructuralEqual eq; - auto ta = a->args[index]->type_as(); - auto tb = b->args[index]->type_as(); - - if (!eq(ta->dtype, tb->dtype) || ta->shape.size() != tb->shape.size()) return false; - - for (size_t i = 0; i < ta->shape.size(); i++) { - if (!eq(ta->shape[i], tb->shape[i])) return false; - } - return true; -} - -Call ParallelOpBatchCombiner::MakeCombinedCallFromFollowingOps(const Expr& data, - const Group& branches, size_t depth, - size_t parent_index) { - Array new_args; - const CallNode* call = branches[0][depth]; - - for (size_t i = 0; i < call->args.size(); i++) { - if (i == parent_index) { - new_args.push_back(data); - continue; - } - - Array tuple; - for (const auto& branch : branches) { - // if the shape of the arg is of shape (j,), - // expand it to (1,j) so it can be properly broadcasted. - Expr arg = branch[depth]->args[i]; - const TensorTypeNode* arg_tensor = arg->type_as(); - if (arg_tensor->shape.size() == 1) { - Expr expanded_arg = MakeExpandDims(arg, 0, 1); - tuple.push_back(expanded_arg); - } else { - tuple.push_back(arg); - } - } - - auto stack = MakeStack(Tuple(tuple), 0); - new_args.push_back(std::move(stack)); - } - - return Call(call->op, new_args, call->attrs, {}); -} - -void ParallelOpBatchCombiner::UpdateGroupOutput(const Expr& data, const Group& branches, - size_t depth, ExprSubstMap* subst_map) { - int index = 0; - auto split = MakeSplit(data, runtime::Int(branches.size()), 0); - for (const auto& branch : branches) { - auto split_data = TupleGetItem(split, index++); - auto squeezed_data = MakeSqueeze(split_data, {0}); - subst_map->insert({GetRef(branch[depth]), squeezed_data}); - } -} - -/*! \brief Combine parallel op into batched op if number of branches >= min_num_branches */ -Expr CombineParallelOpBatch(const Expr& expr, const std::string& op_name, - const std::string& batch_op_name, uint64_t min_num_branches) { - return ParallelOpBatchCombiner(op_name, batch_op_name, min_num_branches).Combine(expr); -} - -namespace transform { - -Pass CombineParallelOpBatch(const String& op_name, const String& batch_op_name, - uint64_t min_num_branches) { - runtime::TypedPackedFunc pass_func = - [=](Function f, IRModule m, PassContext pc) { - return Downcast( - CombineParallelOpBatch(f, op_name, batch_op_name, min_num_branches)); - }; - return CreateFunctionPass(pass_func, 4, "CombineParallelOpBatch", {"InferType"}); -} - -TVM_REGISTER_GLOBAL("relay._transform.CombineParallelOpBatch") - .set_body_typed(CombineParallelOpBatch); - -} // namespace transform - -} // namespace relay -} // namespace tvm diff --git a/src/relay/transforms/combine_parallel_op_batch.h b/src/relay/transforms/combine_parallel_op_batch.h deleted file mode 100644 index b9edafe75494..000000000000 --- a/src/relay/transforms/combine_parallel_op_batch.h +++ /dev/null @@ -1,145 +0,0 @@ -/* - * 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 combine_parallel_op_batch.h - * \brief Combine parallel ops into a single batch op. - */ -#ifndef TVM_RELAY_TRANSFORMS_COMBINE_PARALLEL_OP_BATCH_H_ -#define TVM_RELAY_TRANSFORMS_COMBINE_PARALLEL_OP_BATCH_H_ - -#include -#include -#include -#include -#include -#include - -#include -#include -#include - -#include "./combine_parallel_op.h" -#include "./expr_subst.h" -#include "pattern_utils.h" - -namespace tvm { -namespace relay { - -/* - * Class to find and combine parallel ops and following element-wise - * and broadcast ops into a single batch op. Ops can be combined - * if they have the same input data. Batch op is formed by - * stacking inputs. Final results are retrieved by splitting output. - * For example: - * - * data - * / \ - * dense (2,2) dense (2,2) - * | | - * elemwise/bcast (2,2) elemwise/bcast (2,2) - * - * Would become: - * - * data - * | - * batch_matmul+elemwise/bcast (2,2,2) - */ -class ParallelOpBatchCombiner : public ParallelOpCombiner { - public: - /* - * \brief Constructor. - * \param op_name name of op to combine - * \param batch_op_name name of op that combined branches will be joined into - * \param min_num_branches min number of parallel branches beginning with op - * to start combining - */ - ParallelOpBatchCombiner(const std::string& op_name, const std::string& batch_op_name, - uint64_t min_num_branches); - - protected: - /* - * \brief Checks if node is supported to be combined - * \param n node in question - * \return True by default - */ - virtual bool IsSupportedOp(const CallNode* n); - - /* - * \brief Checks if two ops can be combined - * \param a node a - * \param b node b - * \return True if shapes and dtypes of all args of a and b are the same - */ - virtual bool CanOpsBeCombined(const CallNode* a, const CallNode* b); - - /* - * \brief Makes combined op from parallel ops in branches. This usually involves - * concatenating or stacking inputs, then creating a new call. - * \param branches branches that are to be combined - * \return new call with branches combined as batch op by stacking args - */ - virtual Call MakeCombinedOp(const Group& branches); - - /* - * \brief Checks if argument of op following combined ops are able to be combined - * \param a node a - * \param b node b - * \param index index of argument in question - * \return True if shapes and dtypes of args[index] a and b are the same - */ - bool IsArgCompatible(const CallNode* a, const CallNode* b, size_t index) final; - - /* - * \brief Create combined call from ops that follow the initial combined op at the depth-th level. - * This usually involves concatenating or stacking inputs, then creating a new call. - * Only called if IsArgCompatbile returns true for each arg. - * \param data combined op - * \param branches branches of parallel ops to be combined - * \param depth depth at which to combine ops - * \param parent_index index of arg that corresponds to original input that was shared among - * all combined ops - * \return new combined call as batch op by stacking args - */ - Call MakeCombinedCallFromFollowingOps(const Expr& data, const Group& branches, size_t depth, - size_t parent_index) final; - - /* - * \brief Updates map of expr to substitute with combined expr. This usually involves - * slicing or splitting data. - * \param data combined op - * \param branches branches of parallel ops to be combined - * \param depth depth at which to substitute - * \param subst_map map of Expr to replace with Expr to replace it with - */ - void UpdateGroupOutput(const Expr& data, const Group& branches, size_t depth, - ExprSubstMap* subst_map) final; - - private: - /* \brief name of op to replace combined ops with. for example, - * for combining parallel dense, this will be set to - * nn.batch_matmul - */ - std::string batch_op_name_; -}; - -} // namespace relay -} // namespace tvm - -#endif // TVM_RELAY_TRANSFORMS_COMBINE_PARALLEL_OP_BATCH_H_ diff --git a/src/relay/transforms/compiler_function_utils.cc b/src/relay/transforms/compiler_function_utils.cc deleted file mode 100644 index 653659bb9a89..000000000000 --- a/src/relay/transforms/compiler_function_utils.cc +++ /dev/null @@ -1,302 +0,0 @@ -/* - * 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 src/relay/transforms/compiler_function_utils.cc - * \brief Helper passes for working with functions with the "Compiler" attribute. - */ - -#include "./compiler_function_utils.h" - -#include "tvm/relay/analysis.h" -#include "tvm/relay/expr_functor.h" -#include "tvm/relay/transform.h" - -namespace tvm { -namespace relay { -namespace transform { -namespace { - -/*! - * \brief Returns the \p FunctionNode of if \p expr if it is a "Compiler" function which should - * be processed by a pass using \p compiler_filter. Otherwise returns null. - */ -const FunctionNode* AsFunctionNode(const Expr& expr, const std::string& compiler_filter) { - if (const auto* function_node = expr.as()) { - Optional opt_compiler = function_node->GetAttr(attr::kCompiler); - if (opt_compiler.defined() && - (compiler_filter.empty() || opt_compiler.value() == compiler_filter)) { - return function_node; - } - } - return nullptr; -} - -/*! - * \brief Rewrite calls to inlined and let-bound "Compiler" functions to global functions. The given - * module will be extended with the newly outlined functions. - */ -class Outliner : public MixedModeMutator { - public: - using MixedModeMutator::VisitExpr_; - - Outliner(GlobalSymbolCache* cache, std::string compiler_filter, IRModule mod) - : cache_(cache), compiler_filter_(std::move(compiler_filter)), mod_(std::move(mod)) {} - - Expr VisitExpr_(const LetNode* op) final { - auto pre_visit = [this](const LetNode* op) { - Expr var = this->VisitExpr(op->var); - Expr value = this->VisitExpr(op->value); - - if (AsFunctionNode(value, compiler_filter_)) { - // Inline on-the-fly if the let-bound value is a function of interest. - this->memo_[var] = value; - } - }; - auto post_visit = [this](const LetNode* op) { - // Rely on the Memoizer to cache pre-visit values - Expr value = this->VisitExpr(op->value); - Expr body = this->VisitExpr(op->body); - auto expr = GetRef(op); - - if (AsFunctionNode(value, compiler_filter_)) { - // The let binding is no longer needed since inlined on-the-fly above. - this->memo_[expr] = this->VisitExpr(op->body); - } else { - Var var = Downcast(this->VisitExpr(op->var)); - if (var.same_as(op->var) && value.same_as(op->value) && body.same_as(op->body)) { - this->memo_[expr] = expr; - } else { - this->memo_[expr] = Let(var, value, body); - } - } - }; - ExpandANormalForm(op, pre_visit, post_visit); - return memo_[GetRef(op)]; - } - - Expr Rewrite_(const CallNode* pre, const Expr& post) final { - Call new_call = Downcast(post); - if (const auto* function_node = AsFunctionNode(new_call->op, compiler_filter_)) { - auto function = GetRef(function_node); - DCHECK(FreeVars(function).empty()) << "Function marked with '" << attr::kCompiler - << "' attribute should not have free variables"; - // Ask the cache to supply a unique global var for this function. - GlobalVar global_symbol = cache_->GetGlobalSymbol(function); - // Depending on the cache's implementation, two structurally equal (but not object - // equal) functions may be assigned the same global symbol. If so we'll lift it just - // once, but rewrite all the calls. - if (!mod_->ContainGlobalVar(global_symbol->name_hint)) { - function = - WithAttr(std::move(function), tvm::attr::kGlobalSymbol, global_symbol->name_hint); - mod_->Add(global_symbol, function); - } - // Update the call. - return WithFields(new_call, global_symbol); - } - return post; - } - - private: - /*! - * \brief A cached mapping from functions to global variables. Depending on the implementation - * the cache may generate fresh symbols or require the function to already have a - * "global_symbol" attribute, and may share symbols between structurally equal functions. - */ - GlobalSymbolCache* cache_; - /*! \brief If non-empty, the "Compiler" attribute value to require on functions to outline. */ - std::string compiler_filter_; - /*! \brief Module being rewritten. */ - IRModule mod_; -}; - -/*! - * \brief Inline immediate calls to "Composite" functions. - */ -class InnerInliner : public MixedModeMutator { - public: - InnerInliner() = default; - - private: - using MixedModeMutator::Rewrite_; - - Expr Rewrite_(const CallNode* pre, const Expr& post) final { - Call new_call = Downcast(post); - if (const auto* function_node = new_call->op.as()) { - ICHECK(function_node->GetAttr(attr::kComposite).defined()); - ICHECK_EQ(function_node->params.size(), new_call->args.size()); - Map subst; - for (size_t i = 0; i < new_call->args.size(); ++i) { - subst.Set(function_node->params[i], new_call->args[i]); - } - return Bind(function_node->body, subst); - } - return post; - } -}; - -/*! - * \brief Inline calls to global "Compiler" functions with global var in \p global_vars. - * Both the 'outer' "Compiler" function and any 'inner' "Composite" functions in its body - * are inlined. - */ -class OuterInliner : public MixedModeMutator { - public: - OuterInliner(IRModule mod, Array global_vars_) - : mod_(std::move(mod)), global_vars_(std::move(global_vars_)) {} - - private: - using MixedModeMutator::Rewrite_; - - Expr Rewrite_(const CallNode* pre, const Expr& post) final { - Call new_call = Downcast(post); - if (auto global_var_node = new_call->op.as()) { - auto global_var = global_var_node.value(); - if (std::find(global_vars_.begin(), global_vars_.end(), global_var) != global_vars_.end()) { - BaseFunc base_func = mod_->Lookup(global_var); - const auto* function_node = base_func.as(); - ICHECK(function_node); - ICHECK(function_node->GetAttr(attr::kCompiler).defined()); - ICHECK_EQ(function_node->params.size(), new_call->args.size()); - Map subst; - for (size_t i = 0; i < new_call->args.size(); ++i) { - subst.Set(function_node->params[i], new_call->args[i]); - } - Expr new_body = InnerInliner().VisitExpr(function_node->body); - return Bind(new_body, subst); - } - } - return post; - } - - private: - /*! \brief Original module we are processing. */ - IRModule mod_; - /*! \brief Global vars of functions to inline. */ - Array global_vars_; -}; - -} // namespace - -GlobalSymbolCache::~GlobalSymbolCache() = default; - -GlobalVar ExistingGlobalSymbolCache::GetGlobalSymbol(const Function& function) { - Optional opt_global_symbol = function->GetAttr(tvm::attr::kGlobalSymbol); - ICHECK(opt_global_symbol.defined()) - << "ExistingGlobalSymbolCache requires all functions to already have a '" - << tvm::attr::kGlobalSymbol << "' attribute"; - std::string global_symbol = opt_global_symbol.value(); - auto itr = global_vars_.find(global_symbol); - if (itr != global_vars_.end()) { - return itr->second; - } - // Ok if function does not have a checked_type, but if it does capture it in the global var. - GlobalVar global_var(global_symbol, function->checked_type_, function->span); - global_vars_.emplace(global_symbol, global_var); - return global_var; -} - -tvm::transform::Pass OutlineCompilerFunctions(std::shared_ptr cache, - std::string compiler_filter) { - runtime::TypedPackedFunc pass_func = - [cache = std::move(cache), compiler_filter = std::move(compiler_filter)]( - IRModule mod, transform::PassContext ctx) { - VLOG(1) << "OutlineCompilerFunctions input:" << std::endl << PrettyPrint(mod); - IRModule output_mod = mod->ShallowCopy(); - for (const auto& kv : mod->functions) { - if (const auto* function_node = AsOptimizableFunctionNode(kv.second)) { - Expr new_body = - Outliner(cache.get(), compiler_filter, output_mod).VisitExpr(function_node->body); - Function new_function = - WithFields(GetRef(function_node), /*opt_params=*/{}, new_body); - output_mod->Add(kv.first, new_function); - } - } - VLOG(1) << "OutlineCompilerFunctions result:" << std::endl << PrettyPrint(output_mod); - return output_mod; - }; - - return tvm::transform::CreateModulePass(pass_func, 0, "OutlineCompilerFunctions", {}); -} - -// Any Java programmers in the house? -tvm::transform::Pass OutlineCompilerFunctionsWithExistingGlobalSymbols( - std::string compiler_filter) { - return OutlineCompilerFunctions(std::make_shared(), - std::move(compiler_filter)); -} - -tvm::transform::Pass MarkCompilerFunctionsAsExtern(std::string compiler_filter) { - runtime::TypedPackedFunc pass_func = - [compiler_filter = std::move(compiler_filter)](IRModule mod, transform::PassContext ctx) { - VLOG(1) << "MarkCompilerFunctionsAsExtern input:" << std::endl << PrettyPrint(mod); - IRModule output_mod = mod->ShallowCopy(); - for (const auto& kv : mod->functions) { - if (const auto* function_node = AsFunctionNode(kv.second, compiler_filter)) { - auto new_function = - WithFields(GetRef(function_node), function_node->params, - function_node->body, function_node->ret_type, function_node->type_params, - /* erase attributes */ DictAttrs(Map())); - new_function = WithAttr(std::move(new_function), attr::kExtern, Integer(1)); - output_mod->Update(kv.first, new_function); - } - } - VLOG(1) << "MarkCompilerFunctionsAsExtern result:" << std::endl << PrettyPrint(output_mod); - return output_mod; - }; - - return tvm::transform::CreateModulePass(pass_func, 0, "MarkCompilerFunctionsAsExtern", {}); -} - -tvm::transform::Pass InlineCompilerFunctionsBoundTo(Array global_vars) { - runtime::TypedPackedFunc pass_func = - [global_vars = std::move(global_vars)](IRModule mod, transform::PassContext ctx) { - VLOG(1) << "InlineCompilerFunctionsBoundTo with global_vars: " << PrettyPrint(global_vars); - if (global_vars.empty()) { - return mod; - } - VLOG(1) << "InlineCompilerFunctions input:" << std::endl << PrettyPrint(mod); - IRModule output_mod = mod->ShallowCopy(); - for (const auto& kv : mod->functions) { - if (std::find(global_vars.begin(), global_vars.end(), kv.first) != global_vars.end()) { - output_mod->Remove(kv.first); - } else if (const auto* function_node = AsOptimizableFunctionNode(kv.second)) { - Expr new_body = OuterInliner(mod, global_vars).VisitExpr(function_node->body); - Function new_function = - WithFields(GetRef(function_node), /*opt_params=*/{}, new_body); - output_mod->Add(kv.first, new_function); - } - } - VLOG(1) << "InlineCompilerFunctionsBoundTo result:" << std::endl << PrettyPrint(output_mod); - return output_mod; - }; - - return tvm::transform::CreateModulePass(pass_func, 0, "InlineCompilerFunctionsBoundTo", {}); -} - -TVM_REGISTER_GLOBAL("relay._transform.OutlineCompilerFunctionsWithExistingGlobalSymbols") - .set_body_typed(OutlineCompilerFunctionsWithExistingGlobalSymbols); -TVM_REGISTER_GLOBAL("relay._transform.MarkCompilerFunctionsAsExtern") - .set_body_typed(MarkCompilerFunctionsAsExtern); -TVM_REGISTER_GLOBAL("relay._transform.InlineCompilerFunctionsBoundTo") - .set_body_typed(InlineCompilerFunctionsBoundTo); - -} // namespace transform -} // namespace relay -} // namespace tvm diff --git a/src/relay/transforms/compiler_function_utils.h b/src/relay/transforms/compiler_function_utils.h deleted file mode 100644 index a6cf6c9e7a8f..000000000000 --- a/src/relay/transforms/compiler_function_utils.h +++ /dev/null @@ -1,150 +0,0 @@ -/* - * 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 src/relay/transforms/compiler_function_utils.h - * \brief Helper passes for working with functions with the "Compiler" attribute. - * - * Those wishing to use the "RelayToTIR" custom pass machinery to do IRModule-at-a-time external - * codegen may find the following helpers useful: - * - * - The \p OutlineCompilerFunctionsWithExistingGlobalSymbols pass will lift inline functions with - * a matching "Compiler" attribute to be global functions, using the "global_symbol" attribute - * already assigned. Can be used before custom lowering. - * - * Note that ideally "Compiler" attributed functions would be made global functions as early as - * possible and would stay that way. However, the GraphExecutorCodegen and AOTExecutorCodegen - * assume the entire model can be represented by a single 'main' function, and the Inline pass - * is run to respect that assumption. So this pass is mostly just to undo that Pass after modules - * have passed through the 'codegen' keyhole. - * - * - (The \p OutlineCompilerFunctions pass is a more general version of the above which can use - * a custom cache to both allocate "global_symbol" names and ensure two structurally equal - * functions are assigned the same name, and thus lowered only once. This is used by Collage - * when preparing the optimally partitioned IRModule). - * - * - The \p MarkCompilerFunctionsAsExtern pass will update the attributes of global functions - * with a matching "Compiler" attribute to have just the "Extern" attribute. That will signal - * the function has been dealt with. However calls to such functions will be left unchanged. - * Can be used after lowering to cleanup the IRModule. - * - * - The \p InlineCompilerFunctions pass can selectively inline global functions with a matching - * "Compiler" attribute who's name appears in the given set. Obviously it's more sensible to - * not create that function in the first place, however some external codegen have rules to - * accept or reject partitionings based on the overall partitioned function body. This pass - * can be used do the legwork, and will take care to not only inline the outer "Compiler" - * annotated funcition, but also any "Composite" annotated functions in its body. - */ - -#ifndef TVM_RELAY_TRANSFORMS_COMPILER_FUNCTION_UTILS_H_ -#define TVM_RELAY_TRANSFORMS_COMPILER_FUNCTION_UTILS_H_ - -#include -#include -#include - -#include "tvm/ir/transform.h" -#include "tvm/relay/function.h" - -namespace tvm { -namespace relay { -namespace transform { - -/*! - * \brief Abstract class representing a cache of unique global vars keyed by functions. This can - * be used to ensure structurally equal functions are assigned the same global var object, and - * thus lowered at most once. - */ -class GlobalSymbolCache { - public: - virtual ~GlobalSymbolCache(); - virtual GlobalVar GetGlobalSymbol(const Function& function) = 0; -}; - -/*! - * \brief A \p GlobalSymbolCache that requires every "Compiler" attributed function to already - * have a "global_symbol" attribute. - */ -class ExistingGlobalSymbolCache : public GlobalSymbolCache { - public: - ExistingGlobalSymbolCache() = default; - - GlobalVar GetGlobalSymbol(const Function& function) final; - - private: - /*! \brief Maps already seen global symbol names to their corresponding GlobalVar objects. */ - std::unordered_map global_vars_; -}; - -/*! - * \brief A pass to outline all let-bound and literal functions in direct call positions which have - * a "Compiler" attribute. The given \p GlobalSymbolCache is used to determine a unique global - * symbol for each function, which is also assigned to the "global_symbol" attribute of the new - * global function. - * - * At most one function with the same global symbol is outlined. - * - * If \p compiler_filter is non-empty only functions with that as their attribute value are - * outlined. - */ -tvm::transform::Pass OutlineCompilerFunctions(std::shared_ptr cache, - std::string compiler_filter = ""); - -/*! - * \brief A pass to outline all let-bound and literal functions in direct call positions which have - * a "Compiler" attribute. The functions are bound to unique global vars according to their - * existing "global_symbol" attribute. At most one function with the same global symbol is outlined. - * - * If \p compiler_filter is non-empty only functions with that as their attribute value are - * outlined. - * - * This pass may be useful for external codegen using the "RelayToTIR" custom pass mechanism - * to prepare the IRModule before custom lowering. - */ -tvm::transform::Pass OutlineCompilerFunctionsWithExistingGlobalSymbols( - std::string compiler_filter = ""); - -/*! - * \brief A pass to mark all global functions which have a "Compiler" attribute matching - * compiler_filter as 'extern' by replacing all attributes with a single "Extern" attribute. - * Calls to such functions are not changed. - * - * If \p compiler_filter is non-empty only functions with that as their attribute value are - * outlined. - * - * This pass may be useful for external codegen using the "RelayToTIR" custom pass mechanism to - * cleanup the IRModule after custom lowering. - */ -tvm::transform::Pass MarkCompilerFunctionsAsExtern(std::string compiler_filter = ""); - -/*! - * \brief A pass to inline all global "Compiler" functions which are bound to a global var - * in \p global_vars. Both the global function and any calls to "Composite" functions it its body - * are inlined. - * - * This pass may be useful for external codegen which needs to undo partitioning based on - * properties of the entire partition. - */ -tvm::transform::Pass InlineCompilerFunctionsBoundTo(Array global_vars); - -} // namespace transform -} // namespace relay -} // namespace tvm - -#endif // TVM_RELAY_TRANSFORMS_COMPILER_FUNCTION_UTILS_H_ diff --git a/src/relay/transforms/convert_layout.cc b/src/relay/transforms/convert_layout.cc deleted file mode 100644 index e10be508529e..000000000000 --- a/src/relay/transforms/convert_layout.cc +++ /dev/null @@ -1,171 +0,0 @@ -/* - * 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 convert_op_layout.cc - * \brief Alternate the layouts of operators or replace primitive operators with - other expressions. This pass can be used for computing convolution in - custom layouts or other general weight pre-transformation. - */ -#include -#include -#include -#include -#include - -#include -#include -#include -#include -#include -#include - -#include "pattern_utils.h" -#include "transform_layout.h" - -namespace tvm { -namespace relay { - -namespace convert_op_layout { - -/*! - * \brief Container for the transformations for ConvertLayout. - */ -class ConvertTransformMemorizerNode : public TransformMemorizerNode { - public: - /*! - * \brief Initializes the desired_layout. - * \param desired_layouts Specify mapping of op_name to array of desired layouts for each input. - * For example: Map("nn.conv2d", Array("NHWC", "OHWI")), - * this specifies the desired layout for data then kernel for nn.conv2d. - */ - explicit ConvertTransformMemorizerNode(Map> desired_layouts) - : desired_layouts_(std::move(desired_layouts)) {} - - /*! - * \brief Defines the call transformation for ConvertLayout pass. The new layouts should be the - * desired layout as specified by the user. - * \param ref_call The original call. - * \param new_attrs Updated attributes consistent with new layouts. - * \param new_args The traversed/recursed args to the call. - * \return The new Call after calling the packed func. - */ - Call CallWithNewLayouts(const Call& ref_call, Attrs new_attrs, - const std::vector& new_args) override { - static auto fconvert_layout = Op::GetAttrMap("FTVMConvertOpLayout"); - Op op = Downcast(ref_call->op); - Expr new_e; - bool modified = false; - if (fconvert_layout.count(op)) { - auto desired_layouts = desired_layouts_; - if (desired_layouts.find(op->name) != desired_layouts.end()) { - tvm::Array tinfos; - for (auto& expr : ref_call->args) { - if (expr->checked_type()->IsInstance()) { - auto tuple_ttype_node = expr->type_as(); - for (auto& ttype : tuple_ttype_node->fields) { - auto ttype_node = ttype.as(); - tinfos.push_back(tvm::te::placeholder(ttype_node->shape, ttype_node->dtype)); - } - } else { - auto ttype = expr->type_as(); - tinfos.push_back(tvm::te::placeholder(ttype->shape, ttype->dtype)); - } - } - - Array op_desired_layouts = desired_layouts.at(op->name); - Expr altered_value = fconvert_layout[op](new_attrs, new_args, tinfos, op_desired_layouts); - if (altered_value.defined()) { - new_e = altered_value; - modified = true; - } - } else { - LOG(WARNING) << "Desired layout(s) not specified for op: " << op->name; - } - } - if (!modified) { - new_e = Call(ref_call->op, new_args, new_attrs); - } - - const CallNode* new_call = new_e.as(); - ICHECK(new_call) << "Can only replace the original operator with another call node"; - return Call(new_call->op, new_call->args, new_call->attrs, new_call->type_args, ref_call->span); - } - - Call CallWithNewLayouts(const Call& ref_call, const std::vector& new_args) override { - return CallWithNewLayouts(ref_call, ref_call->attrs, new_args); - } - - /*! \brief A mapping of op_name to array of desired layouts for each input. */ - Map> desired_layouts_; -}; - -/*! - * \brief Container that provides the transformation function for convert layout. - */ -class ConvertTransformMemorizer : public TransformMemorizer { - public: - ConvertTransformMemorizer() = default; - explicit ConvertTransformMemorizer(ObjectPtr n) : TransformMemorizer(n) {} - - ConvertTransformMemorizerNode* operator->() { - return static_cast(get_mutable()); - } - - using ContainerType = ConvertTransformMemorizerNode; -}; - -/*! - * Limitations: - * 1. The altered op should have the same number of arguments as the previous one. - * 2. Do not support nested tuple arguments. - */ -Expr ConvertLayout(const Expr& expr, const Map>& desired_layouts) { - ConvertTransformMemorizer transformMemorizer( - make_object(desired_layouts)); - auto fcontext = [&](const Call& call) -> ObjectRef { return transformMemorizer; }; - - return ForwardRewrite(expr, LayoutRewriter, fcontext); -} - -} // namespace convert_op_layout - -namespace transform { - -Pass ConvertLayout(const Map>& desired_layouts) { - runtime::TypedPackedFunc pass_func = - [=](Function f, IRModule m, PassContext pc) { - return Downcast(relay::convert_op_layout::ConvertLayout(f, desired_layouts)); - }; - return CreateFunctionPass(pass_func, 3, "ConvertLayout", {"InferType", "CanonicalizeOps"}); -} - -TVM_REGISTER_GLOBAL("relay._transform.ConvertLayout").set_body_typed(ConvertLayout); - -TVM_REGISTER_GLOBAL("relay._transform.InferCorrectLayoutOutput") - .set_body_typed([](Array input_layouts, Array output_layouts, Attrs new_attrs) { - return InferCorrectLayoutOutput(input_layouts, output_layouts, new_attrs); - }); - -TVM_REGISTER_NODE_TYPE(InferCorrectLayoutOutputNode); - -} // namespace transform - -} // namespace relay -} // namespace tvm diff --git a/src/relay/transforms/convert_sparse_conv2d.cc b/src/relay/transforms/convert_sparse_conv2d.cc deleted file mode 100644 index 0c6a6bd3a834..000000000000 --- a/src/relay/transforms/convert_sparse_conv2d.cc +++ /dev/null @@ -1,324 +0,0 @@ -/* - * 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 convert_sparse_conv2d.cc - * - * \brief Mutate conv2d operator to sparse conv2d operator - */ -#include -#include -#include -#include -#include -#include -#include - -#include -#include - -namespace tvm { -namespace relay { - -// Search conv2d op weight name from Expr -class Conv2dOpWeightVisitor : private ExprVisitor { - public: - Conv2dOpWeightVisitor() : conv2d_op_(Op::Get("nn.conv2d")) {} - - Array Search(const Expr& expr) { - VisitExpr(expr); - return memo_; - } - - private: - void VisitExpr_(const CallNode* n) final { - if (n->op == conv2d_op_) { - const auto weight = n->args[1].as(); - if (weight) { - memo_.push_back(weight->name_hint()); - } - } - for (const auto& arg : n->args) { - VisitExpr(arg); - } - } - // Cache op - const Op& conv2d_op_; - - Array memo_; -}; // SearchConv2dOpWeight - -Array SearchConv2dOpWeight(const Expr& e) { return Conv2dOpWeightVisitor().Search(e); } - -TVM_REGISTER_GLOBAL("relay.analysis.search_conv2d_op_weight").set_body_typed(SearchConv2dOpWeight); - -// Mutate ```nn.conv2d``` to ```nn.sparse_conv2d``` -class Conv2dToSparseConv2dMutator : public ExprRewriter { - public: - Conv2dToSparseConv2dMutator(const Array& weight_name, - const Array>& weight_shape, const String& layout, - int kernel_size) - : conv2d_op_(Op::Get("nn.conv2d")), sparse_conv2d_op_(Op::Get("nn.sparse_conv2d")) { - ICHECK_EQ(weight_name.size(), weight_shape.size()); - layout_ = layout; - kernel_size_ = kernel_size; - for (size_t i = 0; i < weight_name.size(); ++i) { - ICHECK(weight_name[i]->IsInstance()); - std::string k = weight_name[i].as()->data; - const auto& ws = weight_shape[i]; - std::vector v(ws.size()); - for (size_t j = 0; j < ws.size(); ++j) { - v[j] = ws[j].as()->value; - } - target_weights_.emplace(k, v); - } - } - - Expr Rewrite_(const CallNode* pre, const Expr& post) override { - if (pre->op == conv2d_op_) { - const auto weight = pre->args[1].as(); - if (weight) { - if (target_weights_.count(weight->name_hint())) { - const auto& prefix = weight->name_hint(); - const auto& ws = target_weights_.at(prefix); - const auto data = post.as()->args[0]; - relay::TensorType ws_data_type, ws_indices_type, ws_indptr_type; - if (ws.size() == 5) { - ws_data_type = relay::TensorType({ws.at(0), ws.at(1), ws.at(2)}, DataType::Float(32)); - ws_indices_type = relay::TensorType({ws.at(3)}, DataType::Int(32)); - ws_indptr_type = relay::TensorType({ws.at(4)}, DataType::Int(32)); - } else if (ws.size() == 4) { - ws_data_type = relay::TensorType({ws.at(0), ws.at(1)}, DataType::Float(32)); - ws_indices_type = relay::TensorType({ws.at(2)}, DataType::Int(32)); - ws_indptr_type = relay::TensorType({ws.at(3)}, DataType::Int(32)); - } - Var weight_data(prefix + ".data", ws_data_type); - Var weight_indices(prefix + ".indices", ws_indices_type); - Var weight_indptr(prefix + ".indptr", ws_indptr_type); - auto attrs = make_object(); - attrs->layout = std::move(layout_); - attrs->kernel_size = Array{kernel_size_, kernel_size_}; - return Call(sparse_conv2d_op_, {data, weight_data, weight_indices, weight_indptr}, - Attrs(attrs)); - } - } - } - return post; - } - - private: - // Cached op - const Op& conv2d_op_; - const Op& sparse_conv2d_op_; - std::unordered_map> target_weights_; - String layout_; - int kernel_size_; -}; // class Conv2dToSparseConv2dAlter - -Expr Conv2dToSparse(const Expr& e, const Array& weight_name, - const Array>& weight_shape, const String& layout, - int kernel_size) { - auto rewriter = Conv2dToSparseConv2dMutator(weight_name, weight_shape, layout, kernel_size); - return PostOrderRewrite(e, &rewriter); -} - -template -auto unpack_to_tuple_internal(elemTy* arr, std::index_sequence) { - return std::make_tuple(arr[Is]...); -} - -template -auto unpack_to_tuple(elemTy* arr) { - return unpack_to_tuple_internal(arr, std::make_index_sequence{}); -} - -struct Range { - size_t dim; - explicit Range(size_t d) : dim(d) {} - - struct iterpoint { - size_t val, lim; - iterpoint(size_t v1, size_t v2) : val(v1), lim(v2) {} - - size_t operator*() const { return val; } - - iterpoint operator/(const iterpoint& rhs) const { - return iterpoint(val * rhs.lim + rhs.val, lim * rhs.lim); - } - }; - - struct iterator { - size_t val, lim; - iterator(size_t v1, size_t v2) : val(v1), lim(v2) {} - - bool operator!=(const iterator& rhs) const { return val != rhs.val; } - - void operator++() { ++val; } - - iterpoint operator*() const { return iterpoint(val, lim); } - }; - - iterator begin() { return iterator(0, dim); } - - iterator end() { return iterator(dim, dim); } -}; - -// Mutate ```nn.conv2d``` to ```nn.sparse_conv2d``` -class Conv2dToSparseConv2dMutator2 : public ExprRewriter { - public: - Conv2dToSparseConv2dMutator2(const String& layout, int kernel_size, int blockH, int blockW, - double sparse_thresh) - : sparse_conv2d_op_(Op::Get("nn.sparse_conv2d")), - dev_cpu0_{DLDeviceType::kDLCPU, 0}, - layout_(layout), - kernel_size_(kernel_size), - blockH_(blockH), - blockW_(blockW), - sparse_thresh_(sparse_thresh) {} - - Expr Rewrite_(const CallNode* pre, const Expr& post) override { - // check op type & attrs - const auto pre_attrs = pre->attrs.as(); - if (!pre_attrs || pre_attrs->data_layout != layout_ || - pre_attrs->strides[0].as()->value != 1 || - pre_attrs->kernel_size[0].as()->value != kernel_size_) - return post; - // check constant weight - const auto pre_weight_node = pre->args[1].as(); - if (!pre_weight_node) return post; - - // check weight dtype & shape - auto&& pre_weight = pre_weight_node->data; - auto dtype = pre_weight.DataType(), itype = runtime::DataType::Int(32); - ICHECK(dtype.code() == DataType::kFloat && dtype.bits() == 32); // float32 only - auto pre_weight_shape = unpack_to_tuple<4>(pre_weight.Shape().data()); - int O, I, H, W; - if (layout_ == "NCHW") { - std::tie(O, I, H, W) = pre_weight_shape; - } else { // NHWC - std::tie(H, W, I, O) = pre_weight_shape; - } - int CO = O, CI = H * W * I; - - // copy to vector - std::vector pre_weight_data(CO * CI); - pre_weight.CopyToBytes(pre_weight_data.data(), pre_weight_data.size() * sizeof(float)); - if (layout_ == "NHWC") { - std::vector tmp(pre_weight_data.size()); - for (auto i : Range(CO)) - for (auto j : Range(CI)) tmp[*(i / j)] = pre_weight_data[*(j / i)]; - std::swap(tmp, pre_weight_data); - } - // convert to BSR - std::vector wdata, block(blockH_ * blockW_); - std::vector windices, windptr; - for (auto bh : Range(CO / blockH_)) { - windptr.push_back(windices.size()); - for (auto bw : Range(CI / blockW_)) { - int cntnnz = 0; - for (auto i : Range(blockH_)) - for (auto j : Range(blockW_)) { - auto tmp = pre_weight_data[*(bh / i / bw / j)]; - if (tmp) cntnnz++; - block[*(i / j)] = tmp; - } - if (cntnnz) { - wdata.insert(wdata.end(), block.begin(), block.end()); - windices.push_back(*bw); - } - } - } - windptr.push_back(windices.size()); - double sprate = 1 - 1.0 * wdata.size() / pre_weight_data.size(); - if (sprate < sparse_thresh_) return post; - - // constrct return data - int nnz = windices.size(); - auto weight_data = runtime::NDArray::Empty({nnz, blockH_, blockW_}, dtype, dev_cpu0_); - auto weight_indices = runtime::NDArray::Empty({nnz}, itype, dev_cpu0_); - auto weight_indptr = runtime::NDArray::Empty({CO / blockH_ + 1}, itype, dev_cpu0_); - weight_data.CopyFromBytes(wdata.data(), wdata.size() * sizeof(float)); - weight_indices.CopyFromBytes(windices.data(), windices.size() * sizeof(int32_t)); - weight_indptr.CopyFromBytes(windptr.data(), windptr.size() * sizeof(int32_t)); - - // construct return call - auto args = runtime::Array{post.as()->args[0], Constant(weight_data), - Constant(weight_indices), Constant(weight_indptr)}; - auto attrs = make_object(); - attrs->layout = layout_; - attrs->kernel_size = Array{kernel_size_, kernel_size_}; - return Call(sparse_conv2d_op_, args, Attrs(attrs)); - } - - private: - const Op& sparse_conv2d_op_; - DLDevice dev_cpu0_; - String layout_; - int kernel_size_, blockH_, blockW_; - double sparse_thresh_; -}; // class Conv2dToSparseConv2dMutator2 - -Expr Conv2dToSparse2(const Expr& e, const String& layout, int kernel_size, int blockH, int blockW, - double sparse_thresh) { - auto rewriter = Conv2dToSparseConv2dMutator2(layout, kernel_size, blockH, blockW, sparse_thresh); - return PostOrderRewrite(e, &rewriter); -} - -namespace transform { - -// Convert a model with separate weight info (already sparsified). -Pass Conv2dToSparse(const Array& weight_name, const Array>& weight_shape, - const String& layout, int kernel_size) { - runtime::TypedPackedFunc pass_func = - [=](Function f, IRModule m, PassContext pc) { - // Remove FreeVar warnings - auto f0 = - Downcast(Conv2dToSparse(f, weight_name, weight_shape, layout, kernel_size)); - Array sparse_params = FreeVars(f0); - auto f1 = WithFields(f0, sparse_params); - Array params = FreeVars(f1); - for (const auto& var : sparse_params) { - params.push_back(var); - } - return WithFields(f1, params); - }; - return CreateFunctionPass(pass_func, 4, "Conv2dToSparse", {"DeadCodeElimination"}); -} - -TVM_REGISTER_GLOBAL("relay._transform.Conv2dToSparse").set_body_typed(Conv2dToSparse); - -// Convert a model with freezed params (sparsified in the pass). -Pass Conv2dToSparse2(const String& layout, int kernel_size, int blockH, int blockW, - double sparse_thresh) { - runtime::TypedPackedFunc pass_func = - [=](Function f, IRModule m, PassContext pc) { - auto f0 = Downcast( - Conv2dToSparse2(f, layout, kernel_size, blockH, blockW, sparse_thresh)); - return f0; - }; - return CreateFunctionPass(pass_func, 5, "Conv2dToSparse2", {"DeadCodeElimination"}); -} - -TVM_REGISTER_GLOBAL("relay._transform.Conv2dToSparse2").set_body_typed(Conv2dToSparse2); - -} // namespace transform - -} // namespace relay -} // namespace tvm diff --git a/src/relay/transforms/convert_sparse_dense.cc b/src/relay/transforms/convert_sparse_dense.cc deleted file mode 100644 index 7053f1301cca..000000000000 --- a/src/relay/transforms/convert_sparse_dense.cc +++ /dev/null @@ -1,153 +0,0 @@ -/* - * 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 convert_sparse_dense.cc - * - * \brief Mutate dense operator to sparse dense operator - */ -#include -#include -#include -#include -#include -#include -#include - -#include -#include - -namespace tvm { -namespace relay { - -// Search dense op weight name from Expr -class DenseOpWeightVisitor : private ExprVisitor { - public: - DenseOpWeightVisitor() : dense_op_(Op::Get("nn.dense")) {} - - Array Search(const Expr& expr) { - VisitExpr(expr); - return memo_; - } - - private: - void VisitExpr_(const CallNode* n) final { - if (n->op == dense_op_) { - const auto weight = n->args[1].as(); - if (weight) { - memo_.push_back(weight->name_hint()); - } - } - for (const auto& arg : n->args) { - VisitExpr(arg); - } - } - // Cache op - const Op& dense_op_; - - Array memo_; -}; // SearchDenseOpWeight - -Array SearchDenseOpWeight(const Expr& e) { return DenseOpWeightVisitor().Search(e); } - -TVM_REGISTER_GLOBAL("relay.analysis.search_dense_op_weight").set_body_typed(SearchDenseOpWeight); - -// Mutate ```nn.dense``` to ```nn.sparse_dense``` -class DenseToSparseDenseMutator : public ExprRewriter { - public: - DenseToSparseDenseMutator(const Array& weight_name, - const Array>& weight_shape) - : dense_op_(Op::Get("nn.dense")), sparse_dense_op_(Op::Get("nn.sparse_dense")) { - ICHECK_EQ(weight_name.size(), weight_shape.size()); - for (size_t i = 0; i < weight_name.size(); ++i) { - ICHECK(weight_name[i]->IsInstance()); - std::string k = weight_name[i].as()->data; - const auto& ws = weight_shape[i]; - std::vector v(ws.size()); - for (size_t j = 0; j < ws.size(); ++j) { - v[j] = ws[j].as()->value; - } - target_weights_.emplace(k, v); - } - } - - Expr Rewrite_(const CallNode* pre, const Expr& post) override { - if (pre->op == dense_op_) { - const auto weight = pre->args[1].as(); - if (weight) { - if (target_weights_.count(weight->name_hint())) { - const auto& prefix = weight->name_hint(); - const auto& ws = target_weights_.at(prefix); - const auto data = post.as()->args[0]; - auto ws_data_type = - relay::TensorType({ws.at(0), ws.at(1), ws.at(2)}, DataType::Float(32)); - auto ws_indices_type = relay::TensorType({ws.at(3)}, DataType::Int(32)); - auto ws_indptr_type = relay::TensorType({ws.at(4)}, DataType::Int(32)); - Var weight_data(prefix + ".data", ws_data_type); - Var weight_indices(prefix + ".indices", ws_indices_type); - Var weight_indptr(prefix + ".indptr", ws_indptr_type); - auto attrs = make_object(); - - return Call(sparse_dense_op_, {data, weight_data, weight_indices, weight_indptr}, - Attrs(attrs)); - } - } - } - return post; - } - - private: - // Cached op - const Op& dense_op_; - const Op& sparse_dense_op_; - std::unordered_map> target_weights_; -}; // class DenseToSparseDenseAlter - -Expr DenseToSparse(const Expr& e, const Array& weight_name, - const Array>& weight_shape) { - auto rewriter = DenseToSparseDenseMutator(weight_name, weight_shape); - return PostOrderRewrite(e, &rewriter); -} - -namespace transform { - -Pass DenseToSparse(const Array& weight_name, - const Array>& weight_shape) { - runtime::TypedPackedFunc pass_func = - [=](Function f, IRModule m, PassContext pc) { - // Remove FreeVar warnings - auto f0 = Downcast(DenseToSparse(f, weight_name, weight_shape)); - Array sparse_params = FreeVars(f0); - auto f1 = WithFields(f0, sparse_params); - Array params = FreeVars(f1); - for (const auto& var : sparse_params) { - params.push_back(var); - } - return WithFields(f1, params); - }; - return CreateFunctionPass(pass_func, 4, "DenseToSparse", {"DeadCodeElimination"}); -} - -TVM_REGISTER_GLOBAL("relay._transform.DenseToSparse").set_body_typed(DenseToSparse); - -} // namespace transform - -} // namespace relay -} // namespace tvm diff --git a/src/relay/transforms/de_duplicate.cc b/src/relay/transforms/de_duplicate.cc deleted file mode 100644 index 23e147d5d4c4..000000000000 --- a/src/relay/transforms/de_duplicate.cc +++ /dev/null @@ -1,123 +0,0 @@ -/* - * 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 de_duplicate.cc - * \brief Use a fresh Id for every Var to make the result well-formed. - */ -#include -#include -#include -#include - -#include - -namespace tvm { -namespace relay { - -Expr DeDup(const Expr& e) { - class DeDupMutator : public TypeMutator, public MixedModeMutator, public PatternMutator { - public: - TypeVar Fresh(const TypeVar& tv) { - TypeVar ret = TypeVar(tv->name_hint, tv->kind); - type_rename_[tv] = ret; - return ret; - } - - Var Fresh(const Var& v) { - ICHECK_EQ(rename_.count(v), 0); - ICHECK_EQ(memo_.count(v), 0) << v.as(); - Var ret = Var(v->name_hint(), VisitType(v->type_annotation)); - rename_[v] = ret; - return ret; - } - - Expr DispatchVisitExpr(const Expr& e) final { - auto ret = ExprMutator::VisitExpr(e); - ret->checked_type_ = e->checked_type_; - ret->virtual_device_ = e->virtual_device_; - return ret; - } - - using MixedModeMutator::VisitExpr_; - - Expr VisitExpr_(const VarNode* op) final { - Var v = GetRef(op); - return rename_.count(v) != 0 ? rename_.at(v) : v; - } - - Expr VisitExpr_(const LetNode* op) final { - std::unordered_map new_vars; - auto pre_visit = [this, &new_vars](const LetNode* op) { - Expr expr = GetRef(op); - new_vars[expr] = this->Fresh(op->var); - // Rely on the Memoizer to cache pre-visit values - this->VisitExpr(op->value); - }; - auto post_visit = [this, &new_vars](const LetNode* op) { - Expr expr = GetRef(op); - this->memo_[expr] = - Let(new_vars[expr], this->VisitExpr(op->value), this->VisitExpr(op->body)); - }; - ExpandANormalForm(op, pre_visit, post_visit); - return memo_[GetRef(op)]; - } - - Type VisitType(const Type& t) final { return t.defined() ? TypeMutator::VisitType(t) : t; } - - Expr VisitExpr_(const FunctionNode* func_node) final { - tvm::Array type_params; - for (const TypeVar& type_param : func_node->type_params) { - type_params.push_back(Fresh(type_param)); - } - tvm::Array params; - for (const Var& param : func_node->params) { - params.push_back(Fresh(param)); - } - return WithFields(GetRef(func_node), params, VisitExpr(func_node->body), - VisitType(func_node->ret_type), type_params); - } - - Pattern VisitPattern(const Pattern& p) final { return PatternFunctor::VisitPattern(p); } - - Pattern VisitPattern_(const PatternVarNode* op) final { return PatternVar(Fresh(op->var)); } - - Type VisitType_(const TypeVarNode* op) final { - TypeVar v = GetRef(op); - return type_rename_.count(v) != 0 ? type_rename_.at(v) : v; - } - - Var VisitVar(const Var& v) final { return Fresh(v); } - - private: - std::unordered_map rename_; - std::unordered_map type_rename_; - }; - ICHECK(WellFormed(e)) << AsText(e, false); - Expr ret = DeDupMutator().VisitExpr(e); - ICHECK(WellFormed(ret)); - ICHECK_EQ(FreeVars(e).size(), FreeVars(ret).size()); - return ret; -} // namespace relay - -TVM_REGISTER_GLOBAL("relay._transform.dedup").set_body_typed(DeDup); - -} // namespace relay -} // namespace tvm diff --git a/src/relay/transforms/dead_code.cc b/src/relay/transforms/dead_code.cc deleted file mode 100644 index e2b350a439ea..000000000000 --- a/src/relay/transforms/dead_code.cc +++ /dev/null @@ -1,585 +0,0 @@ -/* - * 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 src/relay/transforms/dead_code.cc - * \brief Elides or inlines let-bindings. - * - * TODO(mbs): Track dead writes into references. - */ - -#include -#include -#include -#include - -#include "../op/call/call.h" - -namespace tvm { -namespace relay { -namespace { - -/*! \brief Maximum depth of calls to analyize. */ -constexpr int kMaxCallDepth = 25; - -/*! - * \brief Captures (an approximation of) the purity for a Relay sub-expression. A pure - * sub-expression is guaranteed never to access or mutate state. Thus the sub-expression - * can safely be elided (if its result is never used), or inlined (which may change the - * number of times and program order for the evaluation.) - */ -struct Purity { - /*! - * \brief True if evaling the sub-expression itself is pure. - */ - bool pure_eval; - /*! - * \brief If the sub-expression is first-order then always true. Otherwise true only if evaling - * a call to the sub-expression is pure. See [RULE A] below. - */ - bool pure_call; -}; - -/*! - * \brief Visits all the global functions in a module and records the purity of every let-bound - * value. - * - * (See also inline.cc for function inlining.) - * - * Generally we track whether evaluation of a sub-expression is definitely pure. However for - * sub-expressions f of higher-order type we also track the 'call purity' of evaling a call to f: - * - [RULE A] If f's result is itself higher-order then f is call-pure only if the result of f is - * also call-pure. - * - [RULE B] Higher-order function arguments are assumed call impure. - * - [RULE C] We assume functions extracted from tuples are call impure. - * - [RULE D] We assume functions extracted from references are call impure. - * - [RULE E] We assume functions extracted from ADTs are call impure. - * - [RULE F] We assume all external Functions and PrimFuncs are call impure. - */ -class PurityVisitor : ExprFunctor { - public: - explicit PurityVisitor(IRModule mod) : mod_(std::move(mod)), current_call_depth_(0) {} - - /*! \brief Visit all the functions in the module. */ - void VisitModule() { - VLOG_CONTEXT << "PurityVisitor"; - // It is safe to visit the global functions in any order. Recursive global functions are - // allowed. - for (const auto& kv : mod_->functions) { - if (const auto* function_node = kv.second.as()) { - if (function_node->HasNonzeroAttr(attr::kPrimitive) || - function_node->HasNonzeroAttr(attr::kExtern)) { - // Ignore primitive and external functions. - continue; - } - // Everything of interest will be recorded in the purity maps so we ignore the result. - (void)VisitGlobalFunction(kv.first, GetRef(function_node)); - } - } - } - - /*! - * \brief Returns a map from every let-bound variable to whether its let-bound value is - * definitely pure. - */ - std::unordered_map GetPurityMap() const { - std::unordered_map result; - for (const auto& kv : var_to_purity_) { - result.emplace(kv.first, kv.second.pure_eval); - } - return result; - } - - private: - Purity VisitExpr(const Expr& expr) final { - auto it = memo_.find(expr.get()); - if (it != this->memo_.end()) { - return it->second; - } else { - Purity result = ExprFunctor::VisitExpr(expr); - memo_[expr.get()] = result; - return result; - } - } - - Purity VisitExpr_(const ConstantNode*) final { return {/*pure_eval=*/true, /*pure_call=*/true}; } - - Purity VisitExpr_(const ConstructorNode*) final { - return {/*pure_eval=*/true, /*pure_call=*/true}; - } - - Purity VisitExpr_(const OpNode* op_node) final { - // Primitive operators are pure unless marked as 'stateful'. - static OpAttrMap attr_map = Op::GetAttrMap("TOpIsStateful"); - bool is_stateful = attr_map.count(GetRef(op_node)) && attr_map[GetRef(op_node)]; - return {/*pure_eval=*/true, /*pure_call=*/!is_stateful}; - } - - Purity VisitExpr_(const GlobalVarNode* global_var_node) final { - auto global_var = GetRef(global_var_node); - ICHECK(mod_->ContainGlobalVar(global_var_node->name_hint)) - << "No definition for '" << global_var_node->name_hint << "'"; - auto func = mod_->Lookup(global_var); - if (const auto* function_node = func.as()) { - if (!function_node->HasNonzeroAttr(attr::kExtern)) { - return VisitGlobalFunction(global_var, GetRef(function_node)); - } - } - // Assume externals and PrimFuncs are call-impure [RULE F]. - // (If they are pure then we should have dealt with them before lowering.) - return {/*pure_eval==*/true, /*pure_call=*/false}; - } - - Purity VisitExpr_(const VarNode* var_node) final { - // The var is bound to a value, but if that value is a function we need to propagate the - // function body's purity. - ICHECK(var_to_purity_.count(var_node)) << PrettyPrint(GetRef(var_node)); - return {/*pure_eval=*/true, /*pure_call=*/var_to_purity_[var_node].pure_call}; - } - - Purity VisitExpr_(const FunctionNode* function_node) final { - for (const auto& param : function_node->params) { - // Any higher-order parameters are assumed to be call-impure [RULE B] - var_to_purity_[param.get()] = {/*pure_eval=*/true, /*pure_call=*/IsFirstOrder(param)}; - } - Purity body_purity = VisitExpr(function_node->body); - // The function itself is a value and thus pure. If the function returns - // a function we'll fold its purity in here [RULE A] - return {/*pure_eval=*/true, /*pure_call=*/body_purity.pure_eval && body_purity.pure_call}; - } - - Purity VisitExpr_(const LetNode* let_node) final { - Expr expr = GetRef(let_node); - bool all_values_pure_eval = true; - while (const auto* inner_let_node = expr.as()) { - // In case the value is a recursive function assume the let-bound variable is call-pure. - var_to_purity_[inner_let_node->var.get()] = {/*pure_eval=*/true, /*pure_call=*/true}; - Purity value_purity = VisitExpr(inner_let_node->value); - // Now revise the variable to it's true purity. - var_to_purity_[inner_let_node->var.get()] = value_purity; - VLOG(2) << (value_purity.pure_eval ? "pure" : "impure") << " expression:" << std::endl - << PrettyPrint(inner_let_node->value) << std::endl - << "let-bound to variable:" << std::endl - << PrettyPrint(inner_let_node->var); - all_values_pure_eval = all_values_pure_eval && value_purity.pure_eval; - expr = inner_let_node->body; - } - Purity body_purity = VisitExpr(expr); - return {/*pure_eval=*/all_values_pure_eval && body_purity.pure_eval, - /*pure_call=*/body_purity.pure_call}; - } - - Purity VisitExpr_(const CallNode* call_node) final { - auto call = GetRef(call_node); - if (current_call_depth_ >= kMaxCallDepth) { - // Assume impure. - VLOG(2) << "assuming call is impure since too deeply nested"; - return {/*pure_eval=*/false, /*pure_call*/ IsFirstOrder(call)}; - } - - ++current_call_depth_; - - // We can work with calls in both pre- and post-lowered form. - Call vanilla_call = GetAnyCall(call_node); - - // Find purity for the callee and the args. - Purity callee_purity = VisitExpr(vanilla_call->op); - bool all_args_pure_eval = true; - for (const auto& arg : vanilla_call->args) { - Purity arg_purity = VisitExpr(arg); - all_args_pure_eval = all_args_pure_eval && arg_purity.pure_eval; - } - - VLOG(2) << (callee_purity.pure_call ? "pure" : "impure") << " call to:" << std::endl - << PrettyPrint(vanilla_call->op); - - ICHECK_GT(current_call_depth_, 0); - --current_call_depth_; - - // If the callee's result is itself a function then by [RULE A] its purity - // is given by callee_purity.pure_call. - return {/*pure_eval=*/all_args_pure_eval && callee_purity.pure_eval && callee_purity.pure_call, - /*pure_call=*/IsFirstOrder(call) || callee_purity.pure_call}; - } - - Purity VisitExpr_(const IfNode* if_node) final { - Purity cond_purity = VisitExpr(if_node->cond); - ICHECK(cond_purity.pure_call); // conditional is first-order - Purity true_purity = VisitExpr(if_node->true_branch); - Purity false_purity = VisitExpr(if_node->false_branch); - return {/*pure_eval=*/cond_purity.pure_eval && true_purity.pure_eval && false_purity.pure_eval, - /*pure_call=*/true_purity.pure_call && false_purity.pure_call}; - } - - Purity VisitExpr_(const TupleNode* tuple_node) final { - bool all_fields_pure = true; - for (const auto& field : tuple_node->fields) { - // The call purity of each tuple field is lost [RULE C]. - Purity field_purity = VisitExpr(field); - if (!field_purity.pure_eval) { - all_fields_pure = false; - } - } - return {/*pure_eval=*/all_fields_pure, /*pure_call=*/true}; - } - - Purity VisitExpr_(const TupleGetItemNode* tuple_get_item_node) final { - Purity tuple_purity = VisitExpr(tuple_get_item_node->tuple); - ICHECK(tuple_purity.pure_call); // tuple is first-order - // We don't track call purity through tuple fields, so if the result is a function type we - // must assume it is call impure [RULE C]. - return {/*pure_eval=*/tuple_purity.pure_eval, - /*pure_call=*/IsFirstOrder(GetRef(tuple_get_item_node))}; - } - - Purity VisitExpr_(const RefCreateNode*) final { - // The creation of the ref itself is unobservable other than via the reads/writes into it. - return {/*pure_eval=*/true, /*pure_call=*/true}; - } - - Purity VisitExpr_(const RefWriteNode* ref_write_node) final { - Purity ref_purity = VisitExpr(ref_write_node->ref); - ICHECK(ref_purity.pure_call); // reference is first-order - // The call purity of the written value is lost [RULE D]. - // (But we must still visit to accumulate purity for any let-bindings within in.) - (void)VisitExpr(ref_write_node->value); - return {/*pure_eval=*/false, /*pure_call=*/true}; - } - - Purity VisitExpr_(const RefReadNode* ref_read_node) final { - Purity ref_purity = VisitExpr(ref_read_node->ref); - ICHECK(ref_purity.pure_call); // reference is first-order - // We don't track call purity through reference values, so if the result is a function - // type we must assume it is call impure [RULE D]. - return {/*pure_eval=*/false, /*pure_call=*/IsFirstOrder(GetRef(ref_read_node))}; - } - - class PurityPatternVisitor : public PatternVisitor { - public: - explicit PurityPatternVisitor(PurityVisitor* outer) : outer_(outer) {} - - private: - void VisitPattern_(const PatternVarNode* pattern_var_node) final { - // We don't track call purity through ADTs, so if var is a function type we must assume - // it is call impure [RULE E]. - outer_->var_to_purity_[pattern_var_node->var.get()] = { - /*pure_eval=*/true, /*pure_call=*/IsFirstOrder(pattern_var_node->var)}; - } - - /*! \brief (Mutable borrow of) the outer visitor. */ - PurityVisitor* outer_; - }; - - Purity VisitExpr_(const MatchNode* match_node) final { - Purity data_purity = VisitExpr(match_node->data); - ICHECK(data_purity.pure_call); // ADT is first order - bool all_clauses_pure_eval = true; - bool all_clauses_pure_call = true; - for (const auto& clause : match_node->clauses) { - PurityPatternVisitor pattern_visitor(this); - pattern_visitor.VisitPattern(clause->lhs); - Purity rhs_purity = VisitExpr(clause->rhs); - all_clauses_pure_eval = all_clauses_pure_eval && rhs_purity.pure_eval; - all_clauses_pure_call = all_clauses_pure_call && rhs_purity.pure_call; - } - return {/*pure_eval=*/data_purity.pure_eval && all_clauses_pure_eval, - /*pure_call=*/all_clauses_pure_call}; - } - - /*! \brief Visits \p func bound to global \p var and returns it's purity. */ - Purity VisitGlobalFunction(const GlobalVar& var, const Function& func) { - VLOG_CONTEXT << "func " << var->name_hint; - VLOG(2) << "visiting"; - auto itr = global_var_to_purity_.find(var.get()); - if (itr != global_var_to_purity_.end()) { - // We've already visited the function body. - return itr->second; - } - // We are entering the body of a possibly-recursive global function. Assume it's body is pure. - global_var_to_purity_[var.get()] = {/*pure_eval=*/true, /*pure_call=*/true}; - // Visit the global function for the first time. - Purity func_purity = VisitExpr(func); - // Update with the true purity. - global_var_to_purity_[var.get()] = func_purity; - return func_purity; - } - - static bool IsFirstOrder(const Expr& expr) { - return expr->checked_type().as() == nullptr; - } - - /*! \brief The module we're analyzing. */ - IRModule mod_; - - /*! - * \brief Maps each let-bound and global variable to the purity of the value it is bound to. - * If the variable is bound to a function then the purity of saturating that function is also - * tracked. - * - * Note that global_var_to_purity_, and all the 'pure_call' fields, are only needed internally - * during the analysis, andonly the var_to_purity_ 'pure_eval' fields are used downstream. - */ - std::unordered_map var_to_purity_; - std::unordered_map global_var_to_purity_; - - /*! \brief The current call depth. We'll just assume deeply nested calls are impure rather than - * spending all that time to check for sure. A deeply nested call is almost certain to be needed - * anyway. - */ - - int current_call_depth_; - - /*! \brief Internal map used for memoization. */ - std::unordered_map memo_; -}; - -/*! - * \brief Accumulate the bound values and usage count for each let-bound variable. - * - * We don't attempt to track the number of calls to local functions, and instead just assume they - * are called at least twice. - */ -class UsageVisitor : public ExprVisitor { - public: - /*! \brief Accumulates the expression bound to every let-bound variable. */ - std::unordered_map let_bound_values_; - /*! \brief Accumulates the usage count for every let-bound variable. */ - std::unordered_map use_map_; - - explicit UsageVisitor(const std::unordered_map* var_to_purity, - bool default_purity) - : var_to_purity_(var_to_purity), default_purity_(default_purity) {} - - void VisitExpr(const Expr& expr) final { - // Once we've seen 2 usages of a variable we know it can be neither elided nor inlined, - // so can stop visiting again. - if (++visit_counter_[expr.get()] <= 2) { - ExprFunctor::VisitExpr(expr); - } - } - - void VisitExpr_(const FunctionNode* function_node) final { - ++current_scope_level_; - ExprVisitor::VisitExpr_(function_node); - ICHECK_GT(current_scope_level_, 0); - --current_scope_level_; - } - - void VisitExpr_(const LetNode* let_node) final { - Expr expr = GetRef(let_node); - while (const auto* inner_let_node = expr.as()) { - ++visit_counter_[inner_let_node]; - let_bound_values_[inner_let_node->var.get()] = inner_let_node->value; - VLOG(2) << "seen let-binding for:" << std::endl << PrettyPrint(inner_let_node->var); - use_map_[inner_let_node->var.get()] = 0; - scope_level_map_[inner_let_node->var.get()] = current_scope_level_; - if (is_pure(inner_let_node->var.get())) { - // We'll defer visiting the let-bound value until we've seen the first use of the let-bound - // variable and thus know it must be evaluated. - // no-op. - } else { - // The let-bound value is impure so must always be evaluated. Visit now. - VisitExpr(inner_let_node->value); - } - expr = inner_let_node->body; - } - VisitExpr(expr); - } - - void VisitExpr_(const VarNode* var_node) final { - if (let_bound_values_.count(var_node)) { - size_t& n = use_map_[var_node]; - ++n; - VLOG(2) << var_node->name_hint() << " = " << n; - if (n == 1 && is_pure(var_node)) { - // Now that we have at least one use of the let-bound var, we know the let-bound - // value is necessary. - VisitExpr(let_bound_values_[var_node]); - } - if (scope_level_map_[var_node] < current_scope_level_) { - // Since the variable was bound outside of the current local function, assume the - // function will be called at least twice. - ++n; - VLOG(2) << var_node->name_hint() << " = " << n << " (bound at level " - << scope_level_map_[var_node] << " but used at level " << current_scope_level_ - << ")"; - } - } - // else: nothing to be done for function parameters or variable in match patterns. - } - - bool is_pure(const VarNode* var_node) const { - auto itr = var_to_purity_->find(var_node); - return itr == var_to_purity_->end() ? default_purity_ : itr->second; - } - - /*! \brief (Immutable borrow of) the already determined purity for every let-bound variable. */ - const std::unordered_map* var_to_purity_; - /*! \brief The default purity for variables which are not in the above map. */ - bool default_purity_; - /*! - * \brief The current scope level. 0 for global functions. Incremented by one within each - * let-bound local function. Necessary so we can avoid inlining an expensive let-bound computation - * into a function which could be called more than once. - */ - int current_scope_level_ = 0; - /*! \brief Accumulates the scope level for every let-bound variable. */ - std::unordered_map scope_level_map_; -}; - -/*! \brief Eliminate/inline let-bound values when sound to do so. */ -class EliminatorMutator : public ExprMutator { - public: - EliminatorMutator(bool inline_once, - const std::unordered_map* let_bound_values, - const std::unordered_map* use_map, - const std::unordered_map* var_to_purity, - bool default_purity) - : inline_once_(inline_once), - let_bound_values_(let_bound_values), - use_map_(use_map), - var_to_purity_(var_to_purity), - default_purity_(default_purity) {} - - private: - enum Action { kElide, kInline, kNoChange }; - - /*! \brief What should we do with let-binding for \p var_node? */ - Action ActionFor(const VarNode* var_node) { - if (let_bound_values_->count(var_node) == 0) { - // Not let-bound var. - return kNoChange; - } - if (!is_pure(var_node)) { - // The let-bound value is impure -- we must leave it exactly where it is. - return kNoChange; - } - switch (use_map_->count(var_node) ? use_map_->at(var_node) : 0) { - case 0: - return kElide; - case 1: - return inline_once_ ? kInline : kNoChange; - default: - return kNoChange; - } - } - - Expr VisitExpr_(const VarNode* var_node) final { - if (ActionFor(var_node) == kInline) { - VLOG(1) << "inlining let-bound variable:" << std::endl << PrettyPrint(GetRef(var_node)); - return VisitExpr(let_bound_values_->at(var_node)); - } else { - return GetRef(var_node); - } - } - - Expr VisitExpr_(const LetNode* op) final { - auto pre_visit = [this](const LetNode* op) { - if (ActionFor(op->var.get()) != kElide) { - (void)VisitExpr(op->value); - } - }; - auto post_visit = [this](const LetNode* op) { - Expr body = VisitExpr(op->body); - auto expr = GetRef(op); - switch (ActionFor(op->var.get())) { - case kElide: - VLOG(1) << "eliding let-bound variable:" << std::endl << PrettyPrint(op->var); - memo_[expr] = body; - break; - case kInline: - // Already inlined at use-side. - memo_[expr] = body; - break; - case kNoChange: - Expr value = VisitExpr(op->value); - memo_[expr] = Let(op->var, value, body); - break; - } - }; - ExpandANormalForm(op, pre_visit, post_visit); - return memo_[GetRef(op)]; - } - - bool is_pure(const VarNode* var_node) const { - auto itr = var_to_purity_->find(var_node); - return itr == var_to_purity_->end() ? default_purity_ : itr->second; - } - - bool inline_once_; - const std::unordered_map* let_bound_values_; - const std::unordered_map* use_map_; - const std::unordered_map* var_to_purity_; - bool default_purity_; -}; - -} // namespace - -namespace transform { - -// Declared in relay/transform.h -Pass DeadCodeElimination(bool inline_once, bool ignore_impurity) { - auto pass_func = [=](IRModule mod, PassContext pc) -> IRModule { - VLOG(1) << "Before:" << std::endl << PrettyPrint(mod); - // Which let bindings are pure and can be safely elided? - std::unordered_map var_to_purity; - if (!ignore_impurity) { - VLOG(1) << "determine purity"; - PurityVisitor purity_visitor(mod); - purity_visitor.VisitModule(); - var_to_purity = purity_visitor.GetPurityMap(); - } - - IRModule result(/*functions=*/{}, mod->type_definitions, mod->Imports(), mod->source_map, - mod->attrs); - for (const auto& kv : mod->functions) { - if (auto opt = kv.second.as()) { - auto function = opt.value(); - - VLOG(1) << "processing " << PrettyPrint(kv.first); - - VLOG(2) << "count usage"; - UsageVisitor usage_visitor(&var_to_purity, /*default_purity=*/ignore_impurity); - usage_visitor.VisitExpr(function); - - // Actually eliminate/inline the let-bindings. - VLOG(2) << "eliminate"; - EliminatorMutator eliminator_mutator(inline_once, &usage_visitor.let_bound_values_, - &usage_visitor.use_map_, &var_to_purity, - /*default_purity=*/ignore_impurity); - result->Add(kv.first, Downcast(eliminator_mutator.VisitExpr(function))); - } else { - // PrimFuncs come across unchanged. - result->Add(kv.first, kv.second); - } - } - VLOG(1) << "After:" << std::endl << PrettyPrint(result); - - return result; - }; - return tvm::transform::CreateModulePass(pass_func, /*opt_level=*/1, "DeadCodeElimination", - {"InferType"}); -} - -TVM_REGISTER_GLOBAL("relay._transform.DeadCodeElimination").set_body_typed(DeadCodeElimination); - -} // namespace transform - -} // namespace relay -} // namespace tvm diff --git a/src/relay/transforms/defunctionalization.cc b/src/relay/transforms/defunctionalization.cc deleted file mode 100644 index 59f94e0cdd86..000000000000 --- a/src/relay/transforms/defunctionalization.cc +++ /dev/null @@ -1,431 +0,0 @@ -/* - * 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 defunctionalization.cc - * - * \brief Defunctionalization for Relay IR - * - * This pass transforms a higher-order program into a first-order program with defunctionalization. - * This means that all higher order functions (i.e functions that take function arguments or return - * functions) should be transformed into a semantically equivalent first order one. - * - * This pass implements a basic typed defunctionalization method. - * All higher order functions are cloned and specialized (so that there are no type params). - * Function type arguments are encoded as datatypes and a helper `apply` function is used - * to "call" them. - * - * For example, take the following higher order program: - * fun map F y = case y of - * Nil => Nil - * | Cons(x, XS) => Cons(F z, map F XS) - * fun addone 1 = map (\x -> \x + 1) 1 - * - * where `addone` is our program. - * When we call the `map` function, we see that it is a higher-order function, - * but we can clone `map ` function and specialize it with the type_params of the call. - * In addition, our function argument `(\x -> \x + 1)` will be encoded as a datatype constructor, - * which we will call `incr`, and all calls to `F` in our specialized map function will use the - * helper `apply` function. - * - * After defunctionalization, we get: - * fun apply encoding arg = case encoding of - * “incr” => incr arg - * fun map’ F y = case y of - * Nil => Nil - * | Cons(x, xs) => Cons(apply F x, map’ F xs) - * fun addone 1 = map’ “incr” 1 - * - * Currently, defunctionalization makes the following assumptions: - * - functions cannot return function values - * - function arguments are in two forms: identifier or a lambda abstraction - * - no functions stored in datatype - * - functions are not let binded - */ - -#include -#include -#include -#include -#include -#include - -#include "../analysis/type_solver.h" -#include "../transforms/pass_utils.h" -namespace tvm { -namespace relay { - -// determine if type contains a FuncType -bool HasFuncType(const Type& t) { - struct FuncTypeVisitor : TypeVisitor { - bool has_func_type; - FuncTypeVisitor() : has_func_type(false) {} - - void VisitType_(const FuncTypeNode* op) { this->has_func_type = true; } - }; - - auto visitor = FuncTypeVisitor(); - visitor.VisitType(t); - return visitor.has_func_type; -} -// determine if FuncType is a higher order type -bool IsHigherOrderFunc(const FuncType& t) { - bool higher_order = false; - for (auto arg : t->arg_types) { - higher_order |= HasFuncType(arg); - } - return higher_order |= HasFuncType(t->ret_type); -} - -/*! - * \brief mutator for driving the Defunctionalization transformation - */ -class DefuncMutator : public ExprMutator { - public: - explicit DefuncMutator(const IRModule& mod) : mod(mod), constructor_counter(0) {} - - Expr VisitExpr_(const CallNode* call) { - if (auto op = call->op.as()) { - ICHECK_EQ(call->type_args.size(), op->checked_type().as()->type_params.size()) - << "all type args must be explicit"; - - auto op_type = InstFuncType(op->checked_type().as(), call->type_args); - ICHECK_EQ(FreeTypeVars(op_type, mod).size(), 0) << "free type vars in instantiated"; - ICHECK(!HasFuncType(op_type->ret_type)) << "returning functions not supported"; - - if (!IsHigherOrderFunc(op_type)) { - // not higher order function - return ExprMutator::VisitExpr_(call); - } - - // first we encode function arguments - Array args; - for (size_t i = 0; i < call->args.size(); i++) { - auto arg = call->args[i]; - auto type = op_type->arg_types[i]; - if (!HasFuncType(type)) { - args.push_back(arg); - } else { - args.push_back(EncodeArg(arg, type)); - } - } - auto name = op->name_hint + TypeToString(op_type); - auto gv = GlobalVar(name); - if (specialized_gv_map.count(name)) { - gv = specialized_gv_map[name]; - } else { - specialized_gv_map[name] = gv; - // clone and specialize with specific type - auto clone = Downcast(DeDup(mod->Lookup(GetRef(op)))); - auto specialized_function = Specialize(clone, call->type_args); - // change var types and change all applications to use `apply` method - auto f = Downcast(FirstifyVars(specialized_function)); - mod->Add(gv, f); - } - return Call(gv, args); - } else if (auto op = call->op.as()) { - // reduction by applying vars - std::unordered_map var_binding_map; - for (size_t i = 0; i < op->params.size(); i++) { - var_binding_map[op->params[i]] = call->args[i]; - } - auto e = Bind(op->body, var_binding_map); - return this->VisitExpr(e); - } else if (auto op = call->op.as()) { - // var node will be encoded as datatype - // so we need to use the `apply` helper method - auto var_original_type = GetUnencodedType(op->type_annotation).as(); - ICHECK(var_original_type) << "var original type not saved in var_save_type map"; - auto op_type = InstFuncType(var_original_type, call->type_args); - - Array args = {GetRef(op)}; - for (auto arg : call->args) { - args.push_back(this->VisitExpr(arg)); - } - - return Call(GetApplyFunction(op_type), args); - } - return ExprMutator::VisitExpr_(call); - } - - private: - // module - IRModule mod; - // gv + str(type) to specialized clone gv - std::unordered_map specialized_gv_map; - // str(func_type) to ADT - std::unordered_map func_encoding; - // str(func_tyoe) to apply gv - std::unordered_map apply_map; - // encoded ADT handle to FuncType - std::unordered_map original_func_type_map; - // gv to (str(func_type) to constructor encoding) - std::unordered_map, ObjectHash, - ObjectEqual> - gv_datatype_map; - // use monotonically increasing integer to represent new constructor_name - uint64_t constructor_counter; - - /*! - * \brief add a constructor to the GlobalTypeVar, creating a new TypeDef if GlobalTypeVar does not - * exist - */ - void AddConstructor(GlobalTypeVar gtv, Constructor c) { - if (!mod->ContainGlobalTypeVar(gtv->name_hint)) { - mod->AddTypeDef(gtv, TypeData(gtv, {}, {c})); - } else { - auto typedata = mod->LookupTypeDef(gtv); - auto constructors = typedata->constructors; - constructors.push_back(c); - mod->UpdateTypeDef(gtv, TypeData(typedata->header, typedata->type_vars, constructors)); - } - } - /*! - * \brief add a case to the apply function, creating the function if it does not exist - * - * \param apply_gv GlobalVar of the apply function - * \param ft is the type functions the apply function handles - * \param c constructor to add a case for - * \param expr calls this expr with the args to the apply_gv - * \param patterns PatterVars to match with the constructor, used for handling free vars in - * functions - */ - void AddApplyCase(GlobalVar apply_gv, FuncType ft, Constructor c, const Expr& expr, - const Array patterns) { - ICHECK(c->inputs.size() == patterns.size()) - << "constructor function and pattern vars have different sizes"; - if (!mod->ContainGlobalVar(apply_gv->name_hint)) { - auto x = Var("x", TypeCall(c->belong_to, {})); - auto vars = Array({x}); - auto args = Array(); - for (auto t : ft->arg_types) { - auto y = Var("y", t); - vars.push_back(y); - args.push_back(y); - } - - auto clauses = Array({Clause(PatternConstructor(c, patterns), Call(expr, args))}); - auto body = Match(x, clauses); - auto f = Function(vars, body, ft->ret_type, {}); - - mod->Add(apply_gv, f); - } else { - auto f = Downcast(mod->Lookup(apply_gv)); - auto body = f->body.as(); - ICHECK(body) << "internal invariant broken; apply function body should be a match node"; - - auto clauses = body->clauses; - auto x = f->params[0]; - auto args = Array(); - for (size_t i = 1; i < f->params.size(); i++) { - args.push_back(f->params[i]); - } - clauses.push_back(Clause(PatternConstructor(c, patterns), Call(expr, args))); - - mod->Add(apply_gv, Function(f->params, Match(x, clauses), f->ret_type, f->type_params), true); - } - } - - Expr EncodeArg(const Expr& arg, const Type& type) { - // we assume arg is either an identifier (var or globalvar) or a function - ICHECK(type.as()) << "assume no nested functions"; - ICHECK(arg.as() || arg.as() || arg.as()) - << "assume all first-order-parameters are identifiers or functions"; - - if (arg.as()) { - // variable with functype will be encoded as datatype in surrounding function - return arg; - } else if (arg.as()) { - return EncodeGlobalVar(Downcast(arg), Downcast(type)); - } else if (auto fn = arg.as()) { - // we handle free vars in anonymous functions by adding arguments to - // the constructor function - auto free_vars = FreeVars(arg); - auto ft = Downcast(type); - - auto arg_types = Array(); - auto pattern_vars = Array(); - auto call_args = Array(); - Map free_var_bind_map; - for (auto free_var : free_vars) { - // free vars are already encoded, can only exist within - // specialized functions - if (free_var->type_annotation.defined()) { - arg_types.push_back(free_var->type_annotation); - } else { - arg_types.push_back(free_var->checked_type()); - } - auto new_var = Var(free_var->name_hint(), free_var->type_annotation); - free_var_bind_map.Set(free_var, new_var); - pattern_vars.push_back(PatternVar(new_var)); - call_args.push_back(free_var); - } - auto gtv = GetFuncEncode(ft); - auto c = Constructor(std::to_string(++constructor_counter), arg_types, gtv); - AddConstructor(gtv, c); - - auto apply_gv = GetApplyFunction(ft); - auto body = this->VisitExpr(Bind(fn->body, free_var_bind_map)); - AddApplyCase(apply_gv, ft, c, WithFields(GetRef(fn), fn->params, body), - pattern_vars); - - return Call(c, call_args); - } - LOG(FATAL) << "EncodeArg failed to cast arg into identifier node or function node"; - } - - /*! - * \brief encode a global var with a specialized type with a datatype - */ - Expr EncodeGlobalVar(const GlobalVar& gv, const FuncType& ft) { - auto map = gv_datatype_map[gv]; - auto type_key = TypeToString(ft); - if (map.count(type_key) == 0) { - auto gtv = GetFuncEncode(ft); - auto c = Constructor(std::to_string(constructor_counter++), {}, gtv); - map[type_key] = c; - AddConstructor(gtv, c); - AddApplyCase(GetApplyFunction(ft), ft, c, gv, {}); - } - return Call(map[type_key], {}); - } - - /*! - * \brief type to string - */ - std::string TypeToString(const Type& t) { - std::ostringstream s; - s << t->GetTypeKey(); - return s.str(); - } - - /*! - * \brief get ADT handle for encoding type t - */ - GlobalTypeVar GetFuncEncode(const Type& t) { - auto adt_name = "Defunc" + TypeToString(t); - if (func_encoding.count(adt_name) == 0) { - func_encoding[adt_name] = GlobalTypeVar(adt_name, TypeKind::kAdtHandle); - } - original_func_type_map[func_encoding[adt_name]] = t; - return func_encoding[adt_name]; - } - - /*! - * \brief get original function type represented by type t - */ - FuncType GetUnencodedType(const Type& t) { - auto tc = t.as(); - ICHECK(tc) << "expected type call when getting original type from encoded type"; - auto gv = tc->func.as(); - ICHECK(gv) << "expected global type var in encoded type"; - auto type = original_func_type_map[GetRef(gv)]; - ICHECK(type.defined()) << "reverse mapping from encoded type to original type not found"; - return Downcast(type); - } - - /*! - * \brief get the apply function for calling datatypes encoding functions of type t - */ - GlobalVar GetApplyFunction(const Type& t) { - auto f_name = "apply" + TypeToString(t); - if (apply_map.count(f_name) == 0) { - apply_map[f_name] = GlobalVar("apply" + TypeToString(t)); - } - return apply_map[f_name]; - } - - /*! - * \brief specialize a function type - */ - FuncType InstFuncType(const FuncTypeNode* fty, const Array type_args) { - ICHECK(fty) << "InstFuncType functype is null"; - ICHECK_EQ(fty->type_params.size(), type_args.size()) - << "size mismatch between function type params and type args"; - auto map = tvm::Map(); - for (size_t i = 0; i < type_args.size(); i++) { - map.Set(fty->type_params[i], type_args[i]); - } - // copy with typevars removed - return Downcast(TypeSubst(FuncType(fty->arg_types, fty->ret_type, {}, {}), map)); - } - - /*! - * \brief specialize a function expression - */ - Function Specialize(const Function& f, const Array type_args) { - ICHECK_EQ(f->type_params.size(), type_args.size()) - << "cannot specialize function with size mismatch between function type params and type " - "args"; - auto map = tvm::Map(); - for (size_t i = 0; i < type_args.size(); i++) { - map.Set(f->type_params[i], type_args[i]); - } - // copy with typevars removed - auto copy = TypeSubst(WithFields(f, {}, {}, {}, /* erase type params */ Array()), map); - return Downcast(copy); - } - - /*! - * \brief transform a function to be first order by transforming arg_types and - * using the `apply` function for applications - */ - Function FirstifyVars(const Function& f) { - ICHECK(f->type_params.size() == 0) << "firstify function has type params"; - - tvm::Map var_bind_map; - Array params; - for (auto var : f->params) { - if (auto var_type = var->type_annotation.as()) { - // first order parameter - auto fop_type = GetRef(var_type); - auto adt = GetFuncEncode(fop_type); - auto new_var = Var(var->name_hint(), TypeCall(adt, {})); - mod->LookupTypeDef(adt); - var_bind_map.Set(var, new_var); - params.push_back(new_var); - } else { - ICHECK(!HasFuncType(var->type_annotation)) - << "nested function type in parameter not supported yet"; - params.push_back(var); - } - } - - auto bind = Downcast(Bind(f, var_bind_map)); - return WithFields(bind, params, this->VisitExpr(bind->body), bind->ret_type, - /* erase type params */ Array()); - } -}; - -Expr Defunctionalization(const Function& f, const IRModule& mod) { - // f is the starting point of the program, all types MUST be known - ICHECK(f->type_params.size() == 0) << "no polymorphism supported for defunctionalization"; - for (const auto& p : f->params) { - ICHECK(!HasFuncType(p->checked_type())) << "program cannot have func type parameters"; - } - ICHECK(!HasFuncType(f->ret_type)) << "return type cannot contain function"; - - return Downcast(DefuncMutator(mod).VisitExpr(f)); -} - -TVM_REGISTER_GLOBAL("relay._transform.Defunctionalization").set_body_typed(Defunctionalization); - -} // namespace relay -} // namespace tvm diff --git a/src/relay/transforms/defuse_ops.cc b/src/relay/transforms/defuse_ops.cc deleted file mode 100644 index 0d97d5a7b75c..000000000000 --- a/src/relay/transforms/defuse_ops.cc +++ /dev/null @@ -1,88 +0,0 @@ -/* - * 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 src/relay/transforms/defuse_ops.cc - * \brief This is an inverse operation of fusion pass. It transforms a fused - * program returned by relay::transform::FuseOps into the program before FuseOps. - * (i.e., x == DefuseOps(FuseOps(x))) - */ - -#include -#include -#include - -#include -#include - -#include "pattern_utils.h" - -namespace tvm { -namespace relay { - -class DefuseOpsMutator : public ExprMutator { - public: - class FuncBodyMutator : public ExprMutator { - public: - explicit FuncBodyMutator(std::unordered_map args) - : ExprMutator(), name_to_args_(std::move(args)) {} - - Expr VisitExpr_(const VarNode* n) { return name_to_args_[n->name_hint()]; } - - private: - std::unordered_map name_to_args_; - }; - - Expr VisitExpr_(const CallNode* n) { - auto new_n = ExprMutator::VisitExpr_(n); - - if (const auto* call = new_n.as()) { - if (const auto* func = call->op.as()) { - std::unordered_map name_to_args; - for (size_t i = 0; i < func->params.size(); ++i) { - const std::string& pname = func->params[i]->name_hint(); - ICHECK(name_to_args.cend() == name_to_args.find(pname)) - << "Found multiple parameters share the same variable name `" << pname - << "` which introduces uncertainty in DefuseOps pass"; - name_to_args[pname] = call->args[i]; - } - return FuncBodyMutator(std::move(name_to_args)).Mutate(func->body); - } - } - return new_n; - } -}; - -Expr DefuseOps(const Expr& expr) { return DefuseOpsMutator().Mutate(expr); } - -namespace transform { - -Pass DefuseOps() { - runtime::TypedPackedFunc pass_func = - [=](Function f, IRModule m, PassContext pc) { return Downcast(DefuseOps(f)); }; - return CreateFunctionPass(pass_func, 3, "DefuseOps", {"InferType"}); -} - -TVM_REGISTER_GLOBAL("relay._transform.DefuseOps").set_body_typed(DefuseOps); - -} // namespace transform - -} // namespace relay -} // namespace tvm diff --git a/src/relay/transforms/device_aware_visitors.cc b/src/relay/transforms/device_aware_visitors.cc deleted file mode 100644 index f3ca1bfa3a9e..000000000000 --- a/src/relay/transforms/device_aware_visitors.cc +++ /dev/null @@ -1,352 +0,0 @@ -/* - * 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 src/relay/transforms/device_aware_visitors.cc - * \brief Visitors which track the device for the current Relay expression. - */ - -#include "./device_aware_visitors.h" - -namespace tvm { -namespace relay { -namespace transform { - -// TODO(mbs): This machinery can be used a) on expressions/modules which have not had -// device planning run, and b) on expressions for which we've not kept track of their -// containing module. For now we'll handle b) by being forgiving as possible when recovering -// the device for an expression, and we'll support a) the same way. But better would be -// to ICHECK fail when, eg, a variable is not in scope or the lexical device stack is empty. - -LexicalOnDeviceMixin::LexicalOnDeviceMixin(const Optional& maybe_mod) { - if (maybe_mod) { - for (const auto& kv : maybe_mod.value()->functions) { - if (const auto* function_node = kv.second.as()) { - VirtualDevice virtual_device = function_node->virtual_device(); - if (!virtual_device->IsFullyUnconstrained()) { - VLOG(2) << "global '" << kv.first->name_hint << "' has virtual device " << virtual_device; - global_var_virtual_devices_.emplace(kv.first, virtual_device); - } - } - } - } -} - -VirtualDevice LexicalOnDeviceMixin::GetVirtualDevice(const Expr& expr) const { - OnDeviceProps props = GetOnDeviceProps(expr); - if (props.body.defined() && props.is_fixed()) { - return props.virtual_device; - } else if (const auto* var_node = expr.as()) { - // Lookup variable binding. - auto itr = var_virtual_devices_.find(GetRef(var_node)); - if (itr != var_virtual_devices_.end()) { - return itr->second; - } - // else: fallthrough to unconstrained - } else if (const auto* global_var_node = expr.as()) { - // Lookup global variable. - auto itr = global_var_virtual_devices_.find(GetRef(global_var_node)); - if (itr != global_var_virtual_devices_.end()) { - return itr->second; - } - // else: fallthrough to unconstrained - } else if (const auto* function_node = expr.as()) { - if (function_node->HasNonzeroAttr(attr::kPrimitive)) { - if (!expr_virtual_devices_.empty()) { - // Use the currently in-scope device type. - return expr_virtual_devices_.back(); - } - // else: fallthrough to unconstrained - } else { - return function_node->virtual_device(); - } - } else { - if (!expr_virtual_devices_.empty()) { - // Use the currently in-scope device type. - return expr_virtual_devices_.back(); - } - // else: fallthrough to unconstrained - } - return VirtualDevice::FullyUnconstrained(); -} - -void LexicalOnDeviceMixin::EnterFunctionBody() { ++function_nesting_; } - -void LexicalOnDeviceMixin::ExitFunctionBody() { - ICHECK_GT(function_nesting_, 0); - --function_nesting_; -} - -void LexicalOnDeviceMixin::PushVirtualDevice(const VirtualDevice& virtual_device) { - expr_virtual_devices_.emplace_back(virtual_device); -} - -void LexicalOnDeviceMixin::PopVirtualDevice() { - if (expr_virtual_devices_.empty()) { - return; - } - expr_virtual_devices_.pop_back(); -} - -void LexicalOnDeviceMixin::PushBoundVar(Var var, const VirtualDevice& virtual_device) { - if (virtual_device->IsFullyUnconstrained()) { - return; - } - ICHECK(var_virtual_devices_.find(var) == var_virtual_devices_.end()); - var_virtual_devices_.emplace(std::move(var), virtual_device); -} - -void LexicalOnDeviceMixin::PopBoundVar(const Var& var) { - auto itr = var_virtual_devices_.find(var); - if (itr == var_virtual_devices_.end()) { - return; - } - var_virtual_devices_.erase(itr); -} - -// TODO(mbs): We'd probably have less tedious code duplication if we redefined the memoizing -// mutator on top of the generic Functor. - -void DeviceAwareExprVisitor::VisitExpr_(const FunctionNode* function_node) { - if (function_node->HasNonzeroAttr(attr::kPrimitive)) { - // No tracking inside primitive functions. - DeviceAwareVisitExpr_(function_node); - } else { - // Function parameters come into scope. - for (auto param : function_node->params) { - PushBoundVar(param, param->virtual_device()); - } - // Entering scope of function body. - PushVirtualDevice(function_node->virtual_device()); - EnterFunctionBody(); - - DeviceAwareVisitExpr_(function_node); - - // Leaving scope of function body. - ExitFunctionBody(); - PopVirtualDevice(); - // Function parameters go out of scope. - for (size_t i = 0; i < function_node->params.size(); ++i) { - PopBoundVar(function_node->params[i]); - } - } -} - -void DeviceAwareExprVisitor::VisitExpr_(const LetNode* let_node) { - PreVisitLetBlock_(let_node); - std::vector bindings; - Expr expr = GetRef(let_node); - while (const auto* inner_let_node = expr.as()) { - // Let-bound var (in pre visited version) goes into scope. - // (We'll just assume this is a letrec). - PushBoundVar(inner_let_node->var, GetVirtualDevice(inner_let_node->value)); - PreVisitLetBinding_(inner_let_node->var, inner_let_node->value); - bindings.emplace_back(inner_let_node); - expr = inner_let_node->body; - } - - VisitExpr(expr); - - for (auto itr = bindings.rbegin(); itr != bindings.rend(); ++itr) { - // Let-bound var goes out of scope. - PopBoundVar((*itr)->var); - PostVisitLet_(*itr); - } - PostVisitLetBlock_(let_node); -} - -void DeviceAwareExprVisitor::VisitExpr_(const CallNode* call_node) { - OnDeviceProps props = GetOnDeviceProps(call_node); - if (props.body.defined() && props.is_fixed()) { - // Entering lexical scope of fixed "on_device" call. - PushVirtualDevice(props.virtual_device); - VisitExpr(props.body); - // Leaving lexical scope of "on_device" call. - PopVirtualDevice(); - } else { - DeviceAwareVisitExpr_(call_node); - } -} - -void DeviceAwareExprVisitor::DeviceAwareVisitExpr_(const FunctionNode* function_node) { - ExprVisitor::VisitExpr_(function_node); -} - -void DeviceAwareExprVisitor::DeviceAwareVisitExpr_(const CallNode* call_node) { - ExprVisitor::VisitExpr_(call_node); -} - -void DeviceAwareExprVisitor::PreVisitLetBlock_(const LetNode* let_node) { - // no-op -} - -void DeviceAwareExprVisitor::PreVisitLetBinding_(const Var& var, const Expr& value) { - VisitExpr(var); - VisitExpr(value); -} - -void DeviceAwareExprVisitor::PostVisitLet_(const LetNode* let_node) { - // no-op -} - -void DeviceAwareExprVisitor::PostVisitLetBlock_(const LetNode* let_node) { - // no-op -} - -Expr DeviceAwareExprMutator::VisitExpr_(const FunctionNode* function_node) { - if (function_node->HasNonzeroAttr(attr::kPrimitive)) { - // No tracking inside primitive functions. - return DeviceAwareVisitExpr_(function_node); - } else { - // Function parameters come into scope. - for (auto param : function_node->params) { - PushBoundVar(param, param->virtual_device()); - } - // Entering scope of function body. - PushVirtualDevice(function_node->virtual_device()); - EnterFunctionBody(); - - Expr result = DeviceAwareVisitExpr_(function_node); - - // Leaving scope of function body. - ExitFunctionBody(); - PopVirtualDevice(); - // Function parameters go out of scope. - for (size_t i = 0; i < function_node->params.size(); ++i) { - PopBoundVar(function_node->params[i]); - } - - return result; - } -} - -Expr DeviceAwareExprMutator::VisitExpr_(const LetNode* let_node) { - PreVisitLetBlock_(let_node); - std::vector> bindings; - Expr expr = GetRef(let_node); - while (const auto* inner_let_node = expr.as()) { - // Let-bound var (in pre visited version) goes into scope. - // (We'll just assume this is a letrec.) - PushBoundVar(inner_let_node->var, GetVirtualDevice(inner_let_node->value)); - std::pair pair = PreVisitLetBinding_(inner_let_node->var, inner_let_node->value); - bindings.emplace_back(pair.first, pair.second, inner_let_node->span, inner_let_node); - expr = inner_let_node->body; - } - - expr = VisitExpr(expr); - - for (auto itr = bindings.rbegin(); itr != bindings.rend(); ++itr) { - // Let-bound var goes out of scope. - const LetNode* pre_let_node = std::get<3>(*itr); - PopBoundVar(pre_let_node->var); - Let post_let = Let(/*var=*/std::get<0>(*itr), /*value=*/std::get<1>(*itr), - /*body=*/expr, /*span=*/std::get<2>(*itr)); - expr = PostVisitLet_(pre_let_node, post_let.get()); - } - return PostVisitLetBlock_(let_node, expr.as()); -} - -Expr DeviceAwareExprMutator::VisitExpr_(const CallNode* call_node) { - OnDeviceProps props = GetOnDeviceProps(call_node); - if (props.body.defined() && props.is_fixed()) { - // Entering lexical scope of fixed "on_device" call. - PushVirtualDevice(props.virtual_device); - Expr expr = VisitExpr(props.body); - // Leaving lexical scope of "on_device" call. - PopVirtualDevice(); - return MaybeOnDeviceWithProps(expr, props); - } else { - return DeviceAwareVisitExpr_(call_node); - } -} - -Expr DeviceAwareExprMutator::DeviceAwareVisitExpr_(const FunctionNode* function_node) { - return ExprMutator::VisitExpr_(function_node); -} - -Expr DeviceAwareExprMutator::DeviceAwareVisitExpr_(const CallNode* call_node) { - return ExprMutator::VisitExpr_(call_node); -} - -void DeviceAwareExprMutator::PreVisitLetBlock_(const LetNode* let_node) { /* no-op */ -} - -std::pair DeviceAwareExprMutator::PreVisitLetBinding_(const Var& var, - const Expr& value) { - return std::make_pair(Downcast(VisitExpr(var)), VisitExpr(value)); -} - -Expr DeviceAwareExprMutator::PostVisitLet_(const LetNode* pre_let_node, - const LetNode* post_let_node) { - if (pre_let_node->var == post_let_node->var && pre_let_node->value == post_let_node->value && - pre_let_node->body == post_let_node->body) { - return GetRef(pre_let_node); - } else { - return GetRef(post_let_node); - } -} - -Expr DeviceAwareExprMutator::PostVisitLetBlock_(const LetNode* pre_let_node, - const LetNode* post_let_node) { - if (pre_let_node->var == post_let_node->var && pre_let_node->value == post_let_node->value && - pre_let_node->body == post_let_node->body) { - return GetRef(pre_let_node); - } else { - return GetRef(post_let_node); - } -} - -std::unordered_map RecoverVirtualDeviceMap(const IRModule& mod, - const Expr& expr) { - class Visitor : public DeviceAwareExprVisitor { - public: - explicit Visitor(const Optional& maybe_mod) : DeviceAwareExprVisitor(maybe_mod) {} - - void VisitExpr(const Expr& expr) final { - if (expr->IsInstance() || expr->IsInstance()) { - // Don't record for ops or constructors since they are 'device polymorphic'. - } else { - map_[expr.get()] = GetVirtualDevice(expr); - } - DeviceAwareExprVisitor::VisitExpr(expr); - } - - std::unordered_map map_; - }; - - Visitor visitor(mod); - visitor.VisitExpr(expr); - return std::move(visitor.map_); -} - -// Export the helper function for testing. -TVM_REGISTER_GLOBAL("relay.transform.RecoverVirtualDeviceMap") - .set_body_typed([](const IRModule& mod, const Expr& expr) { - std::unordered_map raw_map = - RecoverVirtualDeviceMap(mod, expr); - Map map; - for (const auto& kv : raw_map) { - map.Set(GetRef(kv.first), kv.second); - } - return map; - }); - -} // namespace transform -} // namespace relay -} // namespace tvm diff --git a/src/relay/transforms/device_aware_visitors.h b/src/relay/transforms/device_aware_visitors.h deleted file mode 100644 index 8a0166abef9b..000000000000 --- a/src/relay/transforms/device_aware_visitors.h +++ /dev/null @@ -1,363 +0,0 @@ -/* - * 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 src/relay/transforms/device_aware_visitors.h - * \brief Visitors which track the device for the current Relay expression and Relay Vars. - */ - -#ifndef TVM_RELAY_TRANSFORMS_DEVICE_AWARE_VISITORS_H_ -#define TVM_RELAY_TRANSFORMS_DEVICE_AWARE_VISITORS_H_ - -#include -#include -#include -#include - -#include -#include -#include - -#include "../op/annotation/annotation.h" -#include "../op/memory/on_device.h" - -namespace tvm { -namespace relay { -namespace transform { - -/*! - * \brief Helper class for expression transformers which need to keep track of the \p VirtualDevice - * holding the results of expressions. This is recovered from function attributes and "on_device" - * CallNodes added by the PlanDevices pass. - * - * \sa \p DeviceAwareExpr{Functor,Visitor,Mutator}. - */ -class LexicalOnDeviceMixin { - protected: - explicit LexicalOnDeviceMixin(const Optional& maybe_mod); - - /*! - * \brief Returns the \p VirtualDevice on which the result of \p expr should/will be stored, - * assuming {Push,Pop}{VirtualDevice,BoundVar} have been correctly called. May return the - * unconstrained \p VirtualDevice if the device planning pass has not been run. - */ - VirtualDevice GetVirtualDevice(const Expr& expr) const; - - /*! \brief Indicate a function body is being entered. */ - void EnterFunctionBody(); - - /*! \brief Indicate a function body has been processed. */ - void ExitFunctionBody(); - - /*! \brief Push an \p VirtualDevice onto the lexical VirtualDevice stack. Ignore if unconstrained. - */ - void PushVirtualDevice(const VirtualDevice& virtual_device); - - /*! \brief Pop an \p VirtualDevice from the lexical VirtualDevice stack. Ignore if stack is empty. - */ - void PopVirtualDevice(); - - /*! \brief Remember that \p var will be stored at \p virtual_device. Ignore if unconstrained. - * - * CAUTION: Despite the name we don't support re-entering the same function body. - */ - void PushBoundVar(Var var, const VirtualDevice& virtual_device); - - /*! \brief Remove the binding for \p var to its \p VirtualDevice. Ignore if var is not bound. */ - void PopBoundVar(const Var& var); - - /*! - * \brief Returns the number of function definitions wrapping the currently visited expression. - */ - int function_nesting() const { return function_nesting_; } - - private: - /*! - * \brief The number of function bodies entered. Since many transforms need to distinguish global - * functions from local functions this supports the mixin's \p is_global() helper method. - */ - int function_nesting_ = 0; - - /*! - * \brief The stack of lexically enclosing "on_device" \p VirtualDevices, from outermost to - * innermost. When visiting an expression other than a variable we can assume the expression's - * result is to be stored on \p expr_virtual_devices.back(). - */ - std::vector expr_virtual_devices_; - - /*! - * \brief A map from in-scope local variables to their \p VirtualDevices. We may assume the - * variable is only ever bound to a value stored on this \p VirtualDevice at runtime. - * - * Note: We're playing it safe and keying by object refs here just in case the Relay expression - * being rewritten has no module or other global to keep it alive. - */ - std::unordered_map - var_virtual_devices_; - - /*! - * \brief A map from global variables to their \p VirtualDevices, ie the "result_virtual_device" - * of the function they are bound to in the module we are working on. We calculate and store this - * explicitly so that we don't need to hold on to any module, which is often in the process of - * being rewritten. - */ - std::unordered_map - global_var_virtual_devices_; -}; - -template -class DeviceAwareExprFunctor; - -/*! - * \brief ExprFunctor which tracks \p VirtualDevices. We only support 'visitor' style implementation - * with no additional arguments, thus this is equivalent to \p DeviceAwareExprVisitor without - * any memoization. - */ -template <> -class DeviceAwareExprFunctor : public ExprFunctor, - public LexicalOnDeviceMixin { - private: - using TSuper = ExprFunctor; - - public: - explicit DeviceAwareExprFunctor(const Optional& maybe_mod) - : LexicalOnDeviceMixin(maybe_mod) {} - - void VisitExpr_(const FunctionNode* function_node) { - if (function_node->HasNonzeroAttr(attr::kPrimitive)) { - // No tracking inside primitive functions. - return DeviceAwareVisitExpr_(function_node); - } else { - // Function parameters come into scope. - for (auto param : function_node->params) { - PushBoundVar(param, param->virtual_device()); - } - // Entering scope of function body. - VirtualDevice virtual_device = function_node->virtual_device(); - VLOG(2) << "entering " << virtual_device << " for function:" << std::endl - << PrettyPrint(GetRef(function_node)); - PushVirtualDevice(virtual_device); - EnterFunctionBody(); - - DeviceAwareVisitExpr_(function_node); - - // Leaving scope of function body. - ExitFunctionBody(); - PopVirtualDevice(); - VLOG(2) << "leaving " << virtual_device << " for function:" << std::endl - << PrettyPrint(GetRef(function_node)); - // Function parameters go out of scope. - for (size_t i = 0; i < function_node->params.size(); ++i) { - PopBoundVar(function_node->params[i]); - } - } - } - - void VisitExpr_(const LetNode* let_node) { - PreVisitLetBlock_(let_node); - std::vector bindings; - Expr expr = GetRef(let_node); - while (const auto* inner_let_node = expr.as()) { - // Let-bound var (in pre visited version) goes into scope. - // (We'll just assume this is a letrec.) - VirtualDevice virtual_device = GetVirtualDevice(inner_let_node->value); - VLOG(2) << "var '" << inner_let_node->var->name_hint() << "' has virtual device " - << virtual_device; - PushBoundVar(inner_let_node->var, virtual_device); - PreVisitLetBinding_(inner_let_node->var, inner_let_node->value); - bindings.emplace_back(inner_let_node); - expr = inner_let_node->body; - } - - VisitExpr(expr); - - for (auto itr = bindings.rbegin(); itr != bindings.rend(); ++itr) { - // Let-bound var goes out of scope. - const LetNode* visited_let_node = *itr; - PopBoundVar(visited_let_node->var); - PostVisitLet_(visited_let_node); - } - PostVisitLetBlock_(let_node); - } - - void VisitExpr_(const CallNode* call_node) { - OnDeviceProps props = GetOnDeviceProps(call_node); - if (props.body.defined() && props.is_fixed()) { - // Entering lexical scope of "on_device" call. - VLOG(2) << "entering " << props.virtual_device << " for on_device:" << std::endl - << PrettyPrint(GetRef(call_node)); - PushVirtualDevice(props.virtual_device); - VisitExpr(props.body); - // Leaving lexical scope of "on_device" call. - PopVirtualDevice(); - VLOG(2) << "leaving " << props.virtual_device << " for on_device:" << std::endl - << PrettyPrint(GetRef(call_node)); - } else { - DeviceAwareVisitExpr_(call_node); - } - } - - /*! - * \brief These are as for VisitExpr_. \p VirtualDevices for expressions and function parameters - * will be tracked automatically. Default implementation defers to ExprMutator::VisitExpr_. For - * functions the function_nesting count will already include that of \p function_node. - */ - - virtual void DeviceAwareVisitExpr_(const FunctionNode* function_node) { - return TSuper::VisitExpr_(function_node); - } - - virtual void DeviceAwareVisitExpr_(const CallNode* call_node) { - return TSuper::VisitExpr_(call_node); - } - - /*! - * \brief Visit the first let in a chain of let expressions before any let bindings or final - * body has been visited. Default implementation is a no-op. - */ - virtual void PreVisitLetBlock_(const LetNode* let_node) { /* no-op */ - } - - /*! - * \brief Visit a let-bound expression before the let body has been visited. Devices for the - * let-bound variable will be tracked automatically. Default implementation just visits var and - * value. - */ - virtual void PreVisitLetBinding_(const Var& var, const Expr& value) { - VisitExpr(var); - VisitExpr(value); - } - - /*! - * \brief Visit a let expression after the let-bound value and body have been visited. - * Default implementation is a no-op. - */ - virtual void PostVisitLet_(const LetNode* let_node) { /* no-op */ - } - - /*! - * \brief Visit the first let in a chain of let expressions after it has been visited. - * Default implementation is a no-op. - */ - virtual void PostVisitLetBlock_(const LetNode* let_node) {} -}; - -/*! \brief ExprVisitor which tracks \p VirtualDevices. */ -class DeviceAwareExprVisitor : public ExprVisitor, public LexicalOnDeviceMixin { - public: - explicit DeviceAwareExprVisitor(const Optional& maybe_mod) - : LexicalOnDeviceMixin(maybe_mod) {} - - using ExprVisitor::VisitExpr_; - - void VisitExpr_(const FunctionNode* function_node) final; - void VisitExpr_(const LetNode* let_node) final; - void VisitExpr_(const CallNode* call_node) final; - - /*! - * \brief These are as for VisitExpr_. \p VirtualDevices for expressions and function parameters - * will be tracked automatically. Default implementation defers to ExprMutator::VisitExpr_. For - * functions the function_nesting count will already include that of \p function_node. - */ - virtual void DeviceAwareVisitExpr_(const FunctionNode* function_node); - virtual void DeviceAwareVisitExpr_(const CallNode* call_node); - - /*! - * \brief Visit the first let in a chain of let expressions before any let bindings or final - * body has been visited. Default implementation is a no-op. - */ - virtual void PreVisitLetBlock_(const LetNode* let_node); - - /*! - * \brief Visit a let-bound expression before the let body has been visited. \p VirtualDevices for - * the let-bound variable will be tracked automatically. Default implementation just visits var - * and value. - */ - virtual void PreVisitLetBinding_(const Var& var, const Expr& value); - - /*! - * \brief Visit a let expression after the let-bound value and body have been visited. - * Default implementation is a no-op. - */ - virtual void PostVisitLet_(const LetNode* let_node); - - /*! - * \brief Visit the first let in a chain of let expressions after it has been visited. - * Default implementation is a no-op. - */ - virtual void PostVisitLetBlock_(const LetNode* let_node); -}; - -/*! \brief ExprMutator which tracks \p VirtualDevices. */ -class DeviceAwareExprMutator : public ExprMutator, public LexicalOnDeviceMixin { - public: - explicit DeviceAwareExprMutator(const Optional& maybe_mod) - : LexicalOnDeviceMixin(maybe_mod) {} - - Expr VisitExpr_(const FunctionNode* function_node) final; - Expr VisitExpr_(const LetNode* let_node) final; - Expr VisitExpr_(const CallNode* call_node) final; - - /*! - * \brief These are as for VisitExpr_. \p VirtualDevices for expressions and function parameters - * will be tracked automatically. Default implementation defers to ExprMutator::VisitExpr_. For - * functions the function_nesting count will already include that of \p function_node. - */ - virtual Expr DeviceAwareVisitExpr_(const FunctionNode* function_node); - virtual Expr DeviceAwareVisitExpr_(const CallNode* call_node); - - /*! - * \brief Visit the first let in a chain of let expressions before any let bindings or final - * body has been visited. Default implementation is a no-op. - */ - virtual void PreVisitLetBlock_(const LetNode* let_node); - - /*! - * \brief Visit a let-bound expression before the let body has been visited. \p VirtualDevices for - * the let-bound variable will be tracked automatically. Default implementation just visits var - * and value. - */ - virtual std::pair PreVisitLetBinding_(const Var& var, const Expr& value); - - /*! - * \brief Visit a let expression after the let-bound value and body have been visited. - * Default implementation just returns a reference to the post-visited node. - */ - virtual Expr PostVisitLet_(const LetNode* pre_let_node, const LetNode* post_let_node); - - /*! - * \brief Visit the first let in a chain of let expressions after it has been visited. - * Default implementation returns reference to let node. - */ - virtual Expr PostVisitLetBlock_(const LetNode* pre_let_node, const LetNode* post_let_node); -}; - -/*! - * \brief Returs a map from Relay expression node to its virtual device using the annotations - * and \p virtual_device fields of \p expr. The map's lifetime must not exceed that of - * \p expr itself. - */ -std::unordered_map RecoverVirtualDeviceMap(const IRModule& mod, - const Expr& expr); - -} // namespace transform -} // namespace relay -} // namespace tvm - -#endif // TVM_RELAY_TRANSFORMS_DEVICE_AWARE_VISITORS_H_ diff --git a/src/relay/transforms/device_domains.cc b/src/relay/transforms/device_domains.cc deleted file mode 100644 index e2af20022a40..000000000000 --- a/src/relay/transforms/device_domains.cc +++ /dev/null @@ -1,483 +0,0 @@ -/* - * 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 src/relay/analysis/device_domains.cc - * \brief Unification domain for the device planner. - */ - -#include "./device_domains.h" - -#include -#include - -#include "../op/annotation/annotation.h" -#include "../op/call/call.h" -#include "../op/memory/device_copy.h" -#include "../op/memory/on_device.h" - -namespace tvm { -namespace relay { -namespace transform { - -DeviceDomains::DeviceDomains(CompilationConfig config) : config_(std::move(config)) { - host_domain_ = MakeFirstOrderDomain(config_->host_virtual_device); -} - -DeviceDomainPtr DeviceDomains::MakeFirstOrderDomain(const VirtualDevice& virtual_device) { - if (virtual_device->IsFullyConstrained()) { - auto itr = fully_constrained_virtual_device_to_domain_.find(virtual_device); - if (itr != fully_constrained_virtual_device_to_domain_.end()) { - return itr->second; - } - DeviceDomainPtr domain = std::make_shared(virtual_device); - fully_constrained_virtual_device_to_domain_.emplace(virtual_device, domain); - return domain; - } else { - return std::make_shared(virtual_device); - } -} - -DeviceDomainPtr DeviceDomains::MakeDomain(const Type& type, const VirtualDevice& virtual_device) { - if (const auto* func_type_node = type.as()) { - std::vector args_and_result; - args_and_result.reserve(func_type_node->arg_types.size() + 1); - for (const auto& arg_type : func_type_node->arg_types) { - args_and_result.emplace_back(MakeDomain(arg_type, VirtualDevice::FullyUnconstrained())); - } - args_and_result.emplace_back(MakeDomain(func_type_node->ret_type, virtual_device)); - return std::make_shared(std::move(args_and_result)); - } else { - return MakeFirstOrderDomain(virtual_device); - } -} - -DeviceDomainPtr DeviceDomains::ForVirtualDevice(const Type& type, - const VirtualDevice& non_canonical_virtual_device) { - // Generally the virtual device will have come from an annotation so resolve it to ensure we have - // its canonical representation. - VirtualDevice virtual_device = config_->CanonicalVirtualDevice(non_canonical_virtual_device); - ICHECK(!virtual_device->IsFullyUnconstrained()); - return MakeDomain(type, virtual_device); -} - -DeviceDomainPtr DeviceDomains::Lookup(DeviceDomainPtr domain) { - DeviceDomainPtr root = domain; - while (true) { - auto itr = domain_to_equiv_.find(root); - if (itr == domain_to_equiv_.end()) { - break; - } - ICHECK_NE(itr->second, root); - root = itr->second; - ICHECK_NOTNULL(root); - } - // Path compression. - while (domain != root) { - auto itr = domain_to_equiv_.find(domain); - ICHECK(itr != domain_to_equiv_.end()); - domain = itr->second; - ICHECK_NOTNULL(domain); - itr->second = root; - } - return root; -} - -DeviceDomainPtr DeviceDomains::JoinOrNull(const DeviceDomainPtr& lhs, const DeviceDomainPtr& rhs) { - if (lhs == rhs) { - return lhs; - } - ICHECK_EQ(lhs->args_and_result_.size(), rhs->args_and_result_.size()) - << "Device domains:" << std::endl - << ToString(lhs) << std::endl - << "and" << std::endl - << ToString(rhs) << std::endl - << "do not have the same kind and can't be unified."; - if (lhs->args_and_result_.empty()) { - // Directly compare first-order. - if (rhs->virtual_device_->IsFullyUnconstrained()) { - return lhs; - } - if (lhs->virtual_device_->IsFullyUnconstrained()) { - return rhs; - } - Optional joined_virtual_device = - VirtualDevice::Join(lhs->virtual_device_, rhs->virtual_device_); - if (!joined_virtual_device) { - return nullptr; - } - return MakeFirstOrderDomain(config_->CanonicalVirtualDevice(joined_virtual_device.value())); - } else { - // Recurse for higher-order. - std::vector args_and_result; - args_and_result.reserve(lhs->args_and_result_.size()); - for (size_t i = 0; i < lhs->args_and_result_.size(); ++i) { - DeviceDomainPtr joined_domain = - UnifyOrNull(lhs->args_and_result_[i], rhs->args_and_result_[i]); - if (joined_domain == nullptr) { - return nullptr; - } - args_and_result.emplace_back(std::move(joined_domain)); - } - return MakeHigherOrderDomain(std::move(args_and_result)); - } -} - -DeviceDomainPtr DeviceDomains::UnifyOrNull(DeviceDomainPtr lhs, DeviceDomainPtr rhs) { - ICHECK_NOTNULL(lhs); - ICHECK_NOTNULL(rhs); - lhs = Lookup(lhs); - rhs = Lookup(rhs); - DeviceDomainPtr joined_domain = JoinOrNull(lhs, rhs); - if (joined_domain == nullptr) { - return nullptr; - } - if (lhs != joined_domain) { - domain_to_equiv_.emplace(lhs, joined_domain); - } - if (rhs != joined_domain) { - domain_to_equiv_.emplace(rhs, joined_domain); - } - return joined_domain; -} - -bool DeviceDomains::CollapseOrFalse(const DeviceDomainPtr& first_order_domain, - const DeviceDomainPtr& higher_order_domain) { - ICHECK(!first_order_domain->is_higher_order()); - ICHECK(higher_order_domain->is_higher_order()); - for (size_t i = 0; i < higher_order_domain->function_arity(); ++i) { - if (UnifyOrNull(higher_order_domain->function_param(i), first_order_domain) == nullptr) { - return false; - } - } - return UnifyOrNull(higher_order_domain->function_result(), first_order_domain) != nullptr; -} - -bool DeviceDomains::UnifyCollapsedOrFalse(const DeviceDomainPtr& lhs_first_order, - const DeviceDomainPtr& rhs_maybe_higher_order) { - ICHECK(!lhs_first_order->is_higher_order()); - if (rhs_maybe_higher_order->is_higher_order()) { - return CollapseOrFalse(lhs_first_order, rhs_maybe_higher_order); - } else { - return UnifyOrNull(lhs_first_order, rhs_maybe_higher_order) != nullptr; - } -} - -DeviceDomainPtr DeviceDomains::DomainFor(const Expr& expr) { - ICHECK(expr.defined()); - auto itr = expr_to_domain_.find(expr.get()); - if (itr != expr_to_domain_.end()) { - return Lookup(itr->second); - } - auto domain = Free(expr->checked_type()); - expr_to_domain_.emplace(expr.get(), domain); - return domain; -} - -DeviceDomainPtr DeviceDomains::DomainForCallee(const Call& call) { - auto itr = call_to_callee_domain_.find(call.get()); - if (itr != call_to_callee_domain_.end()) { - return Lookup(itr->second); - } - std::vector args_and_result; - - OnDeviceProps on_device_props = GetOnDeviceProps(call.get()); - DeviceCopyProps device_copy_props = GetDeviceCopyProps(call.get()); - CallLoweredProps call_lowered_props = GetCallLoweredProps(call.get()); - - if (call_lowered_props.lowered_func.defined()) { - // Presumably we've already seen the call to the "primitive" Function from which this lowered - // function was derived in an earlier PlanDevices pass. Thus we've already established that - // all the argument and result devices domains must be equal, ignoring memory scopes. - // So at this point we'll let all the arguments and result be free so that memory scopes can - // differ. - // TODO(mbs): As per header comments, need to revisit when can setup sub-virtual device - // constraints. - return DomainFor(call_lowered_props.lowered_func); - } else if (on_device_props.body.defined()) { - // By default: - // on_device(expr, virtual_device=) - // on_device : fn():?x? - // However we'll interpret the constrain_body and constrain_result fields to decide - // on free vs constrained domains for the argument and result respectively. - if (on_device_props.constrain_body) { - args_and_result.emplace_back( - ForVirtualDevice(on_device_props.body->checked_type(), on_device_props.virtual_device)); - } else { - args_and_result.emplace_back(Free(on_device_props.body->checked_type())); - } - if (on_device_props.constrain_result) { - args_and_result.emplace_back( - ForVirtualDevice(on_device_props.body->checked_type(), on_device_props.virtual_device)); - } else { - args_and_result.emplace_back(Free(on_device_props.body->checked_type())); - } - } else if (device_copy_props.body.defined()) { - // device_copy(expr, src_virtual_device=, dst_virtual_device=) - // device_copy: fn(): - args_and_result.emplace_back(ForVirtualDevice(device_copy_props.body->checked_type(), - device_copy_props.src_virtual_device)); - args_and_result.emplace_back(ForVirtualDevice(device_copy_props.body->checked_type(), - device_copy_props.dst_virtual_device)); - } else if (call->op == alloc_storage_op) { - ICHECK_EQ(call->args.size(), 3U); - // alloc_storage(size, shape, alignment, virtual_device=) - // alloc_storage: fn(, , ): - const auto* attrs = call->attrs.as(); - args_and_result.emplace_back(host_domain_); - args_and_result.emplace_back(host_domain_); - args_and_result.emplace_back(host_domain_); - args_and_result.emplace_back(ForVirtualDevice(call->checked_type(), attrs->virtual_device)); - } else if (call->op == alloc_tensor_op) { - ICHECK_EQ(call->args.size(), 3U); - // alloc_tensor(storage, offset, shape) - // alloc_tensor: fn(?x?, , ):?x? - auto free_domain = Free(call->checked_type()); - args_and_result.emplace_back(free_domain); - args_and_result.emplace_back(host_domain_); - args_and_result.emplace_back(host_domain_); - args_and_result.emplace_back(free_domain); - } else if (call->op == shape_of_op) { - ICHECK_EQ(call->args.size(), 1U); - // shape_of(tensor) - // shape_of: fn(?x?): - args_and_result.emplace_back(Free(call->args[0]->checked_type())); - args_and_result.emplace_back(host_domain_); - } else if (call->op == invoke_tvm_op) { - ICHECK_EQ(call->args.size(), 3U); - // invoke_tvm_op(op, inputs, outputs) - // invoke_tvm_op: fn(..., ?x?, ?x?):?x? - // where ... is a free domain appropriate for op's type - auto free_domain = Free(call->checked_type()); - args_and_result.emplace_back(Free(call->args[0]->checked_type())); - args_and_result.emplace_back(free_domain); - args_and_result.emplace_back(free_domain); - args_and_result.emplace_back(free_domain); - } else if (call->op == reshape_tensor_op) { - ICHECK_EQ(call->args.size(), 2U); - // reshape_tensor(data, shape) - // reshape_tensor: fn(?x?, ):?x? - auto free_domain = Free(call->checked_type()); - args_and_result.emplace_back(free_domain); - args_and_result.emplace_back(host_domain_); - args_and_result.emplace_back(free_domain); - } else if (call->op->IsInstance()) { - // (arg1, ..., argn) - // : fn(?x?, ..., ?x?):?x? - // (all args and result must be first-order). - auto free_domain = MakeFirstOrderDomain(VirtualDevice::FullyUnconstrained()); - for (size_t i = 0; i < call->args.size(); ++i) { - args_and_result.emplace_back(free_domain); - } - args_and_result.emplace_back(free_domain); - } else if (call->op->IsInstance()) { - // (arg1, ..., argn) - // : fn(?x1?, ..., ?xn?):?xr? - // where we force all possibly higher-order ?xi? to be collapsed to the first-order ?xr?. - // TODO(mbs): This assumes we've eta-expanded constructors, thus all constructors appear - // in callee positions. - const auto* func_type_node = call->op->checked_type().as(); - ICHECK_NOTNULL(func_type_node); - ICHECK_EQ(func_type_node->arg_types.size(), call->args.size()); - auto result_domain = Free(func_type_node->ret_type); // first-order - for (const auto& arg_type : func_type_node->arg_types) { - auto param_domain = Free(arg_type); // possibly higher-order - bool success = UnifyCollapsedOrFalse(result_domain, param_domain); // collapse if required - ICHECK(success); - args_and_result.emplace_back(param_domain); - } - args_and_result.emplace_back(result_domain); - } else { - // We still need to handle the case where the function / op is not lowered - // because the device planner runs both before and after lowering. - return DomainFor(call->op); - } - auto domain = MakeHigherOrderDomain(std::move(args_and_result)); - call_to_callee_domain_.emplace(call.get(), domain); - return domain; -} - -void DeviceDomains::UnifyExprExact(const Expr& lhs, const Expr& rhs) { - auto lhs_domain = DomainFor(lhs); - auto rhs_domain = DomainFor(rhs); - if (UnifyOrNull(lhs_domain, rhs_domain) == nullptr) { - // TODO(mbs): Proper diagnostics. - LOG(FATAL) << "Incompatible virtual devices for expressions:" << std::endl - << PrettyPrint(lhs) << std::endl - << "with virtual device:" << std::endl - << ToString(lhs_domain) << "and:" << std::endl - << PrettyPrint(rhs) << std::endl - << "with virtual device:" << std::endl - << ToString(rhs_domain); - } -} - -void DeviceDomains::OptionalUnifyExprExact(const Expr& lhs, const Expr& rhs) { - auto lhs_domain = DomainFor(lhs); - auto rhs_domain = DomainFor(rhs); - // Snapshot - std::unordered_map domain_to_equiv_snapshot = domain_to_equiv_; - if (UnifyOrNull(lhs_domain, rhs_domain) == nullptr) { - // Rollback - domain_to_equiv_ = domain_to_equiv_snapshot; - VLOG(2) << "Unable to unify virtual devices for expression:" << std::endl - << PrettyPrint(lhs) << std::endl - << "with virtual device:" << std::endl - << ToString(lhs_domain) << std::endl - << "and expression:" << std::endl - << PrettyPrint(rhs) << std::endl - << "with virtual device:" << std::endl - << ToString(rhs_domain) << std::endl - << ". Leaving virtual devices non-unified."; - } else { - VLOG(2) << "Unified virtual devices for expression:" << std::endl - << PrettyPrint(lhs) << std::endl - << "and expression:" << std::endl - << PrettyPrint(rhs) << std::endl - << "to virtual devices:" << std::endl - << ToString(lhs_domain); - } -} - -void DeviceDomains::UnifyExprExact(const Expr& expr, const DeviceDomainPtr& expected_domain) { - auto actual_domain = DomainFor(expr); - if (UnifyOrNull(actual_domain, expected_domain) == nullptr) { - // TODO(mbs): Proper diagnostics. - LOG(FATAL) << "Incompatible virtual devices for expression:" << std::endl - << PrettyPrint(expr) << std::endl - << "with actual virtual device:" << std::endl - << ToString(actual_domain) << std::endl - << "and expected virtual device:" << std::endl - << ToString(expected_domain); - } -} - -void DeviceDomains::UnifyExprCollapsed(const Expr& expr_first_order, - const DeviceDomainPtr& expected_domain_maybe_higher_order) { - auto actual_domain_first_order = DomainFor(expr_first_order); - if (!UnifyCollapsedOrFalse(actual_domain_first_order, expected_domain_maybe_higher_order)) { - // TODO(mbs): Proper diagnostics. - LOG(FATAL) << "Incompatible virtual devices for expression:" << std::endl - << PrettyPrint(expr_first_order) << std::endl - << "with actual virtual devices:" << std::endl - << ToString(actual_domain_first_order) << std::endl - << "and expected virtual device:" << std::endl - << ToString(expected_domain_maybe_higher_order); - } -} - -bool DeviceDomains::IsFullyConstrained(DeviceDomainPtr domain) { - domain = Lookup(domain); - if (domain->args_and_result_.empty()) { - // First-order. - return domain->virtual_device_->IsFullyConstrained(); - } else { - // Higher-order. - return std::all_of( - domain->args_and_result_.begin(), domain->args_and_result_.end(), - [this](const DeviceDomainPtr& sub_domain) { return IsFullyConstrained(sub_domain); }); - } -} - -void DeviceDomains::SetDefault(DeviceDomainPtr domain, - const VirtualDevice& default_virtual_device) { - ICHECK(!default_virtual_device->IsFullyUnconstrained()); - domain = Lookup(domain); - if (domain->args_and_result_.empty()) { - DeviceDomainPtr default_domain = MakeFirstOrderDomain(config_->CanonicalVirtualDevice( - VirtualDevice::Default(domain->virtual_device_, default_virtual_device))); - DeviceDomainPtr defaulted_domain_ptr = UnifyOrNull(domain, default_domain); - ICHECK(defaulted_domain_ptr != nullptr) << "domain:" << std::endl - << ToString(domain) << std::endl - << "default domain:" << std::endl - << ToString(default_domain); - } else { - for (const auto& sub_domain : domain->args_and_result_) { - SetDefault(sub_domain, default_virtual_device); - } - } -} - -void DeviceDomains::SetResultDefaultThenParams(const DeviceDomainPtr& domain_maybe_higher_order, - const VirtualDevice& default_virtual_device) { - if (domain_maybe_higher_order->args_and_result_.empty()) { - SetDefault(domain_maybe_higher_order, default_virtual_device); - } else { - // First set default for result domain. - SetDefault(ResultDomain(domain_maybe_higher_order), default_virtual_device); - // Then use current result domain as default for everything else. - SetDefault(domain_maybe_higher_order, ResultVirtualDevice(domain_maybe_higher_order)); - } -} - -DeviceDomainPtr DeviceDomains::ResultDomain(DeviceDomainPtr domain) { - domain = Lookup(domain); - while (!domain->args_and_result_.empty()) { - domain = Lookup(domain->args_and_result_.back()); - } - return domain; -} - -std::string DeviceDomains::ToString(DeviceDomainPtr domain) { - domain = Lookup(domain); - std::ostringstream os; - if (domain->args_and_result_.empty()) { - // First-order. - if (!domain->virtual_device_->IsFullyConstrained()) { - os << "?" << static_cast(reinterpret_cast(domain.get())) << "?"; - } - if (!domain->virtual_device_->IsFullyUnconstrained()) { - os << domain->virtual_device_; - } - } else { - // higher-order - os << "fn("; - for (size_t i = 0; i + 1 < domain->args_and_result_.size(); ++i) { - if (i > 0) { - os << ","; - } - os << ToString(domain->args_and_result_[i]); - } - os << "):" << ToString(domain->args_and_result_.back()); - } - return os.str(); -} - -std::string DeviceDomains::ToString() { - std::ostringstream os; - for (const auto& pair : expr_to_domain_) { - os << "expression:" << std::endl - << PrettyPrint(GetRef(pair.first)) << std::endl - << "domain:" << std::endl - << ToString(pair.second) << std::endl - << std::endl; - } - for (const auto& pair : call_to_callee_domain_) { - os << "call:" << std::endl - << PrettyPrint(GetRef(pair.first)) << std::endl - << "callee domain:" << std::endl - << ToString(pair.second) << std::endl - << std::endl; - } - return os.str(); -} - -} // namespace transform -} // namespace relay -} // namespace tvm diff --git a/src/relay/transforms/device_domains.h b/src/relay/transforms/device_domains.h deleted file mode 100644 index 983ecb4b6d5d..000000000000 --- a/src/relay/transforms/device_domains.h +++ /dev/null @@ -1,357 +0,0 @@ -/* - * 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 src/relay/analysis/device_domains.h - * \brief Unification domain for the device planner. - */ - -#ifndef TVM_RELAY_TRANSFORMS_DEVICE_DOMAINS_H_ -#define TVM_RELAY_TRANSFORMS_DEVICE_DOMAINS_H_ - -#include -#include -#include -#include -#include -#include - -#include -#include -#include -#include -#include - -namespace tvm { -namespace relay { -namespace transform { - -class DeviceDomain; -using DeviceDomainPtr = std::shared_ptr; -class DeviceDomains; - -/*! - * \brief Represents the domain over which we collect equality constraints. - * - * \code - * D ::= ?x? -- first order, free - * | -- first order, bound to specific virtual device - * | fn(D1, ..., Dn):Dr -- higher order - * \endcode - * - * We require a function value to be on the same device as its result. To support that we need - * a notion of the 'result domain' of a domain: - * \code - * result_domain(?x?) = ?x? - * result_domain() = - * result_domain(fn(D1, ..., Dn):Dr) = result_domain(Dr) - * \endcode - * - * TODO(mbs): We currently don't allow sub-VirtualDevice constraints. Eg for a function we can - * express that the argument and result VirtualDevices must be exactly equal, but we cannot express - * that though the devices and targets for arguments and results must be equal, it is ok for - * memory scopes to differ. At the moment we can get away with this since we run PlanDevices - * twice: once with all memory scopes unconstrained, then again with just memory scopes as - * the new property to flow. However we're on thin ice here and better would be to allow - * constraints on VirtualDevices to be exploded into their device/target component and their - * memory scope component. Should we fold layout constraints into VirtualDevices then they would - * probably be grouped with memory scopes. - */ -class DeviceDomain { - public: - /*! - * \brief Constructs a first-order domain for \p virtual_device, which may be - * fully free (ie virtual_device is unconstrained), partially free (ie virtual_device has at - * least on of its target, device id or memory scopes known), or fully fixed (ie virtual_device - * has its target, device id and memory scopes set). - * - * CAUTION: Use DeviceDomains::MakeFirstOrderDomain instead of this ctor. - */ - explicit DeviceDomain(VirtualDevice virtual_device) - : virtual_device_(std::move(virtual_device)) {} - - /*! - * \brief Constructs a higher-order domain, where \p args_and_result contain the - * function argument and result domains in order. - * - * CAUTION: Use DeviceDomains::MakeHigherOrderDomain instead of this ctor. - */ - explicit DeviceDomain(std::vector args_and_result) - : virtual_device_(VirtualDevice::FullyUnconstrained()), - args_and_result_(std::move(args_and_result)) {} - - bool is_higher_order() const { return !args_and_result_.empty(); } - - VirtualDevice first_order_virtual_device() const { - ICHECK(args_and_result_.empty()) << "expecting domain to be first-order"; - return virtual_device_; - } - - size_t function_arity() const { - ICHECK(!args_and_result_.empty()) << "expecting domain to be higher-order"; - return args_and_result_.size() - 1UL; - } - - DeviceDomainPtr function_param(size_t i) const { - ICHECK(!args_and_result_.empty()) << "expecting domain to be higher-order"; - ICHECK_LT(i + 1, args_and_result_.size()) << "parameter index is out of range"; - return args_and_result_[i]; - } - - DeviceDomainPtr function_result() const { - ICHECK(!args_and_result_.empty()); - return args_and_result_.back(); - } - - private: - /*! - * \brief If this is a function domain then always fully unconstrained. Otherwise will be - * fully unconstrained (the domain is still completely free), partially constrained - * (for example, the \p target and \p device_type are constrained but the \p virtual_device_id and - * \p memory_scope are still unconstrained), or fully constrained (everything is known). - */ - const VirtualDevice virtual_device_; - - /*! - * \brief If this is a function domain then the sub-domains for each of the function's - * arguments, and the domain for its result. Otherwise empty. - */ - const std::vector args_and_result_; - - friend class DeviceDomains; -}; - -/*! - * \brief Tracks the device domains for a set of expressions w.r.t. an equivalence relation - * built up by calls to \p UnifyOrNull. - */ -class DeviceDomains { - public: - explicit DeviceDomains(CompilationConfig config); - - const CompilationConfig& config() const { return config_; } - - /*! - * \brief Returns the domain representing \p virtual_device. If \p virtual_device is fully - * constrained then the domain will be unique that \p virtual_device. - */ - DeviceDomainPtr MakeFirstOrderDomain(const VirtualDevice& virtual_device); - - /*! - * \brief Returns a higher-order domain with \p args_and_results. - */ - DeviceDomainPtr MakeHigherOrderDomain(std::vector arg_and_results) { - return std::make_shared(std::move(arg_and_results)); - } - - /*! - * \brief Returns a domain appropriate for \p type who's result domain is bound to \p - * virtual_device. If \p type is a function then all parameter domains will be completely free. It - * is valid for \p virtual_device to be fully unconstrained. - */ - DeviceDomainPtr MakeDomain(const Type& type, const VirtualDevice& virtual_device); - - /*! - * \brief Returns a domain with the given result appropriate \p non_canonical_virtual_device, - * which cannot be fully unconstrained. We first canonicalize the virtual device to unsure it has - * a target and is unique. - */ - DeviceDomainPtr ForVirtualDevice(const Type& type, - const VirtualDevice& non_canonical_virtual_device); - - /*! \brief Returns a free domain appropriate for \p type. */ - DeviceDomainPtr Free(const Type& type) { - return MakeDomain(type, VirtualDevice::FullyUnconstrained()); - } - - /*! \brief Returns the domain representing the equivalence class containing \p domain. */ - DeviceDomainPtr Lookup(DeviceDomainPtr domain); - - /*! - * \brief Returns the most constrained domain which agrees with both \p lhs and \p rhs. Returns - * null if no such domain exists, ie some first-order component of \p lhs is constrained - * differently than the corresponding component of \p rhs. - */ - DeviceDomainPtr JoinOrNull(const DeviceDomainPtr& lhs, const DeviceDomainPtr& rhs); - - /*! - * \brief Unifies \p lhs and \p rhs, returning the most-bound of the two. Returns null if - * \p lhs and \p rhs are not unifiable, in which case the constraint system may be left in - * a partially modified state. - */ - // TODO(mbs): I don't think we need an occurs check since the program is well-typed, but - // given we have refs to functions I'm prepared to be surprised. - DeviceDomainPtr UnifyOrNull(DeviceDomainPtr lhs, DeviceDomainPtr rhs); - - /* - * \brief Force all domains in \p higher_order_domain to unify with \p first_order_domain. - * This can be used to handle functions within tuples, references and ADTs since we don't - * attempt to track anything beyond 'the device' for expressions of those first-order types. - * - * Returns false if any unification fails. - */ - bool CollapseOrFalse(const DeviceDomainPtr& first_order_domain, - const DeviceDomainPtr& higher_order_domain); - - /*! - * \brief Unifies \p lhs_first_order and \p rhs_maybe_higher_order. If \p rhs_maybe_higher_order - * is indeed higher-order, require all of its arguments and result to unify with - * \p lhs_first_order. Otherwise same as \p Unify. Returns false if unification is not possible. - * - * In an expression such as: - * \code - * (fn(...) {...}, ...).0 - * \endcode - * we need to force all the devices of the inner function to be the same as the device for the - * overall tuple since the device domain does not understand tuples. Similarly for references - * and ADTs. - */ - bool UnifyCollapsedOrFalse(const DeviceDomainPtr& lhs_first_order, - const DeviceDomainPtr& rhs_maybe_higher_order); - - /*! \brief Returns true if a domain is known for \p expr. */ - bool contains(const Expr& expr) const { return expr_to_domain_.count(expr.get()); } - - /*! \brief Returns the domain representing \p expr. */ - DeviceDomainPtr DomainFor(const Expr& expr); - - /*! - * \brief Returns the domain representing the callee (ie 'op') in \p call expression. If the - * callee is a primitive or special operation we handle it specially. Otherwise defers to \p - * DomainFor(call->op). - * - * This special handling is needed: - * - To handle the "on_device" and "device_copy" ops which constrain devices to the given - * devices. - * - To handle some special ops which constrain devices to the CPU. - * - To allow the same primitive to be called on different devices at different call sites. - * Since each call to the op can have a different domain we index the ops by the call expression - * rather than the op itself. - */ - DeviceDomainPtr DomainForCallee(const Call& call); - - /*! - * \brief Unifies the domains for expressions \p lhs and \p rhs. - * - * Aborts if unification fails. - */ - void UnifyExprExact(const Expr& lhs, const Expr& rhs); - - /*! - * \brief Attempts to unify the domains for expressions \p lhs and \p rhs, however if they - * cannot be unified then returns with no change to the unification system. - */ - void OptionalUnifyExprExact(const Expr& lhs, const Expr& rhs); - - /*! - * \brief Unifies the domain for \p expr with \p expected_domain. - * - * Aborts if unification fails. - */ - void UnifyExprExact(const Expr& expr, const DeviceDomainPtr& expected_domain); - - /*! - * \brief Unifies the domain for \p expr with \p expected_domain. - * If \p expected_domain is higher-order but \p expr is first-order, require all arguments - * and the result of \p expected_domain to have the same domain as for \p expr. - * - * Aborts if unification fails. - */ - void UnifyExprCollapsed(const Expr& expr_first_order, - const DeviceDomainPtr& expected_domain_maybe_higher_order); - - /*! \brief Returns true if \p domain is fully constrainted. */ - bool IsFullyConstrained(DeviceDomainPtr domain); - - /*! \brief Force all \p VirtualDevices in \p domain to default to \p default_virtual_device. */ - void SetDefault(DeviceDomainPtr domain, const VirtualDevice& default_virtual_device); - - /*! - * \brief If \p domain is higher-order default it's result domain to \p default_virtual_device. - * Then force all remaining \p VirtualDevices to the result domain (freshly defaulted or - * original). If \p domain is first-order same as \p SetDefault. - */ - void SetResultDefaultThenParams(const DeviceDomainPtr& domain_maybe_higher_order, - const VirtualDevice& default_virtual_device); - - /*! - * \brief Returns the result domain for \p domain (see defn in DeviceDomain comment). - */ - DeviceDomainPtr ResultDomain(DeviceDomainPtr domain); - - /*! - * \brief Returns the result \p VirtualDevice (possibly unconstrained) for \p domain - * (see defn in DeviceDomain comment). - */ - VirtualDevice ResultVirtualDevice(const DeviceDomainPtr& domain) { - return ResultDomain(domain)->first_order_virtual_device(); - } - - /*! \brief Returns one-line description of \p domain for debugging. */ - std::string ToString(DeviceDomainPtr domain); - - /*! \brief Returns description of entire system of constraints for debugging */ - std::string ToString(); - - private: - /*! \brief Intrinsics we need to handle specially. */ - const Op& alloc_storage_op = Op::Get("memory.alloc_storage"); - const Op& alloc_tensor_op = Op::Get("memory.alloc_tensor"); - const Op& shape_of_op = Op::Get("vm.shape_of"); - const Op& invoke_tvm_op = Op::Get("vm.invoke_tvm_op"); - const Op& reshape_tensor_op = Op::Get("vm.reshape_tensor"); - - CompilationConfig config_; - - /*! - * \brief The domain for first-order expressions of non-tensor type, such as shapes and - * buffer dimensions. Generally this will be a CPU. - */ - DeviceDomainPtr host_domain_; - - /*! \brief Maps expressions to their domains as determined during analysis. */ - std::unordered_map expr_to_domain_; - - /*! - * \brief Maps call expressions to the domains for their callee where the callee is a primitive. - */ - std::unordered_map call_to_callee_domain_; - - /*! \brief Maps device domains to their equivalent domains as determined during unification. */ - std::unordered_map domain_to_equiv_; - - /*! - * \brief Maps fully constrained \p VirtualDevices to their corresponding domains. By sharing - * those domains we can ensure: - * - * \code - * domain0 != domain1 && domain0 fully constrained && domain1 fully constrained - * ==> domain0 and domain1 are incompatible - * \endcode - */ - std::unordered_map - fully_constrained_virtual_device_to_domain_; -}; - -} // namespace transform -} // namespace relay -} // namespace tvm - -#endif // TVM_RELAY_TRANSFORMS_DEVICE_DOMAINS_H_ diff --git a/src/relay/transforms/device_planner.cc b/src/relay/transforms/device_planner.cc deleted file mode 100644 index 80ae66ea9e86..000000000000 --- a/src/relay/transforms/device_planner.cc +++ /dev/null @@ -1,1534 +0,0 @@ -/* - * 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 src/relay/transforms/device_planner.cc - * \brief Determines a unique \p VirtualDevice to hold the result of every Relay sub-expression. - * This pass can be run multiple times, and can be run both before and after lowering. - * - * We say a Relay expression E is 'on device D' if the result of executing E is stored on D. - * We represent D by an \p VirtualDevice, which means we can track anywhere from an arbitrary device - * of some \p DLDeviceType to a specific memory scope on a specific (virtual) \p Device who's - * code is compiled with a specific \p Target. - * - * Note that 'stored on device D' is almost but not quite the same as 'executes on device D', - * see below. - * - * This pass works by collecting and solving device constraints, using defaulting heuristics to - * resolve any remaining undetermined devices, and encoding the results on the output in a form - * that's reasonably friendly to downstream passes. - * - * Specific \p VirtualDevices flow into the constraints from five places: - * - Existing "device_copy" CallNodes (with a \p DeviceCopyAttrs attribute) specify a - * 'src_virtual_device' and 'dst_virtual_device' \p VirtualDevice. Those constrain the argument - * and context of the call respectively. It is ok if source and destination devices are the same, - * such no-op copies will be removed after accounting for the device preference. - * - Existing "on_device" CallNodes (with a \p OnDeviceAttrs attribute) specify an - * 'virtual_device', which constrains the argument of the call, but (usually, see below) leaves the - * context unconstrained. These are called 'annotations' in the rest of the code, have no - * operational significance by themselves, but may trigger the insertion of a new "device_copy" call - * by this pass. In two situations the result of an "on_device" CallNode may also be constrained to - * the given 'virtual_device': - * - The "on_device" call occurs at the top-level of a function body, or occurs as an - * immediately let-bound expression. In this situation the extra degree of freedom in - * the function result and let-binding leads to surprising device copies, so we simply - * force the function result or let-bound variable to the given device. - * - The \p OnDeviceAttrs has an \p is_fixed field of \p true, which indicates we inserted - * it ourselves during an earlier invocation of this pass. This helps make this pass - * idempotent. - * - Some special operators require their arguments or results to be on the 'host' (typcially - * a CPU) \p VirtualDevice, see below. - * - Any \p PrimFuncs in the \p IRModule (if \p LowerTEPass has already run) may constrain their - * argument buffers to have a specific memory scope, which is part of \p VirtualDevice. - * - Annotations left over from a previous run of this pass, such as 'param_virtual_devices' and - * 'result_virtual_device' function attributes we introduce below. This is so the pass is - * idempotent and can be re-run to flow additional memory scope constraints. - * - * We proceed in five phases: - * - * Phase 0 - * ------- - * We rewrite the programs to handle some special cases: - * - "on_device" calls at the top-level of function or immediately let-bound are rewritten - * to have \code is_fixed=true \endcode. - * - We wish to treat \code on_device(expr, device_type=d).0 \endcode as if it were written - * \code on_device(expr.0, device_type_d) \endcode. I.e. we prefer to copy the projection from - * the tuple rather than project from a copy of the tuple. We'll do this by rewriting. - * - We are prepared to insert device_copies on the arguments and result of calls to PrimFuncs, - * on the assumption a) we already ran PlanDevices before lowering so we are not allowing - * any new cross-device copies, but b) after lowering we may have new memory scope constraits - * to deal with. - * - * Phase 1 - * ------- - * We iteratively process the programs and find nodes with conflicting virtual devices. If the - * virtual devices ( \p d1 and \p d2 ) are joinable, they are replaced with a joined device \p d. If - * they are unjoinable, a "device_copy" CallNode is inserted to copy the node output to the second - * device. - * - * Phase 2 - * ------- - * We flow constraints from the "on_device" and "device_copy" calls, PrimFunc buffer memory scopes, - * and some special ops, to all other Relay sub-expressions. - * - * For a primitive such as \code add(e1, e2) \endcode all arguments and results must be on the - * same device. However each call site can use a different device. In other words primitives are - * 'device polymorphic' since we compile and execute them for each required device. ADT constructors - * are similarly polymorphic, but require all constructor args to be on the same device. - * - * For most Relay expressions the device for the overall expression is the same as the device - * for its sub-expressions. E.g. each field of a tuple must be on the same device as the tuple - * itself, the condition and arms of an \p if must all be on the same device as the overall \p if, - * and so on. - * - * Some special ops (or 'dialects') are handled: - * - Relay supports computing the shape of tensors and operators at runtime using "shape_of" - * and "reshape_tensor". Shapes must only be held on the CPU, but the tensors they describe - * may reside on any device. - * - Explicit memory allocation is done using the "alloc_storage" and "alloc_tensor". Again - * shapes reside on the CPU, but the allocated tensors may reside on any device. - * - * Two Relay expression have special handling: - * - For \code let x = e1; e2 \endcode the result of \p e2 must be on the same device as the - * overall let. However the result of \p e1 may be on a different device. - * - For a function \code fn(x, y) { body } \endcode the result of the function must be on the - * same device as \p body. However parameters \p x and \p may be on different devices, even - * different from each other. Every call to the function must use the same choice of parameter - * and result devices -- there is no 'device polymorphism' for Relay functions. - * - * Currently \p PrimFuncs and external functions do not carry over their parameter and result - * devices from their original Relay Function representations. However we know all calls to those - * functions are device-consistent, thus no information is lost. - * - * Phase 3 - * ------- - * After flowing constraints we apply some defaulting heuristics (using a global default \p - * VirtualDevice) to fix the device for any as-yet unconstrained sub-expressions. - * - Unconstrained function result devices default to the global default device. - * - Unconstrained function parameters devices default to the device for the function result. - * - Unconstrained let-bound expression devices default to the device for the overall let. - * TODO(mbs): These are very simple minded heuristics, and ultimately we'd like to treat the - * assignment of the remaining unconstrained sub-expressions as an optimiziation problem in itself. - * This requires a formal notion of 'choicepoint' inside the compiler which can integrate with - * automation. - * - * Phase 4 - * ------- - * Finally, the result of this analysis is reified into the result as: - * - Additional "param_virtual_devices" (an \p Array) and "result_virtual_device" - * (an \p VirtualDevice) attributes for every function (both top-level and local). These describe - * the devices for the function's parameters and the result. - * - Additional "device_copy" CallNodes where a copy is required in order to respect the - * intent of the original "on_device" CallNodes. - * - Additional "on_device" CallNodes where the device type of an expression is not trivially - * implied by the lexically enclosing "on_device" CallNode or function attribute. In practice - * this means "on_device" CallNodes may appear in two places: - * - On let-bound expressions. It is tempting to elide the "on_device" if the let-bound value - * has the same device as the overall let expression. However this would mean passes which - * inline let-bound values, such as FoldConstant and DeadCodeElimination, would need to us - * a DeviceAware visitor which in turn requires the expression to be in ANF to avoid - * deep recursion. To minimize disruption we always include the "on_device" so that it - * can follow the inline. - * - On a call argument if its device differs from the call result. In particular, the - * argument to a "device_copy" call will always be wrapped in an "on_device". (That may - * seem pedantic but simplifies downstream handling.) - * However since we make it easy to track devices for variables we never wrap an "on_device" - * around a var or global var. These uses of "on_device" imply both the argument and result are - * on the same device. We signal this by setting the 'is_fixed' OnDeviceAttrs field to true, - * which helps make this pass idempotent. - * - The buffer maps for called PrimFuncs are updated to capture memory scopes. - * - * Helper visitors (in device_aware_visitors.h) can be used by downstream transforms to recover - * the device for any expression for their own use, e.g. during memory planning. All downstream - * passes must preserve the lexical scoping of the "on_device" CallNodes. E.g. conversion - * to ANF must respect the lexical scoping convention: - * \code - * f(on_device(g(h(a, b), c), virtual_device=CPU)) - * ==> - * let %x0 = on_device(h(a, b), virtual_device=CPU) - * let %x1 = on_device(g(%x0), virtual_device=CPU) - * f(on_device(%x1, virtual_device=CPU)) - * \endcode - * - * This pass can be run before FuseOps so that it can use device-specific fusion rules. - * - * 'Stored on' vs 'Executes on' - * ---------------------------- - * Obviously for a primitive call \code add(x, y) \endcode we can execute the primitive on the - * same device as will hold its result. Thus 'executes on' is the same as 'stored on' for - * primitives. - * - * But what about for arbitrary Relay expressions? Most backends (interpreter, graph, VM) are - * implicitly executed on the 'host' CPU, with only primitive evaluation handed off to specific - * devices, thus the notion of 'executes on' is mute. AOT backends on the other hand need to - * know exactly which device (possibly one of a number of available 'CPU'-like devices) is - * responsible for execution. Currently that's handled independently by the \p AnnotateTargets - * pass, but we'd like to fold that into device planning here to ensure everything is consistent. - * - * Obviously since tensors are passed-by-pointer it's quite possible to execute a Relay - * expression (eg an \p if expression) on one device even though the tensor data resides on - * another. But for AOT that flexibility seems excessive. So we'd like to just take 'executes on' - * to be 'stored on' exactly. In particular, for a Relay function, we'd like to be able to just - * compile the function body for the function's result device. - * - * This works after conversion to ANF provided the compilation for a let expression is prepared - * to make a cross-device call. However we leave it to a downstream transformation to heuristically - * minimize cross-device calls by moving device copies out of functions. E.g.: - * \code - * def @f() { // execute on CPU - * let x = on_device(...GPU computation..., virtual_device=GPU); - * device_copy(...GPU computation..., src_dev_type=GPU, dst_dev_type=CPU) - * } - * def @main() { - * ... call @f() on CPU ... - * } - * \endcode - * could be rewritten to: - * \code - * def @f() { // execute on GPU - * let x = ...GPU computation...; - * ...GPU computation... - * } - * def @main() { - * let x = device_copy(@f(), src_dev_type=GPU, dst_dev_type=CPU) - * ... use x on CPU ... - * } - * \endcode - * - * Higher-order shenanigans - * ------------------------ - * Relay is a 'mostly' higher-order language -- we can let-bind functions, pass functions - * as arguments (even anonymous functions), return functions, evaluate conditional expressions - * over functions, and so on. We handle this during constraint solving using the domain: - * \code - * D ::= -- first-order - * | fn(D,...,D):D -- higher-order - * \endcode - * In this way we can determine the device for all function parameters and results. E.g. for - * \code - * let f = fn(x, y) { ... } - * let g = fn(f, z) { f(z, z) } - * g(f, on_device(..., virtual_device=CPU)) - * \endcode - * the parameters \p x and \p y will be on the CPU. - * - * But now look closely at the call \code e1(e2, e3) \endcode. We know \p e1 must evaluate to a - * function. Our analysis must guarantee that the function's parameters and result devices are - * consistent for \p e2, \p e3, and the context of the call. But: - * - Which device holds the closure result of evaluating \p e1 ? - * - If \p e2 is of function type, what does that mean when we say every function parameter - * is on a device? - * - If \p e1 returns a function, what does that mean when we say every function result is - * on a device? - * - * Since higher-order aspects are later compiled away (by 'defunctionalization' - * aka 'firstification') we'd prefer not to have to answer any of those questions. In particular, - * we really don't want our domain \p D to allow for yet another device for the function closure. - * So we'll just force the 'device for a function' to be the same as the device for the function's - * result using the notion of the 'result domain' for a domain: - * \code - * result_domain() = - * result_domain(fn(D1,...,Dn):Dr) = result_domain(Dr) - * \endcode - * - * Similarly the domain does not have entries for tuples, references, or ADTs. Whenever the - * analysis encounters a function inside one of those it simply forces all argument and result - * devices for the function to match the device for the first-order expression. For example, - * if the tuple \code (fn(x, y) { ... }, 3) \endcode is on the GPU then the inner function - * parameters and result must similarly be on the GPU. - * - * ------- - * | AOR | This pass supports all of Relay. - * ------- - * ^ - * | - * `-- Mark's stamp of completeness :-) - * - * TODO(mbs): Proper diagnostics for unification failure using spans. - * TODO(mbs): We may want some 'device polymorphism' for Relay functions. Eg it's ok for the - * function to be called with params/result on different (virtual) device ids provided the target - * and memory scopes are consistent. - */ - -#include -#include -#include -#include -#include -#include -#include -#include -#include -#include -#include -#include -#include -#include - -#include - -#include "../../tir/analysis/device_constraint_utils.h" -#include "../op/annotation/annotation.h" -#include "../op/memory/device_copy.h" -#include "../op/memory/on_device.h" -#include "./device_domains.h" - -namespace tvm { -namespace relay { -namespace transform { - -namespace { - -/* =============== Phase 0 =============== */ - -/*! - * \brief Rewrites "on_device" calls to handle some special cases. - * - * - Don't let the device for %x remain unconstrained: - * \code - * let %x = on_device(e, virtual_device=d) - * ==> let %x = on_device(e, virtual_device=d, constraint=kBoth) - * \endcode - * - * - Don't let the function result remain unconstrained: - * \code - * fn(%x) { on_device(e, virtual_device=d) } - * ==> fn(%x) { on_device(e, virtual_device=d, constraint=kBoth) - * \endcode - * - * - Project-then-copy rather than copy-then-project: - * \code - * on_device(e).0 - * ==> on_device(e.0) - * \endcode - * - * - Be prepared to copy arguments and results on primitive call boundaries in case memory - * scopes don't line up. We'll use the 'fully unconstrained' version of on_device so that - * we can allow for a device_copy without knowing the specific device for the arguments. - * \code - * call_lowered(@prim, (a, b)) - * ==> copy_ok(call_lowered(@prim, (copy_ok(a), copy_ok(b)))) - * where - * copy_ok(x) = on_device(x, virtual_device=VirtualDevice::FullyUnconstrained, - * constrain_body=False, constrain_result=False) - * \endcode - */ -class RewriteOnDevices : public ExprMutator { - public: - explicit RewriteOnDevices(IRModule mod) : mod_(std::move(mod)) {} - - private: - Expr VisitExpr_(const TupleGetItemNode* tuple_get_item_node) final { - Expr tuple = VisitExpr(tuple_get_item_node->tuple); - OnDeviceProps props = GetOnDeviceProps(tuple); - - Expr tuple_get_item = WithFields(GetRef(tuple_get_item_node), tuple); - if (props.body.defined() && props.is_normal()) { - VLOG(2) << "wrapping tuple get item:" << std::endl - << PrettyPrint(GetRef(tuple_get_item_node)) << std::endl - << "with \"on_device\" for VirtualDevice " << props.virtual_device; - return OnDeviceWithProps(tuple_get_item, props); - } else { - return tuple_get_item; - } - } - - Expr VisitExpr_(const LetNode* let_node) final { - auto expr = GetRef(let_node); - std::vector> bindings; - while (auto opt = expr.as()) { - auto inner_let = opt.value(); - Expr value = VisitExpr(inner_let->value); - OnDeviceProps props = GetOnDeviceProps(value); - if (props.body.defined() && props.is_normal()) { - VLOG(2) << "revising let-bound expression of let:" << std::endl - << PrettyPrint(expr) << std::endl - << "to be fixed to VirtualDevice " << props.virtual_device; - value = MaybeOnDeviceFixed(props.body, props.virtual_device); - } - bindings.emplace_back(inner_let, value); - expr = inner_let->body; - } - expr = VisitExpr(expr); - for (auto itr = bindings.rbegin(); itr != bindings.rend(); ++itr) { - expr = WithFields(/*let=*/std::get<0>(*itr), /*opt_var=*/{}, - /*opt_value=*/std::get<1>(*itr), /*opt_body=*/expr); - } - return expr; - } - - Expr VisitExpr_(const FunctionNode* function_node) final { - Expr body = VisitExpr(function_node->body); - OnDeviceProps props = GetOnDeviceProps(body); - if (props.body.defined() && props.is_normal()) { - VLOG(2) << "revising body of function:" << std::endl - << PrettyPrint(GetRef(function_node)) << std::endl - << "to be fixed to VirtualDevice " << props.virtual_device; - body = MaybeOnDeviceFixed(props.body, props.virtual_device); - } - return WithFields(GetRef(function_node), function_node->params, body); - } - - Expr VisitExpr_(const CallNode* call_node) final { - CallLoweredProps props = GetCallLoweredProps(call_node); - if (props.lowered_func.defined()) { - BaseFunc base_func = mod_->Lookup(props.lowered_func); - if (base_func.as()) { - VLOG(2) << "allowing device_copy on PrimFunc arguments and result"; - Array new_args; - new_args.reserve(props.arguments.size()); - for (const auto& arg : props.arguments) { - Expr new_arg = VisitExpr(arg); - new_args.push_back(OnDeviceCopyOk(std::move(new_arg))); - } - Call new_call = CallLowered(std::move(props.lowered_func), std::move(new_args), props.attrs, - call_node->span); - return OnDeviceCopyOk(std::move(new_call)); - } - } - return ExprMutator::VisitExpr_(call_node); - } - - /*! \brief Module we are rewriting, so we can lookup global definitions. */ - IRModule mod_; -}; - -/* =============== Phase 1 =============== */ - -/*! - * \brief Add "device_copy" calls for nodes that have conflicting virtual devices. - * - * Eg Suppose an IRModule contains the following expr: - * \code - * %0 = add(%a, %b); - * %1 = on_device(%0, virtual_device=d1); - * %2 = add(%b, %c); - * %3 = on_device(%2, virtual_device=d2); - * \endcode - * In the above example, node %b has two possible virtual devices: \p d1 and \p d2. - * - * - If \p d1 and \p d2 are joinable, replace \p d1 and \p d2 with the joined device \p d: - * \code - * %0 = add(%a, %b); - * %1 = on_device(%0, virtual_device=d); - * %2 = add(%b, %c); - * %3 = on_device(%2, virtual_device=d); - * \endcode - * - * - If \p d1 and \p d2 are unjoinable, insert a "device_copy" CallNode to copy \p %b to \p d2: - * \code - * %0 = add(%a, %b); - * %1 = on_device(%0, virtual_device=d); - * %2 = device_copy(%b, src_dev_type=d1, dst_dev_type=d2); - * %3 = add(%2, %c); - * %4 = on_device(%3, virtual_device=d); - * \endcode - */ -struct DeviceContext { - VirtualDevice VirtualDeviceFor(const ExprNode* expr) { - auto itr = expr_to_device.find(expr); - if (itr != expr_to_device.end()) { - return itr->second; - } - auto default_dev = VirtualDevice::FullyUnconstrained(); - expr_to_device.emplace(expr, default_dev); - return default_dev; - } - - bool Update(const ExprNode* expr, VirtualDevice dev) { - bool success = true; - auto pair = expr_to_device.emplace(expr, dev); - if (!pair.second) { - auto replaced_item = pair.first; - auto joined_dev = VirtualDevice::Join(replaced_item->second, dev); - if (joined_dev == nullptr) { - success = false; - } else { - replaced_item->second = joined_dev.value(); - } - } - return success; - } - - bool IsConflicted(const ExprNode* expr) { - auto itr = conflicted_nodes.find(expr); - return itr != conflicted_nodes.end(); - } - - std::unordered_set conflicted_nodes; - std::unordered_map expr_to_device; -}; - -/*! - * \brief Flow the device constraints over the module and find all the conflicted nodes. The - * conflicted nodes only contain nodes that have no explicit constraints. For example, "on_device" - * nodes are not considered as conflicted. - */ -class ConflictedNodeFinder : ExprVisitor { - public: - explicit ConflictedNodeFinder(IRModule mod) - : mod_(std::move(mod)), dev_ctx_(std::make_unique()) {} - - std::unique_ptr Finder() { - VLOG_CONTEXT << "ConflictedNodeFinder"; - for (const auto& kv : mod_->functions) { - if (const auto* function_node = AsOptimizableFunctionNode(kv.second)) { - VisitExpr(GetRef(function_node)); - } - } - for (auto const node : dev_ctx_->conflicted_nodes) { - if (node->IsInstance()) { - auto call = Downcast(GetRef(node)); - // "DeviceCapturer" will insert "device_copy" for "on_device" calls. - // Therefore, "on_device" should not be considered as conflicted. - if (call->op == OnDeviceOp()) { - dev_ctx_->conflicted_nodes.erase(node); - } - } - } - return std::move(dev_ctx_); - } - - private: - void VisitExpr_(const CallNode* call_node) final { - VLOG(2) << "Initial call node: " << std::endl << PrettyPrint(GetRef(call_node)); - auto call_dev = dev_ctx_->VirtualDeviceFor(call_node); - auto body_dev = call_dev; - - auto on_dev_props = GetOnDeviceProps(call_node); - auto dev_cp_props = GetDeviceCopyProps(call_node); - if (call_node->op == OnDeviceOp()) { - if (on_dev_props.constrain_body) { - body_dev = on_dev_props.virtual_device; - } - if (on_dev_props.constrain_result) { - call_dev = on_dev_props.virtual_device; - } - } else if (call_node->op == DeviceCopyOp()) { - body_dev = dev_cp_props.src_virtual_device; - call_dev = dev_cp_props.dst_virtual_device; - } - - if (!dev_ctx_->Update(call_node, call_dev) && call_node->op != OnDeviceOp()) { - LOG(FATAL) << "Mismatched device type after iterating args. Implied device: " << std::endl - << PrettyPrint(call_dev) << "and practial device:" << std::endl - << PrettyPrint(dev_ctx_->VirtualDeviceFor(call_node)) << std::endl - << "With CallNode: " << std::endl - << PrettyPrint(GetRef(call_node)); - } - - for (auto& arg : call_node->args) { - VLOG(3) << "Handle call node arg: " << std::endl << PrettyPrint(arg); - if (!dev_ctx_->Update(arg.get(), body_dev)) { - VLOG(2) << "Conflicted node found:" << std::endl - << PrettyPrint(GetRef(arg.get())) << std::endl - << "With corresponding Callee:" << std::endl - << PrettyPrint(GetRef(call_node)); - dev_ctx_->conflicted_nodes.emplace(arg.get()); - } - } - for (auto& expr : call_node->args) { - VisitExpr(expr); - } - } - - IRModule mod_; - std::unique_ptr dev_ctx_; -}; - -/*! - * \brief Insert "device_copy" CallNode for all the conflicted nodes found by \p - * ConflictedNodeFinder. - */ -class ConflictedNodeRewriter : ExprMutator { - public: - ConflictedNodeRewriter(IRModule mod, CompilationConfig config, - std::unique_ptr dev_ctx) - : mod_(mod), config_(config), dev_ctx_(std::move(dev_ctx)) {} - - IRModule Rewrite() { - VLOG_CONTEXT << "ConflictedNodeRewriter"; - IRModule result(/*functions=*/{}, mod_->type_definitions, mod_->Imports(), mod_->source_map, - mod_->attrs); - for (const auto& kv : mod_->functions) { - if (const auto* function_node = AsOptimizableFunctionNode(kv.second)) { - auto func = Mutate(GetRef(function_node)); - result->Add(kv.first, Downcast(func)); - } else { - result->Add(kv.first, kv.second); - } - } - - return result; - } - - private: - Expr VisitExpr_(const CallNode* call_node) final { - VLOG(3) << "Initial call node:" << std::endl << PrettyPrint(GetRef(call_node)); - auto call = Downcast(ExprMutator::VisitExpr_(call_node)); - tvm::Array call_args; - call_args.reserve(call_node->args.size()); - for (auto arg : call->args) { - if (dev_ctx_->IsConflicted(arg.get())) { - auto src_dev = config_->CanonicalVirtualDevice(dev_ctx_->VirtualDeviceFor(arg.get())); - auto dst_dev = config_->CanonicalVirtualDevice(dev_ctx_->VirtualDeviceFor(call_node)); - call_args.push_back(MaybeDeviceCopy(arg, src_dev, dst_dev)); - VLOG(2) << "Adding DeviceCopy Op: " << std::endl << PrettyPrint(call_args.back()); - } else { - call_args.push_back(arg); - } - } - auto new_call = WithFields(GetRef(call_node), call_node->op, call_args); - VLOG(3) << "Final call node:" << std::endl << PrettyPrint(GetRef(call_node)); - return new_call; - } - - IRModule mod_; - CompilationConfig config_; - std::unique_ptr dev_ctx_; -}; - -/* =============== Phase 2 =============== */ - -/* - * \brief Collects the system of device constraints for all sub-expressions in a module. - * It is possible some devices remain free and will need to be defaulted by \p DeviceDefaulter. - * - * Eg from \code add(%x, %y) \endcode we know \p %x and \p %y must be on the same device. Later, - * from \code on_device(%x, virtual_device=d) \endcode we know \p %x must be on device \p d, and - * thus so must \p %y. - * - * Constraints can flow in interesting ways. E.g. in: - * \code - * let %f = fn(%x, %y) { add(%x, on_device(%y, virtual_device=d)) } - * let %g = fn(%f, %x, %y) { %f(%x, %y) } - * %g(%f, %a, %b) - * \endcode - * we discover \p %b must be on device \p d. - */ -class DeviceAnalyzer : public MixedModeVisitor { - public: - DeviceAnalyzer(IRModule mod, CompilationConfig config) - : mod_(std::move(mod)), domains_(std::make_unique(std::move(config))) {} - - /*! - * \brief Returns the expression-to-device-domain map for all expressions in all the global - * function definitions in the module. Expressions may have free domains, these will be resolved - * by \p DeviceDefaulter below. - */ - std::unique_ptr Analyze() { - VLOG_CONTEXT << "DeviceAnalyzer"; - for (const auto& kv : mod_->functions) { - // The global variable and what it is bound to must obviously agree on domain. - if (const auto* function_node = AsOptimizableFunctionNode(kv.second)) { - VLOG(2) << "collecting constraints from Relay Function '" << kv.first->name_hint << "'"; - domains_->UnifyExprExact(kv.first, kv.second); - VisitExpr(GetRef(function_node)); - } else if (auto prim_func = kv.second.as()) { - VLOG(2) << "collecting constraints from TIR PrimFunc '" << kv.first->name_hint << "'"; - domains_->UnifyExprExact(kv.first, DomainForPrimFunc(kv.first, prim_func.value())); - } else { - VLOG(2) << "skipping '" << kv.first->name_hint << "'"; - } - } - return std::move(domains_); - } - - private: - /*! - * \brief Return the domain representing \p prim_func which, before lowering, had - * the Relay \p type. - */ - DeviceDomainPtr DomainForPrimFunc(const GlobalVar& global_var, const tir::PrimFunc& prim_func) { - // CAUTION: The prim_func->checked_type() is currently w.r.t. the flattened and DPS form - // of the prim func, however here we wish to remain within the Relay view of all functions. - // Thus we'll use the global var who's checked_type is in Relay form. - auto func_domain = domains_->DomainFor(global_var); // higher-order - - // TODO(mbs): We don't visit the body of the function -- there's currently nothing to be done. - const auto* func_type_node = global_var->checked_type().as(); - ICHECK(func_type_node); - ICHECK_EQ(func_domain->function_arity(), func_type_node->arg_types.size()); - - Array virtual_devices = - tir::GetPrimFuncArgAndResultConstraints(prim_func, GetRef(func_type_node)); - - // Build the implied domain (in terms of the function's Relay type) implied by any memory scope - // constrains in the function's buffers, for both arguments and results. - std::vector args_and_result_domains; - args_and_result_domains.reserve(virtual_devices.size()); - for (size_t i = 0; i < func_type_node->arg_types.size(); ++i) { - const VirtualDevice& param_virtual_device = virtual_devices[i]; - VLOG(2) << "param_virtual_device[" << i << "] = " << param_virtual_device; - args_and_result_domains.push_back(domains_->MakeFirstOrderDomain(param_virtual_device)); - } - const VirtualDevice& ret_virtual_device = virtual_devices.back(); - VLOG(2) << "ret_virtual_device = " << ret_virtual_device; - args_and_result_domains.push_back(domains_->MakeFirstOrderDomain(ret_virtual_device)); - - return domains_->MakeHigherOrderDomain(std::move(args_and_result_domains)); - } - - void VisitExpr_(const CallNode* call_node) final { - auto call = GetRef(call_node); - - // We don't care if the call is in pre- or post-lowered form. - auto vanilla_call = GetAnyCall(call_node); - - // Find the higher-order domain for the callee. See DomainForCallee for the special rules - // for primitives. - VisitExpr(vanilla_call->op); - auto func_domain = domains_->DomainForCallee(call); // higher-order - - // Build the domain for the function implied by its arguments and call context. - ICHECK_EQ(func_domain->function_arity(), vanilla_call->args.size()) << PrettyPrint(call); - std::vector args_and_result_domains; - args_and_result_domains.reserve(vanilla_call->args.size() + 1); - for (const auto& arg : vanilla_call->args) { - args_and_result_domains.emplace_back(domains_->DomainFor(arg)); - } - args_and_result_domains.emplace_back(domains_->DomainFor(call)); - auto implied_domain = - domains_->MakeHigherOrderDomain(std::move(args_and_result_domains)); // higher-order - - VLOG(2) << "initial call function domain:" << std::endl - << domains_->ToString(func_domain) << std::endl - << "and implied domain:" << std::endl - << domains_->ToString(implied_domain) << std::endl - << "for call:" << std::endl - << PrettyPrint(call); - - // The above must match. - if (domains_->UnifyOrNull(func_domain, implied_domain) == nullptr) { // higher-order - // TODO(mbs): Proper diagnostics. - LOG(FATAL) - << "Function parameters and result VirtualDevices do not match those of call. Call:" - << std::endl - << PrettyPrint(call) << std::endl - << "with function virtual devices:" << std::endl - << domains_->ToString(func_domain) << std::endl - << "and implied call virtual devices:" << std::endl - << domains_->ToString(implied_domain); - } - - VLOG(2) << "final call function domain:" << std::endl - << domains_->ToString(func_domain) << std::endl - << "for call:" << std::endl - << PrettyPrint(call); - } - - void VisitExpr_(const LetNode* let_node) final { - Expr expr = GetRef(let_node); - // Iteratively visit let nodes to avoid stack overflow. - while (expr->IsInstance()) { - Let let = Downcast(expr); - // Let var must be same device as value it is bound to. - domains_->UnifyExprExact(let->var, let->value); // may be higher-order - // Let body must be same device as overall let. - domains_->UnifyExprExact(let, let->body); // may be higher-order - - VisitExpr(let->var); - VisitExpr(let->value); - - expr = let->body; - } - - // Visit the last body - VisitExpr(expr); - } - - void VisitExpr_(const FunctionNode* function_node) final { - auto function = GetRef(function_node); - auto func_domain = domains_->DomainFor(function); // higher-order - ICHECK_EQ(func_domain->function_arity(), function_node->params.size()); - - VLOG(2) << "initial function domain:" << std::endl - << domains_->ToString(func_domain) << std::endl - << "and function body domain:" << std::endl - << domains_->ToString(domains_->DomainFor(function_node->body)) << std::endl - << "for function:" << std::endl - << PrettyPrint(function); - - // The function body domain must match the function result domain. - domains_->UnifyExprExact(function_node->body, - func_domain->function_result()); // may be higher-order - if (!function_node->virtual_device()->IsFullyUnconstrained()) { - // The function body domain must match any existing virtual device annotation. - domains_->UnifyExprExact(function_node->body, - domains_->ForVirtualDevice(function_node->body->checked_type(), - function_node->virtual_device())); - } - - for (size_t i = 0; i < function_node->params.size(); ++i) { - const auto& param = function_node->params[i]; - // The parameter domain must match the function argument domain. - domains_->UnifyExprExact(param, - func_domain->function_param(i)); // may be higher-order - if (!param->virtual_device()->IsFullyUnconstrained()) { - // The parameter domain must match any existing virtual device annotation. - domains_->UnifyExprExact( - param, domains_->ForVirtualDevice(param->checked_type(), param->virtual_device())); - } - VisitExpr(param); - } - - // No need to step into the body of Primitive functions. - if (!function_node->HasNonzeroAttr(attr::kPrimitive)) { - VisitExpr(function_node->body); - } - - VLOG(2) << "final function domain:" << std::endl - << domains_->ToString(func_domain) << std::endl - << "and function body domain:" << std::endl - << domains_->ToString(domains_->DomainFor(function_node->body)) << std::endl - << "for function:" << std::endl - << PrettyPrint(function); - } - - void VisitExpr_(const TupleNode* tuple_node) final { - Tuple tuple = GetRef(tuple_node); - for (size_t i = 0; i < tuple->fields.size(); i++) { - auto domain = domains_->DomainFor(tuple->fields[i]); // may be higher-order - domains_->UnifyExprCollapsed(tuple, domain); // collapse to first-order if needed - } - } - - void VisitExpr_(const TupleGetItemNode* tuple_get_item_node) final { - TupleGetItem tuple_get_item = GetRef(tuple_get_item_node); - auto domain = domains_->DomainFor(tuple_get_item); // may be higher-order - domains_->UnifyExprCollapsed(tuple_get_item_node->tuple, - domain); // collapse to first-order if needed - } - - class DevicePatternAnalyzer : public PatternVisitor { - public: - DevicePatternAnalyzer(DeviceDomains* domains, const ExprNode* adt_node) - : domains_(domains), adt_node_(adt_node) {} - - private: - void VisitPattern_(const PatternVarNode* pattern_var_node) final { - auto var_domain = domains_->DomainFor(pattern_var_node->var); // may be higher order - domains_->UnifyExprCollapsed(GetRef(adt_node_), - var_domain); // collapse to first-order if needed - } - - /*! \brief (Mutable borrow of) the domains for all expressions processed so far. */ - DeviceDomains* domains_; - /*! \brief The expression for the ADT we are matching over. */ - const ExprNode* adt_node_; - }; - - void VisitPattern(const Pattern& pattern) final {} - - void VisitExpr_(const MatchNode* match_node) final { - // For match node, we unify the value and the rhs of each clause - Match match = GetRef(match_node); - auto match_domain = domains_->DomainFor(match); // may be higher-order - DevicePatternAnalyzer pattern_analyzer(domains_.get(), match->data.get()); - domains_->UnifyExprCollapsed(match->data, match_domain); // collapse to first-order if needed - for (const auto& clause : match->clauses) { - pattern_analyzer.VisitPattern(clause->lhs); - domains_->UnifyExprExact(clause->rhs, match_domain); - VisitExpr(clause->rhs); - } - VisitExpr(match_node->data); - } - - void VisitExpr_(const GlobalVarNode* global_var_node) final { - domains_->DomainFor(GetRef(global_var_node)); - } - - void VisitExpr_(const VarNode* var_node) final { domains_->DomainFor(GetRef(var_node)); } - - void VisitExpr_(const ConstantNode* constant_node) final { - domains_->DomainFor(GetRef(constant_node)); - } - - void VisitExpr_(const ConstructorNode* constructor_node) final { - // no-op, constructors are handled at their call-sites. - // TODO(mbs): Assumes eta-expansion - } - - void VisitExpr_(const IfNode* if_node) final { - auto ife = GetRef(if_node); - auto domain = domains_->DomainFor(ife); // may be higher-order - domains_->UnifyExprCollapsed(if_node->cond, domain); // collapse to first-order if needed - domains_->UnifyExprExact(if_node->true_branch, domain); - domains_->UnifyExprExact(if_node->false_branch, domain); - VisitExpr(if_node->cond); - VisitExpr(if_node->true_branch); - VisitExpr(if_node->false_branch); - } - - void VisitExpr_(const OpNode* op) final { - // no-op, primitive operators are handled at their call-sites. - } - - void VisitExpr_(const RefCreateNode* ref_create_node) final { - auto ref_create = GetRef(ref_create_node); - auto domain = domains_->DomainFor(ref_create_node->value); // may be higher-order - domains_->UnifyExprCollapsed(ref_create, domain); // collapse to first-order if needed - VisitExpr(ref_create_node->value); - } - - void VisitExpr_(const RefReadNode* ref_read_node) final { - auto ref_read = GetRef(ref_read_node); - auto domain = domains_->DomainFor(ref_read); // may be higher-order - domains_->UnifyExprCollapsed(ref_read_node->ref, domain); // collapse to first-order if needed - VisitExpr(ref_read_node->ref); - } - - void VisitExpr_(const RefWriteNode* ref_write_node) final { - auto ref_write = GetRef(ref_write_node); - auto domain = domains_->DomainFor(ref_write->value); // may be higher-order - domains_->UnifyExprCollapsed(ref_write->ref, domain); // collapse to first-order if needed - domains_->UnifyExprCollapsed(ref_write, domain); // collapse to first-order if needed - VisitExpr(ref_write_node->ref); - VisitExpr(ref_write_node->value); - } - - /*! \brief The module we are analyzing. */ - IRModule mod_; - /*! \brief The domains for all expressions processed so far. */ - std::unique_ptr domains_; -}; - -/* =============== Phase 3 =============== */ - -/*! - * \brief Calls to 'free' "on_device" annotations (ie where both constrain_body=false and - * constrain_result=false) indicate a device_copy is allowed if required, but no particular - * device is imposed on the body or the context. At this stage we can attempt to unify the - * body and device contexts. In this way we can avoid the defaulting rules in \p DeviceDefaulter - * from choosing default devices which are only going to induce a device copy. - * - * TODO(mbs): The order in which we encounter the "on_device" calls can influence the final global - * device assignment. However we visit global functions in hash map order. - */ -class FreeOnDeviceDefaulter : public ExprVisitor { - public: - FreeOnDeviceDefaulter(IRModule mod, std::unique_ptr domains) - : mod_(std::move(mod)), domains_(std::move(domains)) {} - - std::unique_ptr Default() { - VLOG_CONTEXT << "FreeOnDeviceDefaulter"; - VLOG(0) << "unifying free on_device annotations"; - for (const auto& kv : mod_->functions) { - if (const auto* function_node = AsOptimizableFunctionNode(kv.second)) { - VLOG(2) << "unifying for '" << kv.first->name_hint << "'"; - VisitExpr(GetRef(function_node)); - } else { - VLOG(2) << "skipping '" << kv.first->name_hint << "'"; - } - } - return std::move(domains_); - } - - private: - void VisitExpr_(const CallNode* call_node) final { - auto call = GetRef(call_node); - OnDeviceProps props = GetOnDeviceProps(call_node); - ExprVisitor::VisitExpr_(call_node); - if (props.body.defined() && !props.constrain_body && !props.constrain_result) { - domains_->OptionalUnifyExprExact(call, props.body); - } - } - - /*! \brief The module we are processing. */ - IRModule mod_; - /*! \brief The domains for all expressions. */ - std::unique_ptr domains_; -}; - -/*! - * \brief Ensures every sub-expression in a module has a device type, using both the global - * default and some local heuristics to avoid unnecessary additional "device_copy" CallNodes. - * - * E.g. in: - * \code - * def @main(%x, %y, %z) { - * let %a = add(%x, %y); - * multiply(%a, on_device(%z, virtual_device=d)) - * } - * \endcode - * we know the parameter \p %z must be on device \p d, but the devices for \p %x and \p %y, - * and the device for the function result, are still 'free'. The global 'default' device type - * is first used to 'fix' \p main's result type, which in turn 'fixes' \p %x and \p %y, which - * in turn 'fixes' the device on which the \p add and \p multiply are executed. - * - * TODO(mbs): I think this is deterministic? We do however visit the top-level defs in hashmap - * order. - */ -class DeviceDefaulter : public ExprVisitor { - public: - DeviceDefaulter(IRModule mod, std::unique_ptr domains) - : mod_(std::move(mod)), domains_(std::move(domains)) {} - - std::unique_ptr Default() { - VLOG_CONTEXT << "DeviceDefaulter"; - VLOG(0) << "defaulting to VirtualDevice " - << domains_->config()->default_primitive_virtual_device; - for (const auto& kv : mod_->functions) { - if (const auto* function_node = AsOptimizableFunctionNode(kv.second)) { - VLOG(2) << "defaulting devices for '" << kv.first->name_hint << "'"; - VisitExpr(GetRef(function_node)); - } else { - VLOG(2) << "skipping '" << kv.first->name_hint << "'"; - } - } - return std::move(domains_); - } - - private: - void VisitExpr_(const FunctionNode* function_node) final { - if (function_node->HasNonzeroAttr(attr::kPrimitive)) { - return; - } - - auto function = GetRef(function_node); - auto func_domain = domains_->DomainFor(function); // higher-order - ICHECK_EQ(func_domain->function_arity(), function_node->params.size()); - if (!domains_->IsFullyConstrained(func_domain)) { - VLOG(2) << "before defaulting function:" << std::endl << domains_->ToString(func_domain); - domains_->SetResultDefaultThenParams(func_domain, - domains_->config()->default_primitive_virtual_device); - VLOG(2) << "after defaulting function:" << std::endl << domains_->ToString(func_domain); - } - VisitExpr(function_node->body); - } - - void VisitExpr_(const CallNode* call_node) final { - auto call = GetRef(call_node); - - // We don't care if the call is pre- or post-lowered. - auto vanilla_call = GetAnyCall(call_node); - - auto func_domain = domains_->DomainForCallee(call); // higher-order - ICHECK_EQ(func_domain->function_arity(), vanilla_call->args.size()); - if (!domains_->IsFullyConstrained(func_domain)) { - // For calls to Relay functions this step is identical to that for VisitExpr_(FunctionNode*) - // above. But for calls to primitives we may still need to force free domains to be - // defaulted. - VLOG(2) << "before defaulting callee:" << std::endl - << PrettyPrint(call_node->op) << std::endl - << "of domain:" << std::endl - << domains_->ToString(func_domain); - domains_->SetResultDefaultThenParams(func_domain, - domains_->config()->default_primitive_virtual_device); - VLOG(2) << "after defaulting callee:" << std::endl - << PrettyPrint(call_node->op) << std::endl - << "of domain:" << std::endl - << domains_->ToString(func_domain); - } - return ExprVisitor::VisitExpr_(call_node); - } - - void VisitExpr_(const LetNode* let_node) final { - Expr expr = GetRef(let_node); - // Iteratively visit let nodes to avoid stack overflow. - while (expr->IsInstance()) { - Let let = Downcast(expr); - // If the let-var device is still free force it to match the overall let. - auto let_domain = domains_->DomainFor(let); // may be higher-order - VirtualDevice let_virtual_device = domains_->ResultVirtualDevice(let_domain); - ICHECK(!let_virtual_device->IsFullyUnconstrained()); - auto let_var_domain = domains_->DomainFor(let->var); // may be higher-order - if (!domains_->IsFullyConstrained(let_var_domain)) { - VLOG(2) << "before defaulting let-var:" << std::endl << domains_->ToString(let_var_domain); - domains_->SetDefault(let_var_domain, let_virtual_device); - VLOG(2) << "after defaulting let-var:" << std::endl << domains_->ToString(let_var_domain); - } - VisitExpr(let->var); - VisitExpr(let->value); - expr = let->body; - } - VisitExpr(expr); - } - - /*! \brief The module we are processing. */ - IRModule mod_; - /*! \brief The domains for all expressions. */ - std::unique_ptr domains_; -}; - -/* =============== Phase 4 =============== */ -/*! - * \brief Inserts missing "device_copy" CallNodes, and ensures the device type of every - * sub-expression in a module can be easily recovered by a later transformation using simple - * lexical scoping rules (e.g. for memory planning). - * - * - Discard any existing "on_device" CallNodes since their job is done. Similarly, discard - * any existing "device_copy" CallNodes which are no-ops. - * - * - The result virtual device for a function is stored in the function's virtual_device_ field - * and the virtual devices of the function's parameters are stored in the parameter's - * virtual_device_ field. - * - * - Additional "device_copy" CallNodes are inserted wherever there's a transition between - * storage device types. Since the DeviceAnalyzer phase succeeded this can only happen - * where the original program explicitly allowed a transition using an "on_device" CallNode. - * That is, we do not not try to 'fix' a program with inconsistent devices. - * - * - Additional "on_device" CallNodes are inserted so that a later transform can discover - * the device for an arbitrary sub-expression by looking only for the lexically enclosing - * "on_device" CallNode or "on_device" function attribute. In particular, since function - * arguments and let-bound expressions can be on a device different from the function - * or let body itself we will insert "on_device" CallNodes to spell out any differences. This - * applies even to the argument to a "device_copy" CallNode, which may look pedantic but - * keeps downstream processing simple. The "on_device" calls should be removed before code gen, - * which is easily done on-the-fly. - * - * - Update memory scopes in PrimFunc buffer maps. - * - * For example, we'll end up with programs that look like: - * \code - * def @main(%x, %y, param_virtual_devices=[...], result_virtual_device=...) { - * let %a = on_device(..., virtual_device=..., is_fixed=True) - * @f(%a, device_copy(on_device(..., virtual_device=..., is_fixed=True), - * src_virtual_device=..., dst_virtual_device=...)) - * } - * \endcode - */ -class DeviceCapturer : public ExprMutator { - public: - DeviceCapturer(IRModule mod, std::unique_ptr domains) - : mod_(std::move(mod)), domains_(std::move(domains)) {} - - IRModule Capture() { - VLOG_CONTEXT << "CaptureDevices"; - IRModule result(/*functions=*/{}, mod_->type_definitions, mod_->Imports(), mod_->source_map, - mod_->attrs); - for (const auto& kv : mod_->functions) { - if (const auto* function_node = AsOptimizableFunctionNode(kv.second)) { - VLOG(2) << "capturing devices for Relay Function '" << kv.first->name_hint << "'"; - result->Add(kv.first, Downcast(Mutate(GetRef(function_node)))); - } else if (auto prim_func = kv.second.as()) { - VLOG(2) << "capturing devices for TIR PrimFunc '" << kv.first->name_hint << "'"; - tir::PrimFunc new_prim_func = UpdatePrimFunc(kv.first, prim_func.value()); - VLOG(2) << "Rewritten prim func:" << std::endl - << PrettyPrint(prim_func) << std::endl - << "to:" << std::endl - << PrettyPrint(new_prim_func); - result->Add(kv.first, std::move(new_prim_func)); - } else { - VLOG(2) << "skipping '" << kv.first->name_hint << "'"; - result->Add(kv.first, kv.second); - } - } - return result; - } - - private: - /*! - * \brief Returns \p prim_func updated to capture any memory scope's implied by its device - * domain. - */ - tir::PrimFunc UpdatePrimFunc(const GlobalVar& global_var, const tir::PrimFunc& prim_func) { - // CAUTION: Same caution as for DeviceAnalyzer::DomainForPrimFunc. - auto func_domain = domains_->DomainFor(global_var); - ICHECK(func_domain->is_higher_order()); - - const auto* func_type_node = global_var->checked_type().as(); - ICHECK(func_type_node); - ICHECK_EQ(func_domain->function_arity(), func_type_node->arg_types.size()); - - std::vector arg_and_result_virtual_devices; - arg_and_result_virtual_devices.reserve(func_type_node->arg_types.size() + 1); - for (size_t i = 0; i < func_type_node->arg_types.size(); ++i) { - VirtualDevice param_virtual_device = - domains_->ResultVirtualDevice(func_domain->function_param(i)); - VLOG(2) << "param_virtual_device[" << i << "] = " << param_virtual_device; - arg_and_result_virtual_devices.push_back(param_virtual_device); - } - VirtualDevice ret_virtual_device = - domains_->ResultVirtualDevice(func_domain->function_result()); - VLOG(2) << "ret_virtual_device = " << ret_virtual_device; - arg_and_result_virtual_devices.push_back(ret_virtual_device); - - return tir::ApplyPrimFuncArgAndResultConstraints(prim_func, GetRef(func_type_node), - arg_and_result_virtual_devices); - } - - // Nothing interesting for VarNode, ConstantNode, GlobalVarNode, OpNode and ConstructorNode - - Expr VisitExpr_(const TupleNode* tuple_node) final { - auto tuple = GetRef(tuple_node); - Array fields; - fields.reserve(tuple_node->fields.size()); - for (const auto& field : tuple_node->fields) { - fields.push_back(VisitChild(tuple, field)); - } - return WithFields(tuple, fields); - } - - Expr VisitExpr_(const FunctionNode* function_node) final { - if (function_node->HasNonzeroAttr(attr::kPrimitive)) { - return GetRef(function_node); - } - - auto function = GetRef(function_node); - auto func_domain = domains_->DomainFor(function); // higher-order - VLOG(2) << "capturing function:" << std::endl - << PrettyPrint(function) << std::endl - << "with domain:" << std::endl - << domains_->ToString(func_domain); - - // Gather the parameter and result device types for the function attributes. - ICHECK_EQ(func_domain->function_arity(), function_node->params.size()); - VirtualDevice result_virtual_device = domains_->ResultVirtualDevice(func_domain); - ICHECK(!result_virtual_device->IsFullyUnconstrained()); - - // Map the function parameters to a new variable annotated with a virtual device so - // we can substitute them later. - Map annotated_bind_map; - Array annotated_params; - annotated_params.reserve(function_node->params.size()); - for (size_t i = 0; i < function_node->params.size(); ++i) { - VirtualDevice param_virtual_device = - domains_->ResultVirtualDevice(func_domain->function_param(i)); - VLOG(4) << "Param: " << function_node->params[i]; - Var annotated_var = WithFields(function_node->params[i], {}, {}, param_virtual_device); - VLOG(4) << "Annotated param: " << annotated_var; - VLOG(4) << "VirtualDevice: " << annotated_var->virtual_device(); - ICHECK(!param_virtual_device->IsFullyUnconstrained()); - annotated_bind_map.Set(function_node->params[i], annotated_var); - annotated_params.push_back(annotated_var); - } - // Eventually we probably want to bind before visiting, but for now this is causing an issue - // with the GetVirtualDevice utility, so leaving as is for now. - - // Rewrite the body. Note that the body may have begun with an "on_device" so - // be prepared to insert a "device_copy". - Expr body = VisitChild( - /*lexical_virtual_device=*/result_virtual_device, - /*expected_virtual_device=*/result_virtual_device, - /*child_virtual_device=*/GetVirtualDevice(function_node->body), function_node->body); - VLOG(4) << "Visited body: " << body; - Function func = WithFields(GetRef(function_node), function_node->params, body); - VLOG(4) << "New function: " << func; - func = SubstituteBoundVars(func, annotated_bind_map); - VLOG(4) << "Func with bound params: " << func; - func->virtual_device_ = result_virtual_device; - VLOG(4) << "Func with bound params & result vid set: " << func; - return std::move(func); - } - - Expr VisitExpr_(const CallNode* call_node) final { - auto call = GetRef(call_node); - - // We don't care if the call is pre- or post-lowered - // (However we'll preserve the form in the result below.) - auto vanilla_call = GetAnyCall(call_node); - - VirtualDevice call_virtual_device = GetVirtualDevice(call); - - auto on_device_props = GetOnDeviceProps(call_node); - if (on_device_props.body.defined()) { - // We're done with the original "on_device" calls and can pinch them out. - // Note that this step has already been simulated by GetDeviceType. - return VisitExpr(on_device_props.body); - } - - DeviceCopyProps device_copy_props = GetDeviceCopyProps(call_node); - if (device_copy_props.body.defined()) { - VirtualDevice src_virtual_device = - domains_->config()->CanonicalVirtualDevice(device_copy_props.src_virtual_device); - VirtualDevice dst_virtual_device = - domains_->config()->CanonicalVirtualDevice(device_copy_props.dst_virtual_device); - ICHECK_EQ(call_virtual_device, dst_virtual_device); - if (src_virtual_device == dst_virtual_device) { - // We can pinch out existing "device_copy" CallNodes if their source and destinations - // match. - return VisitExpr(device_copy_props.body); - } else { - return VisitChild(/*lexical_virtual_device=*/dst_virtual_device, - /*expected_virtual_device=*/dst_virtual_device, - /*child_virtual_device=*/src_virtual_device, device_copy_props.body); - } - } - - // Generic call. - auto func_domain = domains_->DomainForCallee(call); // higher-order - VLOG(2) << "considering call:" << std::endl - << PrettyPrint(call) << std::endl - << "in virtual device " << call_virtual_device - << " with function virtual devices:" << std::endl - << domains_->ToString(func_domain); - VirtualDevice result_virtual_device = domains_->ResultVirtualDevice(func_domain); - ICHECK(!result_virtual_device->IsFullyUnconstrained()); - - // The callee is on the current device. - Expr op = VisitChild( - /*lexical_virtual_device=*/call_virtual_device, - /*expected_virtual_device=*/call_virtual_device, - /*child_virtual_device=*/result_virtual_device, vanilla_call->op); - - // Each argument can be on the device for the corresponding function parameter. However if - // any of those differ from the overall call device then wrap them in an "on_device" to - // help downstream transforms track devices lexically. - Array args; - args.reserve(vanilla_call->args.size()); - ICHECK_EQ(func_domain->function_arity(), vanilla_call->args.size()); - for (size_t i = 0; i < vanilla_call->args.size(); ++i) { - VirtualDevice param_virtual_device = - domains_->ResultVirtualDevice(func_domain->function_param(i)); - ICHECK(!param_virtual_device->IsFullyUnconstrained()) - << "for parameter " << i << " for call:" << std::endl - << PrettyPrint(call); - args.push_back(VisitChild(/*lexical_virtual_device=*/call_virtual_device, - /*expected_virtual_device=*/param_virtual_device, - /*child_virtual_device=*/GetVirtualDevice(vanilla_call->args[i]), - vanilla_call->args[i])); - } - - if (call_node->op == CallLoweredOp()) { - Call new_call = - CallLowered(Downcast(op), args, /*call_lowered_attrs=*/{}, /*span=*/{}); - return WithFields(call, new_call->op, new_call->args); - } else { - return WithFields(call, op, args); - } - } - - Expr VisitExpr_(const LetNode* let_node) final { - Expr expr = GetRef(let_node); - // Iterate through chained lets, provided they all agree on their device type. - VirtualDevice let_virtual_device = GetVirtualDevice(expr); - std::vector> bindings; - while (const auto* inner_let = expr.as()) { - if (GetVirtualDevice(GetRef(inner_let)) != let_virtual_device) { - // We have a device transition which needs to be handled. - break; - } - // The let-bound value can be on a different device than the overall let. - // By using the fully-unconstrained virtual device for the 'lexical' scope we'll force the - // let-bound value to *always* be wrapped by an "on_device" (see introductory comment for - // motivation.) - Expr value = - VisitChild(/*lexical_virtual_device=*/VirtualDevice::FullyUnconstrained(), - /*expected_virtual_device=*/GetVirtualDevice(inner_let->var), - /*child_virtual_device=*/GetVirtualDevice(inner_let->value), inner_let->value); - bindings.emplace_back(inner_let->var, value, inner_let->span); - expr = inner_let->body; - } - Expr body = VisitChild(/*lexical_virtual_device=*/let_virtual_device, - /*expected_virtual_device=*/let_virtual_device, - /*child_virtual_device=*/GetVirtualDevice(expr), expr); - for (auto itr = bindings.rbegin(); itr != bindings.rend(); ++itr) { - body = Let(/*var=*/std::get<0>(*itr), /*value=*/std::get<1>(*itr), body, - /*span=*/std::get<2>(*itr)); - } - return body; - } - - Expr VisitExpr_(const IfNode* if_node) final { - auto ife = GetRef(if_node); - Expr cond = VisitChild(ife, if_node->cond); - Expr true_branch = VisitChild(ife, if_node->true_branch); - Expr false_branch = VisitChild(ife, if_node->false_branch); - return WithFields(ife, cond, true_branch, false_branch); - } - - Expr VisitExpr_(const TupleGetItemNode* tuple_get_item_node) final { - auto tuple_get_item = GetRef(tuple_get_item_node); - Expr tuple = VisitChild(tuple_get_item, tuple_get_item_node->tuple); - return WithFields(tuple_get_item, tuple); - } - - Expr VisitExpr_(const RefCreateNode* ref_create_node) final { - auto ref_create = GetRef(ref_create_node); - Expr value = VisitChild(ref_create, ref_create_node->value); - return WithFields(ref_create, value); - } - - Expr VisitExpr_(const RefReadNode* ref_read_node) final { - auto ref_read = GetRef(ref_read_node); - Expr ref = VisitChild(ref_read, ref_read_node->ref); - return WithFields(ref_read, ref); - } - - Expr VisitExpr_(const RefWriteNode* ref_write_node) final { - auto ref_write = GetRef(ref_write_node); - Expr ref = VisitChild(ref_write, ref_write_node->ref); - Expr value = VisitChild(ref_write, ref_write_node->value); - return WithFields(ref_write, ref, value); - } - - Expr VisitExpr_(const MatchNode* match_node) final { - auto match = GetRef(match_node); - Expr data = VisitChild(match, match_node->data); - Array clauses; - clauses.reserve(match_node->clauses.size()); - for (const auto& clause : match_node->clauses) { - Pattern lhs = VisitPattern(clause->lhs); // actually a no-op, so we're not checking vars - Expr rhs = VisitChild(match, clause->rhs); - clauses.push_back(Clause(lhs, rhs)); - } - return WithFields(match, data, clauses); - } - - VirtualDevice GetVirtualDevice(const Expr& expr) { - // Look through any "on_device" CallNodes, to mimic how we will be pinching them out. - OnDeviceProps props = GetOnDeviceProps(expr); - Expr true_expr = props.body.defined() ? props.body : expr; - ICHECK(domains_->contains(true_expr)); - // If expr is higher order we'll return only the result domain's device. - VirtualDevice virtual_device = domains_->ResultVirtualDevice(domains_->DomainFor(true_expr)); - ICHECK(!virtual_device->IsFullyUnconstrained()) - << "no VirtualDevice was determined for expression:" << std::endl - << PrettyPrint(true_expr); - return std::move(virtual_device); - } - - /*! - * \brief Reconcile the \p child_virtual_device for \p child with both the \p - * expected_virtual_device (as required by the expression context the \p child is in) and the \p - * lexical_virtual_device (as a downstream transform would infer based only on lexically enclosing - * "on_device" CallNodes and function attributes.) Generally \p lexical_virtual_device and \p - * expected_virtual_device are the same by definition, but may differ in arguments to functions - * and let-bound expressions. - * - * If \p child_virtual_device differs from \p expected_virtual_device, wrap it as: - * \code - * device_copy(on_device(child', virtual_device=child_virtual_device), - * src_dev_type=child_virtual_device, dst_dev_type=expected_virtual_device) - * \endcode - * (where child is rewritten to child'). Note the pedantic spelling out of "on_device" on the - * child. - * - * If \p expected_virtual_device differs from \p lexical_virtual_device, then (also) wrap - * the expression as: - * \code - * on_device(..., virtual_device=expected_virtual_device) - * \endcode - * - * TODO(mbs): There's no attempt at sharing here. If usage of child's node could be wrapped - * by a "device_copy", even though those copies will generally all be to the same destination - * device. - */ - Expr VisitChild(const VirtualDevice& lexical_virtual_device, - const VirtualDevice& expected_virtual_device, - const VirtualDevice& child_virtual_device, const Expr& child) { - ICHECK(!expected_virtual_device->IsFullyUnconstrained()); - if (child->IsInstance() || child->IsInstance()) { - // Primitive operators and contructors don't need to be rewritten and can have a - // different domain at each call site. - return child; - } - Expr result = VisitExpr(child); - if (child_virtual_device != expected_virtual_device) { - VLOG(2) << "creating " << DeviceCopyOp()->name << " from virtual device " - << child_virtual_device << " to virtual device " << expected_virtual_device - << " for:" << std::endl - << PrettyPrint(result); - // Also wrap the child in an "on_device" so downstream transforms can track devices - // lexically. - result = MaybeOnDeviceFixed(result, child_virtual_device); - result = DeviceCopy(result, child_virtual_device, expected_virtual_device); - } - if (expected_virtual_device != lexical_virtual_device) { - VLOG(2) << "creating " << OnDeviceOp()->name << " for virtual device " - << expected_virtual_device << " for:" << std::endl - << PrettyPrint(result); - result = MaybeOnDeviceFixed(result, expected_virtual_device); - } - return result; - } - - /*! - * Common case of visiting a direct \p child of \p parent where by default the \p child - * is expected to be on the same device as the \p parent. - */ - Expr VisitChild(const Expr& parent, const Expr& child) { - VirtualDevice expected_virtual_device = GetVirtualDevice(parent); - VirtualDevice child_virtual_device = GetVirtualDevice(child); - return VisitChild(expected_virtual_device, expected_virtual_device, child_virtual_device, - child); - } - - /*! \brief Module we are rewriting, so we can lookup global variables. */ - IRModule mod_; - /*! \brief Device domain for every expression from DeviceAnalyzer. */ - std::unique_ptr domains_; -}; - -/*! \brief Rewrite the "on_device" calls (and implicitly re-type-check). */ -tvm::transform::Pass Rewrite() { - auto pass_func = [](Function f, IRModule m, transform::PassContext ctxt) { - auto attrs = m->attrs; - auto r = Downcast(RewriteOnDevices(std::move(m)).Mutate(f)); - return attrs.defined() ? WithAttrs(r, {attrs->dict}) : r; - }; - return tvm::relay::transform::CreateFunctionPass(pass_func, 0, "PlanDevicesRewrite", {}); -} - -/*! \brief Check the conflicted nodes and add "device_copy" calls. */ -tvm::transform::Pass Check(CompilationConfig config) { - return tvm::transform::CreateModulePass( - [config = std::move(config)](IRModule mod, - tvm::transform::PassContext pass_cnxt) -> IRModule { - auto dev_ctx = ConflictedNodeFinder(mod).Finder(); - return ConflictedNodeRewriter(mod, config, std::move(dev_ctx)).Rewrite(); - }, - /*opt_level=*/0, "PlanDevicesCheckConflicts", {}); -} - -/*! \brief Run the remaining phases. */ -tvm::transform::Pass PlanDevicesCore(CompilationConfig config) { - return tvm::transform::CreateModulePass( - [config = std::move(config)](IRModule mod, - tvm::transform::PassContext pass_cnxt) -> IRModule { - // Collect the system of constraints for every sub-expression using existing "on_device" - // and "device_copy" calls. - std::unique_ptr domains = DeviceAnalyzer(mod, config).Analyze(); - VLOG(3) << "Domains after analysis:" << std::endl << domains->ToString(); - - // Choose sensible default devices for every sub-expression if otherwise unconstrained - // by existing "on_device" or "device_copy" calls. - domains = FreeOnDeviceDefaulter(mod, std::move(domains)).Default(); - domains = DeviceDefaulter(mod, std::move(domains)).Default(); - VLOG(3) << "Domains after defaulting: " << std::endl << domains->ToString(); - - // Insert "device_copy" and "on_device" CallNodes where needed to unambiguously capture - // the above map, and attach additional "param_virtual_devices" and "result_virtual_device" - // attributes to all function definitions. - return DeviceCapturer(mod, std::move(domains)).Capture(); - }, - /*opt_level=*/0, "PlanDevicesCore", {}); -} - -} // namespace - -/* =============== Driver =============== */ - -// This function is declared in the public . -tvm::transform::Pass PlanDevices(CompilationConfig config) { - std::vector passes; - passes.emplace_back(Rewrite()); - passes.emplace_back(Check(config)); - passes.emplace_back(InferType()); - passes.emplace_back(PlanDevicesCore(config)); - return tvm::transform::Sequential(passes, "PlanDevices"); -} - -TVM_REGISTER_GLOBAL("relay._transform.PlanDevices").set_body_typed(PlanDevices); - -} // namespace transform -} // namespace relay -} // namespace tvm diff --git a/src/relay/transforms/div_to_mul.cc b/src/relay/transforms/div_to_mul.cc deleted file mode 100644 index 42983c520682..000000000000 --- a/src/relay/transforms/div_to_mul.cc +++ /dev/null @@ -1,86 +0,0 @@ -/* - * 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 "pattern_utils.h" - -namespace tvm { -namespace relay { - -class DivToMulRewrite : public MixedModeMutator { - Expr Rewrite_(const CallNode* pre, const Expr& post) final { - if (const CallNode* call_node = post.as()) { - if (call_node->op == Op::Get("divide")) { - auto rhs = call_node->args[1].as(); - if (rhs != nullptr) { - auto inv = - runtime::NDArray::Empty(rhs->data.Shape(), rhs->data.DataType(), rhs->data->device); - std::string dtype = DLDataType2String(rhs->data.DataType()); - if (dtype == "float32") { - float rhs_val = static_cast(rhs->data->data)[0]; - // Check for division by zero - if (rhs_val == 0.) { - return post; - } - static_cast(inv->data)[0] = 1. / rhs_val; - } else if (dtype == "float64") { - double rhs_val = static_cast(rhs->data->data)[0]; - // Check for division by zero - if (rhs_val == 0.) { - return post; - } - static_cast(inv->data)[0] = 1. / rhs_val; - } else if (dtype == "float16") { - // Do f16 math in f32 - float rhs_val = __gnu_h2f_ieee(static_cast(rhs->data->data)[0]); - // Check for division by zero - if (rhs_val == 0.) { - return post; - } - static_cast(inv->data)[0] = __gnu_f2h_ieee(1. / rhs_val); - } else { - // Cannot do 1/int because it will truncate - return post; - } - return Multiply(call_node->args[0], Constant(inv)); - } - } - } - return post; - } -}; - -namespace transform { - -Pass DivToMul() { - runtime::TypedPackedFunc pass_func = - [=](Function f, IRModule m, PassContext pc) { - return Downcast(DivToMulRewrite().Mutate(f)); - }; - return CreateFunctionPass(pass_func, 0, "DivToMul", {"InferType", "FoldConstant"}); -} - -TVM_REGISTER_GLOBAL("relay._transform.DivToMul").set_body_typed(DivToMul); - -} // namespace transform -} // namespace relay -} // namespace tvm diff --git a/src/relay/transforms/dynamic_to_static.cc b/src/relay/transforms/dynamic_to_static.cc deleted file mode 100644 index c192097a0b29..000000000000 --- a/src/relay/transforms/dynamic_to_static.cc +++ /dev/null @@ -1,337 +0,0 @@ -/* - * 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 dynamic_to_static.cc - * \brief Rewrite Dynamic Operations to Static operations where possible - */ -#include -#include -#include -#include - -#include "pattern_utils.h" - -namespace tvm { -namespace relay { - -class DynamicToStaticMutator : public MixedModeMutator { - public: - DynamicToStaticMutator(IRModule mod, Function func) : mod_(mod), func_(func) { - op_map_ = { - {Op::Get("dyn.reshape"), - [this](const CallNode* call_node) { - auto args = PrepareArgs(call_node); - if (const ConstantNode* shape = args[1].as()) { - ICHECK_EQ(shape->data->ndim, 1); - return MakeReshape(call_node->args[0], ToVector(shape->data)); - } - return Expr(nullptr); - }}, - {Op::Get("dyn.squeeze"), - [this](const CallNode* call_node) { - auto args = PrepareArgs(call_node); - if (const ConstantNode* axis = args[1].as()) { - ICHECK_EQ(axis->data->ndim, 1); - return MakeSqueeze(call_node->args[0], ToVector(axis->data)); - } - return Expr(nullptr); - }}, - {Op::Get("dyn.tile"), - [this](const CallNode* call_node) { - auto args = PrepareArgs(call_node); - if (const ConstantNode* reps = args[1].as()) { - ICHECK_EQ(reps->data->ndim, 1); - return MakeTile(call_node->args[0], ToVector(reps->data)); - } - return Expr(nullptr); - }}, - {Op::Get("dyn.topk"), - [this](const CallNode* call_node) { - auto args = PrepareArgs(call_node); - if (const ConstantNode* k = args[1].as()) { - const TopKAttrs* param = call_node->attrs.as(); - ICHECK(param); - return MakeTopK(call_node->args[0], static_cast(ToScalar(k->data, 0)), - param->axis, param->ret_type, param->is_ascend, param->dtype); - } - return Expr(nullptr); - }}, - {Op::Get("dyn.broadcast_to"), - [this](const CallNode* call_node) { - auto args = PrepareArgs(call_node); - if (const ConstantNode* shape = args[1].as()) { - ICHECK_EQ(shape->data->ndim, 1); - return MakeBroadCastTo(call_node->args[0], ToVector(shape->data)); - } - return Expr(nullptr); - }}, - {Op::Get("dyn.zeros"), - [this](const CallNode* call_node) { - auto args = PrepareArgs(call_node); - if (const ConstantNode* shape = args[0].as()) { - const InitOpAttrs* param = call_node->attrs.as(); - ICHECK(param); - return MakeZeros(ToVector(shape->data), param->dtype); - } - return Expr(nullptr); - }}, - {Op::Get("dyn.ones"), - [this](const CallNode* call_node) { - auto args = PrepareArgs(call_node); - if (const ConstantNode* shape = args[0].as()) { - const InitOpAttrs* param = call_node->attrs.as(); - ICHECK(param); - return MakeOnes(ToVector(shape->data), param->dtype); - } - return Expr(nullptr); - }}, - {Op::Get("dyn.one_hot"), - [this](const CallNode* call_node) { - auto args = PrepareArgs(call_node); - if (const ConstantNode* depth = args[3].as()) { - const OneHotAttrs* param = call_node->attrs.as(); - ICHECK(param); - return MakeOneHot(call_node->args[0], call_node->args[1], call_node->args[2], - static_cast(ToScalar(depth->data, 0)), param->axis, - param->dtype); - } - return Expr(nullptr); - }}, - {Op::Get("dyn.image.resize2d"), - [this](const CallNode* call_node) { - auto args = PrepareArgs(call_node); - if (const ConstantNode* size = args[1].as()) { - if (const ConstantNode* roi = args[2].as()) { - const Resize2DAttrs* param = call_node->attrs.as(); - ICHECK(param); - auto size_int = ToVector(size->data); - Array size_prim; - for (size_t i = 0; i < size_int.size(); ++i) { - size_prim.push_back(size_int[i]); - } - auto roi_vec = ToFloatVector(roi->data); - Array roi_prim; - for (size_t i = 0; i < roi_vec.size(); ++i) { - roi_prim.push_back(roi_vec[i]); - } - return MakeResize2D(call_node->args[0], size_prim, roi_prim, param->layout, - param->method, param->coordinate_transformation_mode, - param->rounding_method, param->cubic_alpha, param->cubic_exclude, - param->extrapolation_value, param->out_dtype); - } - } - return Expr(nullptr); - }}, - {Op::Get("dyn.full"), - [this](const CallNode* call_node) { - auto args = PrepareArgs(call_node); - if (const ConstantNode* shape = args[1].as()) { - ICHECK_EQ(shape->data->ndim, 1); - const InitOpAttrs* param = call_node->attrs.as(); - ICHECK(param); - return MakeFull(call_node->args[0], ToVector(shape->data), param->dtype); - } - return Expr(nullptr); - }}, - {Op::Get("dyn.nn.upsampling"), - [this](const CallNode* call_node) { - auto args = PrepareArgs(call_node); - const ConstantNode* scale_h = args[1].as(); - const ConstantNode* scale_w = args[2].as(); - if (scale_h && scale_w) { - ICHECK_EQ(scale_h->data->ndim, 0); - ICHECK_EQ(scale_w->data->ndim, 0); - const UpSamplingAttrs* param = call_node->attrs.as(); - ICHECK(param); - return MakeUpSampling(call_node->args[0], ToScalar(scale_h->data), - ToScalar(scale_w->data), param->layout, param->method, - param->align_corners); - } - return Expr(nullptr); - }}, - {Op::Get("dyn.nn.upsampling3d"), - [this](const CallNode* call_node) { - auto args = PrepareArgs(call_node); - const ConstantNode* scale_d = args[1].as(); - const ConstantNode* scale_h = args[2].as(); - const ConstantNode* scale_w = args[3].as(); - if (scale_d && scale_h && scale_w) { - ICHECK_EQ(scale_d->data->ndim, 0); - ICHECK_EQ(scale_h->data->ndim, 0); - ICHECK_EQ(scale_w->data->ndim, 0); - const UpSampling3DAttrs* param = call_node->attrs.as(); - ICHECK(param); - return MakeUpSampling3D(call_node->args[0], ToScalar(scale_d->data), - ToScalar(scale_h->data), ToScalar(scale_w->data), - param->layout, param->method, - param->coordinate_transformation_mode); - } - return Expr(nullptr); - }}, - {Op::Get("dyn.nn.pad"), - [this](const CallNode* call_node) { - auto args = PrepareArgs(call_node); - const ConstantNode* pad_width = args[1].as(); - const ConstantNode* pad_fill = args[2].as(); - if (pad_width && pad_fill) { - ICHECK_EQ(pad_fill->data->ndim, 0); // pad_val is 1d - ICHECK_EQ(pad_width->data->ndim, 2); // pad_width is 2d - - const PadAttrs* param = call_node->attrs.as(); - ICHECK(param); - - Expr pad_value = args[2]; - return MakePad(call_node->args[0], ToMatrix(pad_width->data), pad_value, - param->pad_mode); - } - return Expr(nullptr); - }}, - {Op::Get("dyn.strided_slice"), - [this](const CallNode* call_node) { - auto args = PrepareArgs(call_node); - const ConstantNode* begin = args[1].as(); - const ConstantNode* end = args[2].as(); - const ConstantNode* stride = args[3].as(); - if (begin && end && stride) { - ICHECK_EQ(begin->data->ndim, 1); - ICHECK_EQ(end->data->ndim, 1); - ICHECK_EQ(stride->data->ndim, 1); - const StridedSliceAttrs* param = call_node->attrs.as(); - ICHECK(param); - return MakeStridedSlice(call_node->args[0], ToVector(begin->data), ToVector(end->data), - ToVector(stride->data), param->slice_mode); - } - return Expr(nullptr); - }}, - {Op::Get("dyn.sparse_to_dense"), - [this](const CallNode* call_node) { - auto args = PrepareArgs(call_node); - const ConstantNode* output_shape = args[3].as(); - if (output_shape) { - ICHECK_EQ(output_shape->data->ndim, 1); - return MakeSparseToDense(call_node->args[0], ToVector(output_shape->data), - call_node->args[1], call_node->args[2]); - } - return Expr(nullptr); - }}, - }; - Map vars; - for (auto kv : mod_->functions) { - vars.Set(kv.second, kv.first); - } - gv_ = vars[func_]; - } - - Expr GetCurExpr(const Expr& original_expr) { - if (original_expr.as()) { - return mod_->Lookup(gv_); - } else { - return mod_->Lookup(gv_).as()->body; - } - } - - Expr PrepareInput(const Expr& expr) { - BaseFunc func; - if (auto func_node = expr.as()) { - func = func_node.value(); - } else { - func = relay::Function(relay::FreeVars(expr), expr, Type(), relay::FreeTypeVars(expr, mod_)); - } - mod_->Update(gv_, func); - - mod_ = transform::FoldConstant()(mod_); - transform::InferTypeLocal(GetCurExpr(expr)); - mod_ = transform::FoldConstant()(mod_); - transform::InferTypeLocal(GetCurExpr(expr)); - - Expr out; - if (expr.as()) { - out = mod_->Lookup(gv_); - } else { - out = mod_->Lookup(gv_).as()->body; - } - return out; - } - - std::vector PrepareArgs(const CallNode* call_node) { - std::vector args; - for (auto arg : call_node->args) { - if (arg.as()) { - args.emplace_back(arg); - } else { - args.emplace_back(PrepareInput(arg)); - } - } - return args; - } - - private: - Expr Rewrite_(const CallNode* pre, const Expr& post) override { - if (const CallNode* call_node = post.as()) { - if (op_map_.count(call_node->op)) { - auto out = op_map_[call_node->op](call_node); - if (out.defined()) { - return out; - } - } - } - return post; - } - - Expr DispatchVisitExpr(const Expr& expr) override { - auto post = MixedModeMutator::DispatchVisitExpr(expr); - if (auto op = post.as()) { - return Function(op->params, op->body, NullValue(), op->type_params, op->attrs); - } - return post; - } - - std::unordered_map, ObjectPtrHash, ObjectPtrEqual> - op_map_; - IRModule mod_; - Function func_; - GlobalVar gv_; -}; - -Expr DynamicToStatic(Function f, IRModule m) { - DynamicToStaticMutator mutator(m, f); - Expr expr = mutator.Mutate(f); - Expr out = mutator.PrepareInput(expr); - return out; -} - -namespace transform { - -Pass DynamicToStatic() { - runtime::TypedPackedFunc pass_func = - [=](Function f, IRModule m, PassContext pc) { - return Downcast(DynamicToStatic(f, m)); - }; - return CreateFunctionPass(pass_func, 2, "DynamicToStatic", {}); -} - -TVM_REGISTER_GLOBAL("relay._transform.DynamicToStatic").set_body_typed([]() { - return DynamicToStatic(); -}); - -} // namespace transform -} // namespace relay -} // namespace tvm diff --git a/src/relay/transforms/eliminate_common_subexpr.cc b/src/relay/transforms/eliminate_common_subexpr.cc deleted file mode 100644 index 9de1b86b17e1..000000000000 --- a/src/relay/transforms/eliminate_common_subexpr.cc +++ /dev/null @@ -1,154 +0,0 @@ -/* - * 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 eliminate_common_subexpr.cc - * \brief Combine common subexpressions. - * - * This is an optimization pass that eliminates common subexpressions. During the pass, it tries - * to replace an expression with a previously appeared expression with the same input and - * attributes. The fskip callback argument allows us to skip specific expressions. - */ -#include -#include -#include - -#include - -#include "pattern_utils.h" - -namespace tvm { -namespace relay { - -class CommonSubexprEliminator : public MixedModeMutator { - public: - explicit CommonSubexprEliminator(runtime::TypedPackedFunc fskip) : fskip_(fskip) {} - - Expr Rewrite_(const CallNode* call, const Expr& post) final { - static auto op_stateful = Op::GetAttrMap("TOpIsStateful"); - Expr new_expr = post; - const CallNode* new_call = new_expr.as(); - ICHECK(new_call); - const OpNode* op = new_call->op.as(); - StructuralEqual attrs_equal; - - if (new_call->args.size() == 0 || op == nullptr || op_stateful.get(GetRef(op), false)) { - return new_expr; - } - if (fskip_ != nullptr && fskip_(new_expr)) { - return new_expr; - } - - auto it = expr_map_.find(new_call->op); - if (it != expr_map_.end()) { - for (const Expr& candidate_expr : it->second) { - if (const CallNode* candidate = candidate_expr.as()) { - bool is_equivalent = true; - if (!attrs_equal(new_call->attrs, candidate->attrs)) { - continue; - } - for (size_t i = 0; i < new_call->args.size(); i++) { - if (!IsEquivalent(new_call->args[i], candidate->args[i])) { - is_equivalent = false; - break; - } - } - if (!is_equivalent) continue; - return GetRef(candidate); - } - } - } - expr_map_[new_call->op].push_back(new_expr); - return new_expr; - } - - Expr Rewrite_(const TupleGetItemNode* op, const Expr& post) final { - Expr new_expr = post; - const TupleGetItemNode* new_tuple_item = new_expr.as(); - ICHECK(new_tuple_item); - - if (fskip_ != nullptr && fskip_(new_expr)) { - return new_expr; - } - - auto it = expr_map_.find(new_tuple_item->tuple); - if (it != expr_map_.end()) { - for (const Expr& candidate_expr : it->second) { - if (const TupleGetItemNode* candidate = candidate_expr.as()) { - if (new_tuple_item->index == candidate->index) { - return GetRef(candidate); - } - } - } - } - expr_map_[new_tuple_item->tuple].push_back(new_expr); - return new_expr; - } - - std::unordered_map, ObjectPtrHash, ObjectPtrEqual> expr_map_; - runtime::TypedPackedFunc fskip_; - - private: - bool IsEquivalent(const Expr& arg, const Expr& candidate_arg) { - if (arg->IsInstance() && candidate_arg->IsInstance()) { - const TupleNode* arg_node = arg.as(); - const TupleNode* candidate_arg_node = candidate_arg.as(); - - if (arg_node->fields.size() != candidate_arg_node->fields.size()) { - return false; - } - - for (size_t i = 0; i < arg_node->fields.size(); i++) { - if (!arg_node->fields[i].same_as(candidate_arg_node->fields[i]) && - !IsEqualScalar(arg_node->fields[i], candidate_arg_node->fields[i])) { - return false; - } - } - } else { - if (!arg.same_as(candidate_arg) && !IsEqualScalar(arg, candidate_arg)) { - return false; - } - } - - return true; - } -}; - -Expr EliminateCommonSubexpr(const Expr& expr, PackedFunc callback) { - return CommonSubexprEliminator(callback)(expr); -} - -namespace transform { - -Pass EliminateCommonSubexpr(PackedFunc fskip) { - runtime::TypedPackedFunc pass_func = - [=](Function f, IRModule m, PassContext pc) { - return Downcast(EliminateCommonSubexpr(f, fskip)); - }; - return CreateFunctionPass(pass_func, 3, "EliminateCommonSubexpr", {"InferType"}); -} - -TVM_REGISTER_GLOBAL("relay._transform.EliminateCommonSubexpr") - .set_body_typed(EliminateCommonSubexpr); - -} // namespace transform - -} // namespace relay -} // namespace tvm diff --git a/src/relay/transforms/eta_expand.cc b/src/relay/transforms/eta_expand.cc deleted file mode 100644 index 9759f732df4d..000000000000 --- a/src/relay/transforms/eta_expand.cc +++ /dev/null @@ -1,166 +0,0 @@ -/* - * 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 eta_expand.cc - * - * \brief Add an abstraction over constructors and/or global variables bound to a function. - * - */ -#include -#include -#include -#include - -namespace tvm { -namespace relay { -namespace eta_expand { - -/*! - * \brief mutator to replace type variables with fresh ones, while maintaining alpha equality - */ -class TypeVarReplacer : public TypeMutator { - public: - TypeVarReplacer() : replace_map_({}) {} - - Type VisitType_(const TypeVarNode* type_var_node) final { - const auto type_var = GetRef(type_var_node); - if (replace_map_.find(type_var) == replace_map_.end()) { - replace_map_[type_var] = TypeVar("A", Kind::kType); - } - return replace_map_[type_var]; - } - - private: - /*! \brief variable replacement map to remap old type vars to fresh ones */ - std::unordered_map replace_map_; -}; - -/*! - * \brief mutator to perform eta expansion on all functions in a module - */ -class EtaExpander : public ExprMutator { - public: - explicit EtaExpander(const IRModule& mod, bool expand_constructor, bool expand_global_var) - : mod_(mod), - type_var_replacer_(TypeVarReplacer()), - expand_constructor_(expand_constructor), - expand_global_var_(expand_global_var) { - ICHECK(expand_constructor || expand_global_var) << "must expand at least one language feature"; - } - - IRModule Expand() { - for (GlobalVar global_var : mod_->GetGlobalVars()) { - const BaseFunc base_func = mod_->Lookup(global_var); - if (auto func = base_func.as()) { - const Function new_func = Downcast(VisitExpr(func.value())); - mod_->Update(global_var, new_func); - } - } - return mod_; - } - - Expr VisitExpr_(const CallNode* call) final { - // we don't need to expand constructors when they are being called, so we - // prevent them being visited here - Expr new_op = call->op; - if (!call->op.as()) { - new_op = VisitExpr(new_op); - } - tvm::Array new_args; - for (const auto& arg : call->args) { - new_args.push_back(VisitExpr(arg)); - } - return Call(new_op, new_args, call->attrs, call->type_args); - } - - Expr VisitExpr_(const ConstructorNode* cons_node) final { - Constructor cons = GetRef(cons_node); - if (!expand_constructor_) { - return std::move(cons); - } - // NOTE: we only reach this case if the constructor is not being applied to any arguments - tvm::Array params; - for (const auto& type : cons->inputs) { - Type param_type = type_var_replacer_.VisitType(type); - params.push_back(Var("eta_expand_param", param_type)); - } - tvm::Array type_params; - TypeData adt_def = mod_->LookupTypeDef(cons->belong_to); - for (const auto& type_var : adt_def->type_vars) { - type_params.push_back(type_var_replacer_.VisitType(type_var)); - } - Expr body = Call(cons, params, Attrs()); - Type ret_type = TypeCall(cons->belong_to, type_params); - - return Function(Downcast>(params), body, ret_type, - Downcast>(type_params)); - } - - Expr VisitExpr_(const GlobalVarNode* gvar_node) final { - GlobalVar gvar = GetRef(gvar_node); - if (!expand_global_var_) { - return std::move(gvar); - } - const auto base_func = mod_->Lookup(gvar); - if (auto opt = base_func.as()) { - // handle relay function, skip external functions. - auto func = opt.value(); - tvm::Array params; - tvm::Array args; - for (size_t i = 0; i < func->params.size(); ++i) { - auto var = Var("eta_expand_param", func->params[i]->type_annotation); - params.push_back(var); - args.push_back(var); - } - return WithFields(func, args, Call(gvar, params)); - } else { - return std::move(gvar); - } - } - - private: - /*! \brief reference to module being expanded */ - const IRModule mod_; - /*! \brief type variable replacer */ - TypeVarReplacer type_var_replacer_; - /*! \brief whether to expand constructor nodes */ - bool expand_constructor_; - /*! \brief whether to expand global variable nodes */ - bool expand_global_var_; -}; - -} // namespace eta_expand - -namespace transform { - -Pass EtaExpand(bool expand_constructor, bool expand_global_var) { - runtime::TypedPackedFunc pass_func = [=](IRModule mod, - PassContext pc) { - return eta_expand::EtaExpander(mod, expand_constructor, expand_global_var).Expand(); - }; - return CreateModulePass(pass_func, 1, "EtaExpand", {}); -} - -TVM_REGISTER_GLOBAL("relay._transform.EtaExpand").set_body_typed(EtaExpand); - -} // namespace transform - -} // namespace relay -} // namespace tvm diff --git a/src/relay/transforms/expr_subst.cc b/src/relay/transforms/expr_subst.cc deleted file mode 100644 index 96f139b7cdeb..000000000000 --- a/src/relay/transforms/expr_subst.cc +++ /dev/null @@ -1,55 +0,0 @@ -/* - * 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 expr_subst.h - * \brief Utility functions for substituting expressions. - */ - -#include "./expr_subst.h" - -#include - -namespace tvm { -namespace relay { - -class ExprSubstituter : public ExprMutator { - public: - explicit ExprSubstituter(std::unordered_map subst_map) - : subst_map_(subst_map) {} - - Expr VisitExpr(const Expr& expr) final { - auto it = subst_map_.find(expr); - if (it != subst_map_.end()) { - return ExprMutator::VisitExpr((*it).second); - } - return ExprMutator::VisitExpr(expr); - } - - private: - tvm::Map subst_map_; -}; - -Expr ExprSubst(const Expr& expr, - std::unordered_map subst_map) { - return ExprSubstituter(std::move(subst_map)).Mutate(expr); -} - -} // namespace relay -} // namespace tvm diff --git a/src/relay/transforms/expr_subst.h b/src/relay/transforms/expr_subst.h deleted file mode 100644 index 104ce0be0106..000000000000 --- a/src/relay/transforms/expr_subst.h +++ /dev/null @@ -1,38 +0,0 @@ -/* - * 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 expr_subst.h - * \brief Utility functions for substituting expressions. - */ -#ifndef TVM_RELAY_TRANSFORMS_EXPR_SUBST_H_ -#define TVM_RELAY_TRANSFORMS_EXPR_SUBST_H_ -#include - -#include - -namespace tvm { -namespace relay { - -Expr ExprSubst(const Expr& expr, - std::unordered_map subst_map); - -} // namespace relay -} // namespace tvm -#endif // TVM_RELAY_TRANSFORMS_EXPR_SUBST_H_ diff --git a/src/relay/transforms/fake_quantization_to_integer.cc b/src/relay/transforms/fake_quantization_to_integer.cc deleted file mode 100644 index b767924ecd1f..000000000000 --- a/src/relay/transforms/fake_quantization_to_integer.cc +++ /dev/null @@ -1,587 +0,0 @@ -/* - * 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 src/relay/transforms/quantize_fake_quantization.cc - * \brief A pass for taking fake quantized graphs and converting them - * to actual integer operations. - */ - -#include "fake_quantization_to_integer.h" - -#include -#include -#include -#include -#include -#include -#include - -#include - -#include "../qnn/utils.h" - -namespace tvm { -namespace relay { - -/* Description of FakeQuantizationToInteger - * - * The purpose of this pass is to find regions of the graph that follow - * the general pattern: - * - * x w - * | | - * dq dq - * \ / - * op1 - * | - * op2 - * | - * q - * - * and convert them into subgraphs with actual integer operations on x and w - * - * The pass does this via a multi-pass approach: - * - * The main pass is a MixedModeMutator that traverses the full graph searching for - * quantize operations - * - * The second pass is an ExprVisitor that recursively searches for subgraphs leading to the - * quantize for subtraphs bounded by dequantize operations. This pass extracts the affine - * types of the inputs for later processing, where affine denotes the transformation - * x_real = (x_affine - zero_point) * scale - * - * The third pass is an ExprMutator that recursively rewrites the subgraphs using packed funcs - * registered with the FTVMFakeQuantizationToInteger attribute. These packed funcs rewrite - * the ops based on the affine types of their inputs and then return the affine types of the - * new rewriten ops to pass that information down the stack during rewrite. - * - * After the second and third passes run, the first pass replaces the quantize with the - * rewritten subgraph and the processing continues - * - * - * After that an additional QAT pass can be enabled by use_qat flag. The goal of the pass is to find - * operations in those regions(which were not successfully converted by the main pass) that can - * still be converted into quantized form. The idea is to find and transform operations with - * dequantized inputs one by one individually. Only operations for which all parameters can be - * explicitly calculated are allowed. For example, if on the above general pattern op2 is not - * registered with the FTVMFakeQuantizationToInteger attribute, op1 operation can still be - * converted. Converted pattern below: - * - * x w - * | | - * \ / - * op1 - * | - * dq - * | - * op2 - * | - * q - * - * This pass works in the same multi-pass approach. - */ - -using ExprSet = std::unordered_set; -using ExprMap = std::unordered_map; -using AffineTypeMap = Map; - -using FTVMFakeQuantizationToInteger = - runtime::TypedPackedFunc(const Expr& expr, const AffineTypeMap& map)>; - -const ExprSet SubgraphExtractor::GetSubgraph(const Expr& expr) { - VisitExpr(expr); - ExprSet subgraph; - if (is_fake_quantized_) { - for (auto kv : this->visit_counter_) { - if (auto call_node = GetRef(kv.first).as()) { - if (call_node->op != quantize_op_) { - subgraph.insert(Downcast(GetRef(kv.first))); - } - } - } - } - return subgraph; -} -const AffineTypeMap SubgraphExtractor::GetAffineTypes() { return affine_types_; } -void SubgraphExtractor::VisitExpr(const Expr& expr) { - // When looking for fake quantized subgraphs, we only support data-flow regions of the graph, - // i.e. call nodes/tuples/constants/etc. If we see anything else (like control flow) we - // abort the rewrite. - if (expr.as() == nullptr && expr.as() == nullptr && - expr.as() == nullptr && expr.as() == nullptr && - expr.as() == nullptr) { - DLOG(INFO) << "FakeQuantizationToInteger found a non-dataflow op inside" - << " a fake quantize region, aborting this rewrite"; - is_fake_quantized_ = false; - } else { - ExprVisitor::VisitExpr(expr); - } -} - -void SubgraphExtractor::VisitExpr_(const CallNode* call_node) { - const Op test_op = Downcast(call_node->op); - if (call_node->op == quantize_op_) { - const auto* attrs = call_node->attrs.as(); - ICHECK(attrs != nullptr); - // Only look at arg0 for quantize - VisitExpr(call_node->args[0]); - // Collect type of quantize ops - affine_types_.Set( - GetRef(call_node), - TensorAffineType(call_node->args[1], call_node->args[2], attrs->out_dtype, attrs->axis)); - } else if (call_node->op == dequantize_op_) { - const auto* attrs = call_node->attrs.as(); - ICHECK(attrs != nullptr); - // Collect type of dequantize ops - affine_types_.Set( - GetRef(call_node), - TensorAffineType(call_node->args[1], call_node->args[2], - call_node->args[0]->checked_type().as()->dtype, - attrs->axis)); - } else { - // run normally on everything else. - ExprVisitor::VisitExpr_(call_node); - } -} - -class SubgraphMutator : public ExprMutator { - public: - SubgraphMutator(ExprSet subgraph, AffineTypeMap affine_types, bool hard_fail, - const std::unordered_set& optional_qnn_ops) - : subgraph_(subgraph), - affine_types_(affine_types), - hard_fail_(hard_fail), - optional_qnn_ops_(optional_qnn_ops) {} - - Expr MutateSubgraph(const Expr& expr) { - if (subgraph_.size() == 0) { - return expr; - } - const CallNode* quantize_node = expr.as(); - ICHECK(quantize_node); - ICHECK(quantize_node->op == quantize_op_); - out_type_ = affine_types_[expr]; - static auto fqfq = - Op::GetAttrMap("FTVMFakeQuantizationToInteger"); - static auto opt_fqfq = - Op::HasAttrMap("FTVMOptionalFakeQuantizationToInteger") - ? Op::GetAttrMap("FTVMOptionalFakeQuantizationToInteger") - : fqfq; - for (auto node : subgraph_) { - const Op op = Downcast(node.as()->op); - if (!fqfq.count(Downcast(op)) && - !(optional_qnn_ops_.count(op->name) && opt_fqfq.count(Downcast(op)))) { - // Only modify the subgraph if we have translation - // rules for every op - if (hard_fail_) { - LOG(FATAL) << "Found no rewrite rule for " << AsText(op, false) << std::endl; - } else { - DLOG(INFO) << "Found no rewrite rule for " << AsText(op, false) << std::endl; - return expr; - } - } - } - try { - return Mutate(expr); - } catch (std::exception& e) { - if (hard_fail_) { - LOG(FATAL) << e.what(); - } else { - DLOG(INFO) << "Ran into an error rewriting a subgraph, skipping" << expr << std::endl; - return expr; - } - } - } - - protected: - Expr VisitExpr_(const CallNode* call_node) { - Expr out; - - static auto fqfq = - Op::GetAttrMap("FTVMFakeQuantizationToInteger"); - static auto opt_fqfq = - Op::HasAttrMap("FTVMOptionalFakeQuantizationToInteger") - ? Op::GetAttrMap("FTVMOptionalFakeQuantizationToInteger") - : fqfq; - Op op = Downcast(call_node->op); - if (fqfq.count(op) || (optional_qnn_ops_.count(op->name) && opt_fqfq.count(op))) { - Expr expr; - if (op == dequantize_op_) { - expr = GetRef(call_node); - } else { - expr = ExprMutator::VisitExpr_(call_node); - // Set the current op to the output type, useful if we can't deduce output parameters - // from input parameters - affine_types_.Set(expr, out_type_); - } - // Call the rewrite - Array vals = (fqfq.count(op) ? fqfq : opt_fqfq)[op](expr, affine_types_); - // Save the outputs of the rewrite - ICHECK(vals.size() == 2) - << "got the wrong number of returned arguments from FTVMFakeQuantizationToInteger for " - << AsText(op, false); - out = Downcast(vals[0]); - affine_types_.Set(out, Downcast(vals[1])); - } else { - ICHECK(false) << "When rewriting a fake quantized graph, found an invalid node " - << AsText(GetRef(call_node), false); - } - return out; - } - - Expr VisitExpr_(const TupleNode* node) { - Expr expr = ExprMutator::VisitExpr_(node); - auto new_node = expr.as(); - Array types; - for (Expr field : new_node->fields) { - ICHECK(affine_types_[field].as()); - types.push_back(Downcast(affine_types_[field])); - } - affine_types_.Set(expr, TupleAffineType(types)); - return expr; - } - - Expr VisitExpr_(const TupleGetItemNode* node) { - Expr expr = ExprMutator::VisitExpr_(node); - auto tuple_type = affine_types_[expr.as()->tuple].as(); - affine_types_.Set(expr, tuple_type->types[node->index]); - return expr; - } - - ExprSet subgraph_; - AffineTypeMap affine_types_; - AffineType out_type_; - const bool hard_fail_; - const std::unordered_set& optional_qnn_ops_; - const Op quantize_op_ = Op::Get("qnn.quantize"); - const Op dequantize_op_ = Op::Get("qnn.dequantize"); -}; - -class FakeQuantizationRewriter : public MixedModeMutator { - public: - explicit FakeQuantizationRewriter(bool hard_fail, - const std::unordered_set& optional_qnn_ops) - : hard_fail_(hard_fail), optional_qnn_ops_(optional_qnn_ops) {} - - protected: - Expr Rewrite_(const CallNode* pre, const Expr& post) override { - if (const CallNode* call_node = post.as()) { - if (call_node->op == quantize_op_) { - SubgraphExtractor extractor; - ExprSet subgraph = extractor.GetSubgraph(GetRef(pre)); - AffineTypeMap affine_types = extractor.GetAffineTypes(); - - ExprSet post_subgraph; - AffineTypeMap post_affine_types; - - for (auto kv : affine_types) { - if (pre == kv.first.as()) { - // we havent memoized the current op yet - post_affine_types.Set(post, kv.second); - } else { - post_affine_types.Set(memo_.at(kv.first), kv.second); - } - } - for (auto expr : subgraph) { - post_subgraph.insert(memo_[expr]); - } - Expr out = SubgraphMutator(post_subgraph, post_affine_types, hard_fail_, optional_qnn_ops_) - .MutateSubgraph(post); - return out; - } - } - return post; - } - const Op quantize_op_ = Op::Get("qnn.quantize"); - const bool hard_fail_; - const std::unordered_set& optional_qnn_ops_; -}; - -/* Checks if the operation to convert QAT pass is enabled. - * The following conditions must be satisfied: - * 1. operations registered for FTVMFakeQuantizationToInteger; - * 2. Unary operators or operators with the TensorAffineType calculated during - * FTVMFakeQuantizationToInteger conversion; - * 3. Not one of the "key" operations: requantize,quantize and dequantize(they are at the boundaries - * of regions defined to be quantized). - */ -bool is_op_enabled_for_optional_fq2i(const CallNode* call_node) { - const Op op = Downcast(call_node->op); - static auto fqfq = Op::GetAttrMap("FTVMFakeQuantizationToInteger"); - static std::unordered_set ops = { - Op::Get("broadcast_to"), - Op::Get("clip"), - Op::Get("expand_dims"), - Op::Get("max"), - Op::Get("maximum"), - Op::Get("min"), - Op::Get("minimum"), - Op::Get("nn.avg_pool2d"), - Op::Get("nn.batch_flatten"), - Op::Get("nn.batch_matmul"), - Op::Get("nn.bias_add"), - Op::Get("nn.conv2d"), - Op::Get("nn.conv2d_transpose"), - Op::Get("nn.dense"), - Op::Get("nn.depth_to_space"), - Op::Get("nn.global_avg_pool2d"), - Op::Get("nn.max_pool2d"), - Op::Get("nn.pad"), - Op::Get("nn.relu"), - Op::Get("reshape"), - Op::Get("split"), - Op::Get("squeeze"), - Op::Get("strided_slice"), - Op::Get("transpose")}; - - return ops.find(call_node->op) != ops.end() && fqfq.count(Downcast(op)); -} - -class QATSubgraphExtractor : public ExprVisitor { - public: - const ExprSet GetSubgraph(const Expr& expr) { - expr_call_node_ = expr.as(); - ICHECK(expr_call_node_ != nullptr); - ICHECK(is_op_enabled_for_optional_fq2i(expr_call_node_)); - - VisitExpr(expr); - - ExprSet subgraph; - if (is_fake_quantized_) { - for (auto kv : this->visit_counter_) { - if (auto call_node = GetRef(kv.first).as()) { - if (call_node != expr_call_node_) { - subgraph.insert(Downcast(GetRef(kv.first))); - } - } - } - } - return subgraph; - } - const AffineTypeMap GetAffineTypes() { return affine_types_; } - void VisitExpr(const Expr& expr) override { - // When looking for fake quantized subgraphs, we only support data-flow regions of the graph, - // i.e. call nodes/tuples/constants/etc. If we see anything else (like control flow) we - // abort the rewrite. - if (expr.as() == nullptr && expr.as() == nullptr && - expr.as() == nullptr && expr.as() == nullptr && - expr.as() == nullptr) { - DLOG(INFO) << "FakeQuantizationToInteger found a non - dataflow op inside a fake quantize " - "region, aborting this rewrite"; - is_fake_quantized_ = false; - } else { - ExprVisitor::VisitExpr(expr); - } - } - - protected: - void VisitExpr_(const CallNode* call_node) override { - if (call_node->op == dequantize_op_) { - const auto* attrs = call_node->attrs.as(); - ICHECK(attrs != nullptr); - - affine_types_.Set( - GetRef(call_node), - TensorAffineType( - call_node->args[1], call_node->args[2], - tvm::relay::transform::InferTypeLocal(call_node->args[0]).as()->dtype, - attrs->axis)); - } else if (call_node == expr_call_node_) { - for (auto arg : call_node->args) { - VisitExpr(arg); - } - } else { - // run normally on everything else. - ExprVisitor::VisitExpr_(call_node); - } - } - - const Op dequantize_op_ = Op::Get("qnn.dequantize"); - bool is_fake_quantized_ = true; - AffineTypeMap affine_types_; - const CallNode* expr_call_node_ = nullptr; -}; - -class QATSubgraphMutator : public ExprMutator { - public: - QATSubgraphMutator(ExprSet subgraph, AffineTypeMap affine_types, bool hard_fail, - const std::unordered_set& optional_qnn_ops) - : subgraph_(subgraph), - affine_types_(affine_types), - hard_fail_(hard_fail), - optional_qnn_ops_(optional_qnn_ops) {} - - Expr MutateSubgraph(const Expr& expr) { - if (subgraph_.size() == 0) { - return expr; - } - - quantize_node_ = expr.as(); - ICHECK(quantize_node_); - ICHECK(is_op_enabled_for_optional_fq2i(quantize_node_)); - - for (auto node : subgraph_) { - const Op op = Downcast(node.as()->op); - - if (node.as()->op != dequantize_op_) { - if (hard_fail_) { - LOG(FATAL) << "Not dequantization was found in the input arguments for" - << AsText(op, false) << std::endl; - } else { - DLOG(INFO) << "Not dequantization was found in the input arguments for " - << AsText(op, false) << std::endl; - return expr; - } - } - } - try { - return Mutate(expr); - } catch (std::exception& e) { - if (hard_fail_) { - throw e; - } else { - DLOG(INFO) << "Ran into an error rewriting a subgraph, skipping" << expr << std::endl; - return expr; - } - } - } - - protected: - Expr VisitExpr_(const CallNode* call_node) { - Expr out; - static auto fqfq = - Op::GetAttrMap("FTVMFakeQuantizationToInteger"); - static auto opt_fqfq = - Op::HasAttrMap("FTVMOptionalFakeQuantizationToInteger") - ? Op::GetAttrMap("FTVMOptionalFakeQuantizationToInteger") - : fqfq; - - Op op = Downcast(call_node->op); - if (fqfq.count(op) || (optional_qnn_ops_.count(op->name) && opt_fqfq.count(op))) { - Expr expr; - if (op == dequantize_op_) { - expr = GetRef(call_node); - } else { - expr = ExprMutator::VisitExpr_(call_node); - } - // Call the rewrite - Array vals = (fqfq.count(op) ? fqfq : opt_fqfq)[op](expr, affine_types_); - // Save the outputs of the rewrite - ICHECK(vals.size() == 2) - << "got the wrong number of returned arguments from FTVMFakeQuantizationToInteger for " - << AsText(op, false); - out = Downcast(vals[0]); - - affine_types_.Set(out, Downcast(vals[1])); - - if (call_node == quantize_node_) { - out = qnn::MakeDequantize(out, vals[1].as()->scale, - vals[1].as()->zero_point, - vals[1].as()->axis); - } - } else { - ICHECK(false) << "When rewriting a fake quantized graph, found an invalid node " - << AsText(GetRef(call_node), false); - } - return out; - } - - Expr VisitExpr_(const TupleNode* node) { - Expr expr = ExprMutator::VisitExpr_(node); - auto new_node = expr.as(); - Array types; - for (Expr field : new_node->fields) { - ICHECK(affine_types_[field].as()); - types.push_back(Downcast(affine_types_[field])); - } - affine_types_.Set(expr, TupleAffineType(types)); - return expr; - } - - Expr VisitExpr_(const TupleGetItemNode* node) { - Expr expr = ExprMutator::VisitExpr_(node); - auto tuple_type = affine_types_[expr.as()->tuple].as(); - affine_types_.Set(expr, tuple_type->types[node->index]); - return expr; - } - - ExprSet subgraph_; - AffineTypeMap affine_types_; - const bool hard_fail_; - const std::unordered_set& optional_qnn_ops_; - const Op dequantize_op_ = Op::Get("qnn.dequantize"); - const CallNode* quantize_node_ = nullptr; -}; - -class QATRewriter : public MixedModeMutator { - public: - explicit QATRewriter(bool hard_fail, const std::unordered_set& optional_qnn_ops) - : hard_fail_(hard_fail), optional_qnn_ops_(optional_qnn_ops) {} - - protected: - Expr Rewrite_(const CallNode* pre, const Expr& post) override { - if (const CallNode* call_node = post.as()) { - const Op op = Downcast(call_node->op); - if (is_op_enabled_for_optional_fq2i(call_node)) { - QATSubgraphExtractor extractor; - ExprSet subgraph = extractor.GetSubgraph(post); - AffineTypeMap affine_types = extractor.GetAffineTypes(); - Expr out = QATSubgraphMutator(subgraph, affine_types, hard_fail_, optional_qnn_ops_) - .MutateSubgraph(post); - return out; - } - } - return post; - } - const bool hard_fail_; - const std::unordered_set& optional_qnn_ops_; -}; - -Expr FakeQuantizationToInteger(const Expr& expr, const IRModule& mod, bool hard_fail, bool use_qat, - const Array& optional_qnn_ops) { - const std::unordered_set optional_qnn_ops_(optional_qnn_ops.begin(), - optional_qnn_ops.end()); - auto fq_expr = FakeQuantizationRewriter(hard_fail, optional_qnn_ops_).Mutate(expr); - if (use_qat) { - fq_expr = tvm::relay::InferType(fq_expr); - fq_expr = QATRewriter(hard_fail, optional_qnn_ops_).Mutate(fq_expr); - } - return fq_expr; -} - -namespace transform { - -Pass FakeQuantizationToInteger(bool hard_fail, bool use_qat, - const Array& optional_qnn_ops) { - runtime::TypedPackedFunc pass_func = - [=](Function f, IRModule m, PassContext pc) { - return Downcast( - FakeQuantizationToInteger(f, m, hard_fail, use_qat, optional_qnn_ops)); - }; - return CreateFunctionPass(pass_func, 0, "FakeQuantizationToInteger", {"InferType", "DivToMul"}); -} - -TVM_REGISTER_GLOBAL("relay._transform.FakeQuantizationToInteger") - .set_body_typed(FakeQuantizationToInteger); - -} // namespace transform - -} // namespace relay -} // namespace tvm diff --git a/src/relay/transforms/fake_quantization_to_integer.h b/src/relay/transforms/fake_quantization_to_integer.h deleted file mode 100644 index 1956f94a46b3..000000000000 --- a/src/relay/transforms/fake_quantization_to_integer.h +++ /dev/null @@ -1,54 +0,0 @@ -/* - * 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 src/relay/transforms/fake_quantization_to_integer.h - * \brief Extract subgraph of a fake quantized region. - */ -#ifndef TVM_RELAY_TRANSFORMS_FAKE_QUANTIZATION_TO_INTEGER_H_ -#define TVM_RELAY_TRANSFORMS_FAKE_QUANTIZATION_TO_INTEGER_H_ - -#include -#include - -#include - -namespace tvm { -namespace relay { - -class SubgraphExtractor : public ExprVisitor { - public: - const std::unordered_set GetSubgraph(const Expr& expr); - const Map GetAffineTypes(); - void VisitExpr(const Expr& expr) override; - - protected: - void VisitExpr_(const CallNode* call_node) override; - - private: - const Op quantize_op_ = Op::Get("qnn.quantize"); - const Op dequantize_op_ = Op::Get("qnn.dequantize"); - bool is_fake_quantized_ = true; - Map affine_types_; -}; - -} // namespace relay -} // namespace tvm - -#endif // TVM_RELAY_TRANSFORMS_FAKE_QUANTIZATION_TO_INTEGER_H_ diff --git a/src/relay/transforms/fast_math.cc b/src/relay/transforms/fast_math.cc deleted file mode 100644 index f6da52ebe30c..000000000000 --- a/src/relay/transforms/fast_math.cc +++ /dev/null @@ -1,84 +0,0 @@ -/* - * 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 fast_math.cc - * \brief Replaces non linear activation functions with their fast but approximate counterparts. - */ -#include -#include -#include -#include -#include - -#include "pattern_utils.h" - -namespace tvm { -namespace relay { - -class FastMathMutator : public ExprRewriter { - public: - FastMathMutator() - : exp_op_(Op::Get("exp")), - erf_op_(Op::Get("erf")), - tanh_op_(Op::Get("tanh")), - softmax_op_(Op::Get("nn.softmax")) {} - - Expr Rewrite_(const CallNode* pre, const Expr& post) override { - if (pre->op == exp_op_) { - return FastExp(post.as()->args[0]); - } else if (pre->op == erf_op_) { - return FastErf(post.as()->args[0]); - } else if (pre->op == tanh_op_) { - return FastTanh(post.as()->args[0]); - } else if (pre->op == softmax_op_) { - return FastSoftmax(post.as()->args[0], post.as()->attrs); - } - return post; - } - - private: - // Cache the following ops. They will be used in the passes repeatedly for - // operator equivalence checking so that the registry lookup overhead can be - // reduced. - const Op& exp_op_; - const Op& erf_op_; - const Op& tanh_op_; - const Op& softmax_op_; -}; - -Expr FastMath(const Expr& e) { - auto rewriter = FastMathMutator(); - return PostOrderRewrite(e, &rewriter); -} - -namespace transform { - -Pass FastMath() { - runtime::TypedPackedFunc pass_func = - [=](Function f, IRModule m, PassContext pc) { return Downcast(FastMath(f)); }; - return CreateFunctionPass(pass_func, 4, "FastMath", {"InferType"}); -} - -TVM_REGISTER_GLOBAL("relay._transform.FastMath").set_body_typed(FastMath); - -} // namespace transform - -} // namespace relay -} // namespace tvm diff --git a/src/relay/transforms/first_order_gradient.cc b/src/relay/transforms/first_order_gradient.cc deleted file mode 100644 index f530d61e0d99..000000000000 --- a/src/relay/transforms/first_order_gradient.cc +++ /dev/null @@ -1,325 +0,0 @@ -/* - * 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 first_order_gradient.cc - * \brief First-order Automatic Differentiation in Relay for pure dataflow graphs. - */ -#include -#include -#include -#include -#include -#include -#include - -#include "gradient.h" -#include "let_list.h" -#include "pass_utils.h" -#include "pattern_utils.h" - -namespace tvm { -namespace relay { - -template -Expr MultiFactory(const Type& t, F factory, DiagnosticContext diag_ctx) { - if (auto* tt = t.as()) { - return factory(tt->shape, tt->dtype); - } else if (auto* tt = t.as()) { - std::vector res; - for (size_t i = 0; i < tt->fields.size(); i++) { - res.push_back(MultiFactory(tt->fields[i], factory, diag_ctx)); - } - return Tuple(res); - } else { - diag_ctx.EmitFatal(Diagnostic::Error(t->span) - << "could not build tensors using factory for type " << PrettyPrint(t)); - throw; - } -} - -template -Expr MultiFactoryLike(const Expr& e, const Type& t, F factory, F2 factory_like, - DiagnosticContext diag_ctx) { - if (t.as()) { - return factory_like(e); - } else if (auto* tt = t.as()) { - return MultiFactory(t, factory, diag_ctx); - } else { - diag_ctx.EmitFatal(Diagnostic::Error(t->span) - << "could not build tensors using factory for type " << PrettyPrint(t)); - throw; - } -} - -/*! \brief A fragment of the program being built by the automatic differentation - * pass. - */ -struct ADValueNode { - virtual ~ADValueNode() {} - template - T& get() { - auto ret = dynamic_cast(this); - ICHECK(ret) << "cannot downcast"; - return *ret; - } -}; - -using ADValue = std::shared_ptr; - -/*! \brief AD over a program which generates a tensor output. */ -struct ADTensor : ADValueNode { - Expr forward; - mutable Expr reverse; // must be a variable to avoid duplication - ADTensor(LetList* ll, const Expr& forward, DiagnosticContext diag_ctx) - : forward(ll->Push(forward)), - reverse(ll->Push( - MultiFactoryLike(this->forward, forward->checked_type(), Zeros, ZerosLike, diag_ctx))) { - this->forward->checked_type_ = forward->checked_type(); - } -}; - -/*! \brief A staged representation of the program, we reflect - * Relay functions into a function over fragments of AD. We - * can compute away this function to obtain a reverse mode program. - */ -struct ADFunction : ADValueNode { - // (ad_args, orig) -> ad_ret - using ADFunctionType = ADValue(const std::vector&, const Call&); - std::function func; - explicit ADFunction(const std::function& func) : func(func) {} -}; - -struct FirstOrderReverseAD : ExprFunctor { - const OpAttrMap rev_map = Op::GetAttrMap("FPrimalGradient"); - std::vector> backprop_actions; - // we assume no closure so no need for lexical scoping - std::unordered_map env; - LetList* ll; - DiagnosticContext diag_ctx; - - FirstOrderReverseAD(LetList* ll, DiagnosticContext diag_ctx) : ll(ll), diag_ctx(diag_ctx) {} - - ADValue VisitExpr(const Expr& n) final { - if (env.count(n)) { - return env.at(n); - } - auto ret = ExprFunctor::VisitExpr(n); - env[n] = ret; - return ret; - } - - static Expr LiftedAdd(const Type& t, const Expr& x, const Expr& y, LetList* ll) { - if (t.as()) { - return ll->Push(Add(x, y)); - } else if (auto* tt = t.as()) { - Array fields; - for (size_t i = 0; i < tt->fields.size(); ++i) { - fields.push_back( - LiftedAdd(tt->fields[i], ll->Push(GetField(x, i)), ll->Push(GetField(y, i)), ll)); - } - return ll->Push(Tuple(fields)); - } else { - LOG(FATAL) << "cannot lift addition for type " << PrettyPrint(t); - throw; - } - } - - ADValue VisitExpr_(const OpNode* op) final { - Op op_ref = GetRef(op); - if (!rev_map.count(op_ref)) { - diag_ctx.EmitFatal(Diagnostic::Error(op->span) - << "the operator " << op->name << " does not have a registered gradient."); - } - return std::make_shared([this, op_ref](const std::vector& ad_args, - const Call& orig) { - std::vector orig_args; - for (const ADValue& adval : ad_args) { - orig_args.push_back(adval->get().forward); - } - auto orig_new = Call(op_ref, orig_args, orig->attrs, orig->type_args); - orig_new->checked_type_ = orig->checked_type(); - auto ret = std::make_shared(ll, orig_new, diag_ctx); - backprop_actions.push_back([this, ad_args, orig_new, ret, op_ref](LetList* ll) { - tvm::Array rev = rev_map[op_ref](orig_new, ret->reverse); - if (ad_args.size() != rev.size()) { - diag_ctx.EmitFatal(Diagnostic::Error(op_ref->span) - << "arity mismatch for operator " << op_ref->name - << " and its registered gradient: expected " << ad_args.size() - << " but got " << rev.size() << " gradients."); - } - for (size_t i = 0; i < ad_args.size(); ++i) { - auto& ad_arg = ad_args[i]->get(); - ad_arg.reverse = LiftedAdd(ad_arg.forward->checked_type(), ad_arg.reverse, rev[i], ll); - } - }); - return ret; - }); - } - - ADValue VisitExpr_(const TupleGetItemNode* op) final { - ADValue tup = VisitExpr(op->tuple); - TupleType tt = Downcast(op->tuple->checked_type()); - size_t idx = op->index; - // reconstruct projection using let-bound variable to avoid duplicating input tuple - TupleGetItem orig = TupleGetItem(tup->get().forward, idx); - orig->checked_type_ = op->checked_type(); - auto ret = std::make_shared(ll, orig, diag_ctx); - // for orig = pi(tup, i), pi_grad(tup, i, g) = G where pi(G, i) = g and pi(G, j) = 0 for j != i - backprop_actions.push_back([tup, tt, idx, ret](LetList* ll) { - auto& ad_tup = tup->get(); - std::vector updated_grads; - for (size_t i = 0; i < tt->fields.size(); ++i) { - Expr grad_pre = GetField(ad_tup.reverse, i); - updated_grads.push_back(i != idx ? grad_pre - : LiftedAdd(tt->fields[i], grad_pre, ret->reverse, ll)); - } - ad_tup.reverse = ll->Push(Tuple(updated_grads)); - }); - return ret; - } - - ADValue VisitExpr_(const TupleNode* tuple_node) final { - auto tt = Downcast(tuple_node->checked_type()); - std::vector ad_fields; - Array field_bindings; - field_bindings.reserve(tuple_node->fields.size()); - - for (const auto& f : tuple_node->fields) { - ADValue f_ad = VisitExpr(f); - if (!dynamic_cast(f_ad.get())) { - diag_ctx.EmitFatal(Diagnostic::Error(f->span) - << "first-order AD only supports (nested) tuples of tensors"); - } - ad_fields.push_back(f_ad); - field_bindings.push_back(f_ad->get().forward); - } - // reconstruct tuple using let-bound variables to avoid duplication - auto orig = WithFields(GetRef(tuple_node), field_bindings); - orig->checked_type_ = tt; - auto ret = std::make_shared(ll, orig, diag_ctx); - // for orig = tuple(x1, ..., xn), tuple_grad(x1, ..., xn, G) = [pi(G, 1), ..., pi(G, n)] - backprop_actions.push_back([ad_fields, tt, ret](LetList* ll) { - for (size_t i = 0; i < ad_fields.size(); ++i) { - auto& ad_field = ad_fields[i]->get(); - ad_field.reverse = - LiftedAdd(tt->fields[i], ad_field.reverse, GetField(ret->reverse, i), ll); - } - }); - return ret; - } - - ADValue VisitExpr_(const ConstantNode* op) final { - Expr e = GetRef(op); - return std::make_shared(ll, e, diag_ctx); - } - - ADValue VisitExpr_(const CallNode* op) final { - ADValue f = VisitExpr(op->op); - std::vector args; - for (const auto& arg : op->args) { - args.push_back(VisitExpr(arg)); - } - return f->get().func(args, GetRef(op)); - } - - ADValue VisitExpr_(const FunctionNode* op) final { - Function f = GetRef(op); - // todo: assert no closure - return std::make_shared( - [this, f](const std::vector& ad_args, const Call& orig) { - ICHECK_EQ(f->params.size(), ad_args.size()); - for (size_t i = 0; i < f->params.size(); ++i) { - env[f->params[i]] = ad_args[i]; - } - return VisitExpr(f->body); - }); - } - - // Var will always be in env, handled in VisitExpr (without _), so we don't need - // to implement its VisitExpr_. -}; - -namespace transform { - -Pass FirstOrderGradient() { - runtime::TypedPackedFunc f = [](IRModule mod, PassContext ctx) { - CheckFeature( - mod, FeatureSet({fVar, fConstant, fTuple, fTupleGetItem, fFunction, fOp, fCall, fGraph})); - IRModule ad_mod = GetRef(mod.CopyOnWrite()); - DiagnosticContext diag_ctx = DiagnosticContext::Default(ad_mod); - - if (mod->functions.size() > 1) { - LOG(WARNING) << "IRModule contains multiple global functions: first-order AD will transform " - "them indepedently!"; - } - - for (const auto& pr : mod->functions) { - const FunctionNode* func = pr.second.as(); - if (!func) { - diag_ctx.Emit(Diagnostic::Warning(pr.second->span) - << "AD can only be performed on Relay functions, skipping " - << PrettyPrint(pr.first)); - } - if (func->type_params.size() > 0) { - diag_ctx.EmitFatal(Diagnostic::Error(pr.second->span) - << "first-order AD does not support polymorphism yet."); - } - Expr body = LetList::With([&](LetList* ll) { - FirstOrderReverseAD reverse_ad(ll, diag_ctx); - ADValue rev = reverse_ad(pr.second); - std::vector args; - for (const auto& p : func->params) { - args.push_back(std::make_shared(ll, p, diag_ctx)); - } - Call placeholder = Call(GetRef(func), {}); - placeholder->checked_type_ = func->checked_type().as()->ret_type; - auto grad_call = rev->get().func(args, placeholder); - auto& res = grad_call->get(); - Expr grad_tuple = LetList::With([&](LetList* ll) { - res.reverse = - MultiFactoryLike(res.forward, res.forward->checked_type(), Ones, OnesLike, diag_ctx); - for (auto it = reverse_ad.backprop_actions.rbegin(); - it != reverse_ad.backprop_actions.rend(); ++it) { - (*it)(ll); - } - std::vector grads; - for (const auto& a : args) { - grads.push_back(a->get().reverse); - } - return Tuple(grads); - }); - return Pair(res.forward, grad_tuple); - }); - ad_mod->Update(pr.first, WithFields(GetRef(func), func->params, body, - GradRetType(GetRef(func)), - /* erase type params */ Array())); - } - - return ad_mod; - }; - return CreateModulePass(f, 0, "FirstOrderGradient", {}); -} - -TVM_REGISTER_GLOBAL("relay._transform.FirstOrderGradient").set_body_typed(FirstOrderGradient); - -} // namespace transform - -} // namespace relay -} // namespace tvm diff --git a/src/relay/transforms/flatten_atrous_conv.cc b/src/relay/transforms/flatten_atrous_conv.cc deleted file mode 100644 index 54e0f193cf8b..000000000000 --- a/src/relay/transforms/flatten_atrous_conv.cc +++ /dev/null @@ -1,195 +0,0 @@ -/* - * 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 src/relay/transforms/flatten_atrous_conv.cc - * \brief This transform flattens atrous convolution, which corresponds to the sequence of - * operations: "space_to_batch_nd"->"conv2d"->"batch_to_space_nd". - */ - -#include -#include -#include -#include -#include -#include -#include - -#include -#include -#include - -#include "../qnn/utils.h" -#include "pattern_utils.h" - -namespace tvm { -namespace relay { - -/* Description of FlattenAtrousConv - * - * The purpose of this pass is to find a sequence of space_to_batch_nd-conv2d-batch_to_space_nd - * operations: - * - * x w - * | | - * s2b | - * \ / - * conv2d - * | - * b2s - * - * and convert them into subgraphs with a convolution with the modified "dilation" and - * recalculated "padding" parameters. - */ - -using ExprSet = std::unordered_set; - -class FlattenAtrousConvSubgraphMutator { - public: - Expr MutateSubgraph(const Expr& expr) { - try { - const CallNode* b2s_node_ = expr.as(); - const CallNode* conv2d_node_ = b2s_node_->args[0].as(); - const CallNode* s2b_node_ = conv2d_node_->args[0].as(); - - ICHECK(b2s_node_ != nullptr); - const auto* b2s_attrs = b2s_node_->attrs.as(); - ICHECK(b2s_attrs != nullptr); - - Array dilation = {b2s_attrs->block_shape[0], b2s_attrs->block_shape[1]}; - - ICHECK(conv2d_node_ != nullptr); - const auto* conv2d_attrs = conv2d_node_->attrs.as(); - ICHECK(conv2d_attrs != nullptr); - - Array kernel_shape = conv2d_attrs->kernel_size; - PrimExpr kernel_h = kernel_shape[0]; - PrimExpr kernel_w = kernel_shape[1]; - - ICHECK(s2b_node_ != nullptr); - const auto* s2b_attrs = s2b_node_->attrs.as(); - ICHECK(s2b_attrs != nullptr); - - Expr data = s2b_node_->args[0]; - ICHECK(conv2d_attrs->data_layout == "NHWC"); - Array data_shape = transform::InferTypeLocal(data).as()->shape; - PrimExpr in_h = data_shape[1]; - PrimExpr in_w = data_shape[2]; - - PrimExpr dilation_h = dilation[0]; - PrimExpr dilation_w = dilation[1]; - - PrimExpr dilated_kernel_h = (kernel_h - 1) * dilation_h + 1; - PrimExpr dilated_kernel_w = (kernel_w - 1) * dilation_w + 1; - - Array strides = {1, 1}; - PrimExpr stride_h = strides[0]; - PrimExpr stride_w = strides[1]; - - auto _get_pad_pair = [](PrimExpr input1d, PrimExpr kernel1d, - PrimExpr stride1d) -> Array { - PrimExpr out1d = truncdiv((input1d + stride1d - 1), stride1d); - PrimExpr pad = topi::maximum(((out1d - 1) * stride1d + kernel1d - input1d), 0); - PrimExpr pad_before = truncdiv(pad, 2); - PrimExpr pad_after = pad - pad_before; - return {pad_before, pad_after}; - }; - - Array pad_v = _get_pad_pair(in_h, dilated_kernel_h, stride_h); - Array pad_h = _get_pad_pair(in_w, dilated_kernel_w, stride_w); - - Array padding = {pad_v[0], pad_h[0], pad_v[1], pad_h[1]}; - - Expr weight = conv2d_node_->args[1]; - - if (conv2d_node_->op == Op::Get("nn.conv2d")) { - return Conv2D(data, weight, strides, padding, dilation, conv2d_attrs->groups, - conv2d_attrs->channels, conv2d_attrs->kernel_size, conv2d_attrs->data_layout, - conv2d_attrs->kernel_layout, conv2d_attrs->out_layout, - conv2d_attrs->out_dtype); - } - - if (conv2d_node_->op == Op::Get("qnn.conv2d")) { - Expr input_zero_point = conv2d_node_->args[2]; - Expr kernel_zero_point = conv2d_node_->args[3]; - Expr input_scale = conv2d_node_->args[4]; - Expr kernel_scale = conv2d_node_->args[5]; - return qnn::MakeQnnConv2D(data, weight, input_zero_point, kernel_zero_point, input_scale, - kernel_scale, strides, padding, dilation, conv2d_attrs->groups, - conv2d_attrs->channels, conv2d_attrs->kernel_size, - conv2d_attrs->data_layout, conv2d_attrs->kernel_layout, - conv2d_attrs->out_layout, conv2d_attrs->out_dtype); - } - - DLOG(INFO) << "Ran into an unhandled convolution, skipping " << expr << std::endl; - return expr; - } catch (std::exception& e) { - DLOG(INFO) << "Ran into an error rewriting a subgraph, skipping " << expr << " with " - << e.what() << std::endl; - return expr; - } - } -}; - -class FlattenAtrousConvRewriter : public MixedModeMutator { - protected: - Expr Rewrite_(const CallNode* pre, const Expr& post) override { - if (const CallNode* call_node = post.as()) { - if (ops_[op_iter_].count(call_node->op)) { - ++op_iter_; - if (op_iter_ == ops_.size()) { - op_iter_ = 0; - return FlattenAtrousConvSubgraphMutator().MutateSubgraph(post); - } - } else { - op_iter_ = 0; - } - } - return post; - } - - private: - size_t op_iter_ = 0; - const std::array ops_ = { - ExprSet{Op::Get("nn.space_to_batch_nd")}, - ExprSet{Op::Get("nn.conv2d"), Op::Get("qnn.conv2d")}, - ExprSet{Op::Get("nn.batch_to_space_nd")}, - }; -}; - -Expr FlattenAtrousConv(const Expr& expr, const IRModule& mod) { - return FlattenAtrousConvRewriter().Mutate(expr); -} - -namespace transform { - -Pass FlattenAtrousConv() { - runtime::TypedPackedFunc pass_func = - [=](Function f, IRModule m, PassContext pc) { - return Downcast(FlattenAtrousConv(f, m)); - }; - return CreateFunctionPass(pass_func, 0, "FlattenAtrousConv", {"InferType"}); -} - -TVM_REGISTER_GLOBAL("relay._transform.FlattenAtrousConv").set_body_typed(FlattenAtrousConv); - -} // namespace transform - -} // namespace relay -} // namespace tvm diff --git a/src/relay/transforms/fold_constant.cc b/src/relay/transforms/fold_constant.cc deleted file mode 100644 index df28506c6217..000000000000 --- a/src/relay/transforms/fold_constant.cc +++ /dev/null @@ -1,450 +0,0 @@ -/* - * 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 constant_folding.cc - */ -#include -#include -#include -#include -#include -#include -#include -#include -#include -#include -#include - -#include "../op/memory/on_device.h" -#include "./pattern_utils.h" - -namespace tvm { -namespace relay { -namespace transform { - -namespace { -/*! - * \brief Returns whether \p expr is a literal \p Constant, optionally wrapped by an "on_device" - * annotation CallNode (which serves only to associate an \p VirtualDevice to the constant and has - * no operational effect). - */ -bool IsSimpleConstant(const Expr& expr) { - return AsIgnoringOnDevice(expr) != nullptr; -} - -/*! - * \brief Returns whether \p expr \p IsSimpleConstant directly or is a tuple of - * \p IsComplexConstant expressions. - */ -bool IsComplexConstant(const Expr& expr) { - if (IsSimpleConstant(expr)) { - return true; - } else if (const auto* tuple_node = AsIgnoringOnDevice(expr)) { - return std::all_of(tuple_node->fields.begin(), tuple_node->fields.end(), IsComplexConstant); - } else { - return false; - } -} - -// TODO(tvm-team) consider combine dead-code with constant folder. -// or make a more powerful partial evaluator. -class ConstantFolder : public MixedModeMutator { - public: - explicit ConstantFolder(IRModule module, bool fold_qnn) - : module_(std::move(module)), - fold_qnn_(fold_qnn), - device_copy_op_(Op::Get("device_copy")), - shape_of_op_(Op::Get("shape_of")), - vm_shape_of_op_(Op::Get("vm.shape_of")), - cast_op_(Op::Get("cast")), - ndarray_size_op_(Op::Get("ndarray_size")) {} - - private: - using ExprMutator::VisitExpr_; - - Expr VisitExpr_(const LetNode* let_node) final { - auto pre_visit = [this](const LetNode* op) { - // Rely on the Memoizer to cache pre-visit values - Expr new_value = Mutate(op->value); - if (IsSimpleConstant(new_value)) { - // Inline new value (along with any on_device annotation wrapping it) at all occurrences of - // the variable. - // - // We need to retain any "on_device" annotation so that downstream 'device aware' - // passes can still retrieve the virtual device for the constant in its new position(s). Eg: - // def @f(..., result_virtual_device=D) { - // let %x = on_device(... something we eval to a constant..., virtual_device=E) - // @f(..., %x, ...) - // } - // Here the default virtual device is D, whereas the argument %x to @f is on E (and @f - // expects that). No on_device annotation is required in the call according to the - // convention used by the device-aware visitors. - // - // However once we've inlined the constant we need to insert an on_device, again to - // respect the convention used by the device-aware visitors. - // def @f(..., result_virtual_device=D) { - // @f(..., on_device(...the constant..., virtual_device=E), ...) - // } - VLOG(1) << "Replacing let-binding for " << op->var->name_hint() - << " with constant:" << std::endl - << PrettyPrint(new_value); - memo_[op->var] = new_value; - } else { - this->Mutate(op->var); - } - }; - auto post_visit = [this](const LetNode* op) { - Expr expr = GetRef(op); - // Rely on the Memoizer to cache pre-visit values - Expr new_value = this->Mutate(op->value); - if (IsSimpleConstant(new_value)) { - // The let-bound value has been inlined, drop the let-binding itself. - this->memo_[expr] = Mutate(op->body); - } else { - Var new_var = Downcast(this->Mutate(op->var)); - Expr new_body = this->Mutate(op->body); - if (new_var.same_as(op->var) && new_value.same_as(op->value) && - new_body.same_as(op->body)) { - this->memo_[expr] = expr; - } else { - this->memo_[expr] = Let(new_var, new_value, new_body, op->span); - } - } - }; - ExpandANormalForm(let_node, pre_visit, post_visit); - return memo_[GetRef(let_node)]; - } - - Expr VisitExpr_(const FunctionNode* function_node) final { - if (function_node->HasNonzeroAttr(attr::kPrimitive)) { - ICHECK_EQ(inside_primitive_, false); - inside_primitive_ = true; - auto ret = ExprMutator::VisitExpr_(function_node); - inside_primitive_ = false; - return ret; - } else { - return ExprMutator::VisitExpr_(function_node); - } - } - - Expr Rewrite_(const CallNode* pre_call_node, const Expr& post) final { - Call pre_call = GetRef(pre_call_node); - if (inside_primitive_) { - return std::move(pre_call); - } - - Call post_call = Downcast(post); - - if (post_call->args.empty()) { - // We don't constant fold function with zero arguments. - // This is a heuristic that is useful. - // For example it is harmful to fold ones(shape=(4, 5)). - return std::move(pre_call); - } - - const auto* op_node = post_call->op.as(); - if (op_node == nullptr) { - // Only evaluate primitives. - return std::move(post_call); - } - Op op = GetRef(op_node); - static auto op_stateful = Op::GetAttrMap("TOpIsStateful"); - if (op_stateful.get(op, false)) { - // skip stateful ops. - return std::move(post_call); - } - // Try to evaluate shape_of and ndarray_size ops - // Use the original call rather than new_call here since it still has valid checked_type - // fields. These operators don't care about the value of their argument anyway. - if (Optional opt_result = EvaluateShapeOf(pre_call)) { - return opt_result.value(); - } - // Use the original call rather than new_call here since it still has valid checked_type - // fields. This operator doesn't care about the value of its argument anyway. - if (Optional opt_result = EvaluateNdarraySize(pre_call)) { - return opt_result.value(); - } - static auto fnoncomputational = Op::GetAttrMap("TNonComputational"); - static auto qnn_canonicalize = Op::GetAttrMap("FTVMQnnCanonicalize"); - bool is_no_qnn_canonicalized = !qnn_canonicalize.count(op); - bool is_no_computational = fnoncomputational.count(op) && fnoncomputational[op]; - if (is_no_computational && (is_no_qnn_canonicalized || !fold_qnn_)) { - return std::move(post_call); - } - if (op == device_copy_op_ || op == shape_of_op_ || op == vm_shape_of_op_ || - op == ndarray_size_op_) { - // We should think about potentially constant evaluation over these ops too. - return std::move(post_call); - } - if (!std::all_of(post_call->args.begin(), post_call->args.end(), IsComplexConstant)) { - // At least one non-constant argument. - return std::move(post_call); - } - // During evaluation we have obviously lost all on_device annotations. However any - // on_device wrapping this call will be left in place. - return ConstEvaluate(post_call); - } - - Expr VisitExpr_(const IfNode* if_node) final { - If new_if = Downcast(ExprMutator::VisitExpr_(if_node)); - if (const auto* const_node = AsIgnoringOnDevice(new_if->cond)) { - if (reinterpret_cast(const_node->data->data)[0]) { - return new_if->true_branch; - } else { - return new_if->false_branch; - } - } - return std::move(new_if); - } - - Expr Rewrite_(const TupleGetItemNode* tuple_get_item_node, - const Expr& post_tuple_get_item) final { - const auto* post_tuple_get_item_node = post_tuple_get_item.as(); - if (const auto* tuple_node = AsIgnoringOnDevice(post_tuple_get_item_node->tuple)) { - Expr result = tuple_node->fields[tuple_get_item_node->index]; - OnDeviceProps props = GetOnDeviceProps(post_tuple_get_item_node->tuple); - if (props.body.defined()) { - // (on_device((x, y, z), virtual_device=D).1 ==> on_device(y, virtual_device=D) - return MaybeOnDeviceWithProps(result, props); - } else { - return result; - } - } - return post_tuple_get_item; - } - - // Convert value to expression. - Expr ObjectToExpr(const ObjectRef& value) { - if (value->IsInstance()) { - auto nd_array = Downcast(value); - return Constant(nd_array); - } else if (auto opt = value.as()) { - runtime::ADT adt = opt.value(); - Array fields; - for (size_t i = 0; i < adt.size(); ++i) { - fields.push_back(ObjectToExpr(adt[i])); - } - return Tuple(fields); - } else { - LOG(FATAL) << "Cannot handle " << value->GetTypeKey(); - } - } - - // Constant evaluate an expression. - Expr ConstEvaluate(const Expr& expr) { - VLOG_CONTEXT << "ConstEvaluate"; - VLOG(1) << "Evaluating :" << std::endl << PrettyPrint(expr); - - // We'll invoke the interpreter using the generic CPU device and target. Technically there's - // no guarantee the results will be bitwise equal what we'd get on the true device, however to - // support cross-compilation we don't want to assume the true device is available. - - // Use a fresh build context in case we are already in a build context. - // needed for both execution and creation(due to JIT) - With fresh_build_ctx(transform::PassContext::Create()); - - Map dict = (module_->attrs.defined()) - ? Map(module_->attrs.CopyOnWrite()->dict) - : Map(); - - // always use graph executor with no link-params - dict.Set(tvm::attr::kExecutor, - relay::Executor::Create("graph", {{"link-params", runtime::Bool(false)}})); - Expr result = ObjectToExpr(Eval(expr, module_->type_definitions, module_->Imports(), - eval_cpu_dev_, eval_cpu_target_, dict)); - VLOG(1) << "Evaluated to constant:" << std::endl << PrettyPrint(result); - return result; - } - - /*! - * \brief Returns constant shape result of \p call if it of form \p shape_of(e) and \p e has - * a non-dynamic tensor shape. Returns null otherwise. - */ - Optional EvaluateShapeOf(const Call& call) { - if (call->op != shape_of_op_ && call->op != vm_shape_of_op_) { - return {}; - } - - VLOG(1) << "Evaluating for shape_of:" << std::endl << PrettyPrint(call); - ICHECK_EQ(call->args.size(), 1); - const auto* param = call->attrs.as(); - ICHECK(param != nullptr); - Expr input = call->args[0]; - - tvm::Array ishape; - if (Optional> opt_shape = GetConstantShape(input)) { - ishape = opt_shape.value(); - } else { - return {}; - } - - // Get the constant shape - runtime::NDArray value; - DLDataType cdtype = DataType::Int(32); - if (ishape.empty()) { - value = runtime::NDArray::Empty({}, cdtype, eval_cpu_dev_); - } else { - ICHECK_NE(ishape.size(), 0); - std::vector cshape = {static_cast(ishape.size())}; - value = runtime::NDArray::Empty(cshape, cdtype, eval_cpu_dev_); - auto* dims = static_cast(value->data); - using ::tvm::tir::IntImmNode; - for (size_t i = 0; i < ishape.size(); ++i) { - if (const auto* dim = ishape[i].as()) { - dims[i] = dim->value; - } else { - return {}; - } - } - } - - Constant shape = Downcast(ObjectToExpr(value)); - - if (shape->data.Shape().empty() && GetScalarFromConstant(shape) == 0) { - auto ndarray = runtime::NDArray::Empty({}, cdtype, eval_cpu_dev_); - shape = Constant(ndarray); - } - - return CastValue(shape, param->dtype); - } - - /*! - * \brief Returns the constant NDArray size of result of \p call if it is of the form - * \p ndarray_size(e) and \p e has non-dynamic tensor type. Returns null otherwise. - */ - Optional EvaluateNdarraySize(const Call& call) { - if (call->op != ndarray_size_op_) { - return {}; - } - VLOG(1) << "Evaluating for ndarray_size:" << std::endl << PrettyPrint(call); - ICHECK_EQ(call->args.size(), 1); - Expr input = call->args[0]; - const auto* param = call->attrs.as(); - ICHECK(param != nullptr); - - tvm::Array ishape; - if (Optional> opt_shape = GetConstantShape(input)) { - ishape = opt_shape.value(); - } else { - return {}; - } - - // Get the constant size - runtime::NDArray value; - DLDataType cdtype = DataType::Int(32); - value = runtime::NDArray::Empty({}, cdtype, eval_cpu_dev_); - auto* data = static_cast(value->data); - if (ishape.empty()) { - *data = 0; - } else { - *data = 1; - using ::tvm::tir::IntImmNode; - for (size_t i = 0; i < ishape.size(); ++i) { - if (const auto* dim = ishape[i].as()) { - *data *= dim->value; - } else { - return {}; - } - } - } - - Constant size = Downcast(ObjectToExpr(value)); - return CastValue(size, param->dtype); - } - - Expr CastValue(const Expr& value, DataType dtype) { - // Cast the constant into correct dtype - auto cast_attrs = make_object(); - cast_attrs->dtype = dtype; - Expr ret = Call(cast_op_, {value}, Attrs(cast_attrs), {}); - return ConstEvaluate(ret); - } - - Optional> GetConstantShape(const Expr& input) { - if (const auto* const_node = AsIgnoringOnDevice(input)) { - // TODO(mbs): This is not necessary since we only ever ask for the shapes for - // pre-rewritten expressions which will always have a checked_type. - return const_node->tensor_type()->shape; - } else if (input->checked_type_.defined()) { - return input->checked_type().as()->shape; - } else { - return {}; - } - } - - // Module - IRModule module_; - - // Whether to fold constants for QNN operations. - bool fold_qnn_; - - // The kDLCPU device assumed to be available to the compiler. Used only when evaluating - // sub-expressions. - Device eval_cpu_dev_{kDLCPU, /*device_id=*/0}; - // The target for the above device assumed to be available to the compiler. Used only when - // evaluating sub-expressions. - Target eval_cpu_target_{"llvm"}; - - // Cache the following ops for equivalence checking in this pass. - const Op& device_copy_op_; - const Op& shape_of_op_; - const Op& vm_shape_of_op_; - const Op& cast_op_; - const Op& ndarray_size_op_; - - // True if currently within a "primitive" Relay Function. - bool inside_primitive_ = false; -}; - -} // namespace - -TVM_REGISTER_GLOBAL("relay.analysis.check_constant").set_body_typed(IsComplexConstant); - -Expr FoldConstantExpr(const Expr& expr, const IRModule& mod, bool fold_qnn) { - VLOG_CONTEXT << "FoldConstantExpr"; - VLOG(1) << "folding:" << std::endl << PrettyPrint(expr); - Expr result = ConstantFolder(mod, fold_qnn).VisitExpr(expr); - VLOG(1) << "folded to:" << std::endl << PrettyPrint(result); - return result; -} - -Expr FoldConstantExpr(const Expr& expr, bool fold_qnn) { - auto mod = IRModule::FromExpr(expr); - return FoldConstantExpr(expr, mod, fold_qnn); -} - -TVM_REGISTER_GLOBAL("relay._transform.FoldConstantExpr") - .set_body_typed([](const Expr& expr, const IRModule& mod, bool fold_qnn) { - return FoldConstantExpr(expr, mod, fold_qnn); - }); - -Pass FoldConstant(bool fold_qnn) { - runtime::TypedPackedFunc pass_func = - [=](Function f, IRModule m, PassContext /* pc */) { - return Downcast(FoldConstantExpr(f, m, fold_qnn)); - }; - return CreateFunctionPass(pass_func, 2, "FoldConstant", {}); -} - -TVM_REGISTER_GLOBAL("relay._transform.FoldConstant").set_body_typed(FoldConstant); - -} // namespace transform -} // namespace relay -} // namespace tvm diff --git a/src/relay/transforms/fold_constant.h b/src/relay/transforms/fold_constant.h deleted file mode 100644 index 4f475037d195..000000000000 --- a/src/relay/transforms/fold_constant.h +++ /dev/null @@ -1,55 +0,0 @@ -/* - * 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 fold_constant.h - * \brief Utility functions for folding constants in expressions. - */ -#ifndef TVM_RELAY_TRANSFORMS_FOLD_CONSTANT_H_ -#define TVM_RELAY_TRANSFORMS_FOLD_CONSTANT_H_ - -#include - -namespace tvm { -namespace relay { -namespace transform { - -/*! - * \brief Apply constant folding on an expression. - * - * \param expr The expression to fold. - * \param fold_qnn Whether to fold constants for QNN operations. - * \returns The new folded expression. - */ -Expr FoldConstantExpr(const Expr& expr, bool fold_qnn = true); - -/*! - * \brief Returns \p expr with any constants expressions evaluated and let-bound constants - * inlined. Returns \p expr unchanged if no change. - * - * CAUTION: The importers rely on this function returning \p expr unchanged to preserve sharing - * from their p.o.v. Furthermore, this function can be called before conversion to ANF so - * we must avoid all recursion. - */ -Expr FoldConstantExpr(const Expr& expr, const IRModule& mod, bool fold_qnn); - -} // namespace transform -} // namespace relay -} // namespace tvm -#endif // TVM_RELAY_TRANSFORMS_FOLD_CONSTANT_H_ diff --git a/src/relay/transforms/fold_explicit_padding.cc b/src/relay/transforms/fold_explicit_padding.cc deleted file mode 100644 index 55e1fe854fe3..000000000000 --- a/src/relay/transforms/fold_explicit_padding.cc +++ /dev/null @@ -1,380 +0,0 @@ -/* - * 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 src/relay/transforms/fold_explicit_padding.cc - * \brief A pass for folding explicit pads into other ops. - */ - -#include -#include -#include -#include -#include -#include -#include -#include - -#include -#include -#include - -#include "../op/tensor/transform.h" -#include "pattern_utils.h" - -namespace tvm { -namespace relay { - -/*! - * \brief SimplifyExplicitPad matches a pad followed by a conv/maxpool/avgpool - * with a pad attribute and merges the padding into the kernel. - */ -class SimplifyExplicitPad { - public: - DFPattern pattern() const { return pattern_; } - - SimplifyExplicitPad() { - x_ = IsWildcard(); - pad_ = IsOp("nn.pad")({x_, IsWildcard()}); - - // pad->conv patterns - w_ = IsWildcard(); - conv1d_ = IsOp("nn.conv1d"); - conv2d_ = IsOp("nn.conv2d"); - conv3d_ = IsOp("nn.conv3d"); - contrib_conv2d_nchwc_ = IsOp("nn.contrib_conv2d_NCHWc"); - conv_ = (conv1d_ || conv2d_ || conv3d_ || contrib_conv2d_nchwc_)({pad_, w_}); - - input_zero_point_ = IsWildcard(); - kernel_zero_point_ = IsWildcard(); - input_scale_ = IsWildcard(); - kernel_scale_ = IsWildcard(); - qconv2d_ = IsOp("qnn.conv2d")( - {pad_, w_, input_zero_point_, kernel_zero_point_, input_scale_, kernel_scale_}); - - // pad->pool patterns - avg_pool1d_ = IsOp("nn.avg_pool1d"); - avg_pool2d_ = IsOp("nn.avg_pool2d"); - avg_pool3d_ = IsOp("nn.avg_pool3d"); - max_pool1d_ = IsOp("nn.max_pool1d"); - max_pool2d_ = IsOp("nn.max_pool2d"); - max_pool3d_ = IsOp("nn.max_pool3d"); - max_pool_ = max_pool1d_ || max_pool2d_ || max_pool3d_; - pool_ = (max_pool_ || avg_pool1d_ || avg_pool2d_ || avg_pool3d_)({pad_}); - - pattern_ = conv_ || qconv2d_ || pool_; - } - - template - Array get_combined_padding(const T* old_attrs, Array padding) const { - ICHECK(padding.size() == old_attrs->padding.size()) - << "Number of dimensions to pad and convolution padding attributes should have the same " - "extent"; - - Array combined_padding; - for (size_t i = 0; i < padding.size(); ++i) { - combined_padding.push_back(padding[i] + old_attrs->padding[i]); - } - return combined_padding; - } - - template - Attrs MakeConvAttrs(const PadAttrs* param, const T* old_attrs) const { - // Creates attrs from old_attrs with fields shared by 1D, 2D, 3D conv attrs - ICHECK(old_attrs); - ICHECK(param); - auto padding = get_padding(param, old_attrs->data_layout); - if (!padding) { - return Attrs(); - } - auto combined_padding = get_combined_padding(old_attrs, padding.value()); - - auto new_attrs = make_object(); - new_attrs->strides = old_attrs->strides; - new_attrs->padding = combined_padding; - new_attrs->dilation = old_attrs->dilation; - new_attrs->groups = old_attrs->groups; - new_attrs->channels = old_attrs->channels; - new_attrs->kernel_size = old_attrs->kernel_size; - new_attrs->data_layout = old_attrs->data_layout; - new_attrs->kernel_layout = old_attrs->kernel_layout; - new_attrs->out_layout = old_attrs->out_layout; - new_attrs->out_dtype = old_attrs->out_dtype; - return Attrs(new_attrs); - } - - template - Attrs MakeConv2D3DAttrs(const PadAttrs* param, const T* old_attrs) const { - // Propagate additional Conv2D- and Conv3D-specific attrs - auto attrs = MakeConvAttrs(param, old_attrs); - if (!attrs.defined()) { - return Attrs(); - } - - T* new_attrs = const_cast(attrs.template as()); - new_attrs->auto_scheduler_rewritten_layout = old_attrs->auto_scheduler_rewritten_layout; - new_attrs->meta_schedule_original_shape = old_attrs->meta_schedule_original_shape; - return attrs; - } - - template - Attrs MakePoolAttrs(const PadAttrs* param, const T* old_attrs) const { - // Creates attrs from old_attrs with fields shared by 1D, 2D, 3D pool attrs - ICHECK(old_attrs); - ICHECK(param); - auto padding = get_padding(param, old_attrs->layout); - if (!padding) { - return Attrs(); - } - auto combined_padding = get_combined_padding(old_attrs, padding.value()); - - auto new_attrs = make_object(); - new_attrs->pool_size = old_attrs->pool_size; - new_attrs->strides = old_attrs->strides; - new_attrs->dilation = old_attrs->dilation; - new_attrs->padding = combined_padding; - new_attrs->layout = old_attrs->layout; - new_attrs->out_layout = old_attrs->out_layout; - new_attrs->ceil_mode = old_attrs->ceil_mode; - return Attrs(new_attrs); - } - - template - Attrs MakeAvgPoolAttrs(const PadAttrs* param, const T* old_attrs) const { - // Propagate additional AvgPool-specific attrs - auto attrs = MakePoolAttrs(param, old_attrs); - if (!attrs.defined()) { - return attrs; - } - - T* new_attrs = const_cast(attrs.template as()); - new_attrs->count_include_pad = old_attrs->count_include_pad; - if (!new_attrs->count_include_pad) { - // AvgPool's divisor doesn't include padding, so don't fold the explicit pad - // unless all original pad items are 0. - for (IndexExpr pad : old_attrs->padding) { - const IntImmNode* maybe_int_imm = pad.as(); - if (!maybe_int_imm || maybe_int_imm->value != 0) { - // Return undefined attrs to signal that we don't want to fold explicit pad - return Attrs(); - } - } - // Turn on `count_include_pad` to preserve original pad first, then pool behavior - // where AvgPool's divisor implicitly includes padding. - new_attrs->count_include_pad = true; - } - - return attrs; - } - - static const std::optional> get_padding(const PadAttrs* param, - std::string data_layout) { - // Gets spatial axes padding from the given PadAttrs `param`. If padding - // is non-zero on non-spatial axes, return std::nullopt. - ICHECK(param); - ICHECK(data_layout.size() == param->pad_width.size()) - << "Data Layout and padding attributes should have the same extent"; - - std::set image_dims({'H', 'W', 'D'}); - Array padding; - // If we're padding a non-spatial dimension, don't simplify - // Convolution/Pool can only pad on spatial axes - for (size_t i = 0; i < param->pad_width.size(); ++i) { - if (!image_dims.count(data_layout[i])) { - for (size_t j = 0; j < param->pad_width[i].size(); ++j) { - if (param->pad_width[i][j] != 0) { - return std::nullopt; - } - } - } - } - for (size_t j = 0; j < param->pad_width[0].size(); ++j) { - for (size_t i = 0; i < param->pad_width.size(); ++i) { - if (image_dims.count(data_layout[i])) { - padding.push_back(param->pad_width[i][j]); - } - } - } - return padding; - } - - Expr callback(const Expr& pre, const Expr& post, - const Map>& node_map) const { - const CallNode* call_node = post.as(); - ICHECK(call_node); - auto pad = node_map[pad_][0]; - const CallNode* pad_node = pad.as(); - ICHECK(pad_node); - const PadAttrs* param = pad_node->attrs.as(); - ICHECK(param); - - auto x = node_map[x_][0]; - - const Expr& pv = pad_node->args[1]; - const ConstantNode* pad_value = pv.as(); - - if (node_map.find(qconv2d_) != node_map.end()) { - Attrs attrs = MakeConv2D3DAttrs(param, call_node->attrs.as()); - if (!attrs.defined()) { - return post; - } - auto input_zero_point = node_map[input_zero_point_][0]; - auto kernel_zero_point = node_map[kernel_zero_point_][0]; - auto input_scale = node_map[input_scale_][0]; - auto kernel_scale = node_map[kernel_scale_][0]; - // Fold Padding and QNN Convolution only if pad value == input zero point. - if (IsEqualScalar(input_zero_point, pv)) { - auto w = node_map[w_][0]; - return Call(call_node->op, - {x, w, input_zero_point, kernel_zero_point, input_scale, kernel_scale}, attrs, - call_node->type_args, call_node->span); - } - return post; - } - - if (param->pad_mode == "constant" && pad_value) { - Attrs attrs; - auto pad_scalar = ToScalar(pad_value->data); - if (pad_scalar == 0.0) { - // Fold Padding and Conv/AvgPool only if pad_value == 0. - if (node_map.count(conv_)) { - if (node_map.count(conv1d_)) { - attrs = MakeConvAttrs(param, call_node->attrs.as()); - } else if (node_map.count(conv2d_)) { - attrs = MakeConv2D3DAttrs(param, call_node->attrs.as()); - } else if (node_map.count(conv3d_)) { - attrs = MakeConv2D3DAttrs(param, call_node->attrs.as()); - } - if (!attrs.defined()) { - return post; - } - auto w = node_map[w_][0]; - return Call(call_node->op, {x, w}, attrs, call_node->type_args, call_node->span); - } else if (node_map.count(avg_pool1d_)) { - attrs = MakeAvgPoolAttrs(param, call_node->attrs.as()); - } else if (node_map.count(avg_pool2d_)) { - attrs = MakeAvgPoolAttrs(param, call_node->attrs.as()); - } else if (node_map.count(avg_pool3d_)) { - attrs = MakeAvgPoolAttrs(param, call_node->attrs.as()); - } - } - if (node_map.count(max_pool_)) { - // Fold Padding and MaxPool only if pad_value is the min possible value for the dtype - auto min_value = tvm::min_value(tvm::runtime::DataType(pad_value->data->dtype)); - const FloatImmNode* maybe_min_float = min_value.as(); - const IntImmNode* maybe_min_int = min_value.as(); - - if ((maybe_min_float && pad_scalar == maybe_min_float->value) || - (maybe_min_int && pad_scalar == maybe_min_int->value)) { - if (node_map.count(max_pool1d_)) { - attrs = MakePoolAttrs(param, call_node->attrs.as()); - } else if (node_map.count(max_pool2d_)) { - attrs = MakePoolAttrs(param, call_node->attrs.as()); - } else if (node_map.count(max_pool3d_)) { - attrs = MakePoolAttrs(param, call_node->attrs.as()); - } - } - } - if (!attrs.defined()) { - return post; - } - return Call(call_node->op, {x}, attrs, call_node->type_args, call_node->span); - } - return post; - } - - private: - /*! \brief Pattern for rewriting */ - DFPattern pattern_; - /*! \brief Pattern input */ - DFPattern x_; - /*! \brief Pattern input weight */ - DFPattern w_; - /*! \brief Pattern pad */ - DFPattern pad_; - /*! \brief Pattern conv */ - DFPattern conv_; - DFPattern conv1d_; - DFPattern conv2d_; - DFPattern conv3d_; - DFPattern contrib_conv2d_nchwc_; - DFPattern qconv2d_; - DFPattern input_zero_point_; - DFPattern kernel_zero_point_; - DFPattern input_scale_; - DFPattern kernel_scale_; - /*! \brief Pattern pool */ - DFPattern pool_; - DFPattern avg_pool1d_; - DFPattern avg_pool2d_; - DFPattern avg_pool3d_; - DFPattern max_pool1d_; - DFPattern max_pool2d_; - DFPattern max_pool3d_; - DFPattern max_pool_; -}; - -class SimplifyExplicitPadding { - public: - explicit SimplifyExplicitPadding(IRModule mod) : mod_(mod) { - CreateCallback(SimplifyExplicitPad()); - } - template - void CreateCallback(const T& pattern) { - auto func = [pattern](TVMArgs args, TVMRetValue* rv) { - Expr pre = args[0]; - Expr post = args[1]; - Map> node_map = args[2]; - *rv = pattern.callback(pre, post, node_map); - }; - callbacks_.push_back(DFPatternCallback(pattern.pattern(), PackedFunc(func), true)); - } - - Expr Simplify(const Expr& expr) { return RewritePatterns(callbacks_, expr, mod_); } - - private: - IRModule mod_; - /*! \brief Callbacks for expr simplification */ - Array callbacks_; -}; - -/*! - * \brief FoldExplicitPadding finds explict padding before an op that can - * support implicit padding and fuses them. - */ -Expr FoldExplicitPadding(const Expr& expr, const IRModule& mod) { - return SimplifyExplicitPadding(mod).Simplify(expr); -} - -namespace transform { - -Pass FoldExplicitPadding() { - runtime::TypedPackedFunc pass_func = - [=](Function f, IRModule m, PassContext pc) { - return Downcast(FoldExplicitPadding(f, m)); - }; - return CreateFunctionPass(pass_func, 0, " FoldExplicitPadding", {"InferType"}); -} - -TVM_REGISTER_GLOBAL("relay._transform.FoldExplicitPadding").set_body_typed(FoldExplicitPadding); - -} // namespace transform - -} // namespace relay -} // namespace tvm diff --git a/src/relay/transforms/fold_scale_axis.cc b/src/relay/transforms/fold_scale_axis.cc deleted file mode 100644 index 69e100936839..000000000000 --- a/src/relay/transforms/fold_scale_axis.cc +++ /dev/null @@ -1,1179 +0,0 @@ -/* - * 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 fold_scale_axis.cc - * - * \brief Fold axis scaling into weights of - * conv/dense operators. - */ -#include -#include -#include -#include -#include - -#include "../backend/utils.h" -#include "../op/tensor/transform.h" -#include "pass_utils.h" -#include "pattern_utils.h" - -namespace tvm { -namespace relay { -/*! - * \brief namespace of fold scale axis - * - * Use namespace to reduce potential naming conflict. - */ - -namespace fold_scale_axis { - -using runtime::TypedPackedFunc; - -// FoldScaleAxis algorithm: -// -// The general idea is to transform Expr to tuple of -// (value, axes, scale), where the final result satisfies: -// -// result = value -// for i, k in enumerate(axes): -// k-th dimension of result *= i-th dimension of scale -// -// Then we can propagate this signal along and fold the scale if necessary. -// However, it is possible that certain scale may never be consumed -// if there is no dense/conv2d that follows multiplication. -// -// In order to make sure all the scale we sent out can be consumed eventually, -// we run a backward "preparation phase", which propagates the demand -// of the potential axes scaling back to its input. -// -// Forward folding process is done in two steps: -// - Prepare phase: backward propagation of demand. -// - Transform phase: forward transformation, -// -// Similarly, backward folding process is done in two steps: -// - Prepare phase: forward propagation of demand. -// - Transform phase: transformation by push down the axes scale signal to inputs. -// - -/*! - * \brief sorted array axis, can also be nullptr. - * - * nullptr means no scaling request can be done. - */ -using AxesSet = Array; - -class Message; - -/*! - * \brief Message propogated during the prepare phase. - */ -class MessageNode : public RelayNode { - public: - /*! \brief Axes for scaling */ - AxesSet axes; - /*! - * \brief Whether folding requires the scale to be positive constant. This is necessary if some - * operators (e.g. Relu) is present. - */ - bool require_positive; - - static constexpr const char* _type_key = "relay.pass.fold_scale_axis.Message"; - TVM_DECLARE_FINAL_OBJECT_INFO(MessageNode, RelayNode); -}; - -class Message : public ObjectRef { - public: - /*! - * \brief The constructor - * \param axes Axes for scaling - * \param require_positive If folding requires the scales to be positive - * values. - */ - Message(const AxesSet& axes, bool require_positive); - - TVM_DEFINE_OBJECT_REF_METHODS(Message, ObjectRef, MessageNode); -}; - -Message::Message(const AxesSet& axes, bool require_positive) { - auto n = make_object(); - n->axes = axes; - n->require_positive = require_positive; - data_ = std::move(n); -} - -/*! - * \brief Merge two axis set together by taking - * intersection. - * - * \note The axes in a AxesSet should be sorted. - * - * \param lhs The left axis. - * \param rhs The right axis. - * \return The result of the inersection. - */ -AxesSet Intersect(const AxesSet& lhs, const AxesSet& rhs) { - if (!lhs.defined()) return lhs; - if (!rhs.defined()) return rhs; - // This code relies on axes in a AxesSet to be sorted. - AxesSet ret; - size_t i = 0, j = 0; - while (i < lhs.size() && j < rhs.size()) { - if (lhs[i]->value < rhs[j]->value) { - ++i; - } else if (lhs[i]->value > rhs[j]->value) { - ++j; - } else { - ret.push_back(lhs[i]); - ++i; - ++j; - } - } - return ret; -} - -/*! - * \brief Merge two messages together by taking intersection. - * - * \param lhs The lhs message. - * \param rhs The rhs message. - * \return The result of intersection. - */ -Message Intersect(const Message& lhs, const Message& rhs) { - if (!lhs.defined()) return lhs; - if (!rhs.defined()) return rhs; - auto axes = Intersect(lhs->axes, rhs->axes); - return Message(axes, lhs->require_positive || rhs->require_positive); -} - -/*! - * \brief Preparation function for pass scale forward. - * \param call The call node. - * \param out_message Message from the output containing possible scaling on axes and whether - * positive scale is required. - * \return The message containing the result scaling on axes of the input. - */ -using FForwardPrep = - runtime::TypedPackedFunc(const Call& call, const Message& out_message)>; - -/*! \brief Axis scale tuple. */ -class ScaledExprNode : public TempExprNode { - public: - /*! \brief The value */ - Expr value; - /*! \brief The axes to scale, can be nullptr(means no-scaling) */ - AxesSet axes = NullValue(); - /*! \brief The scaling factor */ - Expr scale = NullValue(); - - Expr Realize() const final { - ICHECK(!axes.defined()) << "outstanding scale"; - return value; - } - - void VisitAttrs(AttrVisitor* v) { - v->Visit("value", &value); - v->Visit("axes", &axes); - v->Visit("scale", &scale); - } - - static constexpr const char* _type_key = "relay.fold_scale_axis.ScaledExpr"; - TVM_DECLARE_FINAL_OBJECT_INFO(ScaledExprNode, TempExprNode); -}; - -using FForwardRewrite = TypedPackedFunc& new_args, - const Message& message)>; - -//---------------------------------------------- -// Generic Visitors for FScaleAxisForward -//---------------------------------------------- -class ForwardPrep : private MixedModeVisitor { - public: - std::unordered_map Prepare(const Expr& body) { - this->Update(body, NullValue()); - this->VisitExpr(body); - // flist is added in the Post-DFS order - // which is a special case of topological order. - // We reversely traverse the list to invoke the lazy functions. - // This act like a backprop of valid scale axis messages - for (auto it = flist_.rbegin(); it != flist_.rend(); ++it) { - (*it)(); - } - // return the created message; - return std::move(message_); - } - - private: - // The invoke list - std::vector> flist_; - // The message on each node. - std::unordered_map message_; - // Update the message stored at node. - void Update(const Expr& node, const Message& message) { - // We run intersection of messages: - // - // %y = multiply(%x, %scale) - // %z1 = conv2d(%y, %w) - // %z2 = exp(%y) - // - // Consider the above code example, - // because %z2 will propagate null to %y, - // the AxesSet on %y is also null, - // and the forward folding won't be triggered. - const Object* key = node.get(); - if (message_.count(key)) { - message_[key] = Intersect(message_[key], message); - } else { - message_[key] = message; - } - } - - // We intended the following overrides on implementations from ExprVisitor. - using MixedModeVisitor::VisitExpr_; - - // Visitor pattern override. - void VisitExpr_(const TupleGetItemNode* op) final { MixedModeVisitor::VisitExpr_(op); } - - void VisitExpr_(const LetNode* op) final { - ExprVisitor::VisitExpr_(op); - // do pass through condition - // by assigning NullValue - // it means fuse signal cannot pass - // through into these subexpressions. - auto flazy = [this, op]() { - this->Update(op->value, NullValue()); - this->Update(op->body, NullValue()); - }; - flist_.push_back(flazy); - } - - void VisitExpr_(const FunctionNode* op) final { - ExprVisitor::VisitExpr_(op); - auto flazy = [this, op] { this->Update(op->body, NullValue()); }; - flist_.push_back(flazy); - } - - void VisitExpr_(const CallNode* call) final { - ExprVisitor::VisitExpr_(call); - // function to be lazily invoked - auto flazy = [this, call]() { - static const auto& fprep = Op::GetAttrMap("FScaleAxisForwardPrep"); - // find the message send to this node. - auto it = message_.find(call); - Message out_message; - if (it != message_.end()) { - out_message = it->second; - } else { - out_message = NullValue(); - } - // pass the message back to all the children it references. - auto f = fprep.get(call->op, nullptr); - if (f != nullptr) { - Array in_messages = f(GetRef(call), out_message); - ICHECK_EQ(in_messages.size(), call->args.size()); - for (size_t i = 0; i < call->args.size(); ++i) { - this->Update(call->args[i], in_messages[i]); - } - } else { - for (size_t i = 0; i < call->args.size(); ++i) { - this->Update(call->args[i], NullValue()); - } - } - }; - flist_.push_back(flazy); - } - - void VisitExpr_(const TupleNode* op) final { - ExprVisitor::VisitExpr_(op); - // do not support pass scale through tuple for now. - auto flazy = [this, op]() { - for (const Expr& field : op->fields) { - this->Update(field, NullValue()); - } - }; - flist_.push_back(flazy); - } - - void VisitExpr_(const IfNode* op) final { - ExprVisitor::VisitExpr_(op); - // do pass through condition - // by assigning NullValue - // it means fuse signal cannot pass - // through into these subexpressions. - auto flazy = [this, op]() { - this->Update(op->cond, NullValue()); - this->Update(op->true_branch, NullValue()); - this->Update(op->false_branch, NullValue()); - }; - flist_.push_back(flazy); - } -}; - -static bool IsIntInArray(const Array& axis, int v) { - for (size_t i = 0; i < axis.size(); i++) { - if (axis[i] == v) return true; - } - return false; -} - -static Expr ReshapeToMatchAxis(Expr scale, const Array& shape, - const Array& axis) { - Array arr; - for (size_t i = 0; i < shape.size(); i++) { - if (IsIntInArray(axis, i)) { - auto node = shape[i].as(); - if (!node) { - // if the shape is not a constant, use normal transform - return Expr(); - } - arr.push_back(node->value); - } else { - arr.push_back(1); - } - } - return MakeReshape(scale, std::move(arr)); -} - -// if only one axis, use expand dim. Else, use reshape -static Expr ReshapeOrExpandToMatchAxis(Expr scale, const Array& shape, - const Array& axis) { - if (axis.size() > 1) { - return ReshapeToMatchAxis(scale, shape, axis); - } else { - return ExpandBiasToMatchAxis(scale, shape.size(), axis); - } -} - -//---------------------------------------------- -// Per operator defs for FScaleAxisForward -//---------------------------------------------- - -// Intermediate operators -Array ReluForwardPrep(const Call& call, const Message& out_message) { - if (out_message.defined()) { - return {Message(out_message->axes, true)}; - } - return {out_message}; -} - -Expr ReluForwardRewrite(const Call& ref_call, const Array& new_args, const Message& message) { - const auto* input = new_args[0].as(); - if (input == nullptr) return Expr(nullptr); - // return transformed conv2d - auto rnode = make_object(); - rnode->value = Call(ref_call->op, {input->value}, ref_call->attrs, ref_call->type_args); - rnode->scale = input->scale; - rnode->axes = input->axes; - return Expr(rnode); -} - -RELAY_REGISTER_OP("nn.relu").set_attr("FScaleAxisForwardPrep", ReluForwardPrep); - -RELAY_REGISTER_OP("nn.relu").set_attr("FScaleAxisForwardRewrite", - ReluForwardRewrite); - -RELAY_REGISTER_OP("nn.leaky_relu").set_attr("FScaleAxisForwardPrep", ReluForwardPrep); - -RELAY_REGISTER_OP("nn.leaky_relu") - .set_attr("FScaleAxisForwardRewrite", ReluForwardRewrite); - -// AddSub -Array AddSubForwardPrep(const Call& call, const Message& out_message) { - const auto* tlhs = call->args[0]->type_as(); - const auto* trhs = call->args[1]->type_as(); - auto none = NullValue(); - if (out_message.defined()) { - if (MatchBroadcastToLeftAxes(tlhs, trhs, out_message->axes)) { - return {out_message, none}; - } else if (MatchBroadcastToLeftAxes(trhs, tlhs, out_message->axes)) { - return {none, out_message}; - } - } - return {none, none}; -} - -Expr AddSubForwardRewrite(const Call& ref_call, const Array& new_args, - const Message& message) { - const auto* slhs = new_args[0].as(); - const auto* srhs = new_args[1].as(); - if (!slhs && !srhs) return Expr(); - const auto* tlhs = ref_call->args[0]->type_as(); - const auto* trhs = ref_call->args[1]->type_as(); - auto rnode = make_object(); - - if (slhs != nullptr) { - ICHECK(srhs == nullptr); - ICHECK(MatchBroadcastToLeftAxes(tlhs, trhs, slhs->axes)); - Expr scale = ReshapeOrExpandToMatchAxis(slhs->scale, tlhs->shape, slhs->axes); - if (!scale.defined()) { - return Expr(); - } - Expr rhs = Divide(new_args[1], scale); - rnode->value = Call(ref_call->op, {slhs->value, rhs}, ref_call->attrs, ref_call->type_args); - rnode->scale = slhs->scale; - rnode->axes = slhs->axes; - } else { - ICHECK(srhs != nullptr); - ICHECK(MatchBroadcastToLeftAxes(trhs, tlhs, srhs->axes)); - Expr scale = ReshapeOrExpandToMatchAxis(srhs->scale, trhs->shape, srhs->axes); - if (!scale.defined()) { - return Expr(); - } - Expr lhs = Divide(new_args[0], scale); - rnode->value = Call(ref_call->op, {lhs, srhs->value}, ref_call->attrs, ref_call->type_args); - rnode->scale = srhs->scale; - rnode->axes = srhs->axes; - } - return Expr(rnode); -} - -RELAY_REGISTER_OP("add").set_attr("FScaleAxisForwardPrep", AddSubForwardPrep); - -RELAY_REGISTER_OP("add").set_attr("FScaleAxisForwardRewrite", - AddSubForwardRewrite); - -RELAY_REGISTER_OP("subtract").set_attr("FScaleAxisForwardPrep", AddSubForwardPrep); - -RELAY_REGISTER_OP("subtract") - .set_attr("FScaleAxisForwardRewrite", AddSubForwardRewrite); - -// Producer operators -// Multiply produces the scale-axis pair. -Expr MultiplyForwardRewrite(const Call& ref_call, const Array& new_args, - const Message& message) { - if (!message.defined()) return Expr(); - const auto& expected_out_axes = message->axes; - ICHECK(expected_out_axes.defined() && expected_out_axes.size()); - // TODO(tvm-team) allow same axes accumulation - // not as important because it is less common in nn. - const auto* slhs = new_args[0].as(); - const auto* srhs = new_args[1].as(); - ICHECK(!slhs && !srhs); - - const auto* tlhs = ref_call->args[0]->type_as(); - const auto* trhs = ref_call->args[1]->type_as(); - Expr lhs = new_args[0]; - Expr rhs = new_args[1]; - auto rnode = make_object(); - - if (MatchBroadcastToLeftAxes(tlhs, trhs, expected_out_axes, &rhs) && - (!message->require_positive || IsAllPositiveConstant(rhs))) { - rnode->value = lhs; - rnode->scale = rhs; - rnode->axes = expected_out_axes; - return Expr(rnode); - } else if (MatchBroadcastToLeftAxes(trhs, tlhs, expected_out_axes, &lhs) && - (!message->require_positive || IsAllPositiveConstant(lhs))) { - rnode->value = rhs; - rnode->scale = lhs; - rnode->axes = expected_out_axes; - return Expr(rnode); - } else { - return Expr(); - } -} - -RELAY_REGISTER_OP("multiply") - .set_attr("FScaleAxisForwardRewrite", MultiplyForwardRewrite); - -// Consumer operators -// Conv send out requirement of axis folding. -template -Array ConvForwardPrep(const Call& call, const ATTRS* param, const Message& out_message) { - // TODO(tvm-team) support general data layout - // by transforming weight - ICHECK(param != nullptr); - Layout data_layout(param->data_layout); - Layout kernel_layout(param->kernel_layout); - int c_big_axis = data_layout.IndexOf(LayoutAxis::Get('C')); - int c_small_axis = data_layout.IndexOf(LayoutAxis::Get('c')); - - ICHECK_GE(c_big_axis, 0); - Message none = NullValue(); - // For now, we only support simple pattern (no folded weight/data) - // More general layout can be supported under the current framework. - // By using a unified layout transformation. - // We only need to change the Prep and Mutate function. - // - // only handle depthwise or full conv2d. - // TODO(tvm-team) handle grouped conv by reshape + bcast - bool is_depthwise_conv = IsDepthwiseConv(call, param, kernel_layout); - if (param->groups == 1 || is_depthwise_conv) { - auto ko_small_axis = kernel_layout.IndexOf(LayoutAxis::Get('o')); - auto ki_small_axis = kernel_layout.IndexOf(LayoutAxis::Get('i')); - if ((ko_small_axis < 0 && ki_small_axis < 0 && c_small_axis < 0) || // simple layout - (ko_small_axis >= 0 && ki_small_axis >= 0 && c_small_axis >= 0)) { // blocked layout - Array arr{c_big_axis}; - if (c_small_axis >= 0) { - arr.push_back(c_small_axis); - } - return {Message(arr, false), none}; - } - } - return {none, none}; -} - -// Conv2D consumes the scale axis during transformation. -template -Expr ConvForwardRewrite(const Call& ref_call, const ATTRS* param, const Array& new_args, - const Message& message) { - // if data do not have scale, normal transform path. - const auto* sdata = new_args[0].as(); - const auto* sweight = new_args[1].as(); - if (sdata == nullptr) return Expr(); - if (sweight != nullptr) return Expr(); - ICHECK(param != nullptr); - Layout data_layout(param->data_layout); - Layout kernel_layout(param->kernel_layout); - int c_big_axis = data_layout.IndexOf(LayoutAxis::Get('C')); - ICHECK_GE(c_big_axis, 0); - int small_ko_axis = kernel_layout.IndexOf(LayoutAxis::Get('o')); - int small_ki_axis = kernel_layout.IndexOf(LayoutAxis::Get('i')); - int big_ki_axis = kernel_layout.IndexOf(LayoutAxis::Get('I')); - int big_ko_axis = kernel_layout.IndexOf(LayoutAxis::Get('O')); - - bool is_simple = (small_ko_axis < 0 && small_ki_axis < 0 && big_ki_axis >= 0); - bool is_blocking = (small_ko_axis >= 0 && small_ki_axis >= 0 && big_ki_axis >= 0); - ICHECK(is_simple || is_blocking); - - // Check it must be depthwise or full conv2d. - bool is_depthwise_conv = IsDepthwiseConv(ref_call, param, kernel_layout); - ICHECK(param->groups == 1 || is_depthwise_conv); - - Expr weight = new_args[1]; - - // match the ic_axis - if (is_depthwise_conv) { - if (is_simple) { - Expr scale = ExpandBiasToMatchAxis(sdata->scale, kernel_layout.ndim(), {big_ko_axis}); - weight = Multiply(weight, scale); - } else { - weight = Multiply(weight, - ReshapeToMatchAxis(sdata->scale, weight->type_as()->shape, - {big_ko_axis, small_ko_axis})); - if (!weight.defined()) return Expr(); - } - - } else { - if (is_simple) { - Expr scale = ExpandBiasToMatchAxis(sdata->scale, kernel_layout.ndim(), {big_ki_axis}); - weight = Multiply(weight, scale); - } else { - weight = Multiply(weight, - ReshapeToMatchAxis(sdata->scale, weight->type_as()->shape, - {big_ki_axis, small_ki_axis})); - if (!weight.defined()) return Expr(); - } - } - // return transformed conv - return Call(ref_call->op, {sdata->value, weight}, ref_call->attrs, ref_call->type_args); -} - -Array PreConvForwardPrep(const Call& call, const Message& out_message) { - if (backend::IsOp(call.as(), "nn.conv2d")) { - const auto* param = call->attrs.as(); - ICHECK(param != nullptr); - return ConvForwardPrep(call, param, out_message); - } - const auto* param = call->attrs.as(); - ICHECK(param != nullptr); - return ConvForwardPrep(call, param, out_message); -} - -Expr PreConvForwardRewrite(const Call& ref_call, const Array& new_args, - const Message& message) { - if (backend::IsOp(ref_call.as(), "nn.conv2d")) { - const auto* param = ref_call->attrs.as(); - ICHECK(param != nullptr); - return ConvForwardRewrite(ref_call, param, new_args, message); - } - const auto* param = ref_call->attrs.as(); - ICHECK(param != nullptr); - return ConvForwardRewrite(ref_call, param, new_args, message); -} - -RELAY_REGISTER_OP("nn.conv2d").set_attr("FScaleAxisForwardPrep", PreConvForwardPrep); - -RELAY_REGISTER_OP("nn.conv2d") - .set_attr("FScaleAxisForwardRewrite", PreConvForwardRewrite); - -RELAY_REGISTER_OP("nn.conv3d").set_attr("FScaleAxisForwardPrep", PreConvForwardPrep); - -RELAY_REGISTER_OP("nn.conv3d") - .set_attr("FScaleAxisForwardRewrite", PreConvForwardRewrite); - -// Dense send out requirement of axis folding. -Array DenseForwardPrep(const Call& call, const Message& out_message) { - return {Message({1}, false), NullValue()}; -} - -// Dense consumes the scale axis during transformation. -Expr DenseForwardRewrite(const Call& ref_call, const Array& new_args, - const Message& message) { - const auto* sdata = new_args[0].as(); - const auto* sweight = new_args[1].as(); - if (sdata == nullptr) return Expr(); - if (sweight != nullptr) return Expr(); - - Expr weight = Multiply(new_args[1], sdata->scale); - return Call(ref_call->op, {sdata->value, weight}, ref_call->attrs, ref_call->type_args); -} - -RELAY_REGISTER_OP("nn.dense").set_attr("FScaleAxisForwardPrep", DenseForwardPrep); - -RELAY_REGISTER_OP("nn.dense") - .set_attr("FScaleAxisForwardRewrite", DenseForwardRewrite); - -Expr ForwardFoldScaleAxis(const Expr& data) { - auto message = ForwardPrep().Prepare(data); - for (const auto& m : message) { - if (m.second.defined()) { - // run optimization - auto fcontext = [&](const Call& call) -> ObjectRef { - auto it = message.find(call.get()); - if (it != message.end()) { - return it->second; - } else { - return ObjectRef(nullptr); - } - }; - return ForwardRewrite(data, "FScaleAxisForwardRewrite", fcontext); - } - } - // no messages - no optimization - return data; -} - -//---------------------------------------- -// Implement backward transformations. -//---------------------------------------- -class BackwardTransformer; - -/*! - * \brief Preparation function for pass scale backward. - * \param call The call node. - * \param in_messages Messages from the input containing allowed input scaling and whether - * positive scale is required. - * \return Message containing the result scaling on axes of the input. - */ -using FBackwardPrep = TypedPackedFunc& in_messages)>; - -using FBackwardTransform = - TypedPackedFunc; - -//---------------------------------------------- -// Generic Visitors for FScaleAxisBackward -//---------------------------------------------- - -class BackwardPrep : private MixedModeVisitor { - public: - // The message on each node. - std::unordered_map Prepare(const Expr& body) { - ref_counter_ = GetExprRefCount(body); - this->VisitExpr(body); - return std::move(message_); - } - - private: - // The message on each node. - std::unordered_map message_; - // reference counter of an internal expr - std::unordered_map ref_counter_; - // Visit the expression. - void VisitExpr_(const CallNode* call) { - ExprVisitor::VisitExpr_(call); - static const auto& fprep = Op::GetAttrMap("FScaleAxisBackwardPrep"); - auto f = fprep.get(call->op, nullptr); - if (f == nullptr) return; - auto rit = ref_counter_.find(call); - ICHECK(rit != ref_counter_.end()); - // We only allow propagation of scale backward - // if the expression is only referred by a single parent. - if (rit->second != 1) return; - Array in_messages = GetInMessages(call); - Message out_message = f(GetRef(call), in_messages); - if (out_message.defined()) { - message_[call] = out_message; - } - } - - Array GetInMessages(const CallNode* call) { - Array in_messages; - for (Expr arg : call->args) { - auto it = message_.find(arg.get()); - if (it != message_.end()) { - in_messages.push_back(it->second); - } else { - in_messages.push_back(NullValue()); - } - } - return in_messages; - } -}; - -/* - * Hybrid apporach is used with the transformation - * itself is recursive but the traversal is non-recursive - */ -class BackwardTransformerNode : public Object, private MixedModeMutator { - public: - using MixedModeMutator::Mutate; - // Run forward transform. - Expr Fold(Expr expr) { - message_ = BackwardPrep().Prepare(expr); - for (const auto& m : message_) { - if (m.second.defined()) { - // run optimization - return this->Mutate(expr); - } - } - // no messages - no optimization - return expr; - } - - /*! - * \brief Transform the expr to consider the scaling. - */ - Expr Transform(const Expr& expr, Message message, Expr scale); - /*! - * \brief Get the message propogated to the expr. - * \param expr The expresison. - * \return The message containing the expected axes and whether positive scale is required. - */ - Message GetMessage(const Expr& expr) const { - auto it = message_.find(expr.get()); - if (it != message_.end()) return it->second; - return NullValue(); - } - - // solver is not serializable. - void VisitAttrs(tvm::AttrVisitor* v) {} - - static constexpr const char* _type_key = "relay.fold_scale_axis.FBackwardTransformer"; - TVM_DECLARE_FINAL_OBJECT_INFO(BackwardTransformerNode, Object); - - private: - // Valid axes on each node. - std::unordered_map message_; - // Override mutation of call. - Expr Rewrite_(const CallNode* call_node, const Expr& post) final { - return Transform(GetRef(call_node), NullValue(), NullValue()); - } - - public: - Expr NormalCallTransform(const CallNode* call_node) { return ExprMutator::VisitExpr_(call_node); } -}; - -class BackwardTransformer : public ObjectRef { - public: - BackwardTransformer() {} - explicit BackwardTransformer(::tvm::ObjectPtr<::tvm::Object> n) : ObjectRef(n) {} - BackwardTransformerNode* operator->() const { - return static_cast(get_mutable()); - } - using ContainerType = BackwardTransformerNode; -}; - -/*! - * \brief Transform the expr to consider the scaling. - * - * \param expr The input expression. - * \param message The axes to scale. - * \param scale The scale applied to the axes. - * \return The result of transformation. - */ -Expr BackwardTransformerNode::Transform(const Expr& expr, Message message, Expr scale) { - if (const CallNode* call_node = expr.as()) { - static const auto& ftransform = - Op::GetAttrMap("FScaleAxisBackwardTransform"); - auto f = ftransform.get(call_node->op, nullptr); - const Call call = GetRef(call_node); - // ignore if there is a message - if (!message.defined()) { - const auto it = memo_.find(call); - if (it != memo_.end()) { - return it->second; - } - } - Expr new_expr = NullValue(); - if (f != nullptr) { - new_expr = f(call, message, scale, GetRef(this)); - } else { - ICHECK(!message.defined()) << "outstanding scale"; - new_expr = NormalCallTransform(call.operator->()); - } - memo_[call] = new_expr; - return new_expr; - } else { - ICHECK(!message.defined()) << "outstanding scale"; - return this->Mutate(expr); - } -} - -//---------------------------------------------- -// Per operator defs for FScaleAxisForward -//---------------------------------------------- - -// Intermediate operators -Message ReluBackwardPrep(const Call& call, const Array& in_messages) { - if (in_messages[0].defined()) { - return Message(in_messages[0]->axes, true); - } - return in_messages[0]; -} - -Expr ReluBackwardTransform(const Call& call, const Message& message, const Expr& scale, - const BackwardTransformer& transformer) { - if (!message.defined()) { - return transformer->NormalCallTransform(call.operator->()); - } - Expr input = transformer->Transform(call->args[0], message, scale); - return Call(call->op, {input}, call->attrs, call->type_args); -} - -RELAY_REGISTER_OP("nn.relu").set_attr("FScaleAxisBackwardPrep", ReluBackwardPrep); - -RELAY_REGISTER_OP("nn.relu").set_attr("FScaleAxisBackwardTransform", - ReluBackwardTransform); - -RELAY_REGISTER_OP("nn.leaky_relu") - .set_attr("FScaleAxisBackwardPrep", ReluBackwardPrep); - -RELAY_REGISTER_OP("nn.leaky_relu") - .set_attr("FScaleAxisBackwardTransform", ReluBackwardTransform); - -// AddSub -Message AddSubBackwardPrep(const Call& call, const Array& in_messages) { - const auto* tlhs = call->args[0]->type_as(); - const auto* trhs = call->args[1]->type_as(); - StructuralEqual equal; - if (in_messages[0].defined() && MatchBroadcastToLeftAxes(tlhs, trhs, in_messages[0]->axes)) { - return in_messages[0]; - } else if (in_messages[1].defined() && - MatchBroadcastToLeftAxes(trhs, tlhs, in_messages[1]->axes)) { - return in_messages[1]; - } else if (in_messages[0].defined() && in_messages[1].defined() && - equal(in_messages[0]->axes, in_messages[1]->axes) && equal(tlhs->shape, trhs->shape)) { - // add of two elements. - return in_messages[0]; - } else { - auto res = NullValue(); - return res; - } -} - -Expr AddSubBackwardTransform(const Call& call, const Message& message, const Expr& scale, - const BackwardTransformer& transformer) { - const auto* tlhs = call->args[0]->type_as(); - const auto* trhs = call->args[1]->type_as(); - if (!message.defined()) { - return transformer->NormalCallTransform(call.operator->()); - } - - Message lhs_message = transformer->GetMessage(call->args[0]); - Message rhs_message = transformer->GetMessage(call->args[1]); - StructuralEqual equal; - - if (lhs_message.defined() && rhs_message.defined()) { - ICHECK(equal(lhs_message->axes, rhs_message->axes)); - ICHECK(equal(message->axes, lhs_message->axes)); - Expr lhs = transformer->Transform(call->args[0], message, scale); - Expr rhs = transformer->Transform(call->args[1], message, scale); - return Call(call->op, {lhs, rhs}, call->attrs, call->type_args); - } else if (lhs_message.defined()) { - ICHECK(equal(message->axes, lhs_message->axes)); - Expr lhs = transformer->Transform(call->args[0], message, scale); - Expr rhs = transformer->Transform(call->args[1], NullValue(), NullValue()); - Expr rhs_scale = ReshapeOrExpandToMatchAxis(scale, tlhs->shape, message->axes); - if (!rhs_scale.defined()) { - return transformer->NormalCallTransform(call.operator->()); - } - rhs = Multiply(rhs, rhs_scale); - return Call(call->op, {lhs, rhs}, call->attrs, call->type_args); - } else if (rhs_message.defined()) { - ICHECK(equal(message->axes, rhs_message->axes)); - Expr lhs = transformer->Transform(call->args[0], NullValue(), NullValue()); - Expr rhs = transformer->Transform(call->args[1], message, scale); - Expr lhs_scale = ReshapeOrExpandToMatchAxis(scale, trhs->shape, message->axes); - if (!lhs_scale.defined()) { - return transformer->NormalCallTransform(call.operator->()); - } - lhs = Multiply(lhs, lhs_scale); - return Call(call->op, {lhs, rhs}, call->attrs, call->type_args); - } else { - LOG(FATAL) << "outstanding scale"; - } -} - -RELAY_REGISTER_OP("add").set_attr("FScaleAxisBackwardPrep", AddSubBackwardPrep); - -RELAY_REGISTER_OP("add").set_attr("FScaleAxisBackwardTransform", - AddSubBackwardTransform); - -RELAY_REGISTER_OP("subtract").set_attr("FScaleAxisBackwardPrep", AddSubBackwardPrep); - -RELAY_REGISTER_OP("subtract") - .set_attr("FScaleAxisBackwardTransform", AddSubBackwardTransform); - -// Producer operators -// Multiply produces the scale-axis pair. -Expr MultiplyBackwardTransform(const Call& call, const Message& message, const Expr& scale, - const BackwardTransformer& transformer) { - ICHECK(!message.defined()) << "outstanding scale"; - const auto* tlhs = call->args[0]->type_as(); - const auto* trhs = call->args[1]->type_as(); - Message lhs_message = transformer->GetMessage(call->args[0]); - Message rhs_message = transformer->GetMessage(call->args[1]); - if (lhs_message.defined()) { - ICHECK(lhs_message->axes.defined() && lhs_message->axes.size()); - // NOTE we won't recursively call mutating on scale part. - // since there won't be scale chance within scale part. - Expr rhs = call->args[1]; - if (MatchBroadcastToLeftAxes(tlhs, trhs, lhs_message->axes, &rhs) && - (!lhs_message->require_positive || IsAllPositiveConstant(rhs))) { - return transformer->Transform(call->args[0], lhs_message, rhs); - } - } else if (rhs_message.defined()) { - ICHECK(rhs_message->axes.defined() && rhs_message->axes.size()); - Expr lhs = call->args[0]; - if (MatchBroadcastToLeftAxes(trhs, tlhs, rhs_message->axes, &lhs) && - (!rhs_message->require_positive || IsAllPositiveConstant(lhs))) { - return transformer->Transform(call->args[1], rhs_message, lhs); - } - } - return transformer->NormalCallTransform(call.operator->()); -} - -RELAY_REGISTER_OP("multiply") - .set_attr("FScaleAxisBackwardTransform", MultiplyBackwardTransform); - -// Consumer operators -// Conv send out requirement of axis folding. -template -Message ConvBackwardPrep(const Call& call, const ATTRS* param, const Array& in_messages) { - ICHECK(param != nullptr); - Layout kernel_layout(param->kernel_layout); - Layout out_layout(param->out_layout == "" ? param->data_layout : param->out_layout); - int c_big_axis = out_layout.IndexOf(LayoutAxis::Get('C')); - int c_small_axis = out_layout.IndexOf(LayoutAxis::Get('c')); - - ICHECK_GE(c_big_axis, 0); - // For now, we only support simple pattern (no folded weight/data) - // More general layout can be supported under the current framework. - // By using a unified layout transformation. - // We only need to change the Prep and Mutate function. - // - // only handle depthwise or full conv. - // TODO(tvm-team) handle grouped conv by reshape + bcast - bool is_depthwise_conv = IsDepthwiseConv(call, param, kernel_layout); - if (param->groups == 1 || is_depthwise_conv) { - auto ko_small_axis = kernel_layout.IndexOf(LayoutAxis::Get('o')); - auto ki_small_axis = kernel_layout.IndexOf(LayoutAxis::Get('i')); - if ((ko_small_axis < 0 && ki_small_axis < 0 && c_small_axis < 0) || // simple layout - (ko_small_axis >= 0 && ki_small_axis >= 0 && c_small_axis >= 0)) { // blocked layout - Array arr{c_big_axis}; - if (c_small_axis >= 0) { - arr.push_back(c_small_axis); - } - return Message(arr, false); - } - } - return NullValue(); -} - -// Conv consumes the scale axis during transformation. -template -Expr ConvBackwardTransform(const Call& call, const ATTRS* param, const Message& message, - const Expr& scale, const BackwardTransformer& transformer) { - if (!message.defined()) { - return transformer->NormalCallTransform(call.operator->()); - } - ICHECK(param != nullptr); - Layout kernel_layout(param->kernel_layout); - Layout out_layout(param->out_layout == "" ? param->data_layout : param->out_layout); - int c_big_axis = out_layout.IndexOf(LayoutAxis::Get('C')); - ICHECK_GE(c_big_axis, 0); - // For now, we only support simple pattern (no folded weight/data) - // TODO(tvm-team) support general data layout - int small_ko_axis = kernel_layout.IndexOf(LayoutAxis::Get('o')); - int small_ki_axis = kernel_layout.IndexOf(LayoutAxis::Get('i')); - int big_ki_axis = kernel_layout.IndexOf(LayoutAxis::Get('I')); - int big_ko_axis = kernel_layout.IndexOf(LayoutAxis::Get('O')); - // Check it must be depthwise or full conv. - bool is_depthwise_conv = IsDepthwiseConv(call, param, kernel_layout); - ICHECK(param->groups == 1 || is_depthwise_conv); - bool is_simple = (small_ko_axis < 0 && small_ki_axis < 0 && big_ki_axis >= 0); - bool is_blocking = (small_ko_axis >= 0 && small_ki_axis >= 0 && big_ki_axis >= 0); - ICHECK(is_simple || is_blocking); - - Expr data = transformer->Transform(call->args[0], NullValue(), NullValue()); - Expr weight = transformer->Transform(call->args[1], NullValue(), NullValue()); - // scale on input for deptwise. - Expr wscale; - if (is_simple) { - wscale = ExpandBiasToMatchAxis(scale, kernel_layout.ndim(), {big_ko_axis}); - } else { - wscale = ReshapeToMatchAxis(scale, weight->type_as()->shape, - {big_ko_axis, small_ko_axis}); - if (!wscale.defined()) { - return transformer->NormalCallTransform(call.operator->()); - } - } - weight = Multiply(weight, wscale); - return Call(call->op, {data, weight}, call->attrs, call->type_args); -} - -Message PreConvBackwardPrep(const Call& call, const Array& in_messages) { - if (backend::IsOp(call.as(), "nn.conv2d")) { - const auto* param = call->attrs.as(); - ICHECK(param != nullptr); - return ConvBackwardPrep(call, param, in_messages); - } - const auto* param = call->attrs.as(); - ICHECK(param != nullptr); - return ConvBackwardPrep(call, param, in_messages); -} - -Expr PreConvBackwardTransform(const Call& call, const Message& message, const Expr& scale, - const BackwardTransformer& transformer) { - if (backend::IsOp(call.as(), "nn.conv2d")) { - const auto* param = call->attrs.as(); - ICHECK(param != nullptr); - return ConvBackwardTransform(call, param, message, scale, transformer); - } - const auto* param = call->attrs.as(); - ICHECK(param != nullptr); - return ConvBackwardTransform(call, param, message, scale, transformer); -} - -RELAY_REGISTER_OP("nn.conv2d") - .set_attr("FScaleAxisBackwardPrep", PreConvBackwardPrep); - -RELAY_REGISTER_OP("nn.conv2d") - .set_attr("FScaleAxisBackwardTransform", PreConvBackwardTransform); - -RELAY_REGISTER_OP("nn.conv3d") - .set_attr("FScaleAxisBackwardPrep", PreConvBackwardPrep); - -RELAY_REGISTER_OP("nn.conv3d") - .set_attr("FScaleAxisBackwardTransform", PreConvBackwardTransform); - -Message BiasAddBackwardPrep(const Call& call, const Array& in_messages) { - const BiasAddAttrs* attrs = call->attrs.as(); - ICHECK(attrs); - if (in_messages[0].defined() && in_messages[0]->axes.size() == 1 && - attrs->axis == static_cast(in_messages[0]->axes[0]->value)) { - return in_messages[0]; - } else { - return NullValue(); - } -} - -Expr BiasAddBackwardTransform(const Call& call, const Message& message, const Expr& scale, - const BackwardTransformer& transformer) { - if (!message.defined()) { - return transformer->NormalCallTransform(call.operator->()); - } - Message lhs_message = transformer->GetMessage(call->args[0]); - Message rhs_message = transformer->GetMessage(call->args[1]); - StructuralEqual equal; - - if (lhs_message.defined()) { - ICHECK(equal(message->axes, lhs_message->axes)); - Expr lhs = transformer->Transform(call->args[0], message, scale); - Expr rhs = transformer->Transform(call->args[1], NullValue(), NullValue()); - rhs = Multiply(rhs, scale); - return Call(call->op, {lhs, rhs}, call->attrs, call->type_args); - } else { - LOG(FATAL) << "outstanding scale"; - } -} - -RELAY_REGISTER_OP("nn.bias_add") - .set_attr("FScaleAxisBackwardPrep", BiasAddBackwardPrep); - -RELAY_REGISTER_OP("nn.bias_add") - .set_attr("FScaleAxisBackwardTransform", BiasAddBackwardTransform); - -// Dense send out requirement of axis folding. -Message DenseBackwardPrep(const Call& call, const Array& in_messages) { - return Message({1}, false); -} - -// Dense consumes the sacle axis during trasformation. -Expr DenseBackwardTransform(const Call& call, const Message& message, const Expr& scale, - const BackwardTransformer& transformer) { - if (!message.defined()) { - return transformer->NormalCallTransform(call.operator->()); - } - Expr data = transformer->Transform(call->args[0], NullValue(), NullValue()); - Expr weight = transformer->Transform(call->args[1], NullValue(), NullValue()); - Expr wscale = ExpandBiasToMatchAxis(scale, 2, {0}); - weight = Multiply(weight, wscale); - return Call(call->op, {data, weight}, call->attrs, call->type_args); -} - -RELAY_REGISTER_OP("nn.dense").set_attr("FScaleAxisBackwardPrep", DenseBackwardPrep); - -RELAY_REGISTER_OP("nn.dense") - .set_attr("FScaleAxisBackwardTransform", DenseBackwardTransform); - -Expr BackwardFoldScaleAxis(const Expr& data) { - return make_object()->Fold(data); -} - -} // namespace fold_scale_axis - -namespace transform { - -Pass ForwardFoldScaleAxis() { - runtime::TypedPackedFunc pass_func = - [=](Function f, IRModule m, PassContext pc) { - return Downcast(relay::fold_scale_axis::ForwardFoldScaleAxis(f)); - }; - return CreateFunctionPass(pass_func, 3, "ForwardFoldScaleAxis", {"InferType"}); -} - -TVM_REGISTER_GLOBAL("relay._transform.ForwardFoldScaleAxis").set_body_typed(ForwardFoldScaleAxis); - -Pass BackwardFoldScaleAxis() { - runtime::TypedPackedFunc pass_func = - [=](Function f, IRModule m, PassContext pc) { - return Downcast(relay::fold_scale_axis::BackwardFoldScaleAxis(f)); - }; - return CreateFunctionPass(pass_func, 3, "BackwardFoldScaleAxis", {"InferType"}); -} - -TVM_REGISTER_GLOBAL("relay._transform.BackwardFoldScaleAxis").set_body_typed(BackwardFoldScaleAxis); - -Pass FoldScaleAxis() { - // FoldScaleAxis pass contains the following three passes. Therefore, we can - // register it as a sequential pass. - Pass pass = Sequential({BackwardFoldScaleAxis(), ForwardFoldScaleAxis(), FoldConstant()}, - "FoldScaleAxis"); - return pass; -} - -TVM_REGISTER_GLOBAL("relay._transform.FoldScaleAxis").set_body_typed(FoldScaleAxis); - -} // namespace transform - -} // namespace relay -} // namespace tvm diff --git a/src/relay/transforms/forward_rewrite.cc b/src/relay/transforms/forward_rewrite.cc deleted file mode 100644 index 857e0e8c91e0..000000000000 --- a/src/relay/transforms/forward_rewrite.cc +++ /dev/null @@ -1,187 +0,0 @@ -/* - * 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 forward_rewrite.cc - * \brief Apply rewriting rules in a forward fashion. - */ -#include -#include -#include -#include - -#include "pass_utils.h" - -namespace tvm { -namespace relay { - -// Realizer class that realizes the expression -// Note that we can take benefit of its internal memo -// so that calling realize repeatively won't hurt perf. -class TempRealizer : private MixedModeMutator { - public: - Expr Realize(Expr expr) { return Mutate(expr); } - - private: - Expr DispatchVisitExpr(const Expr& expr) final { - Expr res; - if (const auto* temp = expr.as()) { - res = temp->Realize(); - } else { - res = MixedModeMutator::DispatchVisitExpr(expr); - } - return res; - } -}; - -class ForwardRewriter : private MixedModeMutator { - public: - ForwardRewriter(const OpAttrMap* rewrite_map, - std::function fcontext, - std::function fmulti_ref_trigger) - : rewrite_map_(rewrite_map), fcontext_(fcontext), fmulti_ref_trigger_(fmulti_ref_trigger) {} - - ForwardRewriter(const FForwardRewrite* rewrite_func, - std::function fcontext, - std::function fmulti_ref_trigger) - : rewrite_func_(rewrite_func), fcontext_(fcontext), fmulti_ref_trigger_(fmulti_ref_trigger) {} - - // Transform expression. - Expr Rewrite(const Expr& expr) { - if (fmulti_ref_trigger_ != nullptr) { - ref_counter_ = GetExprRefCount(expr); - } - return realizer_.Realize(this->VisitExpr(expr)); - } - - private: - // The rewrite rule. - const OpAttrMap* rewrite_map_{nullptr}; - const FForwardRewrite* rewrite_func_{nullptr}; - // The context.const - std::function fcontext_{nullptr}; - // The multiple reference trigger - std::function fmulti_ref_trigger_{nullptr}; - // Internal ref counter - std::unordered_map ref_counter_; - // internal realizer - TempRealizer realizer_; - - // Visit and allow non-realized version. - Expr GetTempExpr(const Expr& expr, const Expr& post) { - if (fmulti_ref_trigger_ != nullptr) { - Expr ret = post; - auto it = ref_counter_.find(expr.get()); - ICHECK(it != ref_counter_.end()); - if (it->second > 1) { - ret = fmulti_ref_trigger_(ret); - } - return ret; - } else { - return post; - } - } - - // Automatic fold TupleGetItem. - Expr Rewrite_(const TupleGetItemNode* op, const Expr& post) final { - Expr tuple = this->GetTempExpr(op->tuple, post.as()->tuple); - if (const auto* ptuple = tuple.as()) { - return ptuple->fields[op->index]; - } else { - if (tuple.same_as(op->tuple)) { - return GetRef(op); - } else { - return TupleGetItem(tuple, op->index); - } - } - } - - Expr Rewrite_(const TupleNode* tuple_node, const Expr& post) final { - tvm::Array fields; - fields.reserve(tuple_node->fields.size()); - - const auto* post_tuple_node = post.as(); - for (size_t i = 0; i < tuple_node->fields.size(); ++i) { - fields.push_back(this->GetTempExpr(tuple_node->fields[i], post_tuple_node->fields[i])); - } - - return WithFields(GetRef(tuple_node), fields); - } - - Expr Rewrite_(const CallNode* call_node, const Expr& post) final { - const Call& ref_call = GetRef(call_node); - PackedFunc frewrite; - if (rewrite_func_) { - frewrite = *rewrite_func_; - } else { - ICHECK(rewrite_map_); - frewrite = rewrite_map_->get(call_node->op, nullptr); - } - const auto* post_node = post.as(); - auto new_op = post_node->op; - if (new_op->IsInstance()) { - new_op = realizer_.Realize(new_op); - } - bool unchanged = call_node->op.same_as(new_op); - - Array call_args; - for (size_t i = 0; i < call_node->args.size(); ++i) { - Expr new_arg = this->GetTempExpr(call_node->args[i], post_node->args[i]); - if (frewrite == nullptr) { - new_arg = realizer_.Realize(new_arg); - } - unchanged &= new_arg.same_as(call_node->args[i]); - call_args.push_back(new_arg); - } - // try to rewrite. - if (frewrite != nullptr) { - Expr res = frewrite(ref_call, call_args, - fcontext_ != nullptr ? fcontext_(ref_call) : ObjectRef(nullptr)); - if (res.defined()) return res; - // abort, use old rule - for (size_t i = 0; i < call_args.size(); ++i) { - Expr arg = call_args[i]; - Expr new_arg = realizer_.Realize(arg); - if (!arg.same_as(new_arg)) { - call_args.Set(i, new_arg); - unchanged = false; - } - } - } - if (unchanged) return ref_call; - return Call(new_op, call_args, call_node->attrs, call_node->type_args, call_node->span); - } -}; - -Expr ForwardRewrite(const Expr& expr, const String& rewrite_map_name, - std::function fcontext, - std::function fmulti_ref_trigger) { - auto rewrite_map = Op::GetAttrMap(rewrite_map_name); - return ForwardRewriter(&rewrite_map, fcontext, fmulti_ref_trigger).Rewrite(expr); -} - -Expr ForwardRewrite(const Expr& expr, const FForwardRewrite& rewrite_func, - std::function fcontext, - std::function fmulti_ref_trigger) { - return ForwardRewriter(&rewrite_func, fcontext, fmulti_ref_trigger).Rewrite(expr); -} - -} // namespace relay -} // namespace tvm diff --git a/src/relay/transforms/fuse_ops.cc b/src/relay/transforms/fuse_ops.cc deleted file mode 100644 index ee005aa17052..000000000000 --- a/src/relay/transforms/fuse_ops.cc +++ /dev/null @@ -1,590 +0,0 @@ -/* - * 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 src/relay/transforms/fuse_ops.cc - * - * \brief This is a backend-aware optimization pass. - * Fuse necessary ops into a single one. - */ -#include -#include -#include -#include -#include -#include - -#include "../../support/arena.h" -#include "../analysis/graph_partitioner.h" -#include "../op/annotation/annotation.h" -#include "./pass_utils.h" -#include "./pattern_utils.h" - -namespace tvm { -namespace relay { - -/* - Note on Fusing algorithm: - - The main challenge of general fusor is to handle possible diamond shape branches, - in the following graph, conv2d can be fused to elemwise add. - - conv2d - / | \ - / | \ - op op op - \ | / - \ | / - elemwise add - | - - However, at the point of conv2d we do not necessarily know that all the future paths - will merge at the elemwise add. The fusion algorithm applies post-dominator analysis. - - The immediate post-dominator of a node defined by the closest node where all the future path goes - into. In the above case, the elemwise add is the post-dominator of conv2d. The general algorithm - is as follows: - - - Construct a DAG of dataflow graph for dominator analysis - - Construct a post-dominator tree which gives immediate post dominator of each node. - - Run fusion algorithm with the given post-dominator information. - - Note that, because we run analysis on a DAG, we use a single pass post-dominator - tree construction algorithm via LCA, which is simpler than the full version that handles cycles. - - The fusion algorithm traverses from each node and checks if it can be fused to its - immediate post dominator. It has to check the following things: - - - CheckPath: check all the path between a node and its immediate post-dominator - satisfies the fuse condition. - - Note that these intermediate node can already be fused with another nodes, the algorithm - will still run correctly. - - CommitFuse: mark all the nodes between source and post-dominator as the same group. - - We use an Union-Find data structure to manage the groups. -*/ -using support::LinkedList; -using support::LinkNode; - -constexpr uint32_t kMaxFusedOps = 256; - -static const Op& stop_fusion_op = Op::Get("annotation.stop_fusion"); - -TVM_REGISTER_PASS_CONFIG_OPTION("relay.FuseOps.max_depth", Integer); -TVM_REGISTER_PASS_CONFIG_OPTION("relay.FuseOps.link_params", Bool); - -// Creator of post dominator tree of the dataflow -class IndexedForwardGraphCreator : private ExprVisitor { - public: - static IndexedForwardGraph Create(support::Arena* arena, const Expr& body) { - IndexedForwardGraphCreator creator(arena); - return creator.Prepare(body); - } - - private: - explicit IndexedForwardGraphCreator(support::Arena* arena) : arena_(arena) {} - - IndexedForwardGraph Prepare(const Expr& body) { - this->Update(body, nullptr, kOpaque); - this->VisitExpr(body); - return std::move(graph_); - } - - private: - /*! \brief allocator of all the internal node object */ - support::Arena* arena_; - // The output. - IndexedForwardGraph graph_; - // attribute equal comparator - StructuralEqual attr_equal_; - // Update the message stored at the node. - void Update(const Expr& node, IndexedForwardGraph::Node* parent, OpPatternKind pattern) { - const tvm::Object* key = node.get(); - IndexedForwardGraph::Node* current; - auto it = graph_.node_map.find(key); - if (it != graph_.node_map.end()) { - current = it->second; - } else { - current = arena_->make(); - graph_.node_map[key] = current; - } - if (parent != nullptr) { - auto* link = arena_->make>(); - link->value.node = parent; - link->value.pattern = pattern; - current->outputs.Push(link); - } else { - current->extern_ref = true; - } - } - - void AddNode(const tvm::Object* key) { - auto it = graph_.node_map.find(key); - ICHECK(it != graph_.node_map.end()) << "Cannot find node " << GetRef(key); - IndexedForwardGraph::Node* node = it->second; - ICHECK(node->ref == nullptr); - node->ref = key; - node->index = graph_.post_dfs_order.size(); - graph_.post_dfs_order.push_back(node); - } - - // Post order tree - void VisitExpr_(const FunctionNode* op) final { - // Skip the function that should be handled by external codegen. - if (op->GetAttr(attr::kCompiler).defined()) return; - - for (auto param : op->params) { - this->Update(param, nullptr, kOpaque); - } - this->Update(op->body, nullptr, kOpaque); - ExprVisitor::VisitExpr_(op); - } - - void VisitExpr_(const ConstantNode* op) final { - this->AddNode(op); - IndexedForwardGraph::Node* node = graph_.node_map.at(op); - DataType dtype = DataType(op->data->dtype); - // This rule must be consistent with code generator. - bool is_simple_const = - (dtype == DataType::Int(32) || dtype == DataType::Int(64) || dtype == DataType::Float(32) || - dtype == DataType::Float(64) || dtype == DataType::Bool()); - if (op->is_scalar() && is_simple_const) { - node->pattern = kElemWise; - } else { - // for now, mark non-scalar constant - // as opaque, we will not choose to fuse it. - node->pattern = kOpaque; - } - } - - void VisitExpr_(const CallNode* call) final { - ICHECK(graph_.node_map.count(call)); - IndexedForwardGraph::Node* node = graph_.node_map.at(call); - static auto fpattern = Op::GetAttrMap("TOpPattern"); - // Now we set the pattern of this call. - // - // If we see a call mentioning an operator we should mark it with its - // annotated pattern. - // - // If the pattern is not annotated we will default to opaque. - // - // Finally if the operator position is not a call node we will - // need to call Update, as it may be an arbitrary expression. - OpPatternKind op_pattern = kOpaque; - if (auto optional = call->op.as()) { - auto op = optional.value(); - if (IsDynamic(call->checked_type()) && IsDataDependent(call)) { - // output of a shape func can't be fed to a data-dependent shape func - op_pattern = kOpaque; - } else { - op_pattern = static_cast(fpattern[op]); - } - } else { - this->Update(call->op, node, kOpaque); - } - - node->pattern = op_pattern; - this->Update(call->op, nullptr, kOpaque); - const auto* rtype = call->checked_type().as(); - // pass the analysis back to all the children it references. - for (size_t i = 0; i < call->args.size(); ++i) { - const auto* arg_type = call->args[i]->checked_type().as(); - // specifically check if result type is the same as arguments type - OpPatternKind edge_pattern = op_pattern; - if (edge_pattern == kBroadcast && arg_type != nullptr && rtype != nullptr && - attr_equal_(rtype->shape, arg_type->shape)) { - edge_pattern = kElemWise; - } - this->Update(call->args[i], node, edge_pattern); - } - ExprVisitor::VisitExpr_(call); - this->AddNode(call); - } - - void VisitExpr_(const TupleNode* op) final { - ICHECK(graph_.node_map.count(op)); - IndexedForwardGraph::Node* tuple_node = graph_.node_map.at(op); - tuple_node->pattern = kTuple; - for (const Expr& field : op->fields) { - if (field->checked_type().as()) { - this->Update(field, tuple_node, kInjective); - } else { - this->Update(field, nullptr, kOpaque); - } - } - ExprVisitor::VisitExpr_(op); - this->AddNode(op); - } - - void VisitExpr_(const TupleGetItemNode* op) final { - auto tuple_type = op->tuple->checked_type().as(); - ICHECK(tuple_type); - // When TVM lowers a fused function, it expects all arguments to be a Tensor or - // a tuple containing only Tensors. But this tuple may contain a reference or - // another tuple. To avoid modifying codegen logic, we do not allow fusing through this node - // if the tuple contains such non Tensor fields. However, all fields will be recursively - // visited via call to ExprVisitor::VisitExpr_(op) below and corresponding visitor methods. - bool has_non_tensor = false; - for (auto ty : tuple_type->fields) { - if (!ty.as()) { - has_non_tensor = true; - break; - } - } - if (has_non_tensor) { - this->Update(op->tuple, nullptr, kOpaque); - } else { - ICHECK(graph_.node_map.count(op)); - IndexedForwardGraph::Node* node = graph_.node_map.at(op); - node->pattern = kInjective; - this->Update(op->tuple, node, kInjective); - } - ExprVisitor::VisitExpr_(op); - this->AddNode(op); - } - - void VisitExpr_(const VarNode* op) final { this->AddNode(op); } - - void VisitExpr_(const LetNode* op) final { - // do not fuse through let. - auto pre_visit = [this](const LetNode* op) { - // Rely on the Memoizer to cache pre-visit values - this->Update(op->var, nullptr, kOpaque); - this->Update(op->value, nullptr, kOpaque); - this->Update(op->body, nullptr, kOpaque); - this->VisitExpr(op->var); - this->VisitExpr(op->value); - }; - auto post_visit = [this](const LetNode* op) { - this->VisitExpr(op->body); - this->visit_counter_[op] += 1; - this->AddNode(op); - }; - ExpandANormalForm(op, pre_visit, post_visit); - } - - void VisitExpr_(const IfNode* op) final { - // do not fuse through if. - this->Update(op->cond, nullptr, kOpaque); - this->Update(op->true_branch, nullptr, kOpaque); - this->Update(op->false_branch, nullptr, kOpaque); - ExprVisitor::VisitExpr_(op); - this->AddNode(op); - } - - void VisitExpr_(const RefCreateNode* op) final { - this->Update(op->value, nullptr, kOpaque); - ExprVisitor::VisitExpr_(op); - this->AddNode(op); - } - - void VisitExpr_(const RefReadNode* op) final { - this->Update(op->ref, nullptr, kOpaque); - ExprVisitor::VisitExpr_(op); - this->AddNode(op); - } - - void VisitExpr_(const RefWriteNode* op) final { - this->Update(op->ref, nullptr, kOpaque); - this->Update(op->value, nullptr, kOpaque); - ExprVisitor::VisitExpr_(op); - this->AddNode(op); - } - - void VisitExpr_(const MatchNode* op) final { - this->Update(op->data, nullptr, kOpaque); - for (const Clause& c : op->clauses) { - this->Update(c->rhs, nullptr, kOpaque); - } - ExprVisitor::VisitExpr_(op); - this->AddNode(op); - } -}; - -class FuseMutator : private MixedModeMutator { - public: - FuseMutator(int fuse_opt_level, size_t max_fuse_depth, size_t max_function_args, bool link_params) - : fuse_opt_level_(fuse_opt_level), - max_fuse_depth_(max_fuse_depth), - max_function_args_(max_function_args), - link_params_(link_params) {} - - // Run the transform - Expr Transform(const Expr& body) { - return Transform(body, fuse_opt_level_, max_fuse_depth_, link_params_); - } - - protected: - // Run the transform - Expr Transform(const Expr& body, int fuse_opt_level, size_t max_fuse_depth, bool link_params) { - // setup the group map. - auto graph = IndexedForwardGraphCreator::Create(&arena_, body); - auto groups = GraphPartitioner(&arena_, fuse_opt_level, max_fuse_depth, max_function_args_) - .Partition(graph); - for (size_t nid = 0; nid < graph.post_dfs_order.size(); ++nid) { - ICHECK(graph.post_dfs_order[nid]->ref != nullptr); - gmap_[graph.post_dfs_order[nid]->ref] = groups[nid]; - } - // The following line can be used for debug. - // this->DebugDumpGroup(body); - return this->Mutate(body); - } - - private: - int fuse_opt_level_; - size_t max_fuse_depth_; - size_t max_function_args_; - bool link_params_; - - using MixedModeMutator::VisitExpr_; - - /*! \brief Temporary information from each group. */ - struct GroupInfo { - public: - // The parameters of the function. - Array params; - // The arguments to call the functions. - Array arguments; - // Get a new parameter or allocate an old one - Var GetOrAllocParam(const Expr& expr, const Type& type) { - // run linear scan as most fused groups contain only a few inputs. - for (size_t i = 0; i < arguments.size(); ++i) { - if (expr.same_as(arguments[i])) return params[i]; - } - // create a new parameter. - std::ostringstream os; - os << "p" << params.size(); - auto var = Var(os.str(), type); - params.push_back(var); - arguments.push_back(expr); - return var; - } - }; - /*! \brief Internal arena. */ - support::Arena arena_; - /*! \brief The group assignment map. */ - std::unordered_map gmap_; - /* \brief Internal group information map. */ - std::unordered_map ginfo_; - - // Skip primitive function. - Expr VisitExpr_(const FunctionNode* fn_node) { - if (fn_node->HasNonzeroAttr(attr::kPrimitive)) { - return GetRef(fn_node); - } else { - return ExprMutator::VisitExpr_(fn_node); - } - } - - // Transform calls. - Expr Rewrite_(const CallNode* call, const Expr& post) { - if (call->op.as()) { - static auto fnoncomputational = Op::GetAttrMap("TNonComputational"); - static auto fqnncanonicalize = Op::GetAttrMap("FTVMQnnCanonicalize"); - - Op op = Downcast(call->op); - if (fnoncomputational.get(op, false) && !fqnncanonicalize.count(op)) { - return ExprMutator::VisitExpr_(call); - } - - // If it is a primitive op call - // then we must have a group assignment for it already. - ICHECK(gmap_.count(call)); - if (call->op == stop_fusion_op) { - return ExprMutator::VisitExpr(call->args[0]); - } - auto* ret_group = gmap_.at(call)->FindRoot(); - Array new_args = GetNewArguments(call->args, ret_group); - - auto new_call = Call(call->op, new_args, call->attrs, call->type_args, call->span); - - if (ret_group->root_ref == call) { - // This is the root of the group - // create the new call node. - return MakeNewFunction(ret_group, call->checked_type(), new_call); - } else { - // This is an intermediate node of a fused function - // simply return the new call. - return std::move(new_call); - } - } else { - return ExprMutator::VisitExpr_(call); - } - } - - Expr Rewrite_(const TupleNode* tuple_node, const Expr& post) { - auto* ret_group = gmap_.at(tuple_node)->FindRoot(); - if (ret_group->root_ref == tuple_node) { - return ExprMutator::VisitExpr_(tuple_node); - } - // This tuple is an intermediate node in the group - Array new_fields = GetNewArguments(tuple_node->fields, ret_group); - return WithFields(GetRef(tuple_node), new_fields); - } - - Expr Rewrite_(const TupleGetItemNode* tuple_get, const Expr& post) { - auto* ret_group = gmap_.at(tuple_get)->FindRoot(); - auto new_tuple = GetNewArguments({tuple_get->tuple}, ret_group)[0]; - auto new_node = TupleGetItem(new_tuple, tuple_get->index); - if (ret_group->root_ref == tuple_get) { - if (gmap_.at(tuple_get->tuple.get())->FindRoot() != ret_group) { - // Isolated. This case occurs when tuple is created by an Opaque op - // e.g. multibox_transform_loc - return ExprMutator::VisitExpr_(tuple_get); - } - // A new function whose output is a tuple field access - return MakeNewFunction(ret_group, tuple_get->checked_type(), new_node); - } - // This is an intermediate node in the group - return std::move(new_node); - } - - Expr VisitExpr_(const LetNode* op) final { - auto pre_visit = [this](const LetNode* op) { - // Rely on the Memoizer to cache pre-visit values - this->VisitExpr(op->var); - this->VisitExpr(op->value); - }; - auto post_visit = [this](const LetNode* op) { - // Rely on the Memoizer to cache pre-visit values - Var var = Downcast(this->VisitExpr(op->var)); - Expr value = this->VisitExpr(op->value); - // Visit body and cache the op - Expr body = this->VisitExpr(op->body); - auto expr = GetRef(op); - if (var.same_as(op->var) && value.same_as(op->value) && body.same_as(op->body)) { - this->memo_[expr] = expr; - } else { - this->memo_[expr] = Let(var, value, body); - } - }; - ExpandANormalForm(op, pre_visit, post_visit); - return memo_[GetRef(op)]; - } - - Expr MakeNewFunction(GraphPartitioner::Group* group, Type ret_type, Expr body) { - // Quickly check special properties of the fused function. - // A pass to check if the fused op contains only reshape ops. - class CheckReshapeOnly : public ExprVisitor { - public: - void VisitExpr_(const CallNode* cn) final { - this->has_call = true; - static auto freshape_op = Op::GetAttrMap("TReshapeOp"); - - if (!freshape_op.get(cn->op, false)) { - this->reshape_only = false; - } - - if (!this->reshape_only) return; - ExprVisitor::VisitExpr_(cn); - } - - void VisitExpr_(const VarNode* vn) final { - if (!vn->type_annotation.defined() || !vn->type_annotation->IsInstance()) { - this->reshape_only = false; - } - } - - bool reshape_only = true; - bool has_call = false; - } visitor; - - visitor(body); - const GroupInfo& ginfo = ginfo_[group]; - auto func = Function(ginfo.params, body, ret_type, {}); - func = WithAttr(std::move(func), attr::kPrimitive, tvm::Integer(visitor.has_call)); - // TODO(mbs): "reshape" cleanup. - if (visitor.has_call && visitor.reshape_only) { - func = WithAttr(std::move(func), attr::kReshapeOnly, tvm::Integer(visitor.reshape_only)); - } - return Call(func, ginfo.arguments, Attrs()); - } - - Array GetNewArguments(const tvm::Array& args, - GraphPartitioner::Group* current_group) { - Array new_args; - for (auto arg : args) { - auto* arg_group = gmap_.at(arg.get())->FindRoot(); - auto type = arg->checked_type(); - Expr new_arg = this->Mutate(arg); - if (current_group != arg_group) { - if (!link_params_ || new_arg.as() == nullptr) { - Var param = ginfo_[current_group].GetOrAllocParam(new_arg, type); - new_args.push_back(param); - } else { - new_args.push_back(new_arg); - } - } else { - new_args.push_back(new_arg); - } - } - return new_args; - } - - // Debug function, dump the group assignment in text. - void DebugDumpGroup(const Expr& body) { - std::string text = AsText(body, false, [this](const ObjectRef& expr) -> std::string { - auto it = gmap_.find(expr.get()); - if (it == gmap_.end()) return ""; - std::ostringstream os; - auto* group = it->second->FindRoot(); - os << " /* group=" << group << " */"; - return os.str(); - }); - LOG(INFO) << "Dump of group info:\n" << text; - } -}; - -Expr FuseOps(const Expr& expr, int fuse_opt_level, size_t max_fuse_depth, size_t max_function_args, - bool link_params, const IRModule& module) { - return FuseMutator(fuse_opt_level, max_fuse_depth, max_function_args, link_params) - .Transform(expr); -} - -namespace transform { - -Pass FuseOps(int fuse_opt_level) { - runtime::TypedPackedFunc pass_func = - [=](Function f, IRModule m, PassContext pc) { - bool link_params = false; - Executor executor = - m->GetAttr(tvm::attr::kExecutor).value_or(NullValue()); - link_params = executor.defined() - ? executor->attrs.GetAttr("link-params").value_or(Bool(link_params)) - : link_params; - link_params = pc->GetConfig("relay.FuseOps.link_params", Bool(link_params)).value(); - int opt_level = fuse_opt_level == -1 ? pc->opt_level : fuse_opt_level; - auto max_fuse_depth = pc->GetConfig("relay.FuseOps.max_depth", Integer(kMaxFusedOps)); - auto target = Target::Current(); - size_t max_function_args = - (target.defined()) - ? target->GetAttr("max_function_args", Integer(0)).value().IntValue() - : 0; - return Downcast(FuseOps(f, opt_level, max_fuse_depth.value().IntValue(), - max_function_args, link_params, m)); - }; - return CreateFunctionPass(pass_func, 0, "FuseOps", {"InferType"}); -} - -TVM_REGISTER_GLOBAL("relay._transform.FuseOps").set_body_typed(FuseOps); - -} // namespace transform - -} // namespace relay -} // namespace tvm diff --git a/src/relay/transforms/gradient.h b/src/relay/transforms/gradient.h deleted file mode 100644 index 2e6ffbcc7c9e..000000000000 --- a/src/relay/transforms/gradient.h +++ /dev/null @@ -1,54 +0,0 @@ -/* - * 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 gradient.h - * \brief Utility functions for Automatic Differentiation in Relay. - */ -#ifndef TVM_RELAY_TRANSFORMS_GRADIENT_H_ -#define TVM_RELAY_TRANSFORMS_GRADIENT_H_ - -#include -#include - -#include - -namespace tvm { -namespace relay { - -inline Type GradRetType(const Function& f) { - // if type annotations are provided, we will construct a ret type; - // otherwise, leave it to be inferred - if (!f->ret_type.defined()) { - return Type(); - } - std::vector vt; - for (const auto& p : f->params) { - if (!p->type_annotation.defined()) { - return Type(); - } - vt.push_back(p->type_annotation); - } - - return TupleType({f->ret_type, TupleType(vt)}); -} - -} // namespace relay -} // namespace tvm -#endif // TVM_RELAY_TRANSFORMS_GRADIENT_H_ diff --git a/src/relay/transforms/higher_order_gradient.cc b/src/relay/transforms/higher_order_gradient.cc deleted file mode 100644 index da7a8f6420cd..000000000000 --- a/src/relay/transforms/higher_order_gradient.cc +++ /dev/null @@ -1,466 +0,0 @@ -/* - * 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 higher_order_gradient.cc - * \brief Higher-order Automatic Differentiation in Relay IR, for non-graph programs. - */ -#include -#include -#include -#include -#include -#include - -#include "gradient.h" -#include "let_list.h" -#include "pass_utils.h" -#include "pattern_utils.h" - -namespace tvm { -namespace relay { - -/*! What is automatic differentiation(AD) and why is it important? - * By AD, we roughly mean, given a term which denotes some mathematical function, - * derive a term which denotes the derivative of that mathematical function. - * Such a method can be compile-time, which is a macro on completely known function. - * Formally speaking, such requirement mean that the input function is a closed expression - - * that is, it only refer to local variable that is it's parameter, or defined inside it. - * Every top level definition satisfy this criteria. - * AD can also be run-time, which mean it is merely a function term of AD : (Float[] -> Float[]) -> - * (Float[] -> Float[]). In relay we currently only support compile-time AD, but it should be enough - * for a lot of use case. - * - * In deep learning, the most common way to train a deep neural network is by gradient descent or - * some of it's variant. Such optimization method require us to input the gradient of neural - * network, which can be obtained easily using AD. In fact, back propagation is essentially - * reverse-mode automatic differentiation, a kind of AD! - */ - -/*! In relay, automatic differentiation(AD) is a macro, - * that transform closed expr(expr without free variable/free type variable) of type - * (x0, x1, x2, ...) -> Float[] to - * (x0, x1, x2, ...) -> (Float[], (x0, x1, x2, ...)), - * When x0, x1, x2... are Float of different shape. - * the return value is a pair, with left hand side as the original value, and right hand side as - * gradient of the input. WithGradientType will take the type of input, and produce the type of - * output. There are multiple implementation of AD in relay, with different characteristic. However, - * they all transform the input expr according to WithGradientType. - */ -Type WithGradientType(const Type& t) { - // TODO(@M.K.): stricter checking - auto ty = t.as(); - ICHECK(ty) << "input should be a function"; - return FuncType(ty->arg_types, TupleType({ty->ret_type, TupleType(ty->arg_types)}), {}, {}); -} - -//! \brief if the expression is a GlobalVar, transform to it's expression. -Expr DeGlobal(const Optional& mod, const Expr& e) { - const auto* x = e.as(); - - if (mod.defined() && x) { - BaseFunc base_func = mod.value()->Lookup(GetRef(x)); - if (auto func = base_func.as()) { - return func.value(); - } else { - return e; - } - } else { - return e; - } -} - -static Type bpt = RelayRefType(FuncType({}, TupleType(Array()), {}, {})); - -struct ReverseADType : TypeMutator { - Type VisitType_(const TensorTypeNode* ttn) final { - Type t = GetRef(ttn); - return TupleType({t, RelayRefType(t)}); - } - - Type VisitType_(const FuncTypeNode* ftn) final { - std::vector arg_types; - for (const auto& t : ftn->arg_types) { - arg_types.push_back(VisitType(t)); - } - arg_types.push_back(bpt); - return FuncType(arg_types, ftn->ret_type, ftn->type_params, ftn->type_constraints); - } -}; - -Type ReverseType(const Type& t) { return ReverseADType()(t); } - -/*! \brief Lift a function that transform Tensor to a function that also transform more type - * by doing a structure preserving map. - */ -Expr LiftTensor(const std::function& f, - const std::function& tf, const Type& forward_type, const Expr& e, - LetList* ll) { - ICHECK(IsAtomic(e)) << e; - if (forward_type.as()) { - auto ret = ll->Push(f(e)); - ret->checked_type_ = tf(forward_type); - return std::move(ret); - } else if (auto* tt = forward_type.as()) { - tvm::Array fields; - tvm::Array types; - for (size_t i = 0; i < tt->fields.size(); ++i) { - auto field = LiftTensor(f, tf, tt->fields[i], ll->Push(GetField(e, i)), ll); - fields.push_back(field); - types.push_back(field->checked_type_); - } - auto ret = ll->Push(Tuple(fields)); - ret->checked_type_ = TupleType(types); - return std::move(ret); - } else { - LOG(FATAL) << "unsupported input/output type: " << tt; - throw; - } -} - -/*! \brief Transfers the gradients from an Expr to a deep duplication of the Expr, - * by stitching the references in the AD values. - */ -void TransferGrads(const Type& forward_type, const Expr& from, const Expr& to, LetList* ll) { - ICHECK(IsAtomic(from)) << from; - ICHECK(IsAtomic(to)) << to; - if (forward_type.as()) { - auto from_ref = TupleGetItem(from, 1); - auto to_ref = TupleGetItem(to, 1); - ll->Push(RefWrite(to_ref, RefRead(from_ref))); - } else if (auto* tt = forward_type.as()) { - for (size_t i = 0; i < tt->fields.size(); ++i) { - TransferGrads(tt->fields[i], ll->Push(TupleGetItem(from, i)), ll->Push(TupleGetItem(to, i)), - ll); - } - } else { - LOG(FATAL) << "Unsupported input/output type: " << forward_type; - throw; - } -} - -// TODO(@M.K.): why take Expr? -/*! \brief t -> ReverseType(t). Transform to Reverse Mode Value. */ -Expr GetRev(const Type& forward_type, const Expr& e, LetList* ll) { - auto rev = [&](const Expr& e) { return Pair(e, RefCreate(ZerosLike(e))); }; - auto rev_type = [&](const Type& forward_type) { return ReverseType(forward_type); }; - return LiftTensor(rev, rev_type, forward_type, e, ll); -} - -/*! \brief ReverseType(t) -> t. Get the original value. */ -Expr GetValue(const Type& forward_type, const Expr& e, LetList* ll) { - auto val = [&](const Expr& e) { return GetField(e, 0); }; - auto val_type = [&](const Type& forward_type) { return forward_type; }; - return LiftTensor(val, val_type, forward_type, e, ll); -} - -/*! \brief ReverseType(t) -> t. Get the gradient. */ -Expr GetGrad(const Type& forward_type, const Expr& e, LetList* ll) { - auto grad = [&](const Expr& e) { return RefRead(GetField(e, 1)); }; - auto grad_type = [&](const Type& forward_type) { return forward_type; }; - return LiftTensor(grad, grad_type, forward_type, e, ll); -} - -void UpdateGrad(const Type& t, const Expr& arg, const Expr& grad, LetList* ll) { - if (t.as()) { - ll->Push(RefWrite(GetField(arg, 1), Add(RefRead(GetField(arg, 1)), grad))); - } else if (auto* tt = t.as()) { - for (size_t i = 0; i < tt->fields.size(); ++i) { - UpdateGrad(tt->fields[i], ll->Push(GetField(arg, i)), ll->Push(GetField(grad, i)), ll); - } - } else { - LOG(FATAL) << "unsupported arg type of operator: " << t; - throw; - } -} - -Expr BPEmpty() { - Expr unitF = Function({}, Tuple(tvm::Array({})), TupleType::Empty(), {}); - return RefCreate(unitF); -} - -struct ReverseAD : ExprMutator { - using ADVarMap = std::unordered_map; - using ADGlobalVarMap = std::unordered_map; - Optional mod; - // TODO(@M.K.) refactor AD to always use mod. - Var bp; - std::shared_ptr ad_vars; - std::shared_ptr ad_gvars; - const OpAttrMap rev_map = Op::GetAttrMap("FPrimalGradient"); - - explicit ReverseAD(const Optional& mod, const Var& bp, - const std::shared_ptr& ad_vars, - const std::shared_ptr& ad_gvars) - : mod(mod), bp(bp), ad_vars(ad_vars), ad_gvars(ad_gvars) {} - - Expr VisitExpr_(const OpNode* op) final { - LOG(FATAL) << "op should only be inside call"; - throw; - } - - Expr Remap(const Expr& e) { - struct Remapper : ExprMutator { - std::shared_ptr ad_vars; - LetList* ll; - Remapper(const std::shared_ptr& ad_vars, LetList* ll) : ad_vars(ad_vars), ll(ll) {} - Expr VisitExpr_(const VarNode* var) final { - // memoize Var -> ADVar so we don't end up with free Vars when checkpointing - auto var_ref = GetRef(var); - if (ad_vars->count(var_ref) == 0) { - return std::move(var_ref); - } else { - return GetValue(var_ref->checked_type(), ad_vars->at(var_ref), ll); - } - } - }; - return LetList::With([&](LetList* ll) { return Remapper(ad_vars, ll)(e); }); - } - - Expr VisitCheckpoint(const CallNode* call) { - auto optional = call->op.as(); - ICHECK(optional) << "expected op in call"; - Op op_ref = optional.value(); - ICHECK(op_ref->name == "annotation.checkpoint") << "expected checkpoint annotation"; - auto x = call->args[0]; - return LetList::With([&](LetList* ll) { - auto x_var = ll->Push(Remap(x)); - auto ret = ll->Push(GetRev(call->checked_type(), x_var, ll)); - auto bpv = ll->Push(RefRead(bp)); - Expr nbp = Function({}, LetList::With([&](LetList* ll) { - // we need a new ReverseAD visitor to avoid clobbering the bp local var - auto dup_bp = ll->Push(BPEmpty()); - auto dup_ad = - ll->Push(ReverseAD(mod, dup_bp, ad_vars, ad_gvars)(DeDup(x))); - TransferGrads(call->checked_type(), ret, dup_ad, ll); - ll->Push(Call(RefRead(dup_bp), {})); - return Call(bpv, {}); - }), - TupleType::Empty(), {}); - ll->Push(RefWrite(bp, nbp)); - return ret; - }); - } - - Expr VisitExpr_(const CallNode* call) final { - if (auto optional = call->op.as()) { - Op op_ref = optional.value(); - - if (op_ref->name == "annotation.checkpoint") { - return VisitCheckpoint(call); - } - - ICHECK(rev_map.count(op_ref)) << op_ref->name << " does not have reverse mode defined"; - return LetList::With([&](LetList* ll) { - std::vector args; - for (const auto& arg : call->args) { - args.push_back(ll->Push(VisitExpr(arg))); - } - std::vector orig_args; - for (size_t i = 0; i < args.size(); i++) { - orig_args.push_back(GetValue(call->args[i]->checked_type(), args[i], ll)); - } - Expr orig = Call(call->op, orig_args, call->attrs, call->type_args); - orig->checked_type_ = call->checked_type(); - Var orig_var = ll->Push(orig); - orig_var->checked_type_ = call->checked_type(); - auto ret = ll->Push(GetRev(call->checked_type(), orig_var, ll)); - auto bpv = ll->Push(RefRead(bp)); - Expr nbp_body = LetList::With([&](LetList* ll) { - tvm::Array rev = rev_map[op_ref](orig, GetGrad(call->checked_type(), ret, ll)); - ICHECK(args.size() == rev.size()); - for (size_t i = 0; i < args.size(); ++i) { - UpdateGrad(call->args[i]->checked_type(), args[i], rev[i], ll); - } - return Call(bpv, {}); - }); - Expr nbp = Function({}, nbp_body, TupleType::Empty(), {}); - ll->Push(RefWrite(bp, transform::ToANormalForm(nbp))); - // TODO(@M.K.): ToANF should be called on rev. Enhance ToANF for that. - return ret; - }); - } else if (call->op.as()) { - return ExprMutator::VisitExpr_(call); - } else { - std::vector args; - for (const auto& arg : call->args) { - args.push_back(VisitExpr(arg)); - } - args.push_back(bp); - return Call(VisitExpr(call->op), args); - } - } - - Expr VisitExpr_(const ConstantNode* op) final { - return LetList::With([&](LetList* ll) { - Expr e = ll->Push(GetRef(op)); - return Pair(e, RefCreate(ZerosLike(e))); - }); - } - - Expr VisitExpr_(const IfNode* op) final { - return If(TupleGetItem(VisitExpr(op->cond), 0), VisitExpr(op->true_branch), - VisitExpr(op->false_branch)); - } - - Expr VisitExpr_(const VarNode* var) final { - // memoize Var -> ADVar so we don't end up with free Vars when checkpointing - auto var_ref = GetRef(var); - if (ad_vars->count(var_ref) == 0) { - auto res = Downcast(ExprMutator::VisitExpr_(var)); - (*ad_vars)[var_ref] = res; - } - - return ad_vars->at(var_ref); - } - - Expr VisitExpr_(const GlobalVarNode* op) final { - // todo: concatenating string to add attribute seems like a brittle hack. - // maybe get module indexed by a rose tree of string? - ICHECK(mod.defined()); - auto orig_gv = GetRef(op); - if (ad_gvars->count(orig_gv) == 0) { - GlobalVar gv(op->name_hint + "_grad"); - (*ad_gvars)[orig_gv] = gv; - Function orig_f = Downcast(DeDup(mod.value()->Lookup(orig_gv))); - Array params; - for (const auto& p : orig_f->params) { - params.push_back(Downcast(VisitExpr(p))); - } - params.push_back(bp); - Function f = WithFields(orig_f, params, VisitExpr(orig_f->body), VisitType(orig_f->ret_type)); - std::cout << "gv " << op->name_hint << ": " << AsText(f, false) << std::endl; - mod.value()->Add(gv, f); - } - return ad_gvars->at(orig_gv); - } - - Expr VisitExpr_(const FunctionNode* func_node) final { - Array params; - for (const auto& var : func_node->params) { - params.push_back(Downcast(VisitExpr(var))); - } - auto new_bp = Var("bp", bpt); - params.push_back(new_bp); - return WithFields(GetRef(func_node), params, - ReverseAD(mod, new_bp, ad_vars, ad_gvars)(func_node->body), - VisitType(func_node->ret_type)); - } - - Type VisitType(const Type& t) final { return t.defined() ? ReverseType(t) : t; } -}; - -bool MissingGrad(const Expr& e) { - struct MGVisitor : ExprVisitor { - const OpAttrMap rev_map = Op::GetAttrMap("FPrimalGradient"); - std::unordered_set op_names; - - void VisitExpr_(const OpNode* op) final { - Op op_ref = GetRef(op); - if (op_ref->name != "annotation.checkpoint" && !rev_map.count(op_ref)) { - op_names.insert(op_ref->name); - } - ExprVisitor::VisitExpr_(op); - } - }; - - MGVisitor mg; - mg.VisitExpr(e); - - if (mg.op_names.size() > 0) { - LOG(WARNING) << "found operators with missing gradients:"; - for (const auto& op : mg.op_names) { - LOG(WARNING) << " " << op; - } - return true; - } - - return false; -} - -Expr Gradient(const Expr& re, const Optional& mod) { - CheckFeature(re, FeatureSet::All() - fGraph); - if (mod.defined()) { - CheckFeature(mod.value(), FeatureSet::All() - fGraph); - } - auto e = DeGlobal(mod, re); - auto f = e.as(); - ICHECK(f) << "input need to be a function"; - ICHECK(f->type_params.size() == 0) << "no polymorphism supported for now"; - for (const auto& p : f->params) { - ICHECK(p->checked_type().as()) << "input parameters need to be tensor"; - } - ICHECK(!MissingGrad(e)) << "input has operators with missing gradients"; - Expr body = LetList::With([&](LetList* ll) { - Var bp = ll->Push(BPEmpty(), bpt); - Expr rev = ReverseAD(mod, bp, std::make_shared(), - std::make_shared())(e); - std::vector normal_args, args; - for (const auto& p : f->params) { - auto x = ll->Push(Pair(p, RefCreate(ZerosLike(p)))); - normal_args.push_back(x); - args.push_back(x); - } - args.push_back(bp); - auto c = ll->Push(Call(rev, args)); - std::function init_grad; - init_grad = [&](const Expr& e, const Type& t) { - if (t.as()) { - ll->Push(RefWrite(GetField(e, 1), OnesLike(GetField(e, 0)))); - } else if (auto tt = t.as()) { - ICHECK_GT(tt->fields.size(), 0); - init_grad(ll->Push(GetField(e, 0)), tt->fields[0]); - } else { - LOG(FATAL) << "unhandled type " << t; - throw; - } - }; - init_grad(c, f->body->checked_type()); - ll->Push(Call(RefRead(bp), {})); - std::vector ret; - for (const auto& a : normal_args) { - ret.push_back(RefRead(GetField(a, 1))); - } - std::function get_final_result; - get_final_result = [&](const Expr& e, const Type& t) -> Expr { - if (t.as()) { - return GetField(e, 0); - } else if (auto tt = t.as()) { - tvm::Array fields; - for (size_t i = 0; i < tt->fields.size(); ++i) { - fields.push_back(get_final_result(ll->Push(GetField(e, i)), tt->fields[i])); - } - return Tuple(fields); - } else { - LOG(FATAL) << "unhandled type " << t; - throw; - } - }; - return Pair(get_final_result(c, f->body->checked_type()), Tuple(ret)); - }); - Function ret = WithFields(GetRef(f), f->params, body, GradRetType(GetRef(f)), - /* erase type params */ Array()); - CheckFeature(ret, FeatureSet::All() - fGraph); - return std::move(ret); -} - -TVM_REGISTER_GLOBAL("relay._transform.gradient").set_body_typed(Gradient); - -} // namespace relay -} // namespace tvm diff --git a/src/relay/transforms/infer_layout_utils.cc b/src/relay/transforms/infer_layout_utils.cc deleted file mode 100644 index 984e23ad15f1..000000000000 --- a/src/relay/transforms/infer_layout_utils.cc +++ /dev/null @@ -1,265 +0,0 @@ -/* - * 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 "infer_layout_utils.h" - -#include -#include -#include - -#include -#include -#include -#include -#include - -#include "pattern_utils.h" -#include "tvm/runtime/logging.h" - -namespace tvm { -namespace relay { - -Layout AdjustSubordinateFactors(const Layout& src_layout, const Layout& old_layout, - const Array& old_shape) { - // For each subordinate axis - // 1) Find the corresponding dual axis. - // 2) Find the Index of this dual axis in old_layout. - // 3) Find the shape of the that axis in old_shape. - // 4) a) Adjust factor to 1, if that shape is 1. b) Else retain the factor. - DLOG(INFO) << "AdjustSubordinateFactors" - << "src_layout: " << src_layout << " old_layout: " << old_layout - << " old_shape: " << old_shape << std::endl; - std::string new_layout; - for (auto axis : src_layout->axes) { - if (!LayoutAxis::Get(axis).IsPrimal()) { - bool is_shape_one = false; - // 1) Find the corresponding dual axis - const auto& dual_axis = LayoutAxis::Get(axis).ToPrimal(); - - // 2) Find the index of this dual axis in old_layout - int old_axis = old_layout.IndexOf(dual_axis); - - if (old_axis == -1) { - new_layout += "1"; - is_shape_one = true; - } else { - // 3) Find the shape of this index in old_shape - auto shape_val = old_shape[old_axis]; - - // 4) a) Check if this shape element is 1. - if (auto* shape_int = shape_val.as()) { - // We can treat 1 as broadcast only if axis was not split before - if (shape_int->value == 1 && old_layout.IndexOf(LayoutAxis::Get(axis)) == -1) { - new_layout += "1"; - is_shape_one = true; - } - } - } - - // 4) b) If shape is not 1, retain the factor. - if (!is_shape_one) { - auto new_shape_val = src_layout.FactorOf(dual_axis); - new_layout += std::to_string(new_shape_val); - } - } - new_layout += LayoutAxis::Get(axis).name(); - } - return new_layout != "" ? Layout(new_layout) - : Layout("H").SubLayout(0, 0); // hack to create a scalar layout -} - -bool Isomorphic(const Layout& lhs, const Layout& rhs) { - DLOG(INFO) << "Isomorphic: " - << "lhs: " << lhs << " rhs: " << rhs << std::endl; - ICHECK(lhs.defined()); - ICHECK(rhs.defined()); - if (lhs->axes.size() != rhs->axes.size()) return false; - std::map map_to, map_back; - for (size_t i = 0; i < lhs->axes.size(); ++i) { - auto& lhs_axis = LayoutAxis::Get(lhs->axes[i]); - auto& rhs_axis = LayoutAxis::Get(rhs->axes[i]); - std::string name_lhs = lhs_axis.name(); - std::string name_rhs = rhs_axis.name(); - if (lhs_axis.IsPrimal() != rhs_axis.IsPrimal()) return false; - - auto it = map_to.find(name_lhs); - if (it == map_to.end()) - map_to[name_lhs] = name_rhs; - else if (it->second != name_rhs) - return false; - - it = map_back.find(name_rhs); - if (it == map_back.end()) - map_back[name_rhs] = name_lhs; - else if (it->second != name_lhs) - return false; - if (!lhs_axis.IsPrimal() && lhs.FactorOf(lhs_axis) != rhs.FactorOf(rhs_axis)) return false; - } - return true; -} - -Layout TryTransformLike(const Layout& old, const Layout& ref_old, const Layout& ref_new) { - DLOG(INFO) << "transform_layout: old = " << old << ", ref_new = " << ref_new - << ", ref_old = " << ref_old << std::endl; - ICHECK(ref_old.defined()); - ICHECK(ref_new.defined()); - ICHECK(old.defined()); - - { // check if old and ref_old are similar enough such that it's - // compatible for the transform ref_old -> ref_new - const Layout& large = ref_old.ndim() > old.ndim() ? ref_old : old; - const Layout& small = large == ref_old ? old : ref_old; - Layout large_sublayout = large.SubLayout(large.ndim() - small.ndim(), small.ndim()), - rest_sublayout = large.SubLayout(0, large.ndim() - small.ndim()); - bool orthorgonal = true; - for (auto i : rest_sublayout->axes) - if (large_sublayout.IndexOf(LayoutAxis::Get(i).ToPrimal()) != -1 || - large_sublayout.IndexOf(LayoutAxis::Get(i).ToSubordinate()) != -1) { - orthorgonal = false; - break; - } - if (!orthorgonal || !Isomorphic(large_sublayout, small)) - return Layout::Undef(); // For now this case is not supported. - } - - // `old` is compatible. Now learn the axis name mapping between `old` and `ref_old` - if (old.ndim() == 0) return old; // an optmization for scalar: no-op - std::vector mapping(26, -1); - std::vector used(26, false); - - auto find_unused = [&](char preference) -> char { - if (!used[preference - 'A']) return preference; // preference unused - for (int i = 0; i < 26; ++i) - if (!used[i]) return 'A' + i; - LOG(FATAL) << "All letters are used"; - }; - - for (int j = old->axes.size() - 1, i = ref_old->axes.size() - 1; j >= 0; --i, --j) { - char name_ref = LayoutAxis::Get(ref_old->axes[i]).ToPrimal().name()[0]; - char name = LayoutAxis::Get(old->axes[j]).ToPrimal().name()[0]; - mapping[name_ref - 'A'] = name - 'A'; - used[name - 'A'] = true; - } - - for (int i = ref_old->axes.size() - 1; i >= 0; --i) { - char name_ref = LayoutAxis::Get(ref_old->axes[i]).ToPrimal().name()[0]; - int name = mapping[name_ref - 'A']; - if (name == -1) { - mapping[name_ref - 'A'] = find_unused(name_ref) - 'A'; - used[mapping[name_ref - 'A']] = true; - } - } - - // apply the mapping to rename `ref_new` - std::string new_layout; - for (auto c : std::string(ref_new->name)) { - if (c >= 'A' && c <= 'Z') { - ICHECK(mapping[c - 'A'] != -1); - new_layout += mapping[c - 'A'] + 'A'; - } else if (c >= 'a' && c <= 'z') { - ICHECK(mapping[c - 'a'] != -1); - new_layout += mapping[c - 'a'] + 'a'; - } else { - new_layout += c; - } - } - - DLOG(INFO) << "new_layout = " << new_layout << std::endl; - return Layout(new_layout); -} - -std::pair, Array> BinaryBroadcastLayoutHelper( - const Attrs& attrs, const Array& new_in_layouts, const Array& old_in_layouts, - const Array& old_in_types) { - // Two steps. Step (2) only executes if the function is called after rewrite. - // (1) infer input layouts before rewrite - // (2) if some input layouts are changed by its producer after rewrite, rewrite the other - // layout to make sure it's changed in the same way, so that they are still broadcastable. - Array layouts; - Array> old_in_shapes; - for (auto old_in_t : old_in_types) { - ICHECK(old_in_t.as()); - old_in_shapes.push_back(old_in_t.as()->shape); - } - int old_large_idx = old_in_shapes[0].size() >= old_in_shapes[1].size() ? 0 : 1; - - layouts.Assign(old_in_layouts.begin(), old_in_layouts.end()); - // always operate on the original layouts first for consistency - - std::pair, Array> out, - out_default{{Layout::Undef(), Layout::Undef()}, {Layout::Undef()}}; - - if (!layouts[0].defined() && !layouts[1].defined()) { - // both undefined, infer fails - out = out_default; - } else if (!layouts[0].defined() || !layouts[1].defined()) { - // only one is defined, use shape information to help infer - int defined_idx = layouts[0].defined() ? 0 : 1; - int undef_idx = 1 - defined_idx; - - if (old_in_shapes[defined_idx].size() >= old_in_shapes[undef_idx].size()) { - // TODO(lazycal): handle the case when the sublayout contains subcoordinate of factor one but - // the other tensor has the corresponding dimension size other than one. - // E.g. defined's shape = [x, x, x, x, 1] in NCHW1c and undefined's shape = [3] - layouts.Set(undef_idx, layouts[defined_idx].SubLayout(old_in_shapes[defined_idx].size() - - old_in_shapes[undef_idx].size(), - old_in_shapes[undef_idx].size())); - out = {layouts, {layouts[defined_idx]}}; - } else { - // only know the tensor with smaller dimensions, - // so we cannot infer the final broadcasted output. - // fails in this case. - out = out_default; - } - } else { - // when both are defined, return the larger one - out = {layouts, {layouts[old_large_idx]}}; - } - if (!new_in_layouts.defined()) return out; - // Step (2) rewrite the layouts to make them broadcastable again. - Layout ret = new_in_layouts[old_large_idx]; - int large_idx = new_in_layouts[0].ndim_primal() >= new_in_layouts[1].ndim_primal() ? 0 : 1; - int small_idx = 1 - large_idx; - // start adjusting - - // Apply a greedy strategy that always transform the small layout in the same way as the - // large layout is transformed, if possible. - Layout tgt_layout = - TryTransformLike(layouts[small_idx], layouts[large_idx], new_in_layouts[large_idx]); - if (!tgt_layout.defined()) return out_default; // fallback - - // Support scenarios where original operands were of type [N, H, W, C] and [N, H, W, 1] - // In this case, we might have NCHW16c coming for 1 operand. However, the other operand does - // not have enough C dimension. To reuse broadcasting, we would want to use NCHW1c for the - // second operand. The following section of code walks through the layouts and shapes to - // perform that operation. - // a in NCHWC16c - // b in NHW1 - // b = layout_transform(b) from NHW1 -> NCHW1c - // add(a, b) - auto old_small_shape = old_in_shapes[small_idx]; - auto old_small_layout = layouts[small_idx]; - auto new_small_layout = AdjustSubordinateFactors(tgt_layout, old_small_layout, old_small_shape); - layouts.Set(large_idx, new_in_layouts[large_idx]); - layouts.Set(small_idx, new_small_layout); - return {layouts, {ret}}; -} - -} // namespace relay -} // namespace tvm diff --git a/src/relay/transforms/infer_layout_utils.h b/src/relay/transforms/infer_layout_utils.h deleted file mode 100644 index 3b1cb29951e4..000000000000 --- a/src/relay/transforms/infer_layout_utils.h +++ /dev/null @@ -1,159 +0,0 @@ -/* - * 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 infer_layout_utils.h - * \brief Utility functions to alter the layouts of operators or replace primitive operators with - other expressions. This pass can be used for computing convolution in - custom layouts or other general weight pre-transformation. - */ - -#ifndef TVM_RELAY_TRANSFORMS_INFER_LAYOUT_UTILS_H_ -#define TVM_RELAY_TRANSFORMS_INFER_LAYOUT_UTILS_H_ - -#include -#include -#include - -#include -#include -#include - -#include "pattern_utils.h" - -namespace tvm { -namespace relay { - -/*! - * \brief Returns a new layout where the subordinate factors are adjusted based on the tensor - * shape. - * \param old_layout The old layout before any transformation. - * \param old_shape The shape of the original tensor. - * \return The adjusted Layout. - */ -Layout AdjustSubordinateFactors(const Layout& src_layout, const Layout& old_layout, - const Array& old_shape); - -bool Isomorphic(const Layout& lhs, const Layout& rhs); - -/*! - * \brief Try transforming `old` in as the smae way as how`ref_old` is transformed to `ref_new`. - * `old` and `ref_old` are expected to describe two broadcastable tensors. Layout with fewer rank - * will be expanded. For example, - * if old = 'NW', ref_old = 'NC', ref_new = 'NC1c', then the result is 'NW1w'; - * if old = 'W', ref_old = 'NC', ref_new = 'NC1c', then the result is 'NW1w'. - * When `old` and `ref_old` are isomorphic (same structure, only differ in naming), the transform - * is guaranteed to succeed, in which case the function is simply renaming the axes of `ref_new` - * to conform to `old`'s naming. - * \param old The layout to be transformed. - * \param ref_old The reference layout before transform. - * \param ref_new The reference layout after transform. - * \return The transformed layout. - */ -Layout TryTransformLike(const Layout& old, const Layout& ref_old, const Layout& ref_new); - -/* - * \brief An output structure to hold results from FInferCorrectLayout calls. - * \tparam input_layouts Inferred input layouts. - * \tparam output_layouts Inferred output layouts. - * \tparam new_attrs Updated attributes consistent with inferred layouts. - */ -class InferCorrectLayoutOutputNode : public Object { - public: - Array input_layouts; - Array output_layouts; - Attrs new_attrs; - - void VisitAttrs(tvm::AttrVisitor* v) { - v->Visit("input_layouts", &input_layouts); - v->Visit("output_layouts", &output_layouts); - v->Visit("new_attrs", &new_attrs); - } - - TVM_DECLARE_BASE_OBJECT_INFO(InferCorrectLayoutOutputNode, Object); - - static constexpr const char* _type_key = "relay._transform.InferCorrectLayoutOutput"; -}; - -class InferCorrectLayoutOutput : public ObjectRef { - public: - InferCorrectLayoutOutput(Array input_layouts, Array output_layouts, - Attrs new_attrs) { - auto n = make_object(); - n->input_layouts = std::move(input_layouts); - n->output_layouts = std::move(output_layouts); - n->new_attrs = std::move(new_attrs); - data_ = n; - } - TVM_DEFINE_OBJECT_REF_METHODS(InferCorrectLayoutOutput, ObjectRef, InferCorrectLayoutOutputNode); -}; - -/*! - * \brief Infer & correct function of node layout. See \p Layout for layout convention - * \param attrs The attribute of the node. - * \param new_in_layouts The layouts of input arguments after alter_op_layout. - * This can be undefined, which means we call this function before alternating - * any operators. - * \param old_in_layouts The layouts of input arguments before alter_op_layout. - * \param old_in_types The types of old input arguments. - * \return infer_layout_output Inferred layouts and updated attributes stored in - * InferCorrectLayoutOutput above. - */ -using FInferCorrectLayout = runtime::TypedPackedFunc& new_in_layouts, const Array& old_in_layouts, - const Array& old_in_types)>; - -inline InferCorrectLayoutOutput ElemwiseArbitraryLayout( - const Attrs& attrs, const Array& new_in_layouts, const Array& old_in_layouts, - const Array& old_in_types) { - Layout ret; - - if (new_in_layouts.defined()) { - ICHECK_GE(new_in_layouts.size(), 1); - ret = new_in_layouts[0]; - } else { - for (size_t i = 0; i < old_in_layouts.size(); ++i) { - if (old_in_layouts[i].defined()) { - ret = old_in_layouts[i]; - break; - } - } - } - - return InferCorrectLayoutOutput(Array(old_in_layouts.size(), ret), {ret}, attrs); -} - -std::pair, Array> BinaryBroadcastLayoutHelper( - const Attrs& attrs, const Array& new_in_layouts, const Array& old_in_layouts, - const Array& old_in_types); - -/*! \brief Infer layout for binary broadcast operators */ -inline InferCorrectLayoutOutput BinaryBroadcastLayout(const Attrs& attrs, - const Array& new_in_layouts, - const Array& old_in_layouts, - const Array& old_in_types) { - auto inferred_layout = - BinaryBroadcastLayoutHelper(attrs, new_in_layouts, old_in_layouts, old_in_types); - return InferCorrectLayoutOutput(inferred_layout.first, inferred_layout.second, attrs); -} - -} // namespace relay -} // namespace tvm - -#endif // TVM_RELAY_TRANSFORMS_INFER_LAYOUT_UTILS_H_ diff --git a/src/relay/transforms/inline.cc b/src/relay/transforms/inline.cc deleted file mode 100644 index 564c0daef70f..000000000000 --- a/src/relay/transforms/inline.cc +++ /dev/null @@ -1,229 +0,0 @@ -/* - * 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 src/relay/transforms/inline.cc - * \brief Global function inliner. It contains the following steps: - * - * - Preprocessing: eligibility checking. Only inline the functions that can - * be inlined. We currently only use simple rules to make the decision. No - * profitibility analysis is available for now. - * - * - Inline: replace the call with a function or the function body depending on - * the attribute of the callee function. For example, we return the function - * node when it doesn't use default compiler, i.e. llvm. This is because these - * functions are packed to be offloaded to external codegen. - * - * - Postprocessing: remove the replaced functions that have no reference. - */ - -#include -#include -#include -#include -#include - -#include -#include - -#include "../analysis/call_graph.h" -#include "../op/call/call.h" - -using namespace tvm::runtime; - -namespace tvm { -namespace relay { - -class Inliner : ExprMutator { - public: - explicit Inliner(CallGraphEntry* cur_node, CallGraphNode* call_graph) - : cur_node_(cur_node), call_graph_(call_graph) {} - - Expr VisitExpr_(const CallNode* call_node) final { - // We can work with calls in both pre- and post-lowered form. - Call vanilla_call = GetAnyCall(call_node); - - const auto* global_var_node = vanilla_call->op.as(); - if (global_var_node) { - GlobalVar gv = GetRef(global_var_node); - auto* cg_node = (*call_graph_)[gv->name_hint]; - if (CanInline(cg_node)) { - Array new_args; - new_args.reserve(vanilla_call->args.size()); - for (auto arg : vanilla_call->args) { - new_args.push_back(VisitExpr(arg)); - } - // TODO(mbs): Does not handle multiple calls to the same global function. - cur_node_->RemoveCallTo(gv); - return MakeNewExpr(gv, new_args, GetRef(call_node)); - } - // else: fallthrough - } - // else: fallthrough - - // If not calling a global function then nothing to inline. - return ExprMutator::VisitExpr_(call_node); - } - - Expr VisitExpr_(const GlobalVarNode* gvn) final { - GlobalVar gv = GetRef(gvn); - auto* cg_node = (*call_graph_)[gv->name_hint]; - if (CanInline(cg_node)) { - cur_node_->RemoveCallTo(gv); - return MakeNewExpr(gv, {}, GetRef(gvn)); - } - return ExprMutator::VisitExpr_(gvn); - } - - Function Inline(const Function& func) { - return WithFields(func, func->params, VisitExpr(func->body)); - } - - private: - bool CanInline(const CallGraphEntry* cg_node) { - // The node must be a leaf node and it cannot be recursive. - if (!cg_node->empty() || cg_node->IsRecursive()) return false; - - auto base_func = call_graph_->GetGlobalFunction(cg_node->GetGlobalVar()); - const auto* function_node = base_func.as(); - if (!function_node) { - // Can't inline PrimFuncs! - return false; - } - // The body of a global functions must be defined. - if (!function_node->body.defined()) return false; - - // The function must be annotated with the inline attribute. - // (Note that partitioned functions and external functions do not have this attribute!) - if (!function_node->HasNonzeroAttr(attr::kInline)) return false; - - // The function is not able to be inlined if any callee under the CallGraph - // of this function cannot be inlined. - for (const auto& it : *cg_node) { - if (!CanInline(it.second)) { - return false; - } - } - - return true; - } - - // Make a new Relay expression to replace \p expr. - Expr MakeNewExpr(const GlobalVar& global, const Array& args, const Expr& expr) { - ICHECK(expr->IsInstance() || expr->IsInstance()); - auto base_func = call_graph_->GetGlobalFunction(global); - const auto* fn = base_func.as(); - ICHECK(fn) << "Expected to work on a Relay function."; - - // There is an inconsistency here, the function itself gets shallow-copied but the body is not - // shallow-copied. - auto func = Function(fn->params, fn->body, fn->ret_type, fn->type_params, fn->attrs); - // Inline the function body to the caller if this function uses default - // compiler, i.e. no external codegen is needed. - if (!func->GetAttr(attr::kCompiler).defined() && !func->HasNonzeroAttr(attr::kExtern)) { - ICHECK_EQ(func->params.size(), args.size()) - << "Mismatch found in the number of parameters and call args"; - // Bind the parameters with call args. - Map bind_map; - for (size_t i = 0; i < args.size(); i++) { - bind_map.Set(fn->params[i], args[i]); - } - if (const auto* gvn = expr.as()) { - auto ret_type = gvn->checked_type(); - // Cannot replace TensorType/TensorTupleType with FuncType. Therefore, - // we simply inline the function as a closure instead of directly using - // its body when the global var returns FuncType. - return ret_type->IsInstance() ? std::move(func) : func->body; - } else { - ICHECK(expr->IsInstance()); - return Bind(func->body, bind_map); - } - } else if (const auto* call_node = expr.as()) { - return Call(func, args, call_node->attrs, call_node->type_args); - } else { - return std::move(func); - } - } - - /*! - * \brief The current call graph entry that is being handled. Each entry - * contains a global function. - */ - CallGraphEntry* cur_node_; - /*! \brief The call graph that is used for global function lookup. */ - const CallGraphNode* call_graph_; -}; - -IRModule Inline(const IRModule& module) { - CallGraph cg(module); - auto topo = cg->TopologicalOrder(); - // Get the reverse topological order of the global functions. - std::reverse(topo.begin(), topo.end()); - // Cache the functions that are originally entries. These functions will - // remain in the module after inlining. - std::unordered_set original_entry; - - for (auto* it : topo) { - if (it->GetRefCount() == 0) original_entry.emplace(it); - // Skip the leaf calls and the recursive calls that don't call other - // functions. - if (it->empty() || (it->IsRecursive() && it->size() == 1)) continue; - auto base_func = module->Lookup(it->GetNameHint()); - if (auto func = base_func.as()) { - auto new_func = Inliner(it, cg.operator->()).Inline(func.value()); - // TODO(zhiics) Maybe move this to CallGraph, but updating function from - // CallGraph arbitarily may lead to incorrect CallGraph. - cg->module->Update(it->GetGlobalVar(), new_func); - } - } - - // Clean up the functions that are inlined and have no reference. - for (auto* cgn : topo) { - // Skip recursive functions and entry functions even if they are marked as - // `inline`. - if (cgn->IsRecursive() || original_entry.count(cgn)) continue; - auto base_func = cg->GetGlobalFunction(cgn->GetGlobalVar()); - // Skip calls to PrimFuncs since they can't be inlined. - if (const auto* func = base_func.as()) { - if (func->HasNonzeroAttr(attr::kInline)) { - ICHECK_EQ(cgn->GetRefCount(), 0U) - << cgn->GetNameHint() << " is marked as inline but not inlined."; - cgn->CleanCallGraphEntries(); - cg->RemoveGlobalVarFromModule(cgn, /*update_call_graph*/ true); - } - } - } - - return cg->module; -} - -namespace transform { - -Pass Inline() { - runtime::TypedPackedFunc pass_func = - [=](IRModule m, PassContext pc) { return relay::Inline(m); }; - return CreateModulePass(pass_func, 1, "InlineGlobals", {}); -} - -TVM_REGISTER_GLOBAL("relay._transform.Inline").set_body_typed(Inline); - -} // namespace transform - -} // namespace relay -} // namespace tvm diff --git a/src/relay/transforms/label_ops.cc b/src/relay/transforms/label_ops.cc deleted file mode 100644 index bc75d555c9af..000000000000 --- a/src/relay/transforms/label_ops.cc +++ /dev/null @@ -1,143 +0,0 @@ -/* - * 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 - -namespace tvm { -namespace relay { -namespace transform { - -namespace { - -/*! \brief Collect all attributes whose name contains "layout". - */ -struct CollectAttrs : public AttrVisitor { - void Visit(const char* key, std::string* value) final { - if (std::string(key).find("layout") != std::string::npos) { - attrs[key] = String(*value); - } - } - void Visit(const char* key, double* value) final {} - void Visit(const char* key, uint64_t* value) final {} - void Visit(const char* key, int* value) final {} - void Visit(const char* key, int64_t* value) final {} - void Visit(const char* key, bool* value) final {} - void Visit(const char* key, runtime::NDArray* value) final {} - void Visit(const char* key, ObjectRef* value) final { - if (std::string(key).find("layout") != std::string::npos) { - attrs[key] = *value; - } - } - void Visit(const char* key, DataType* value) final {} - void Visit(const char* key, void** value) final {} - std::unordered_map attrs; -}; -} // namespace - -/*! \brief Visitor to add structural hash and layout information to `Function` - * nodes. Sets the "hash" field on the attr to the structural hash of the - * function. Propogates any attributes with "layout" in their name from call - * nodes in the Function to the Function's attrs. - */ -class LabelOpsMutator : public MixedModeMutator { - private: - using MixedModeMutator::VisitExpr_; - std::unordered_map body_attrs; - Expr VisitExpr_(const FunctionNode* op) final { - if (op->GetAttr("hash").defined()) { - // Already labelled. - return ExprMutator::VisitExpr_(op); - } - - // body_attrs collects attrs from Calls in the body of this Function. Reset - // it so we only get attrs from this Function. - body_attrs = {}; - auto updated = ExprMutator::VisitExpr_(op); - size_t hash = StructuralHash()(updated); - - // format hash as fixed length hex string so it is easier to read - std::stringstream s; - s << std::setfill('0') << std::setw(sizeof(size_t) * 2) << std::hex << hash; - - Function f = WithAttr(Downcast(updated), "hash", String(s.str())); - for (auto p : body_attrs) { - f = WithAttr(f, p.first, p.second); - } - return std::move(f); - } - - Expr VisitExpr_(const LetNode* op) final { - auto pre_visit = [this](const LetNode* op) { - this->Mutate(op->var); - this->Mutate(op->value); - }; - auto post_visit = [this](const LetNode* op) { - Var var = Downcast(this->Mutate(op->var)); - auto value = this->Mutate(op->value); - auto body = this->Mutate(op->body); - auto expr = GetRef(op); - if (var.same_as(op->var) && value.same_as(op->value) && body.same_as(op->body)) { - this->memo_[expr] = expr; - } else { - this->memo_[expr] = Let(var, value, body); - } - }; - ExpandANormalForm(op, pre_visit, post_visit); - return memo_[GetRef(op)]; - } - - Expr Rewrite_(const CallNode* op, const Expr& post) final { - auto updated = MixedModeMutator::Rewrite_(op, post); - if (op->attrs.defined()) { - CollectAttrs collect; - const_cast(op->attrs.get())->VisitAttrs(&collect); - for (auto p : collect.attrs) { - if (body_attrs.find(p.first) != body_attrs.end() && p.second == body_attrs[p.first]) { - LOG(WARNING) << "LabelOps found two call sites with different values for " << p.first - << " (" << p.second << " vs " << body_attrs[p.first] - << "). Only the first will be recorded."; - } - body_attrs[p.first] = p.second; - } - } - return updated; - } -}; - -/*! \brief Add structural hash and layout information to Function nodes. This - * information is used later by the profiler. - * - * The hash and layout information is added to the attrs field of the Function. - * The key "hash" contains the structural hash of the node. Any attributes with - * "layout" in their name are also added to attrs (for example, - * `attrs["src_layout"]` contains the `src_layout` attribute of the TVM op - * corresponding to this function call). - */ -Pass LabelOps() { - runtime::TypedPackedFunc pass_func = - [=](Function f, IRModule m, PassContext pc) { - return Downcast(LabelOpsMutator().Mutate(f)); - }; - return CreateFunctionPass(pass_func, 1, "LabelOps", {}); -} - -} // namespace transform -} // namespace relay -} // namespace tvm diff --git a/src/relay/transforms/lazy_gradient_init.cc b/src/relay/transforms/lazy_gradient_init.cc deleted file mode 100644 index 548951f19404..000000000000 --- a/src/relay/transforms/lazy_gradient_init.cc +++ /dev/null @@ -1,282 +0,0 @@ -/* - * 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 lazy_gradient_init.cc - * - * \brief Lazily instantiate 0-filled or 1-filled tensors. - * This pass should be used after reverse-mode ad so that gradient tensors - * are not instantiated until after the forward pass. - * - * This pass delays or removes memory allocation by converting tensors into - * GradCell, an algebraic data type defined in gradient.rly. - * - * This will delay or decrease memory usage. All calls to - * ones, ones_like, zeros, zeros_like will call the One or Zero constructor - * of GradCell, which will not instantiate in memory until needed. All other cases result - * in using the Raw constructor which means the tensor is instantiated in memory. - * - * It also overloads + and * operation which can increase performance when doing - * operations involving tensors with values of only 0 or 1. - * - * Note: this pass can only be used with functions where the input/output types are - * a combination of TupleTypes and TensorTypes - * - * This pass optimizes 6 ops: - * - add - * - multiply - * - ones - * - ones_like - * - zeros - * - zeros_like - * - * This pass makes use of three visitor. The most important one visits the entire function, - * one is used for wrap inputs and one to unwrap outputs. - * - * For example: - * fn: TensorType[(10,10), float32] -> TensorType[(10,10), float32] - * - * After this pass - * fn: GradCell[TensorType[(10,10), float32]] -> GradCell[TensorType[(10,10), float32]] - * - * Thus, it is necessary to wrap this outer function so that the input/output types remain the same - */ - -#include -#include -#include -#include -#include -#include - -#include "let_list.h" - -namespace tvm { -namespace relay { - -class LazyGradientInitializer : public ExprMutator, public TypeMutator { - public: - explicit LazyGradientInitializer(IRModule module) : module_(module) { - module_->ImportFromStd("gradient.rly"); - } - - Expr WrapExpr(const Var& var, const Type& type, LetList* ll) { - if (type.as()) { - return Call(module_->GetConstructor("GradCell", "Raw"), {var}, Attrs(), {type}); - } else if (auto* type_anno = type.as()) { - tvm::Array fields; - for (size_t i = 0; i < type_anno->fields.size(); i++) { - const Type& t = type_anno->fields[i]; - fields.push_back(WrapExpr(ll->Push(TupleGetItem(var, i)), t, ll)); - } - Expr tuple = Tuple(fields); - return tuple; - } - - return var; - } - - Expr UnwrapExpr(const Var& var, const Type& type, LetList* ll) { - if (auto* type_call = type.as()) { - if (type_call->func.same_as(module_->GetGlobalTypeVar("GradCell"))) { - return Call(module_->GetGlobalVar("FromGradCell"), {var}); - } - return var; - } else if (auto* type_anno = type.as()) { - tvm::Array fields; - for (size_t i = 0; i < type_anno->fields.size(); i++) { - const Type& t = type_anno->fields[i]; - fields.push_back(UnwrapExpr(ll->Push(TupleGetItem(var, i)), t, ll)); - } - Expr tuple = Tuple(fields); - return tuple; - } - - return var; - } - - // Turn off memo for constant node. - Expr VisitExpr(const Expr& e) final { - if (e.as()) { - return ExprFunctor::VisitExpr(e); - } else { - return ExprMutator::VisitExpr(e); - } - } - - /*! - * \brief apply LazyGradientInit transformation and wrap function - * so that function type stays the same - * - * input/output types should only be a combination of TupleTypes and TensorTypes - */ - Expr Transform(const Expr& e) { - auto* f = e.as(); - auto* transformed = this->Mutate(e).as(); - - ICHECK(f); - ICHECK(transformed); - - if (e.same_as(GetRef(transformed))) { - return GetRef(transformed); - } - - auto tensorOutput = LetList::With([&](LetList* ll) { - // wrap inputs of Tensor type using InputVisitor class - tvm::Array args; - for (const Var& var : f->params) { - args.push_back(WrapExpr(var, var->checked_type(), ll)); - } - Expr transformedExpr = Call(GetRef(transformed), args); - // unwrap outputs of GradCell type into Tensor type using OutputVisitor class - return UnwrapExpr(ll->Push(transformedExpr), transformed->ret_type, ll); - }); - return Function(f->params, tensorOutput, f->ret_type, Array()); - } - - Expr VisitExpr_(const ConstantNode* op) final { - return Call(module_->GetConstructor("GradCell", "Raw"), {GetRef(op)}, Attrs(), - {op->checked_type()}); - } - - Expr VisitExpr_(const CallNode* call_node) final { - if (auto op = call_node->op.as()) { - Expr op_expr = op.value(); - - if (op_expr == Op::Get("add")) { - return CallGradCellFunction(call_node, module_->GetGlobalVar("AddGradCell")); - } - - if (op_expr == Op::Get("multiply")) { - return CallGradCellFunction(call_node, module_->GetGlobalVar("MultiplyGradCell")); - } - - if (op_expr == Op::Get("ones") || op_expr == Op::Get("zeros")) { - // ones and zeros need TensorType input - Expr result = CallPrimitiveOp(call_node); - Expr func = Function({}, result, {call_node->checked_type()}, Array()); - // call appropriate GradCell constructor - std::string constructor_name = op_expr == Op::Get("ones") ? "One" : "Zero"; - return Call(module_->GetConstructor("GradCell", constructor_name), {func}, Attrs(), - {call_node->checked_type()}); - } - - if (op_expr == Op::Get("ones_like") || op_expr == Op::Get("zeros_like")) { - // ones_like and zeros_like need TensorType input - Expr result = CallPrimitiveOp(call_node); - // fn() -> T, function returns result of operation - Expr func = Function({}, result, {call_node->checked_type()}, Array()); - // call appropriate GradCell constructor - std::string constructor_name = op_expr == Op::Get("ones_like") ? "One" : "Zero"; - return Call(module_->GetConstructor("GradCell", "One"), {func}, Attrs(), - {call_node->checked_type()}); - } - - // handle all other ops - Expr result = CallPrimitiveOp(call_node); - // wrap result with Raw constructor - return Call(module_->GetConstructor("GradCell", "Raw"), {result}, Attrs(), - {call_node->checked_type()}); - } - // not an op - return ExprMutator::VisitExpr_(call_node); - } - - Type VisitType(const Type& t) final { return TypeMutator::VisitType(t); } - - Type VisitType_(const TensorTypeNode* op) { - GlobalTypeVar gradCell = module_->GetGlobalTypeVar("GradCell"); - tvm::Array args; - args.push_back(GetRef(op)); - return TypeCall(gradCell, args); - } - - private: - // Module - IRModule module_; - - /*! - * \brief Convert call_node to add/multiply op to use overloaded functions for GradCell type - */ - Expr CallGradCellFunction(const CallNode* call_node, GlobalVar overloaded_op) { - // can only use overloaded functions if 2 arguments of same type - if (call_node->args.size() != 2 || - !tvm::StructuralEqual()(call_node->args[0]->checked_type(), - call_node->args[1]->checked_type())) { - Expr result = CallPrimitiveOp(call_node); - return Call(module_->GetConstructor("GradCell", "Raw"), {result}, Attrs(), - {call_node->checked_type()}); - } - - tvm::Array args; - // create "fallback" function for overloaded function - Type paramType = call_node->args[0]->checked_type(); - tvm::Array params = {Var("lhs", paramType), Var("rhs", paramType)}; - // use primitive op in this case - Expr callOp = Call(call_node->op, {params[0], params[1]}); - Expr func = Function(params, callOp, paramType, Array()); - - // pass "fallback" function and tensors as arguments - args.push_back(func); - for (Expr expr : call_node->args) { - args.push_back(VisitExpr(expr)); - } - // return new call to overloaded function - return Call(overloaded_op, args, Attrs(), {paramType}); - } - - /*! - * \brief Convert calls to other ops by converting args into TensorType - * \return call expr returning result of op - */ - Expr CallPrimitiveOp(const CallNode* call_node) { - const auto fromFunc = module_->GetGlobalVar("FromGradCell"); - tvm::Array args; - // use FromGradCell to convert args to Tensor - for (Expr expr : call_node->args) { - args.push_back(Call(fromFunc, {VisitExpr(expr)}, Attrs(), {expr->checked_type()})); - } - // result of operation - return Call(call_node->op, args, call_node->attrs); - } -}; - -Expr LazyGradientInit(const Expr& e, IRModule mod) { - CheckFeature(e, mod, FeatureSet::All() - fGraph); - auto ret = LazyGradientInitializer(mod).Transform(e); - CheckFeature(ret, mod, FeatureSet::All() - fGraph); - return ret; -} - -namespace transform { -Pass LazyGradientInit() { - runtime::TypedPackedFunc pass_func = - [=](Function f, IRModule m, PassContext pc) { - return Downcast(LazyGradientInit(f, m)); - }; - return CreateFunctionPass(pass_func, 2, "LazyGradientInit", {}); -} - -TVM_REGISTER_GLOBAL("relay._transform.LazyGradientInit").set_body_typed(LazyGradientInit); - -} // namespace transform - -} // namespace relay -} // namespace tvm diff --git a/src/relay/transforms/legalize.cc b/src/relay/transforms/legalize.cc deleted file mode 100644 index 7daa028bbcf3..000000000000 --- a/src/relay/transforms/legalize.cc +++ /dev/null @@ -1,112 +0,0 @@ -/* - * 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 legalize.cc - * \brief Converts an expr to another expr. This pass can be used to transform an op based on its - * shape, dtype or layout to another op or a sequence of ops. - */ - -#include -#include -#include -#include - -namespace tvm { -namespace relay { - -namespace legalize { - -// Call registered FTVMLegalize of an op -// Returns the legalized expression -class Legalizer : public ExprRewriter { - public: - explicit Legalizer(const std::string& legalize_map_attr_name) - : legalize_map_attr_name_{legalize_map_attr_name} {} - - Expr Rewrite_(const CallNode* call_node, const Expr& post) override { - // Get the new_call node without any changes to current call node. - Call new_call = Downcast(post); - - // Check if the string is registered. - if (!Op::HasAttrMap(legalize_map_attr_name_)) { - return post; - } - - // Collect the registered legalize function. - auto fop_legalize = Op::GetAttrMap(legalize_map_attr_name_); - auto call_op = call_node->op; - if (call_op.as()) { - Op op = Downcast(call_node->op); - - if (fop_legalize.count(op)) { - // Collect the new_args. - tvm::Array call_args = new_call->args; - - // Collect input and output dtypes to pass on to Legalize API. - tvm::Array types; - for (auto arg : call_node->args) { - types.push_back(arg->checked_type()); - } - types.push_back(call_node->checked_type()); - - // Transform the op by calling the registered legalize function. - Expr legalized_value = fop_legalize[op](call_node->attrs, call_args, types); - - // Return the new expr if the transformation succeeded. - if (legalized_value.defined()) { - // Check that the returned Expr from legalize is CallNode. - const CallNode* legalized_call_node = legalized_value.as(); - ICHECK(legalized_call_node) - << "Can only replace the original operator with another call node"; - return legalized_value; - } - } - } - - return post; - } - - private: - std::string legalize_map_attr_name_; -}; - -Expr Legalize(const Expr& expr, const std::string& legalize_map_attr_name) { - auto rewriter = Legalizer(legalize_map_attr_name); - return PostOrderRewrite(expr, &rewriter); -} - -} // namespace legalize - -namespace transform { - -Pass Legalize(const String& legalize_map_attr_name) { - runtime::TypedPackedFunc pass_func = - [=](Function f, IRModule m, PassContext pc) { - return Downcast(relay::legalize::Legalize(f, legalize_map_attr_name)); - }; - return CreateFunctionPass(pass_func, 1, "Legalize", {"InferType"}); -} - -TVM_REGISTER_GLOBAL("relay._transform.Legalize").set_body_typed(Legalize); - -} // namespace transform - -} // namespace relay -} // namespace tvm diff --git a/src/relay/transforms/let_list.h b/src/relay/transforms/let_list.h deleted file mode 100644 index f908fbcee514..000000000000 --- a/src/relay/transforms/let_list.h +++ /dev/null @@ -1,154 +0,0 @@ -/* - * 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 let_list.h - * \brief LetList record let binding and insert let expression implicitly. - * using it, one can treat AST as value instead of expression, - * and pass them around freely without fear of AST explosion (or effect duplication). - * for example, if one write 'b = a + a; c = b + b; d = c + c', the AST will contain 8 'a'. - * if one instead write 'b = ll.Push(a + a); c = ll.Push(b + b); d = ll.Get(c + c);', - * the AST will contain 2 'a', as b and c are now variables. - */ -#ifndef TVM_RELAY_TRANSFORMS_LET_LIST_H_ -#define TVM_RELAY_TRANSFORMS_LET_LIST_H_ - -#include -#include - -#include -#include -#include -#include - -#include "tvm/relay/type.h" - -namespace tvm { -namespace relay { - -/*! - * \brief LetList allow you to transform expression into variables, so you can copy them around. - * one can insert into the LetList by calling Push, and wrap an expression with bindings with Get. - * additionally, there is the 'With' function, which automatically call Get. - */ -class LetList { - public: - ~LetList() { - if (lets_.size() > 0 && !used_) { - LOG(WARNING) << "letlist not used"; - } - } - /*! - * \brief insert a binding. - * - * \param pv the var of the binding. - * - * \param expr the value of the binding. - * - * \return a Var that hold the inserted expr. - */ - Var Push(Var pv, Expr expr) { - ICHECK(!used_); - ICHECK(WellFormed(expr)) << "expression:" << std::endl << PrettyPrint(expr); - lets_.emplace_back(std::make_pair(pv, expr)); - return pv; - } - - /*! - * \brief insert a binding. - * - * \param expr the value of the binding. - * - * \param ty the type of the binding. - * - * \return a Var that hold the inserted expr. - */ - Var Push(Expr expr, Type ty) { return Push(Var::GenSym(ty), expr); } - - /*! - * \brief insert a binding. - * - * \param expr the value of the binding. - * - * \return a Var that hold the inserted expr. - */ - Var Push(Expr expr) { return Push(expr, Type()); } - - /*! - * \brief wrap an expr around the LetList. - * - * \param body the Expression to be wrapped around. - * - * \return the wrapped expr. - */ - Expr Get(const Expr& body) { - ICHECK(!used_); - Expr ret = body; - for (auto rit = lets_.rbegin(); rit != lets_.rend(); ++rit) { - ret = Let(std::get<0>(*rit), std::get<1>(*rit), ret); - } - used_ = true; - return ret; - } - - /*! \brief get the number of let bindings in the let list. - * - * \return the let list size. - */ - size_t size() const { return lets_.size(); } - - /*! \brief generate an LetList and wrap the result automatically. - * - * \param f a function that generate the unwrapped Expr. - * - * \code - * // Example code that generate `16 * a` using 4 plus instead of 15 plus. - * Expr mult_sixteen(const Var& a) { - * Op plus = Op::Get("plus"); - * // Automatically call Get with LetList::With - * return LetList::With([&](LetList* ll) { - * // Turn a call to plus into a variable to avoid duplication of code - * Var b = ll->Push(Call(plus, {a, a})); - * Var c = ll->Push(Call(plus, {b, b})); - * Var d = ll->Push(Callplus, {c, c})); - * return Call(plus, {d, d}); - * }); - * } - * \endcode - * - * \return the wrapped Expr. - */ - template - static Expr With(F&& f) { - LetList ll; - return ll.Get(f(&ll)); - } - - static Expr LetBind(const Expr& e, const std::function& f) { - return With([&](LetList* ll) { return f(ll->Push(e)); }); - } - - private: - std::vector> lets_; - bool used_ = false; -}; - -} // namespace relay -} // namespace tvm -#endif // TVM_RELAY_TRANSFORMS_LET_LIST_H_ diff --git a/src/relay/transforms/memory_alloc.cc b/src/relay/transforms/memory_alloc.cc deleted file mode 100644 index fcf8a784a9e7..000000000000 --- a/src/relay/transforms/memory_alloc.cc +++ /dev/null @@ -1,454 +0,0 @@ -/* - * 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 src/relay/transforms/memory_alloc.cc - * \brief A pass for manifesting explicit memory allocations. - */ - -#include -#include -#include -#include -#include -#include -#include -#include -#include -#include -#include -#include -#include - -#include -#include -#include -#include -#include - -#include "../backend/te_compiler.h" -#include "../backend/te_compiler_cache.h" -#include "../op/annotation/annotation.h" -#include "../op/call/call.h" -#include "../op/memory/device_copy.h" -#include "../op/memory/memory.h" -#include "../op/vm/vm.h" -#include "./device_aware_visitors.h" -#include "./let_list.h" -#include "./pass_utils.h" -#include "./pattern_utils.h" - -using namespace tvm::runtime; - -namespace tvm { -namespace relay { - -class DialectRewriter : public transform::DeviceAwareExprMutator { - public: - DialectRewriter(IRModule mod, VirtualDevice host_virtual_device) - : transform::DeviceAwareExprMutator(mod), - mod_(std::move(mod)), - host_virtual_device_(std::move(host_virtual_device)) {} - - Function Rewrite(const Function& expr) { return Downcast(Mutate(expr)); } - - private: - using ExprMutator::VisitExpr_; - - Expr VisitExpr_(const TupleNode* tuple_node) final { - LetList& scope = scopes_.back(); - Array new_fields; - new_fields.reserve(tuple_node->fields.size()); - - for (auto field : tuple_node->fields) { - auto new_field = Mutate(field); - if (const auto* op = new_field.as()) { - DataType dtype(op->data->dtype); - bool is_simple_const = (dtype == DataType::Int(32) || dtype == DataType::Int(64) || - dtype == DataType::Float(32) || dtype == DataType::Float(64) || - dtype == DataType::Bool()); - if (!op->is_scalar() || !is_simple_const) { - VirtualDevice virtual_device = GetVirtualDevice(field); - ICHECK(!virtual_device->IsFullyUnconstrained()); - Var const_var("const", Type(nullptr)); - new_field = scope.Push(const_var, MaybeOnDeviceFixed(new_field, virtual_device)); - } - } - new_fields.push_back(new_field); - } - return WithFields(GetRef(tuple_node), new_fields); - } - - void PreVisitLetBlock_(const LetNode* let_node) final { scopes_.emplace_back(); } - - std::pair PreVisitLetBinding_(const Var& var, const Expr& value) final { - Expr new_value = Mutate(value); - VirtualDevice virtual_device = GetVirtualDevice(value); - ICHECK(!virtual_device->IsFullyUnconstrained()); - scopes_.back().Push(var, MaybeOnDeviceFixed(new_value, virtual_device)); - // Since we always need a let block on which to bind sub-expressions the rewritten bindings - // are tracked in the current scopes. But return the rewritten binding anyway. - return {var, new_value}; - } - - Expr PostVisitLetBlock_(const LetNode* pre_let_node, const LetNode* post_let_node) final { - // The current scope has captured all the rewritten let-binding, as well as any additional - // bindings we needed to add. All we need is the rewritted body. - Expr new_body = post_let_node->body; - while (const auto* inner_let_node = new_body.as()) { - new_body = inner_let_node->body; - } - auto ret = scopes_.back().Get(new_body); - scopes_.pop_back(); - return ret; - } - - Expr DeviceAwareVisitExpr_(const CallNode* call_node) final { - DeviceCopyProps device_copy_props = GetDeviceCopyProps(call_node); - CallLoweredProps call_lowered_props = GetCallLoweredProps(call_node); - - if (device_copy_props.body.defined()) { - // Special case: device_copy calls remain in their original (and functional) form. - // TODO(mbs): device_copy cleanup. - return transform::DeviceAwareExprMutator::DeviceAwareVisitExpr_(call_node); - } - - if (!call_lowered_props.lowered_func.defined()) { - // This is a call to a user-defined Relay functinon, which will be handled directly by - // the VM and does not need conversion to DPS. - return transform::DeviceAwareExprMutator::DeviceAwareVisitExpr_(call_node); - } - - Call call = GetRef(call_node); - VLOG(1) << "converting lowered call to DPS:" << std::endl << PrettyPrint(call); - - VirtualDevice virtual_device = GetVirtualDevice(call); - ICHECK(!virtual_device->IsFullyUnconstrained()); - ICHECK(!scopes_.empty()) - << "Calls out of a let block are not supported, do you forget to transform " - << "with ToANormalForm or set opt_level >= 1 in the pass context?"; - LetList& scope = scopes_.back(); - - std::vector new_args; - for (const auto& arg : call_lowered_props.arguments) { - new_args.push_back(Mutate(arg)); - } - Tuple ins(new_args); - Type ret_type = call_node->checked_type_; - std::vector out_types = FlattenTupleType(ret_type); - - // Handle reshape. - // Original: - // reshape(body, ) - // dyn.reshape(body, shape, ) - // After FuseOps: - // let %f = fn(x, primitive=1, relay.reshape_only=1) { reshape(x, ) } - // %f(body) - // After LowerTEPass: - // call_lowered(@xxx_reshape, (body), dict[relay.reshape_only] = 1) - // -OR- - // call_lowered(@xxx_dyn_reshape, (body, shape), ) - // where @reshape_xxx is bound as a PrimFunc. - // (the name is irrelevant, only the relay.reshape_only attribute matters) - // After this pass: - // vm.reshape_tensor(body, shape, ) - if (IsReshapeOnly(call_lowered_props)) { - return EmitReshapeTensor(&scope, ins, call_lowered_props.attrs, ret_type); - } - - // At this point we could be calling a PrimFunc or an 'external' and already compiled primitive. - // The calling conventions are identical. - - // Handle 'dynamic' calls, ie to PrimFuncs whose result shape must be first computed - // by a companion shape function. - if (IsDynamic(ret_type)) { - return DynamicInvoke(&scope, call_lowered_props.lowered_func, ins, call_lowered_props.attrs, - out_types, ret_type, virtual_device); - } - - // Handle ordinary primitive calls. - Array outputs; - for (size_t i = 0; i < out_types.size(); ++i) { - outputs.push_back( - MakeStaticAllocation(&scope, out_types[i], virtual_device, std::to_string(i))); - } - Tuple outs(outputs); - Expr invoke = - InvokeTVMOp(call_lowered_props.lowered_func, ins, outs, - Downcast(call_lowered_props.attrs.metadata.at("relay_attrs"))); - scope.Push(MaybeOnDeviceFixed(invoke, virtual_device)); - return ToTupleType(ret_type, std::vector(outputs.begin(), outputs.end())); - } - - /*! - * \brief Returns the Relay Constant representing the 1d tensor with \p value. - * - * CAUTION: Make sure the constant ends up on the correct device. - */ - inline Constant MakeConstant(const std::vector& value) { - return MakeConstantTensor(DataType::Int(64), {static_cast(value.size())}, value); - } - - /*! Returns an \p alloc_tensor call for a tensor of \p shape and \p dtype over \p storage. */ - inline Expr AllocTensor(const Expr& storage, tvm::relay::Expr shape, DataType dtype, - Array assert_shape) { - Expr offset = - MaybeOnDeviceFixed(MakeConstantScalar(DataType::Int(64), 0), host_virtual_device_); - return tvm::relay::AllocTensor(storage, std::move(offset), std::move(shape), dtype, - assert_shape); - } - - Expr ComputeAlignment(const DataType& dtype) const { - int64_t align = dtype.bits() / 8 * dtype.lanes(); - if (align < 64) { - align = 64; - } - return MakeConstantScalar(DataType::Int(64), align); - } - - Expr ComputeStorageInRelay(const Expr& shape, const TensorType& type) const { - auto dtype = DataType(type->dtype); - Expr els = Prod(shape, Array(nullptr), false, false); - Expr num = MakeConstantScalar(DataType::Int(64), dtype.bits() * dtype.lanes()); - Expr add = Add(num, MakeConstantScalar(DataType::Int(64), 7)); - Expr div = MakeConstantScalar(DataType::Int(64), 8); - Expr ret = Multiply(els, Divide(add, div)); - return std::move(ret); - } - - Expr ComputeStorage(const TensorType& type) { - int64_t size = 1; - for (auto it : type->shape) { - auto val = it.as(); - CHECK(val); - size *= val->value; - } - size *= (type->dtype.bits() * type->dtype.lanes() + 7) / 8; - return std::move(MakeConstantScalar(DataType::Int(64), size)); - } - - // Allocate a tensor with a statically known shape. - Var MakeStaticAllocation(LetList* scope, const TensorType& type, - const VirtualDevice& virtual_device, String name_hint) { - std::vector int_shape; - for (auto it : type->shape) { - const auto* imm = it.as(); - CHECK(imm) << "expect static int shape"; - int_shape.push_back(imm->value); - } - Expr shape = MaybeOnDeviceFixed(MakeConstant(int_shape), host_virtual_device_); - Expr size = MaybeOnDeviceFixed(ComputeStorage(type), host_virtual_device_); - // Alignment is directly captured in the instruction rather than calculated, so we - // don't want to wrap it with an "on_device". - Expr alignment = ComputeAlignment(type->dtype); - // Run type inference later to get the correct type. - Var var("storage_" + name_hint, Type(nullptr)); - Expr value = AllocStorage(size, shape, alignment, virtual_device, type->dtype); - auto sto = scope->Push(var, MaybeOnDeviceFixed(value, virtual_device)); - - // TODO(@jroesch): There is a bug with typing based on the constant shape. - auto tensor = AllocTensor(sto, shape, type->dtype, /*assert_shape=*/type->shape); - Var tensor_var("tensor_" + name_hint, Type(nullptr)); - return scope->Push(tensor_var, MaybeOnDeviceFixed(tensor, virtual_device)); - } - - /*! - * \brief Appends to \p scope the computation necessary to call the shape function given - * in \p tir_call_attrs and bind the resulting result shapes into \p scope. The result - * shapes are for a call to a primitive with \p ins arguments. Some combinationn of the - * data and/or shapes of \p ins will be needed by the shape function. - */ - Array EmitShapeFunc(LetList* scope, const Tuple& ins, const CallLoweredAttrs& attrs) { - ICHECK(attrs.metadata.count("prim_shape_fn_states")); - Array input_states = - Downcast>(attrs.metadata.at("prim_shape_fn_states")); - ICHECK(attrs.metadata.count("prim_shape_fn_var")); - auto prim_fn_var = Downcast(attrs.metadata.at("prim_shape_fn_var")); - - const auto* func_type_node = prim_fn_var->checked_type().as(); - ICHECK(func_type_node); - - // Establish the arguments to the shape function. - Array shape_func_ins; - int input_pos = 0; - ICHECK_EQ(ins->fields.size(), input_states.size()); - for (size_t i = 0; i < ins->fields.size(); ++i) { - const Expr& arg = ins->fields[i]; - Type ty; - if (const auto* vn = arg.as()) { - ty = vn->type_annotation; - } else { - ty = arg->checked_type(); - } - int64_t state = input_states[i]->value; - // Pass Shapes - if (state == tec::kNeedInputShape) { - std::vector exprs = FromTupleType(ty, arg); - for (size_t j = 0; j < exprs.size(); ++j) { - Expr sh_of = Mutate(ShapeOf(exprs[j])); - Var in_shape_var("in_shape_" + std::to_string(input_pos + j), Type(nullptr)); - shape_func_ins.push_back( - scope->Push(in_shape_var, MaybeOnDeviceFixed(sh_of, host_virtual_device_))); - input_pos++; - } - } else if (state == tec::kNeedInputData) { - auto new_arg = Mutate(arg); // already accounts for device - VirtualDevice arg_virtual_device = GetVirtualDevice(arg); - ICHECK(!arg_virtual_device->IsFullyUnconstrained()); - // The dynamic shape function is expecting its data on the host/CPU, so insert a - // device_copy otherwise. (We'll need to fuse & lower these copies in the same way - // we fuse & lower other operators we insert for, eg, dynamic tensor size calculation.) - new_arg = MaybeDeviceCopy(MaybeOnDeviceFixed(new_arg, arg_virtual_device), - arg_virtual_device, host_virtual_device_); - Var in_shape_var("in_shape_" + std::to_string(input_pos), Type(nullptr)); - shape_func_ins.push_back( - scope->Push(in_shape_var, MaybeOnDeviceFixed(new_arg, host_virtual_device_))); - input_pos++; - } else { - // TODO(@jroesch): handle kNeedBoth - LOG(FATAL) << "unsupported shape function input state"; - } - } - ICHECK_EQ(shape_func_ins.size(), func_type_node->arg_types.size()); - - // Establish the result shapes. - const auto* res_tuple_node = func_type_node->ret_type.as(); - ICHECK(res_tuple_node); - - Array out_shapes; - for (size_t i = 0; i < res_tuple_node->fields.size(); ++i) { - const auto* tensor_type_node = res_tuple_node->fields[i].as(); - ICHECK(tensor_type_node); - // Put the shape func on the host. This also ensures that everything between - // shape_of and shape_func is similarly on the host. - Var alloc = MakeStaticAllocation(scope, GetRef(tensor_type_node), - host_virtual_device_, "out_shape_" + std::to_string(i)); - out_shapes.push_back(alloc); - } - - // Represent the call in DPS form. - auto shape_call = InvokeTVMOp(prim_fn_var, Tuple(shape_func_ins), Tuple(out_shapes), - Downcast(attrs.metadata.at("relay_attrs"))); - Var shape_func_var("shape_func", Type(nullptr)); - scope->Push(shape_func_var, MaybeOnDeviceFixed(shape_call, host_virtual_device_)); - return out_shapes; - } - - // Generate the code for invoking the TVM primitive \p func who's results have dynamic shapes. - Expr DynamicInvoke(LetList* scope, const Expr& func, const Tuple& ins, - const CallLoweredAttrs& attrs, const std::vector& out_types, - const Type& ret_type, const VirtualDevice& virtual_device) { - Array out_shapes = EmitShapeFunc(scope, ins, attrs); - std::vector storages; - CHECK_EQ(out_shapes.size(), out_types.size()); - for (size_t i = 0; i < out_shapes.size(); ++i) { - auto out_shape = out_shapes[i]; - auto out_type = out_types[i]; - auto size = - MaybeOnDeviceFixed(ComputeStorageInRelay(out_shape, out_type), host_virtual_device_); - // Alignment is directly captured in the instruction so don't wrap in "on_device". - auto alignment = ComputeAlignment(out_type->dtype); - Var sto_var("storage_" + std::to_string(i), Type(nullptr)); - auto val = AllocStorage(size, out_shape, alignment, virtual_device, out_type->dtype); - storages.push_back(scope->Push(sto_var, MaybeOnDeviceFixed(val, virtual_device))); - } - - Array outs; - for (size_t i = 0; i < storages.size(); ++i) { - auto out_shape = out_shapes[i]; - auto out_type = out_types[i]; - auto storage = storages[i]; - auto alloc = AllocTensor(storage, out_shape, out_type->dtype, out_type->shape); - Var out_var("out_" + std::to_string(i), Type(nullptr)); - outs.push_back(scope->Push(out_var, MaybeOnDeviceFixed(alloc, virtual_device))); - } - - Tuple tuple_outs(outs); - auto call = - InvokeTVMOp(func, ins, tuple_outs, Downcast(attrs.metadata.at("relay_attrs"))); - scope->Push(MaybeOnDeviceFixed(call, virtual_device)); - return ToTupleType(ret_type, - std::vector(tuple_outs->fields.begin(), tuple_outs->fields.end())); - } - - Expr EmitReshapeTensor(LetList* scope, const Tuple& ins, const CallLoweredAttrs& attrs, - const Type& ret_type) { - ICHECK_GE(ins->fields.size(), 1); // static reshape - ICHECK_LE(ins->fields.size(), 2); // dynamic reshape, second arg is shape - TensorType ret_ty = Downcast(ret_type); - Expr shape_expr; - if (IsDynamic(ret_type)) { - // Even though the desired output shape has been passed as the second argument to - // the dyn.reshape primitive, we'll still call that primitive's shape function. Go figure. - Array out_shapes = EmitShapeFunc(scope, ins, attrs); - ICHECK_EQ(out_shapes.size(), 1); - shape_expr = out_shapes[0]; - } else { - std::vector shape; - for (const auto& it : ret_ty->shape) { - const auto* imm = it.as(); - CHECK(imm) << "expect static int shape"; - shape.push_back(imm->value); - } - shape_expr = MaybeOnDeviceFixed(MakeConstant(shape), host_virtual_device_); - } - return ReshapeTensor(ins->fields[0], shape_expr, ret_ty->shape); - } - - private: - const Op& device_copy_op_ = Op::Get("device_copy"); - runtime::DataType compute_dtype_ = runtime::DataType::Int(64); - IRModule mod_; - VirtualDevice host_virtual_device_; - - std::vector scopes_; -}; - -namespace transform { - -Pass ManifestAllocImportStorage() { - auto pass_func = [](IRModule mod, tvm::transform::PassContext pass_cnxt) { - mod.CopyOnWrite(); - mod->ImportFromStd("core.rly"); - return mod; - }; - return tvm::transform::CreateModulePass(pass_func, /*opt_level=*/0, "ManifestAllocImportStorage", - /*required=*/{}); -} - -Pass ManifestAllocImpl(VirtualDevice host_virtual_device) { - auto pass_func = [host_virtual_device](Function func, IRModule mod, PassContext ctxt) { - return DialectRewriter(mod, host_virtual_device).Rewrite(func); - }; - return CreateFunctionPass(pass_func, 0, "ManifestAllocImpl", {}); -} - -Pass ManifestAlloc(VirtualDevice cpu_virtual_device) { - std::vector passes = {ManifestAllocImportStorage(), InferType(), - ManifestAllocImpl(std::move(cpu_virtual_device)), InferType()}; - return Sequential(passes, "ManifestAlloc"); -} - -TVM_REGISTER_GLOBAL("relay.transform.ManifestAlloc").set_body_typed(ManifestAlloc); - -} // namespace transform - -} // namespace relay -} // namespace tvm diff --git a/src/relay/transforms/merge_compiler_regions.cc b/src/relay/transforms/merge_compiler_regions.cc deleted file mode 100644 index 92e881fa6100..000000000000 --- a/src/relay/transforms/merge_compiler_regions.cc +++ /dev/null @@ -1,230 +0,0 @@ -/* - * 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 src/relay/transforms/merge_compiler_regions.cc - * - * \brief After operators have been annotated with the targets that support - * them, this pass creates regions of the operators for each target. It - * is guaranteed that the regions will have a topological ordering so that - * no data dependency issues exist. - * - * This pass only introduces annotations to indicate the regions. - * partition_graph must subsequently be called to lift these regions out - * as external functions. - */ - -#include -#include -#include -#include -#include -#include - -#include -#include -#include -#include - -#include "../analysis/annotated_region_set.h" -#include "pass_utils.h" - -namespace tvm { -namespace relay { -namespace merge_compiler_region { - -class RegionMerger : public MixedModeVisitor { - public: - explicit RegionMerger(AnnotatedRegionSet regions) : regions_(regions) {} - - void find_control_flow_regions( - const Expr op, - std::unordered_set& correlative_regions) { - // Find correlative restriction regions from control flow. - - // In IfNode, find from condition, true_branch and false branch. - const IfNode* if_node = op.as(); - if (if_node) { - auto cond_region = regions_->GetRegion(if_node->cond); - auto true_branch_region = regions_->GetRegion(if_node->true_branch); - auto false_branch_region = regions_->GetRegion(if_node->false_branch); - if (cond_region.defined()) { - correlative_regions.insert(cond_region); - } else { - find_control_flow_regions(if_node->cond, correlative_regions); - } - if (true_branch_region.defined()) { - correlative_regions.insert(true_branch_region); - } else { - find_control_flow_regions(if_node->true_branch, correlative_regions); - } - if (false_branch_region.defined()) { - correlative_regions.insert(false_branch_region); - } else { - find_control_flow_regions(if_node->false_branch, correlative_regions); - } - } - } - - void VisitExpr_(const CallNode* call) final { - if (call->op == CompilerEndOp()) { - auto region = regions_->GetRegion(GetRef(call)); - - // Skip this region if it has been merged to the other region. - if (merged_regions_.find(region->GetID()) != merged_regions_.end()) { - return; - } - - // Check the region target. - auto compiler_attrs = call->attrs.as(); - ICHECK_EQ(region->GetTarget(), compiler_attrs->compiler); - - // Visit the unmerged parent regions. - for (const auto& arg : region->GetInputs()) { - // Region inputs must be begin annotation, and the region of - // the begin annotation's argument is the parent region. - auto begin = Downcast(arg); - ICHECK_EQ(begin->op, CompilerBeginOp()); - auto parent_region = regions_->GetRegion(begin->args[0]); - - // Skip this region if it has been merged. - if (!parent_region.defined()) { - continue; - } else if (merged_regions_.find(parent_region->GetID()) == merged_regions_.end()) { - VisitExpr(begin->args[0]); - } - } - - // Collect unmerged parent regions. - std::unordered_set mergeable_regions; - // Collect correlative regions to propagate restrictions. - std::unordered_set correlative_regions; - for (const auto& arg : region->GetInputs()) { - auto begin = Downcast(arg); - ICHECK_EQ(begin->op, CompilerBeginOp()); - auto parent_region = regions_->GetRegion(begin->args[0]); - if (parent_region.defined()) { - mergeable_regions.insert(parent_region); - correlative_regions.insert(parent_region); - } else { - find_control_flow_regions(begin->args[0], correlative_regions); - } - } - - // Propogate all the parent restrictions to the current region. - auto& region_restrictions = region_restrictions_[region->GetID()]; - for (const auto& parent_region : correlative_regions) { - auto parent_restrictions = region_restrictions_[parent_region->GetID()]; - region_restrictions.insert(parent_restrictions.begin(), parent_restrictions.end()); - } - - for (const auto& parent_region : mergeable_regions) { - // Skip the parent region with a different target. - if (parent_region->GetTarget() != compiler_attrs->compiler) { - region_restrictions.insert(parent_region->GetID()); - continue; - } - - // Skip the parent region if it is in the restriction set. - if (region_restrictions.find(parent_region->GetID()) != region_restrictions.end()) { - continue; - } - - // Merge the parent region to the current one. - regions_->MergeRegions(parent_region, region); - - // Replace the parent region ID with the current region for all - // other regions' restriction sets. - for (const auto& r : regions_) { - auto& restrictions = region_restrictions_[r->GetID()]; - if (restrictions.find(parent_region->GetID()) != restrictions.end()) { - restrictions.erase(parent_region->GetID()); - restrictions.insert(region->GetID()); - } - } - } - merged_regions_.insert(region->GetID()); - } - } - - private: - AnnotatedRegionSet regions_; - std::unordered_set merged_regions_; - std::unordered_map> region_restrictions_; -}; - -class MergeAnnotations : public ExprRewriter { - public: - explicit MergeAnnotations(AnnotatedRegionSet regions) : regions_(regions) {} - - Expr Rewrite_(const CallNode* call, const Expr& post) final { - // Merge annotations which are now internal to a region. - // This happens if we see a compiler begin next to a - // compiler end and they're both in the same region. - if (call->op == CompilerBeginOp() && call->args[0]->IsInstance()) { - auto arg = Downcast(call->args[0]); - if (arg->op == CompilerEndOp()) { - auto region1 = regions_->GetRegion(GetRef(call)); - auto region2 = regions_->GetRegion(arg); - if (region1 == region2) { - auto post_arg = post.as()->args[0]; - return post_arg.as()->args[0]; - } - } - } - return post; - } - - private: - AnnotatedRegionSet regions_; -}; - -Expr MergeCompilerRegions(const Expr& expr) { - // Create regions using the annotations. - AnnotatedRegionSet regions = AnnotatedRegionSet::Create(expr, CompilerBeginOp(), CompilerEndOp()); - - // Analyze the graph to explore the opportunities of merging regions. - RegionMerger merger(regions); - merger.VisitExpr(expr); - - // Remove annotations that are not in the region boundaries. - MergeAnnotations merge_anno(regions); - return PostOrderRewrite(expr, &merge_anno); -} - -} // namespace merge_compiler_region - -namespace transform { - -Pass MergeCompilerRegions() { - runtime::TypedPackedFunc part_func = - [=](Function f, IRModule m, PassContext pc) { - return Downcast(merge_compiler_region::MergeCompilerRegions(f)); - }; - auto merged = CreateFunctionPass(part_func, 0, "MergeCompilerRegions", {}); - return Sequential({merged, InferType()}); -} - -TVM_REGISTER_GLOBAL("relay._transform.MergeCompilerRegions") - .set_body_typed(transform::MergeCompilerRegions); - -} // namespace transform - -} // namespace relay -} // namespace tvm diff --git a/src/relay/transforms/merge_composite.cc b/src/relay/transforms/merge_composite.cc deleted file mode 100644 index 51f1387fd9ca..000000000000 --- a/src/relay/transforms/merge_composite.cc +++ /dev/null @@ -1,89 +0,0 @@ -/* - * 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 src/relay/transforms/merge_composite.cc - * \brief Merges expressions matching patterns into functions marked - * as 'composite'. This is primarily intended to be used alongside the - * external codegen infrastructure to support the case where multiple - * Relay operators map to a single external operator. - */ - -#include -#include -#include -#include -#include -#include - -namespace tvm { -namespace relay { -namespace merge_composite { - -Function InferType(const Function& expr, const IRModule& m) { - IRModule mod(m); - mod->Update(mod->GetGlobalVar("main"), expr); - mod = transform::InferType()(mod); - return Downcast(mod->Lookup("main")); -} - -Expr MergeComposite(const Function& func, const Array& pattern_names, - const Array& patterns, const std::vector& checks, - const IRModule& m) { - ICHECK_EQ(pattern_names.size(), patterns.size()); - Function merged_func = func; - // merge the patterns one-by-one in order - for (size_t i = 0; i < patterns.size(); i++) { - Map attrs; - attrs.Set("Composite", pattern_names[i]); - merged_func = Downcast(PartitionPattern(patterns[i], merged_func, attrs, checks[i])); - merged_func = InferType(merged_func, m); - } - return std::move(merged_func); -} - -} // namespace merge_composite - -namespace transform { - -Pass MergeComposite(const tvm::Array& pattern_names, - const tvm::Array& patterns, const std::vector& checks) { - runtime::TypedPackedFunc pass_func = - [=](Function f, IRModule m, PassContext pc) { - return Downcast( - relay::merge_composite::MergeComposite(f, pattern_names, patterns, checks, m)); - }; - auto func_pass = CreateFunctionPass(pass_func, 0, "MergeComposite", {}); - return func_pass; -} - -TVM_REGISTER_GLOBAL("relay._transform.MergeComposite").set_body([](TVMArgs args, TVMRetValue* rv) { - tvm::Array pattern_names = args[0]; - tvm::Array patterns = args[1]; - std::vector checks; - for (int i = 2; i < args.size(); i++) { - checks.push_back(args[i]); - } - *rv = MergeComposite(pattern_names, patterns, checks); -}); - -} // namespace transform - -} // namespace relay -} // namespace tvm diff --git a/src/relay/transforms/meta_schedule_layout_rewrite.cc b/src/relay/transforms/meta_schedule_layout_rewrite.cc deleted file mode 100644 index 1ae6a62629dc..000000000000 --- a/src/relay/transforms/meta_schedule_layout_rewrite.cc +++ /dev/null @@ -1,181 +0,0 @@ -/* - * 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 "./meta_schedule_layout_rewrite.h" - -#include -#include -#include -#include - -#include -#include -#include - -#include "../backend/te_compiler.h" - -namespace tvm { -namespace relay { - -class LayoutIndexQueue { - public: - static LayoutIndexQueue* Global() { - static LayoutIndexQueue inst; - return &inst; - } - - void Clear() { - std::lock_guard lock(mutex_); - queue_.clear(); - } - - private: - friend class MetaScheduleLayoutRewriter; - std::mutex mutex_; - std::deque queue_; -}; - -void MetaScheduleLayoutRewriter::LayoutQueuePush(const tir::IndexMap& index_map) { - LayoutIndexQueue* self = LayoutIndexQueue::Global(); - { - std::lock_guard lock(self->mutex_); - self->queue_.push_back(index_map); - } -} - -bool IsSupportedOp(const OpNode* op) { - static std::vector target_ops{ - "nn.conv2d", // - "nn.contrib_conv2d_winograd_without_weight_transform", - "nn.conv3d", - "nn.matmul", - "nn.dense", - "nn.batch_matmul", - }; - return std::find(target_ops.begin(), target_ops.end(), op->name) != target_ops.end(); -} - -#define TVM_RELAY_LAYOUT_WITH_ORIGINAL_SHAPE(Attr, AttrType, OriginalShape, Result) \ - if (const AttrType* ptr = Attr.as()) { \ - ObjectPtr n = make_object(*ptr); \ - n->meta_schedule_original_shape = OriginalShape; \ - Result = Attrs(n); \ - } - -// Mutate ops in a function -class MetaScheduleFuncMutator : public ExprMutator { - public: - explicit MetaScheduleFuncMutator(std::deque&& layout_queue) - : layout_queue_(std::move(layout_queue)) {} - - Expr VisitExpr_(const CallNode* call) { - Expr expr = ExprMutator::VisitExpr_(call); - if (layout_queue_.empty()) { - return expr; - } - if (const auto* call = expr.as()) { - if (const auto* op = call->op.as()) { - if (IsSupportedOp(op)) { - ICHECK_EQ(call->args.size(), 2); - tir::IndexMap index_map = layout_queue_.front(); - layout_queue_.pop_front(); - Array shape; - if (call->args[1]->IsInstance()) { - Var var = Downcast(call->args[1]); - shape = Downcast(var->type_annotation)->shape; - } else if (const ConstantNode* cnst = call->args[1].as()) { - shape = cnst->tensor_type()->shape; - } else { - LOG(FATAL) << "Unexpected input " << call->args[1]; - } - Attrs attrs{nullptr}; - TVM_RELAY_LAYOUT_WITH_ORIGINAL_SHAPE(call->attrs, Conv2DAttrs, shape, attrs); - TVM_RELAY_LAYOUT_WITH_ORIGINAL_SHAPE(call->attrs, Conv2DWinogradAttrs, shape, attrs); - TVM_RELAY_LAYOUT_WITH_ORIGINAL_SHAPE(call->attrs, Conv3DAttrs, shape, attrs); - TVM_RELAY_LAYOUT_WITH_ORIGINAL_SHAPE(call->attrs, MatmulAttrs, shape, attrs); - TVM_RELAY_LAYOUT_WITH_ORIGINAL_SHAPE(call->attrs, DenseAttrs, shape, attrs); - TVM_RELAY_LAYOUT_WITH_ORIGINAL_SHAPE(call->attrs, BatchMatmulAttrs, shape, attrs); - ICHECK(attrs.defined()) << "TypeError: Unknown attribute: " << call->attrs; - expr = Call(call->op, - {call->args[0], MakeMetaScheduleLayoutTransform(call->args[1], index_map)}, - attrs); - } - } - } - return expr; - } - - private: - std::deque layout_queue_; -}; - -#undef TVM_RELAY_LAYOUT_WITH_ORIGINAL_SHAPE - -Expr MetaScheduleLayoutRewriter::VisitExpr_(const CallNode* call) { - Expr expr = ExprMutator::VisitExpr_(call); - call = expr.as(); - if (call != nullptr) { - if (const auto* func = call->op.as()) { - LayoutIndexQueue* self = LayoutIndexQueue::Global(); - self->queue_.clear(); - tec::PrimFuncFor(GetRef(func), Target::Current()); - if (!self->queue_.empty()) { - std::deque queue = std::move(self->queue_); - self->queue_.clear(); - return MetaScheduleFuncMutator(std::move(queue)).VisitExpr(expr); - } - } - } - return expr; -} - -namespace transform { - -Pass MetaScheduleLayoutRewrite() { - runtime::TypedPackedFunc pass_func = - [=](Function f, IRModule m, PassContext pc) -> Function { - return Downcast(MetaScheduleLayoutRewriter().Mutate(std::move(f))); - }; - return CreateFunctionPass(pass_func, 3, "MetaScheduleLayoutRewrite", {"InferType"}); -} - -#define TVM_RELAY_META_SCHEDULE_LAYOUT_REWRITE_GET_ORIGINAL_SHAPE(Attrs, AttrType) \ - if (const auto* p = Attrs.as()) { \ - return p->meta_schedule_original_shape; \ - } - -TVM_REGISTER_GLOBAL("relay.attrs.get_meta_schedule_original_shape") - .set_body_typed([](const Attrs& attrs) -> Array { - TVM_RELAY_META_SCHEDULE_LAYOUT_REWRITE_GET_ORIGINAL_SHAPE(attrs, Conv2DAttrs); - TVM_RELAY_META_SCHEDULE_LAYOUT_REWRITE_GET_ORIGINAL_SHAPE(attrs, Conv2DWinogradAttrs); - TVM_RELAY_META_SCHEDULE_LAYOUT_REWRITE_GET_ORIGINAL_SHAPE(attrs, Conv3DAttrs); - TVM_RELAY_META_SCHEDULE_LAYOUT_REWRITE_GET_ORIGINAL_SHAPE(attrs, MatmulAttrs); - TVM_RELAY_META_SCHEDULE_LAYOUT_REWRITE_GET_ORIGINAL_SHAPE(attrs, DenseAttrs); - TVM_RELAY_META_SCHEDULE_LAYOUT_REWRITE_GET_ORIGINAL_SHAPE(attrs, BatchMatmulAttrs); - LOG(FATAL) << "TypeError: Unknown attribute: " << attrs; - throw; - }); -TVM_REGISTER_GLOBAL("relay._transform.MetaScheduleLayoutRewrite") - .set_body_typed(MetaScheduleLayoutRewrite); - -#undef TVM_RELAY_META_SCHEDULE_LAYOUT_REWRITE_GET_ORIGINAL_SHAPE - -} // namespace transform -} // namespace relay -} // namespace tvm diff --git a/src/relay/transforms/meta_schedule_layout_rewrite.h b/src/relay/transforms/meta_schedule_layout_rewrite.h deleted file mode 100644 index f60df9b3e2ee..000000000000 --- a/src/relay/transforms/meta_schedule_layout_rewrite.h +++ /dev/null @@ -1,38 +0,0 @@ -/* - * 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. - */ -#ifndef TVM_RELAY_TRANSFORMS_META_SCHEDULE_LAYOUT_REWRITE_H_ -#define TVM_RELAY_TRANSFORMS_META_SCHEDULE_LAYOUT_REWRITE_H_ - -#include -#include - -namespace tvm { -namespace relay { - -class MetaScheduleLayoutRewriter : public ExprMutator { - public: - Expr VisitExpr_(const CallNode* n) final; - - static void LayoutQueuePush(const tir::IndexMap& index_map); -}; - -} // namespace relay -} // namespace tvm - -#endif // TVM_RELAY_TRANSFORMS_META_SCHEDULE_LAYOUT_REWRITE_H_ diff --git a/src/relay/transforms/partial_eval.cc b/src/relay/transforms/partial_eval.cc deleted file mode 100644 index c574a5772c16..000000000000 --- a/src/relay/transforms/partial_eval.cc +++ /dev/null @@ -1,1199 +0,0 @@ -/* - * 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 partial_eval.cc - * - * \brief Perform known computation in compile time. - * - * The partial evaluator try to do computation at compile time, - * so it can generate code that do less work. - * Additionally, it might open more chance for further optimization, - * since the high level, structural part of the code (closure, reference, control flow) - * might get partially evaluated away, and the subsequent optimization (for example, kernel fusion) - * can reason across those structural code as it got removed. - * In the extreme case, partial evaluation can even turn the whole program - * into pure first order computation with no control flow. - * In such a case, we can compile the whole computation onto SIMD Instruction/GPU/FPGA, - * and get huge speedup. - * - * It works by making the following modifications to the standard relay interpreter: - * - * 0: The values become partially static value. - * Since we cannot know the value of every term at compile time, - * Term might get partially evaluated to 'Unknown Value'. - * Every partially static value is, hence, - * a static fragment that might not be there (partially static), - * and a dynamic fragment that is semantically equivalent to the original term, - * so the unknown part will be computed at runtime, using the dynamic fragment. - * - * 1: The interpreter holds a LetList, which preserves A Normal Form for the generated code. - * More specifically, we require that all dynamic is an atom. - * This avoids code duplication (which is both inefficient and incorrect), as atom has constant size - * and allow us to not handle capture-avoidance substitution (as atom has no binder). - * - * 2: The map of References to partially static values is reified, as described below. - * Instead of Reference having mutable field, Reference only has an unique identifier. - * There will be a mutable mapping of id to partially static value, called the store. - * This allow us to rollback the store: - * when a path may or may not be executed (as in a conditional), we copy the store, - * recurse with the copy, and reinstate the original when the call returns - * so that the effects of the computation are not preserved. - * We do this in if else, pattern matching, and in function, - * as, when we see a function, we partially evaluate it with all the argument as dynamic, - * to generate efficient dynamic for that function. - * - * 3: The generated code reuses bindings (although they are not shadowed), - * so we have to deduplicate them. - * - * 4: In the generated code, as it call TypeSubst, multiple VarNode might have same Id. - * While it is permitted, most pass use ObjectPtrHash for Var, - * and having multiple VarNode for same Id break them. - * Thus we remap them to a single Id for now. - * - * Also, It will also generate lots of dead code, - * so it is a good idea to feed it through the dead code eliminator after partial evaluation. - * - * The partial evaluator makes several assumptions, so there is room for improvement: - * - * 0: Every time an unknown effect happened, we clear the whole store. - * It is too conservative: if a local reference is created (and do not get passed outside), - * An unknown global function call/global reference write can not modify it. - * We can pair PE with escape analysis/alias analysis. - * - * 1: We assume all unknown code has effect. Doing effect analysis can make the store more precise. - * - * 2: When doing pattern matching, we can simplify the match even for dynamic case. - * Right now it is all or nothing: either a complete match, or the original dynamic code. - * Instead, we can get a match tree, pair it with the data and evaluate it to a normal form. - * We then can reify the result. - * - * 3: Every time a function is called, its code will get expanded and partially evaluated. - * We can do a binding time analysis to cache the result and avoid re-partial evaluation. - * - * These assumptions do not affect the correctness of the algorithm, however. - */ -#include -#include -#include -#include -#include -#include -#include - -#include "let_list.h" -#include "pass_utils.h" - -namespace tvm { -namespace relay { -namespace partial_eval { - -using namespace runtime; - -/*! \brief Hash Var by it's id. - * Different VarNode might has same vid, and they are considered to be the same var in such case. - * Use VarHash to hash Var by id. - */ -struct VarHash { - size_t operator()(const Var& v) const { return ObjectPtrHash()(v->vid); } -}; - -/*! \brief Compare Var by it's id. - * Different VarNode might has same vid, and they are considered to be the same var in such case. - * Use VarEqual to compare Var by id. - */ -struct VarEqual { - bool operator()(const Var& l, const Var& r) const { return l->vid.get() == r->vid.get(); } -}; - -Expr PostProcess(const Expr&); - -/*! \brief A StaticNode contains some static data that the Partial Evaluator can use. */ -class StaticNode : public RelayNode { - public: - static constexpr const char* _type_key = "relay.Static"; - TVM_DECLARE_BASE_OBJECT_INFO(StaticNode, RelayNode); -}; - -class Static : public ObjectRef { - public: - Static() {} - explicit Static(ObjectPtr n) : ObjectRef(n) {} - const StaticNode* operator->() const { return static_cast(get()); } - - using ContainerType = StaticNode; -}; - -using Time = size_t; - -struct PStaticNode : Object { - static Time time() { - static Time time_ = 0; - Time ret = time_; - time_++; - return ret; - } - Static pstatic; // may be null - Expr dynamic; - Time created_time; - PStaticNode(const Static& pstatic, const Expr& dynamic) - : pstatic(pstatic), dynamic(dynamic), created_time(time()) {} - explicit PStaticNode(const Expr& dynamic) : PStaticNode(Static(), dynamic) {} - static constexpr const char* _type_key = "relay.PStatic"; - TVM_DECLARE_FINAL_OBJECT_INFO(PStaticNode, Object); -}; - -class PStatic : public ObjectRef { - public: - TVM_DEFINE_OBJECT_REF_METHODS(PStatic, ObjectRef, PStaticNode); -}; - -struct STupleNode : StaticNode { - std::vector fields; - explicit STupleNode(const std::vector& fields) : fields(fields) {} - static constexpr const char* _type_key = "relay.STuple"; - TVM_DECLARE_FINAL_OBJECT_INFO(STupleNode, StaticNode); -}; - -class STuple : public Static { - public: - TVM_DEFINE_OBJECT_REF_METHODS(STuple, Static, STupleNode); -}; - -Static MkSTuple(const std::vector& fields) { - return Static(make_object(fields)); -} - -struct STensorNode : StaticNode { - runtime::NDArray data; - explicit STensorNode(const NDArray& data) : data(data) {} - static constexpr const char* _type_key = "relay.STensor"; - TVM_DECLARE_FINAL_OBJECT_INFO(STensorNode, StaticNode); -}; - -class STensor : public Static { - public: - TVM_DEFINE_OBJECT_REF_METHODS(STensor, Static, STensorNode); -}; - -Static MkSTensor(const NDArray& data) { return Static(make_object(data)); } - -struct SConstructorNode : StaticNode { - Constructor constructor; - std::vector fields; - SConstructorNode(const Constructor& constructor, const std::vector& fields) - : constructor(constructor), fields(fields) {} - static constexpr const char* _type_key = "relay.SConstructor"; - TVM_DECLARE_FINAL_OBJECT_INFO(SConstructorNode, StaticNode); -}; - -class SConstructor : public Static { - public: - TVM_DEFINE_OBJECT_REF_METHODS(SConstructor, Static, SConstructorNode); -}; - -Static MkSConstructor(const Constructor& constructor, const std::vector& fields) { - return Static(make_object(constructor, fields)); -} - -struct SRefNode : StaticNode { - static constexpr const char* _type_key = "relay.SRef"; - // we will use the address as the guid for hashing - TVM_DECLARE_FINAL_OBJECT_INFO(SRefNode, StaticNode); -}; - -class SRef : public Static { - public: - TVM_DEFINE_OBJECT_REF_METHODS(SRef, Static, SRefNode); -}; - -Static MkSRef() { return Static(make_object()); } - -using Func = std::function&, const Attrs&, - const Array&, LetList*)>; - -struct SFuncNode : StaticNode { - Func func; - explicit SFuncNode(const Func& func) : func(func) {} - static constexpr const char* _type_key = "relay.SFunc"; - TVM_DECLARE_FINAL_OBJECT_INFO(SFuncNode, StaticNode); -}; - -class SFunc : public Static { - public: - TVM_DEFINE_OBJECT_REF_METHODS(SFunc, Static, SFuncNode); -}; - -Static MkSFunc(const Func& func) { return Static(make_object(func)); } - -class FuelNode; -/*! \brief A meet-semilattice with finite descending chain. - * It means that we can meet two element to get an element, - * and for every element, there is only a finite amount of meet before getting back the same - * element. - * - * Every time we recurse, we do a meet and require that progress must be made. - * This ensures we do not recurse infinitely in the Partial Evaluator. - */ -class Fuel : public ObjectRef { - public: - Fuel() {} - explicit Fuel(ObjectPtr n) : ObjectRef(n) {} - const FuelNode* operator->() const; - - using ContainerType = FuelNode; -}; - -class FuelNode : public RelayNode { - public: - virtual ~FuelNode() {} - // Please implement one of the following function or there will be infinite loop. - /*! \brief return the new Fuel, and whether progress is made. - * - * Note that progress is not symmetric - it only measure progress for (*this). - * - * Thus, if the generated is smaller then the argument of Meet, - * and the generated is not smaller then (*this), - * progress should be false. - */ - virtual std::tuple Meet(const Fuel& f) const { - bool progress = false; - auto ret = Meet(f, &progress); - return std::make_tuple(ret, progress); - } - /*! \brief return the new Fuel, and write (*progress | is progress made) to *progress. */ - virtual Fuel Meet(const Fuel& f, bool* progress) const { - ICHECK(progress); - auto ret = Meet(f); - *progress |= std::get<1>(ret); - return std::get<0>(ret); - } - static constexpr const char* _type_key = "relay.Fuel"; - TVM_DECLARE_BASE_OBJECT_INFO(FuelNode, RelayNode); -}; - -const FuelNode* Fuel::operator->() const { return static_cast(get()); } - -Fuel MkFSeq(const std::vector& fuels); -struct FSeqNode : FuelNode { - std::vector fuels; - Fuel Meet(const Fuel& f, bool* progress) const final { - auto x = f.as(); - ICHECK(x); - ICHECK_EQ(fuels.size(), x->fuels.size()); - std::vector new_fuels; - for (size_t i = 0; i < fuels.size(); ++i) { - new_fuels.push_back(fuels[i]->Meet(x->fuels[i], progress)); - } - return MkFSeq(new_fuels); - } - explicit FSeqNode(const std::vector& fuels) : fuels(fuels) {} - static constexpr const char* _type_key = "relay.FSeq"; - TVM_DECLARE_FINAL_OBJECT_INFO(FSeqNode, FuelNode); -}; - -class FSeq : public Fuel { - public: - TVM_DEFINE_OBJECT_REF_METHODS(FSeq, Fuel, FSeqNode); -}; - -Fuel MkFSeq(const std::vector& fuels) { return Fuel(make_object(fuels)); } - -Fuel MkFTime(Time time); -struct FTimeNode : FuelNode { - Time time; - std::tuple Meet(const Fuel& f) const final { - auto x = f.as(); - ICHECK(x); - Time new_time = std::min(time, x->time); - return std::make_tuple(MkFTime(new_time), new_time < time); - } - explicit FTimeNode(Time time) : time(time) {} - static constexpr const char* _type_key = "relay.FTime"; - TVM_DECLARE_FINAL_OBJECT_INFO(FTimeNode, FuelNode); -}; - -class FTime : public Fuel { - public: - TVM_DEFINE_OBJECT_REF_METHODS(FTime, Fuel, FTimeNode); -}; - -Fuel MkFTime(Time time) { return Fuel(make_object(time)); } - -Fuel MkFTValue(size_t tvalue); -/*! \brief If the pstatic is hold a positive integer scalar, that number, else 0. */ -struct FTValueNode : FuelNode { - size_t tvalue; - std::tuple Meet(const Fuel& f) const final { - auto x = f.as(); - ICHECK(x); - size_t new_tvalue = std::min(tvalue, x->tvalue); - return std::make_tuple(MkFTValue(new_tvalue), new_tvalue < tvalue); - } - explicit FTValueNode(size_t tvalue) : tvalue(tvalue) {} - static constexpr const char* _type_key = "relay.FTValue"; - TVM_DECLARE_FINAL_OBJECT_INFO(FTValueNode, FuelNode); -}; - -class FTValue : public Fuel { - public: - TVM_DEFINE_OBJECT_REF_METHODS(FTValue, Fuel, FTValueNode); -}; - -Fuel MkFTValue(size_t tvalue) { return Fuel(make_object(tvalue)); } - -/*! \brief Initially every element has Fuel of FTop. It is the largest element. - * - * Note that it is illegal to has FTop inside some other Fuel - - * doing so break the finite descending chain property. - */ -struct FTopNode : FuelNode { - std::tuple Meet(const Fuel& f) const final { - return std::make_tuple(f, !f.as()); - } - static constexpr const char* _type_key = "relay.FTop"; - TVM_DECLARE_FINAL_OBJECT_INFO(FTopNode, FuelNode); -}; - -class FTop : public Fuel { - public: - TVM_DEFINE_OBJECT_REF_METHODS(FTop, Fuel, FTopNode); -}; - -Fuel MkFTop() { return Fuel(make_object()); } - -/*! - * \brief A stack frame in the Relay interpreter. - * - * Contains a mapping from relay::Var to relay::Object. - */ -struct Frame { - /*! \brief The set of local variables and arguments for the frame. */ - std::unordered_map locals; - Frame() = default; -}; - -class Environment { - public: - Environment() : env_({Frame()}) {} - Environment(const Environment&) = delete; - - template - T Extend(const std::function& body) { - FrameContext fc(this); - return body(); - } - - void Insert(const Var& v, const PStatic& ps) { - ICHECK(ps.defined()); - ICHECK_GT(env_.size(), 0); - ICHECK_EQ(env_.back().locals.count(v), 0); - env_.back().locals[v] = ps; - } - - PStatic Lookup(const Var& v) { - auto rit = env_.rbegin(); - while (rit != env_.rend()) { - if (rit->locals.find(v) != rit->locals.end()) { - return rit->locals.find(v)->second; - } - ++rit; - } - LOG(FATAL) << "Unknown Variable: " << v; - throw; - } - - private: - std::list env_; - - struct FrameContext { - Environment* env_; - explicit FrameContext(Environment* env) : env_(env) { env_->env_.push_back(Frame()); } - ~FrameContext() { env_->env_.pop_back(); } - }; -}; - -/*! - * \brief As our store require rollback, we implement it as a frame. - * - * Every time we need to copy the store, a new frame is insert. - * Every time we roll back, a frame is popped. - */ -struct StoreFrame { - std::unordered_map store; - /*! - * \brief On unknown effect, history_valid is set to true to signal above frame is outdated. - * - * It only outdate the frame above it, but not the current frame. - */ - bool history_valid = true; - explicit StoreFrame(const std::unordered_map& store) : store(store) {} - StoreFrame() = default; -}; - -class Store { - public: - Store() : store_({StoreFrame()}) {} - Store(const Store&) = delete; - - template - T Extend(const std::function& body) { - StoreFrameContext sfc(this); - return body(); - } - - void Insert(const SRefNode* r, const PStatic& ps) { - ICHECK(r); - store_.back().store[r] = ps; - } - - // return null if not found - PStatic Lookup(const SRefNode* r) { - auto rit = store_.rbegin(); - while (rit != store_.rend()) { - if (rit->store.find(r) != rit->store.end()) { - return rit->store.find(r)->second; - } - if (!rit->history_valid) { - return PStatic(); - } - ++rit; - } - return PStatic(); - } - - void Invalidate() { - StoreFrame sf; - sf.history_valid = false; - store_.push_back(sf); - } - - private: - std::list store_; - - struct StoreFrameContext { - Store* store_; - explicit StoreFrameContext(Store* store) : store_(store) { - store_->store_.push_back(StoreFrame()); - } - ~StoreFrameContext() { - // push one history valid frame off. - while (!store_->store_.back().history_valid) { - store_->store_.pop_back(); - } - store_->store_.pop_back(); - } - }; -}; - -PStatic HasStatic(const Static& stat, const Expr& dynamic) { - ICHECK(stat.defined()); - return PStatic(make_object(stat, dynamic)); -} - -PStatic NoStatic(const Expr& dynamic) { return PStatic(make_object(dynamic)); } - -enum struct MatchStatus { Match, NoMatch, Unknown }; - -bool StatefulOp(const Expr& e) { - static auto op_stateful = Op::GetAttrMap("TOpIsStateful"); - struct StatefulOpVisitor : ExprVisitor { - bool stateful = false; - void VisitExpr_(const OpNode* op) { - stateful = stateful || op_stateful.get(GetRef(op), false); - } - }; - StatefulOpVisitor sov; - sov(e); - return sov.stateful; -} - -using FInterpreter = runtime::TypedPackedFunc; - -Target CPUTarget() { return Target("llvm"); } - -Device CPUDevice() { - Device dev; - dev.device_type = kDLCPU; - dev.device_id = 0; - return dev; -} - -using FuncId = int; - -/*! - * \brief Annotate a function with a FuncId. - */ -struct WithFuncIdAttrs : public tvm::AttrsNode { - FuncId fid; - - TVM_DECLARE_ATTRS(WithFuncIdAttrs, "relay.attrs.WithFuncIdAttrs") { - TVM_ATTR_FIELD(fid).describe("The FuncId that an function is annotated with.").set_default(-1); - } -}; - -TVM_REGISTER_NODE_TYPE(WithFuncIdAttrs); - -RELAY_REGISTER_OP("annotation.with_funcid") - .describe(R"code(Annotate a function with a funcid.)code" TVM_ADD_FILELINE) - .set_num_inputs(1) - .add_argument("func", "Function", "The input data."); - -// Cache with_funcid op to reduce lookup overhead during traversal. -static const Op& with_funcid_op = Op::Get("annotation.with_funcid"); - -Expr MkWithFuncId(const Expr& expr, FuncId fid) { - auto attrs = make_object(); - attrs->fid = fid; - return Call(with_funcid_op, {expr}, Attrs(attrs), {}); -} - -Expr StripWithFuncId(const Expr& e); - -Function AsFunc(const Expr& e) { - if (e.as()) { - return Downcast(e); - } else if (const CallNode* c = e.as()) { - ICHECK(c->op == with_funcid_op); - ICHECK_EQ(c->args.size(), 1); - return AsFunc(c->args[0]); - } else { - LOG(FATAL) << "Unknown case"; - throw; - } -} - -class PartialEvaluator : public ExprFunctor, - public PatternFunctor { - public: - PartialEvaluator(const IRModule& mod) : mod_(mod) {} - - PStatic VisitExpr(const Expr& e, LetList* ll) final { - PStatic ret = ExprFunctor::VisitExpr(e, ll); - ICHECK(IsAtomic(ret->dynamic)) << ret->dynamic; - return ret; - } - - PStatic VisitExpr(const Expr& e, LetList* ll, const Var& name) { - if (const CallNode* c = e.as()) { - if (c->op == with_funcid_op) { - ICHECK_EQ(c->args.size(), 1); - return VisitExpr(c->args[0], ll, name); - } - } - PStatic ret = - e.as() ? VisitFunc(Downcast(e), ll, name) : VisitExpr(e, ll); - ICHECK(IsAtomic(ret->dynamic)) << ret->dynamic; - return ret; - } - - PStatic VisitExpr_(const ConstantNode* op, LetList* ll) final { - return HasStatic(MkSTensor(op->data.CopyTo(device_)), ll->Push(GetRef(op))); - } - - PStatic VisitExpr_(const TupleNode* op, LetList* ll) final { - std::vector value; - tvm::Array expr; - for (const Expr& e : op->fields) { - PStatic ps = VisitExpr(e, ll); - value.push_back(ps); - expr.push_back(ps->dynamic); - } - // Note: The partial evaluator seems to do some weird stuff with sharing. Changing Tuple(expr) - // to WithFields(op, expr) causes failures in the partial evaluator tests. - return HasStatic(MkSTuple(value), ll->Push(Tuple(expr))); - } - - PStatic VisitExpr_(const TupleGetItemNode* op, LetList* ll) final { - PStatic ps = VisitExpr(op->tuple, ll); - if (ps->pstatic.defined()) { - return Downcast(ps->pstatic)->fields[op->index]; - } else { - return NoStatic(ll->Push(TupleGetItem(ps->dynamic, op->index))); - } - } - - PStatic VisitExpr_(const VarNode* op, LetList* ll) final { return env_.Lookup(GetRef(op)); } - - PStatic VisitGlobalVar(const GlobalVar& gv) { - ICHECK(mod_.defined()); - if (gv_map_.count(gv) == 0) { - BaseFunc base_func = mod_->Lookup(gv); - if (auto opt = base_func.as()) { - auto func = opt.value(); - InitializeFuncId(func); - Func f = VisitFuncStatic(func, gv); - gv_map_.insert({gv, HasStatic(MkSFunc(f), gv)}); - func = AsFunc(PostProcess(VisitFuncDynamic(func, f, gv))); - mod_->Update(gv, func); - return gv_map_.at(gv); - } else { - return NoStatic(gv); - } - } - return gv_map_.at(gv); - } - - PStatic VisitExpr_(const GlobalVarNode* op, LetList* ll) final { - return VisitGlobalVar(GetRef(op)); - } - - PStatic VisitExpr_(const LetNode* op, LetList* ll) final { - env_.Insert(op->var, VisitExpr(op->value, ll, op->var)); - return VisitExpr(op->body, ll); - } - - PStatic VisitExpr_(const IfNode* op, LetList* ll) final { - PStatic c = VisitExpr(op->cond, ll); - if (c->pstatic.defined()) { - NDArray cpu_array = Downcast(c->pstatic)->data.CopyTo(CPUDevice()); - ICHECK_EQ(DataType(cpu_array->dtype), DataType::Bool()); - if (reinterpret_cast(cpu_array->data)[0]) { - return VisitExpr(op->true_branch, ll); - } else { - return VisitExpr(op->false_branch, ll); - } - } else { - Expr t = store_.Extend([&]() { - return LetList::With([&](LetList* ll) { return VisitExpr(op->true_branch, ll)->dynamic; }); - }); - Expr f = store_.Extend([&]() { - return LetList::With([&](LetList* ll) { return VisitExpr(op->false_branch, ll)->dynamic; }); - }); - store_.Invalidate(); - return NoStatic(ll->Push(If(c->dynamic, t, f))); - } - } - - PStatic VisitExpr_(const RefCreateNode* op, LetList* ll) final { - PStatic ps = VisitExpr(op->value, ll); - Static r = MkSRef(); - store_.Insert(r.as(), ps); - return HasStatic(r, ll->Push(RefCreate(ps->dynamic))); - } - - PStatic VisitExpr_(const RefWriteNode* op, LetList* ll) final { - PStatic r = VisitExpr(op->ref, ll); - PStatic v = VisitExpr(op->value, ll); - if (r->pstatic.defined()) { - store_.Insert(r->pstatic.as(), v); - } else { - store_.Invalidate(); - } - return HasStatic(MkSTuple({}), ll->Push(RefWrite(r->dynamic, v->dynamic))); - } - - PStatic VisitExpr_(const RefReadNode* op, LetList* ll) final { - PStatic r = VisitExpr(op->ref, ll); - if (r->pstatic.defined()) { - PStatic ret = store_.Lookup(r->pstatic.as()); - if (ret.defined()) { - return ret; - } - } - return NoStatic(ll->Push(RefRead(r->dynamic))); - } - - PStatic VisitExpr_(const CallNode* op, LetList* ll) final { - if (op->op == with_funcid_op) { - ICHECK_EQ(op->args.size(), 1); - return VisitExpr(op->args[0], ll); - } - PStatic f = VisitExpr(op->op, ll); - std::vector x; - tvm::Array x_dyn; - for (const Expr& e : op->args) { - PStatic ps = VisitExpr(e, ll); - x.push_back(ps); - x_dyn.push_back(ps->dynamic); - } - if (f->pstatic.defined()) { - return Downcast(f->pstatic)->func(f, x, op->attrs, op->type_args, ll); - } else { - store_.Invalidate(); - return NoStatic(ll->Push(Call(f->dynamic, x_dyn, op->attrs, op->type_args))); - } - } - - struct FuelFrame { - PartialEvaluator* pe_; - FuncId fid_; - Fuel old_fuel; - FuelFrame(PartialEvaluator* pe, FuncId fid, const Fuel& new_fuel) : pe_(pe), fid_(fid) { - ICHECK_GT(pe_->fuel_map_.count(fid_), 0); - old_fuel = pe_->fuel_map_[fid_]; - pe_->fuel_map_[fid_] = new_fuel; - } - ~FuelFrame() { pe_->fuel_map_[fid_] = old_fuel; } - }; - - size_t GetFTValue(const PStatic& ps) { - if (ps->pstatic.defined()) { - if (auto* st = ps->pstatic.as()) { - if (st->data.Shape().empty()) { - NDArray cpu_array = st->data.CopyTo(CPUDevice()); - DataType dtype = DataType(cpu_array->dtype); - if (dtype == DataType::Int(32)) { - return std::max(0, *static_cast(cpu_array->data)); - } else if (dtype == DataType::Int(64)) { - return std::max(0, *static_cast(cpu_array->data)); - } - } - } - } - return 0; - } - - Fuel GetFuel(const PStatic& ps) { - std::vector fuels; - fuels.push_back(MkFTime(ps->created_time)); - fuels.push_back(MkFTValue(GetFTValue(ps))); - return MkFSeq(fuels); - } - - Func VisitFuncStatic(const Function& func, const Expr& var) { - ICHECK(IsAtomic(var)); - if (func->HasNonzeroAttr(attr::kPrimitive)) { - return ConstEvaluateFunc(func); - } - std::vector> free_vars; - for (const auto& v : FreeVars(func)) { - if (v != var) { - free_vars.push_back(std::pair(v, env_.Lookup(v))); - } - } - return [=](const PStatic& self, const std::vector& pv, const Attrs& attrs, - const tvm::Array& type_args, LetList* ll) { - return env_.Extend([&]() { - ICHECK_EQ(pv.size(), func->params.size()); - ICHECK_GT(func_map_.count(func), 0); - FuncId fid = func_map_.at(func); - if (fuel_map_.count(fid) == 0) { - fuel_map_.insert({fid, MkFTop()}); - } - std::vector args_fuel; - for (const auto& v : pv) { - args_fuel.push_back(GetFuel(v)); - } - auto meet_res = fuel_map_[fid]->Meet(MkFSeq(args_fuel)); - if (std::get<1>(meet_res)) { - FuelFrame tf(this, fid, std::get<0>(meet_res)); - Expr dedup_func = RegisterFuncId(DeDup(AnnotateFuncId(func))); - Function func = AsFunc(dedup_func); - if (var.as()) { - env_.Insert(Downcast(var), self); - } - for (size_t i = 0; i < pv.size(); ++i) { - env_.Insert(func->params[i], pv[i]); - } - for (const auto& p : free_vars) { - env_.Insert(p.first, p.second); - } - tvm::Map subst; - for (size_t i = 0; i < type_args.size(); ++i) { - subst.Set(func->type_params[i], type_args[i]); - } - for (size_t i = type_args.size(); i < func->type_params.size(); ++i) { - subst.Set(func->type_params[i], IncompleteType(kType)); - } - return VisitExpr(RegisterFuncId(TypeSubst(AnnotateFuncId(func->body), subst)), ll); - } else { - std::vector dyn; - for (const auto& v : pv) { - dyn.push_back(v->dynamic); - } - return NoStatic(ll->Push(Call(var, dyn, attrs, type_args))); - } - }); - }; - } - - Expr VisitFuncDynamic(const Function& func, const Func& f, const Expr& self) { - return store_.Extend([&]() { - store_.Invalidate(); - return WithFields( - func, func->params, LetList::With([&](LetList* ll) { - std::vector pv; - for (const auto& v : func->params) { - pv.push_back(NoStatic(v)); - } - tvm::Array type_args; - for (const auto& tp : func->type_params) { - type_args.push_back(tp); - } - return f(HasStatic(MkSFunc(f), self), pv, Attrs(), type_args, ll)->dynamic; - })); - }); - } - - PStatic VisitFunc(const Function& func, LetList* ll, const Var& name) { - Func f = VisitFuncStatic(func, name); - Function u_func = AsFunc(RegisterFuncId(DeDup(AnnotateFuncId(func)))); - // TODO(@M.K.): we seems to reduce landin knot into letrec. - // restore letrec support across whole relay. - return HasStatic(MkSFunc(f), ll->Push(name, VisitFuncDynamic(u_func, f, name))); - } - - PStatic VisitExpr_(const FunctionNode* op, LetList* ll) final { - return VisitFunc(GetRef(op), ll, Var::GenSym()); - } - - struct ReflectError : Error { - ReflectError() : Error("static value not found") {} - }; - - Expr Reflect(const PStatic& st) { - if (!st->pstatic.defined()) { - throw ReflectError(); - } else if (const STensorNode* op = st->pstatic.as()) { - return Constant(op->data); - } else if (const STupleNode* op = st->pstatic.as()) { - tvm::Array fields; - for (const PStatic& field : op->fields) { - fields.push_back(Reflect(field)); - } - return Tuple(fields); - } else { - LOG(FATAL) << "Unknown case: " << st->dynamic; - throw; - } - } - - PStatic Reify(const ObjectRef& v, LetList* ll) const { - if (v->IsInstance()) { - auto nd_array = Downcast(v); - return HasStatic(MkSTensor(nd_array), ll->Push(Constant(nd_array))); - } else if (auto opt = v.as()) { - std::vector fields; - tvm::Array fields_dyn; - auto adt = opt.value(); - for (size_t i = 0; i < adt.size(); ++i) { - PStatic ps = Reify(adt[i], ll); - fields.push_back(ps); - fields_dyn.push_back(ps->dynamic); - } - return HasStatic(MkSTuple(fields), ll->Push(Tuple(fields_dyn))); - } else { - LOG(FATAL) << "Unknown case"; - throw; - } - } - - // Constant evaluate an expression. - PStatic ConstEvaluate(const Expr& expr, LetList* ll) { - // use a fresh build context in case we are already in a build context. - With fresh_build_ctx(transform::PassContext::Create()); - return Reify(Eval(expr, mod_->type_definitions, mod_->Imports(), CPUDevice(), CPUTarget()), ll); - } - - Func ConstEvaluateFunc(const Expr& expr) { - ICHECK_EQ(FreeVars(expr).size(), 0); - return [=](const PStatic& self, const std::vector& pv, const Attrs& attrs, - const tvm::Array& type_args, LetList* ll) { - tvm::Array ns_args; - for (const PStatic& ps : pv) { - ns_args.push_back(ps->dynamic); - } - auto ns = [&]() { return NoStatic(ll->Push(Call(expr, ns_args, attrs, type_args))); }; - if (StatefulOp(expr)) { - return ns(); - } - try { - tvm::Array args; - for (const PStatic& ps : pv) { - args.push_back(Reflect(ps)); - } - return ConstEvaluate(Call(expr, args, attrs, type_args), ll); - } catch (const ReflectError&) { - return ns(); - } - }; - } - - PStatic VisitExpr_(const OpNode* op, LetList* ll) final { - return HasStatic(MkSFunc(ConstEvaluateFunc(GetRef(op))), GetRef(op)); - } - - PStatic VisitExpr_(const ConstructorNode* op, LetList* ll) final { - Constructor c = GetRef(op); - Func f = [=](const PStatic& self, const std::vector& pv, const Attrs& attrs, - const tvm::Array& type_args, LetList* ll) { - tvm::Array dyn; - for (const PStatic& ps : pv) { - dyn.push_back(ps->dynamic); - } - return HasStatic(MkSConstructor(c, pv), ll->Push(Call(c, dyn))); - }; - return HasStatic(MkSFunc(f), GetRef(op)); - } - - PStatic VisitExpr_(const MatchNode* op, LetList* ll) final { - PStatic ps = VisitExpr(op->data, ll); - return env_.Extend([&]() { - for (const Clause& c : op->clauses) { - switch (VisitPattern(c->lhs, ps)) { - case MatchStatus::Match: - return VisitExpr(c->rhs, ll); - case MatchStatus::NoMatch: - continue; - case MatchStatus::Unknown: - return [&]() { - tvm::Array clauses; - for (const Clause& c : op->clauses) { - Expr expr = store_.Extend([&]() { - return LetList::With([&](LetList* ll) { - for (const Var& v : BoundVars(c->lhs)) { - env_.Insert(v, NoStatic(v)); - } - return VisitExpr(c->rhs, ll)->dynamic; - }); - }); - clauses.push_back(Clause(c->lhs, expr)); - } - store_.Invalidate(); - return NoStatic(ll->Push(Match(ps->dynamic, clauses, op->complete))); - }(); - default: - LOG(FATAL) << "Unknown MatchStatus"; - throw; - } - } - LOG(FATAL) << "No case Match"; - throw; - }); - } - - MatchStatus VisitPattern_(const PatternWildcardNode* op, const PStatic& ps) final { - return MatchStatus::Match; - } - - MatchStatus VisitPattern_(const PatternVarNode* op, const PStatic& ps) final { - env_.Insert(op->var, ps); - return MatchStatus::Match; - } - - MatchStatus VisitPattern_(const PatternConstructorNode* op, const PStatic& ps) final { - if (ps->pstatic.defined()) { - SConstructor scn = Downcast(ps->pstatic); - ICHECK_NE(op->constructor->tag, -1); - ICHECK_NE(scn->constructor->tag, -1); - if (op->constructor->tag == scn->constructor->tag) { - ICHECK_EQ(op->patterns.size(), scn->fields.size()); - MatchStatus current_match_status = MatchStatus::Match; - for (size_t i = 0; i < op->patterns.size(); ++i) { - MatchStatus ms = VisitPattern(op->patterns[i], scn->fields[i]); - switch (ms) { - case MatchStatus::Match: - continue; - case MatchStatus::NoMatch: - return MatchStatus::NoMatch; - case MatchStatus::Unknown: - current_match_status = MatchStatus::Unknown; - } - } - return current_match_status; - } - return MatchStatus::NoMatch; - } else { - return MatchStatus::Unknown; - } - } - - MatchStatus VisitPattern_(const PatternTupleNode* op, const PStatic& ps) final { - if (ps->pstatic.defined()) { - STuple stn = Downcast(ps->pstatic); - ICHECK_EQ(op->patterns.size(), stn->fields.size()); - MatchStatus current_match_status = MatchStatus::Match; - for (size_t i = 0; i < op->patterns.size(); ++i) { - MatchStatus ms = VisitPattern(op->patterns[i], stn->fields[i]); - switch (ms) { - case MatchStatus::Match: - continue; - case MatchStatus::NoMatch: - return MatchStatus::NoMatch; - case MatchStatus::Unknown: - current_match_status = MatchStatus::Unknown; - } - } - return current_match_status; - } else { - return MatchStatus::Unknown; - } - } - - void InitializeFuncId(const Expr& e) { - struct InitializeFuncIdVisitor : ExprVisitor, PatternVisitor { - PartialEvaluator* pe; - explicit InitializeFuncIdVisitor(PartialEvaluator* pe) : pe(pe) {} - - void VisitExpr_(const FunctionNode* op) final { - Function f = GetRef(op); - ICHECK_EQ(pe->func_map_.count(f), 0); - pe->func_map_.insert({f, pe->func_map_.size()}); - VisitExpr(f->body); - } - - void VisitPattern(const Pattern& p) final { PatternVisitor::VisitPattern(p); } - }; - InitializeFuncIdVisitor(this).VisitExpr(e); - } - - Expr RegisterFuncId(const Expr& e) { - struct RegisterFuncIdVisitor : ExprVisitor, PatternVisitor { - PartialEvaluator* pe; - explicit RegisterFuncIdVisitor(PartialEvaluator* pe) : pe(pe) {} - - void VisitExpr_(const CallNode* op) final { - if (op->op == with_funcid_op) { - ICHECK_EQ(op->args.size(), 1); - ICHECK(op->attrs.defined()); - ICHECK(op->attrs.as()); - Function f = AsFunc(op->args[0]); - FuncId fid = op->attrs.as()->fid; - if (pe->func_map_.count(f) != 0) { - ICHECK_EQ(pe->func_map_.at(f), fid); - } - pe->func_map_.insert({f, fid}); - } - ExprVisitor::VisitExpr_(op); - } - - void VisitExpr_(const FunctionNode* op) final { - Function f = GetRef(op); - ICHECK_GT(pe->func_map_.count(f), 0); - ExprVisitor::VisitExpr_(op); - } - - void VisitPattern(const Pattern& p) final { PatternVisitor::VisitPattern(p); } - }; - RegisterFuncIdVisitor(this).VisitExpr(e); - return e; - } - - Expr AnnotateFuncId(const Expr& e) { - struct AnnotateFuncIdMutator : ExprMutator, PatternMutator { - PartialEvaluator* pe; - explicit AnnotateFuncIdMutator(PartialEvaluator* pe) : pe(pe) {} - - Expr VisitExpr_(const FunctionNode* op) final { - Function f = GetRef(op); - ICHECK_GT(pe->func_map_.count(f), 0); - return MkWithFuncId(ExprMutator::VisitExpr_(op), pe->func_map_.at(f)); - } - - Pattern VisitPattern(const Pattern& p) final { return PatternMutator::VisitPattern(p); } - - Var VisitVar(const Var& v) final { return v; } - }; - return AnnotateFuncIdMutator(this).VisitExpr(e); - } - - private: - Environment env_; - IRModule mod_; - std::unordered_map gv_map_; - /*! Termination checking is done as follows: - * We have finitely many FunctionIds. - * Each FunctionId maps to a class of semantically equivalent function (ignoring type), - * as both TypeSubst and DeDup create semantically equivalent function. - * We partially map each FunctionId to a Fuel. - * Every time we try to inline a Function, - * we make sure it either does not have a Fuel, - * or we meet the existing fuel with the fuel calculated from the argument. - * If no progress is made, we do not inline. - * In both case, we remap the mapping to the new Fuel - * when we PE inside the Function body. - * Termination is guaranteed because Fuel is finitely descending - there can only be so many - * meet. - */ - std::unordered_map func_map_; - std::unordered_map fuel_map_; - Store store_; - Device device_ = CPUDevice(); -}; - -/*! \brief Remap multiple Var sharing the same Id into the same Var. */ -Expr Remap(const Expr& e) { - class RemapMutator : public ExprMutator, public PatternMutator { - Expr VisitExpr_(const VarNode* op) final { - Var v = GetRef(op); - if (remap_.count(v) == 0) { - remap_.insert({v, v}); - } - return remap_.at(v); - } - - Var VisitVar(const Var& v) final { return Downcast(VisitExpr(v)); } - - private: - std::unordered_map remap_; - }; - return RemapMutator().VisitExpr(e); -} - -Expr StripWithFuncId(const Expr& e) { - struct StripWithFuncIdMutator : ExprMutator, PatternMutator { - Expr VisitExpr_(const CallNode* op) final { - if (op->op == with_funcid_op) { - ICHECK_EQ(op->args.size(), 1); - return VisitExpr(op->args[0]); - } else { - return ExprMutator::VisitExpr_(op); - } - } - - Pattern VisitPattern(const Pattern& p) final { return PatternMutator::VisitPattern(p); } - - Var VisitVar(const Var& v) final { return v; } - }; - return StripWithFuncIdMutator().VisitExpr(e); -} - -Expr PostProcess(const Expr& e) { return StripWithFuncId(DeDup(Remap(e))); } - -} // namespace partial_eval - -IRModule PartialEval(const IRModule& m) { - CheckFeature(m, FeatureSet::All() - fGraph); - relay::partial_eval::PartialEvaluator pe(m); - std::vector gvs; - for (const auto& p : m->functions) { - gvs.push_back(p.first); - } - for (const auto& gv : gvs) { - pe.VisitGlobalVar(gv); - } - CheckFeature(m, FeatureSet::All() - fGraph); - return m; -} - -namespace transform { - -Pass PartialEval() { - runtime::TypedPackedFunc pass_func = - [=](IRModule m, PassContext pc) { return relay::PartialEval(m); }; - return CreateModulePass(pass_func, 1, "PartialEval", {}); -} - -TVM_REGISTER_GLOBAL("relay._transform.PartialEvaluate").set_body_typed(PartialEval); - -} // namespace transform - -} // namespace relay -} // namespace tvm diff --git a/src/relay/transforms/partition_graph.cc b/src/relay/transforms/partition_graph.cc deleted file mode 100644 index 0be68872dd9c..000000000000 --- a/src/relay/transforms/partition_graph.cc +++ /dev/null @@ -1,619 +0,0 @@ -/* - * 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 src/relay/transforms/partition_graph.cc - * - * \brief Partition an input function into multiple functions according based - * on the inserted annotation nodes (i.e. compiler_begin and compiler_end). - * These nodes are used as boundaries to partition the Relay function into - * multiple regions that can be offloaded to different accelerators/backends. - * - * Each of these paritioned functions, a.k.a regions, will be viewed as - * external functions, and they will use the provided compiler for codegen. - */ - -#include -#include -#include -#include -#include -#include -#include -#include - -#include -#include -#include -#include - -#include "../analysis/annotated_region_set.h" -#include "../backend/name_transforms.h" -#include "../backend/utils.h" -#include "pass_utils.h" - -namespace tvm { -namespace relay { - -namespace partitioning { - -/*! \brief This struct maintains the required metadata for a region to generate a corresponding - * global function and function call. Global function will be passed to the target specific codegen - * and function call will be used in the transform Relay graph to invoke the function in runtime. - */ -struct RegionFuncMetadata { - /*! \brief The call node of the generated global function for this region. */ - Call func_call; - - /*! \brief A list of argument pairs. Each pair includes (var, expr). var is used - * as a function node argument; input expression is used as a function call parameter. - */ - std::vector> args; - - /*! \brief Map from each region output expr (compiler end) node to - * the corresponding function output expr. - */ - std::unordered_map region_func_out; - - /*! \brief Map from each region input expression (compiler begin) to - * the corresponding function input variable. This cache is used to make sure - * a region function will not have duplicated inputs even if it refers to - * the same expr multiple times. - */ - std::unordered_map region_func_in; -}; - -/*! \brief This class partitions the expr labeled with begin and end annotations - * into function containing multiple regions. Each region is labeled with - * a compiler attribute so that it will be handled by any compilers that are not - * in the TVM stack. - * - * Input : A Relay module that has functions with disjoint annotated regions - * using compiler_begin and compiler_end. There could be multiple - * outputs. - * - * Output : A Relay module with global functions for such disjoint annotated - * regions with calls inserted at the respective location - * - * Dependencies : AnnotatedRegionSet Utility class. - * - * Methodology : - * 1) The AnnotatedRegionSet utility class is able to construct a collection - * of nodes that are bound by a given annotation -- here we use - * compiler_begin and compiler_end - * 2) Initially, for each function in the module RegionSets are populated. - * 3) Then, Vistor pass is traversed until a compiler_end node is encountered - * that belongs to a "region". - * 4) When the first compiler_end of a given annotated region is found, - * a function is formed and inserted. - * a) if the region has multiple outputs, a Tuple node (capturing - * all outputs) is returned. - * 5) Thereafter, if we encounter an another output of the same annotated - * region, it is important to note that the function is already formed. - * Therefore, it will lookup the function and add a TupleGetItemNode. - * a) We will use the location index of "rets" of each Region" of - * AnnotatedRegionSet as TupleGetItemNode index. - * 6) Therefore, functions will be created for all annotated regions. - * The name for each global function is created using "Region" id and - * the compiler name. - */ - -class Partitioner : public MixedModeMutator { - public: - Partitioner(const IRModule& module, bool bind_constants) - : module_(module), bind_constants_(bind_constants) { - std::set func_names; - for (auto f : module->functions) { - GlobalVar f_var = f.first; - BaseFunc f_func = f.second; - std::string f_name = f_var.as()->name_hint; - while (func_names.find(f_name) != func_names.end()) { - f_name += "_a"; - } - func_names.insert(f_name); - - // Creating regionset per function in the module. - auto region_set = - AnnotatedRegionSet::Create(f_func, CompilerBeginOp(), CompilerEndOp(), f_name); - regions_sets_[region_set] = f_func; - } - } - - Expr Rewrite_(const CallNode* call, const Expr& post) final { - auto op_node = call->op.as(); - if (op_node == nullptr || call->attrs.as() == nullptr) { - return post; - } else if (call->op == CompilerBeginOp()) { - // The annotation node is inserted on edge so it must have only one argument. - ICHECK_EQ(call->args.size(), 1U); - - // Traverse the rest graph. - Expr parent = call->args[0]; - auto input_expr = Downcast(post)->args[0]; - - // Backtrace the parent to find the first ancestor node that is not a begin or end op - while (const auto* parent_call = parent.as()) { - if (parent_call->op == CompilerBeginOp() || parent_call->op == CompilerEndOp()) { - parent = parent_call->args[0]; - } else { - break; - } - } - - AnnotatedRegion sg = GetRegion(GetRef(call)); - int index = GetArgIdx(sg, GetRef(call)); - ICHECK_NE(index, -1); - - if (region_func_meta_[sg].region_func_in.count(parent)) { - return region_func_meta_[sg].region_func_in[parent]; - } else { - // The type of the created variable is the same as the compiler_begin - // node. - std::string target = call->attrs.as()->compiler; - std::string varname = - target + "_" + std::to_string(sg->GetID()) + "_i" + std::to_string(index); - auto var = Var(varname, GetRef(call)->checked_type_); - - std::pair cand = std::make_pair(var, input_expr); - - if (std::find(region_func_meta_[sg].args.begin(), region_func_meta_[sg].args.end(), cand) == - region_func_meta_[sg].args.end()) { - region_func_meta_[sg].args.push_back(cand); - } - region_func_meta_[sg].region_func_in[parent] = var; - return std::move(var); - } - } else { - ICHECK_EQ(call->op, CompilerEndOp()); - // The annotation node is inserted on edge so it must have only one - // argument. - ICHECK_EQ(call->args.size(), 1U); - - AnnotatedRegion region = GetRegion(GetRef(call)); - - // TODO(@manupa-arm) : need to use the parent function (to which region - // belongs to) name/key for the functions that are created - BaseFunc f = GetFunc(GetRef(call)); - - // Traverse subgraph inputs. - auto input = Downcast(post)->args[0]; - ICHECK(region.defined()) << "Region not defined for " << GetRef(call); - // functions are created for each annotated regions, - // when their first output is encountered. - // If multiple outputs are there, a tuple node is inserted at the end. - - if (!region_func_meta_[region].func_call.defined()) { - // First time this region is encountered in the traversal. Creating the function. - CreateFunction(region, call); - } - - // Retrieve this particular output of function. - Expr region_out_expr = Downcast(GetRef(call))->args[0]; - ICHECK(region_func_meta_[region].region_func_out.count(region_out_expr)); - return region_func_meta_[region].region_func_out[region_out_expr]; - } - } - - IRModule Partition() { - auto glob_funcs = module_->functions; - for (const auto& pair : glob_funcs) { - if (auto opt = pair.second.as()) { - Function func = opt.value(); - func = WithFields(func, func->params, VisitExpr(func->body)); - module_->Update(pair.first, func); - module_ = transform::InferType()(module_); - } - } - return module_; - } - - private: - /*! - * \brief Get the region an expression belongs to - * if its in a region. - */ - AnnotatedRegion GetRegion(const Expr& e) { - for (auto sg_set_it : regions_sets_) { - auto sg_set = sg_set_it.first; - AnnotatedRegion sg = sg_set->GetRegion(e); - if (sg.defined()) { - return sg; - } - } - return AnnotatedRegion(nullptr); - } - - /*! - * \brief Get the function an expression belongs to - * if its in a region. - */ - BaseFunc GetFunc(const Expr& e) { - for (auto sg_set_it : regions_sets_) { - auto sg_set = sg_set_it.first; - auto func = sg_set_it.second; - - AnnotatedRegion sg = sg_set->GetRegion(e); - if (sg.defined()) { - return func; - } - } - return BaseFunc(nullptr); - } - - /*! - * \brief Get the index of the argument; - * this is to be used as tuplegetitem idx - */ - int GetArgIdx(AnnotatedRegion sg, const Expr& arg) { - int idx = 0; - for (auto arg_ : sg->GetInputs()) { - if (arg == arg_) { - return idx; - } - idx++; - } - return -1; - } - - /*! - * \brief Check if an expr is a constant or a tuple that only contain constants. - */ - bool IsConstant(const Expr& expr) const { - if (expr->IsInstance()) return true; - if (!expr->IsInstance()) return false; - const auto* tn = expr.as(); - return std::all_of(tn->fields.begin(), tn->fields.end(), - [](const Expr& e) { return e->IsInstance(); }); - } - - /*! - * \brief Create a call to the function that represents a region. - * \note The customized optimization pipeline will be invoked as well to - * optimize each function that is handled by external codegen. - */ - Call CreateRegionCall(AnnotatedRegion region, const Array& fields, - const CallNode* end_node) { - Array params; - Array param_expr; - Map params_bind; - for (auto pair : region_func_meta_[region].args) { - params.push_back(pair.first); - if (bind_constants_ && IsConstant(pair.second)) { - params_bind.Set(pair.first, pair.second); - } else { - param_expr.push_back(pair.second); - } - } - - Function global_region_func; - if (fields.size() == 1) { - // If there are only a single output; no need to add a tuple - global_region_func = - Function(params, fields[0], end_node->args[0]->checked_type_, {}, DictAttrs()); - } else { - auto tuple = Tuple(fields); - global_region_func = Function(params, tuple, tuple->checked_type_, {}, DictAttrs()); - } - - std::string target = end_node->attrs.as()->compiler; - std::string name = target + "_" + region->GetName() + "_" + std::to_string(region->GetID()); - - // Constant propagation - if (!params_bind.empty()) { - global_region_func = Downcast(relay::Bind(global_region_func, params_bind)); - } - std::string ext_opt = "relay.ext." + target + ".optimize"; - auto pf = tvm::runtime::Registry::Get(ext_opt); - if (pf != nullptr) { - auto mod = IRModule::FromExpr(global_region_func); - mod = transform::InferType()(mod); - mod = (*pf)(mod); - global_region_func = Downcast(mod->Lookup("main")); - } - - global_region_func = - WithAttr(std::move(global_region_func), tvm::attr::kGlobalSymbol, runtime::String(name)); - global_region_func = WithAttr(std::move(global_region_func), attr::kPrimitive, tvm::Integer(1)); - global_region_func = - WithAttr(std::move(global_region_func), attr::kCompiler, tvm::runtime::String(target)); - global_region_func = WithAttr(std::move(global_region_func), attr::kInline, tvm::Integer(1)); - - GlobalVarSupply global_var_supply = GlobalVarSupply(module_); - GlobalVar glob_func = global_var_supply->FreshGlobal(name, false); - ICHECK(!module_->ContainGlobalVar(glob_func->name_hint)) - << "Global function " << glob_func->name_hint << " already exists"; - // Create a global function and add it to the IRModule for the region. - // This way we lift the functions that should be handled by external - // codegen to the module scope and rely on the pass manager to prevent - // relay function level passes (i.e. simplify inference and fusion) - // optimizing it. - module_->Add(glob_func, global_region_func); - module_ = relay::transform::InferType()(module_); - - // Create a call node for the function. - auto call = Call(glob_func, param_expr); - region_func_meta_[region].func_call = call; - - return call; - } - - /*! - * \brief Create a function and its function call for the given region. If the function has - * multiple outputs, a Tuple will be formed to aggregate all outputs, and TupleGetItem nodes - * will be created to serve output consumers. - */ - void CreateFunction(AnnotatedRegion region, const CallNode* end_node) { - // Create fields which is a unique list of outputs. - Array fields; - std::unordered_map out_expr_to_idx; - int out_idx = 0; - for (auto region_end_node : region->GetOutputs()) { - auto ret_node = Downcast(region_end_node)->args[0]; - // Don't duplicate outputs. - if (!out_expr_to_idx.count(ret_node)) { - auto ret_expr = MixedModeMutator::VisitExpr(ret_node); - fields.push_back(ret_expr); - out_expr_to_idx[ret_node] = out_idx++; - } - } - - Call call = CreateRegionCall(region, fields, end_node); - - // Create output expr(s) for the function call. - if (out_expr_to_idx.size() == 1) { - // Single output direcly uses the call node as the output expr. - region_func_meta_[region].region_func_out[out_expr_to_idx.begin()->first] = call; - } else { - // Multiple outptus need to create TupleGetItem nodes as output exprs. - for (auto pair : out_expr_to_idx) { - Expr region_out_expr = pair.first; // The arg of a compiler end node of this region. - int idx = pair.second; // Corresponding function output tuple index. - auto tuple_get_item = TupleGetItem(call, idx); - tuple_get_item->checked_type_ = region_out_expr->checked_type_; - region_func_meta_[region].region_func_out[region_out_expr] = tuple_get_item; - } - } - } - - /*! \brief Map from each region to its metadata of the generated function. */ - std::unordered_map - region_func_meta_; - - /*! \brief Each region set is associated with a function in the module. - * This map maintains the mapping between regionsets and the function it - * belongs to - */ - std::unordered_map regions_sets_; - - /*!\brief The IRModule used for partitioning. */ - IRModule module_; - - /*!\brief Whether or not to bind constants in partitioned subgraphs. */ - bool bind_constants_{false}; -}; - -IRModule RemoveDefaultAnnotations(IRModule module) { - class DefaultRemover : public ExprRewriter { - public: - DefaultRemover() = default; - - Expr Rewrite_(const CallNode* call, const Expr& post) final { - auto attrs = call->attrs.as(); - if (attrs != nullptr && attrs->compiler == "default") { - return Downcast(post)->args[0]; - } - return post; - } - }; - - auto glob_funcs = module->functions; - // module is mutable, hence, we make a copy of it. - module.CopyOnWrite(); - for (const auto& pair : glob_funcs) { - if (auto opt = pair.second.as()) { - auto func = opt.value(); - DefaultRemover remover; - auto removed = PostOrderRewrite(func->body, &remover); - func = WithFields(func, func->params, removed); - module->Update(pair.first, func); - module = relay::transform::InferType()(module); - } - } - return module; -} - -/*! \brief There can be regions with multiple outputs where each output - * could be a tuple output. Such tuple outputs needs to be flattened - * otherwise the function would create tuples of tuples. Moreover, tuple - * of tuples are valid relay, however they are not currently supported by - * graph executor or relay VM. - */ - -// New annotations would be required to be added for each flattened output -static const PackedFunc* make_end_op = - runtime::Registry::Get("relay.op.annotation._make.compiler_end"); - -IRModule FlattenTupleOutputs(IRModule module) { - class TupleOutFlattener : public ExprRewriter { - public: - TupleOutFlattener() = default; - - Expr Rewrite_(const CallNode* call, const Expr& post) final { - if (call->op == CompilerEndOp()) { - std::string target = call->attrs.as()->compiler; - // Arguments of annotation ops should be 1 - ICHECK_EQ(call->args.size(), 1U); - auto annotated_op = Downcast(post)->args[0]; - if (const auto* tuple_node = annotated_op.as()) { - Array new_fields; - new_fields.reserve(tuple_node->fields.size()); - - // Here each input of the tuple will be annotated with compiler_ends - for (auto& tn_arg : tuple_node->fields) { - new_fields.push_back((*make_end_op)(tn_arg, target)); - } - - // Return a tuple of compiler_ends in the place of the tuple that was - // annotated with a compiler_end. - return WithFields(GetRef(tuple_node), new_fields); - } - } - return post; - } - }; - - auto glob_funcs = module->functions; - // module is mutable, hence, we make a copy of it. - module.CopyOnWrite(); - for (const auto& pair : glob_funcs) { - if (auto opt = pair.second.as()) { - Function func = opt.value(); - TupleOutFlattener to_flattener; - auto removed = PostOrderRewrite(func->body, &to_flattener); - func = WithFields(func, func->params, removed); - module->Update(pair.first, func); - module = relay::transform::InferType()(module); - } - } - return module; -} - -class NameMangleExtFuncs : public MixedModeMutator { - public: - explicit NameMangleExtFuncs(const IRModule& module, std::function mangle_fn) - : module_(module), mangle_fn_(mangle_fn) {} - - IRModule Run() { - auto glob_funcs = module_->functions; - - // Collect function names to be mangled and create - // global mangled variables - for (const auto& pair : glob_funcs) { - if (auto opt = pair.second.as()) { - auto func = opt.value(); - if (func->GetAttr(attr::kCompiler).defined()) { - auto fn_name_mangled = tvm::runtime::SanitizeName(mangle_fn_(pair.first->name_hint)); - GlobalVar gvar = GlobalVar(fn_name_mangled); - mangled_gvars_[pair.first->name_hint] = gvar; - } - } - } - - // Walk the tree and mangle the functions. Then replace compiler functions - // with mangled functions in the module - IRModule new_module = module_->ShallowCopy(); - new_module->functions = {}; - - for (const auto& pair : glob_funcs) { - if (auto opt = pair.second.as()) { - auto func = opt.value(); - - if (func->GetAttr(attr::kCompiler).defined()) { - auto new_dict = func->attrs->dict; - new_dict.Set(tvm::attr::kGlobalSymbol, - String(tvm::runtime::SanitizeName(mangle_fn_(pair.first->name_hint)))); - func = WithFields(func, func->params, VisitExpr(func->body), func->ret_type, - func->type_params, DictAttrs(new_dict)); - - new_module->Add(mangled_gvars_[pair.first->name_hint], func); - } else { - func = WithFields(func, func->params, VisitExpr(func->body)); - new_module->Add(pair.first, func); - } - } - } - - return new_module; - } - - private: - Expr Rewrite_(const CallNode* call, const Expr& post) final { - Expr new_expr = post; - const CallNode* new_call = new_expr.as(); - auto op_node = new_call->op.as(); - if (op_node == nullptr || mangled_gvars_.find(op_node->name_hint) == mangled_gvars_.end()) { - return new_expr; - } else { - return Call(mangled_gvars_[op_node->name_hint], new_call->args, new_call->attrs, - new_call->type_args, new_call->span); - } - } - - /*!\brief The IRModule used for partitioning. */ - IRModule module_; - /*!\brief The function used to mangle operators name */ - std::function mangle_fn_; - /*!\brief Tabled used to store (unmangled_var_name, mangled_gvar) pairs*/ - std::unordered_map mangled_gvars_; -}; - -} // namespace partitioning - -namespace transform { - -Pass PartitionGraph(String mod_name, bool bind_constants) { - runtime::TypedPackedFunc flatten_tuples = [=](IRModule m, - PassContext pc) { - // There could be compiler_end annotations on tuples - // If the corresponding region is having multiple compiler_ends, - // this would lead to creation of tuples of tuples. - // Thus, we flatten the tuples by transfering the compiler_end to - // the tuple inputs. - return partitioning::FlattenTupleOutputs(m); - }; - - runtime::TypedPackedFunc remove_defaults = [=](IRModule m, - PassContext pc) { - // TODO(@comaniac, @zhiics): We should also handle the annotation with "default" attribute - // by treating them as un-annotated, but we don't have it yet. This workaround pass removes - // all "default" annotations and should be deleted in the future. - return partitioning::RemoveDefaultAnnotations(m); - }; - - runtime::TypedPackedFunc part_func = [=](IRModule m, - PassContext pc) { - return partitioning::Partitioner(m, bind_constants).Partition(); - }; - - auto name_mangling_fn = [mod_name](String name) { - return runtime::get_name_mangled(mod_name, name); - }; - - runtime::TypedPackedFunc name_mangling_func = - [=](IRModule m, PassContext pc) { - return partitioning::NameMangleExtFuncs(m, name_mangling_fn).Run(); - }; - - auto flatten_tuples_pass = CreateModulePass(flatten_tuples, 0, "FlattenNestedTuples", {}); - auto remove_default_pass = CreateModulePass(remove_defaults, 0, "RemoveDefaultAnnotations", {}); - auto partition_pass = CreateModulePass(part_func, 0, "PartitionGraph", {}); - auto name_mangling_pass = CreateModulePass(name_mangling_func, 0, "NameMangleExtFuncs", {}); - return Sequential( - {flatten_tuples_pass, remove_default_pass, partition_pass, name_mangling_pass, InferType()}); -} - -TVM_REGISTER_GLOBAL("relay._transform.PartitionGraph") - .set_body_typed([](String mod_name, bool bind_constants) { - return transform::PartitionGraph(mod_name, bind_constants); - }); - -} // namespace transform - -} // namespace relay -} // namespace tvm diff --git a/src/relay/transforms/pass_utils.h b/src/relay/transforms/pass_utils.h deleted file mode 100644 index b14a93f02b55..000000000000 --- a/src/relay/transforms/pass_utils.h +++ /dev/null @@ -1,236 +0,0 @@ -/* - * 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 tvm/relay/_transforms/pass_utils.h - * \brief Utilities for writing passes - */ -#ifndef TVM_RELAY_TRANSFORMS_PASS_UTILS_H_ -#define TVM_RELAY_TRANSFORMS_PASS_UTILS_H_ - -#include -#include -#include -#include - -#include -#include -#include -#include - -#include "../analysis/dependency_graph.h" -#include "../op/annotation/annotation.h" -#include "../op/memory/on_device.h" -#include "./let_list.h" - -namespace tvm { -namespace relay { - -/*! - * \brief Check if expr is positive constant. - * \param expr The expression to be checked. - * \return Whether all elements of expr is positive constant. - */ -bool IsAllPositiveConstant(const Expr& expr); - -/*! - * \brief Substitute var with subst. - * \param type The type to be substituted. - * \param tvar The type variable to be substituted. - * \param subst The target of substitution. - * \return The substituted result. - */ -Type TypeSubst(const Type& type, const TypeVar& tvar, const Type& subst); - -/*! - * \brief Substitute var with subst. - * \param expr The expr to be substituted. - * \param tvar The type variable to be substituted. - * \param subst The target of substitution. - * \return The substituted result. - */ -Expr TypeSubst(const Expr& expr, const TypeVar& tvar, const Type& subst); - -/*! - * \brief Substitute type vars in type. - * \param type The type to be substituted. - * \param subst_map The map of substitution. - * \return The substituted result. - */ -Type TypeSubst(const Type& type, const tvm::Map& subst_map); - -/*! - * \brief Substitute type vars in type. - * \param expr The expr to be substituted. - * \param subst_map The map of substitution. - * \return The substituted result. - */ -Expr TypeSubst(const Expr& expr, const tvm::Map& subst_map); - -/*! - * \brief Check if type is dynamic. - * \param ty The type to be checked. - * \return Whether the type is dynamic. - */ -bool IsDynamic(const Type& ty); - -/*! - * \brief Check if call is data dependent. - * \param call The call to be checked. - * \return Whether the call is data dependent. - */ -bool IsDataDependent(const CallNode* call); - -/*! - * \brief Make arbitrary transformation preserve the out most function. - * \param func The transformation. - * \param e The expression - * \return the transformed expression. If e is a function the return is also a function. - */ -inline Expr TransformF(const std::function& func, const Expr& e) { - if (const FunctionNode* f = e.as()) { - return WithFields(GetRef(f), f->params, func(f->body)); - } else { - return func(e); - } -} - -/*! - * \brief Decide whether the expression atomic or not? - * \param e the expression - * \return - * is it atomic? - * if so, the compute cost of the expression is bounded so it can be copy without graph mode. - */ -inline bool IsAtomic(const Expr& expr) { - Expr true_expr = IgnoreOnDevice(expr); - return true_expr.as() || true_expr.as() || true_expr.as() || - true_expr.as() || - true_expr.as(); // Constant is always by reference. -} - -/*! - * \brief Cache the compiler_begin annotation op to reduce registry lookup overhead - * \param void - * \return compiler_begin op - */ -inline const Op& CompilerBeginOp() { - static auto op = Op::Get("annotation.compiler_begin"); - return op; -} - -/*! - * \brief Cache the compiler_end annotation op to reduce registry lookup overhead - * \param void - * \return compiler_end op - */ -inline const Op& CompilerEndOp() { - static auto op = Op::Get("annotation.compiler_end"); - return op; -} - -template -struct TreeNode { - typedef std::shared_ptr> pointer; - virtual ~TreeNode() {} -}; - -template -struct TreeLeafNode : TreeNode { - using TreeObjectPtr = typename TreeNode::pointer; - - Expr body; - - explicit TreeLeafNode(Expr body) : body(body) {} - - static TreeObjectPtr Make(Expr body) { return std::make_shared(body); } - - ~TreeLeafNode() {} -}; - -template -struct TreeLeafFatalNode : TreeNode { - using TreeObjectPtr = typename TreeNode::pointer; - - TreeLeafFatalNode() = default; - - static TreeObjectPtr Make() { return std::make_shared(); } - - ~TreeLeafFatalNode() {} -}; - -template -struct TreeBranchNode : TreeNode { - using TreeObjectPtr = typename TreeNode::pointer; - - ConditionObjectPtr cond; - TreeObjectPtr then_branch; - TreeObjectPtr else_branch; - - TreeBranchNode(ConditionObjectPtr cond, TreeObjectPtr then_branch, TreeObjectPtr else_branch) - : cond(cond), then_branch(then_branch), else_branch(else_branch) {} - - static TreeObjectPtr Make(ConditionObjectPtr cond, TreeObjectPtr then_branch, - TreeObjectPtr else_branch) { - return std::make_shared(cond, then_branch, else_branch); - } - - ~TreeBranchNode() {} -}; - -struct ScopeNode; -using Scope = std::shared_ptr; -using NodeScopeMap = std::unordered_map; -using ExprSet = std::unordered_set; - -/* Invariant: when parent is null level is 0 - * Invariant: when parent is not null level is 1 + parent->level - */ -struct ScopeNode { - // the level of the scope - size_t level; - // the parent scope - Scope parent; - // the corresponding let list which holds all let bindings in the scope - std::shared_ptr let_list = std::make_shared(); - explicit ScopeNode(const Scope& parent) : level(1 + parent->level), parent(parent) {} - ScopeNode() : level(0) {} -}; - -/*! \brief Calculate the scope of nodes in the dependency graph by least common ancestor. - * - * \param dg the input dependency graph - * \param expr_scope the output node -> scope mapping for all nodes. - * \param lifted_exprs the output set of expressions whose scope is lifted due to dependency - */ -std::pair CalcScope(const DependencyGraph& dg); - -/*! \brief find the least common ancestor of lhs scope and rhs scope. - */ -Scope LCA(Scope lhs, Scope rhs); - -// For basic block normal form. -Expr ToBasicBlockNormalFormAux(const Expr& e); - -// ToANormalForm for expressions and as a Pass are declared in transform.h - -} // namespace relay -} // namespace tvm -#endif // TVM_RELAY_TRANSFORMS_PASS_UTILS_H_ diff --git a/src/relay/transforms/pattern_utils.h b/src/relay/transforms/pattern_utils.h deleted file mode 100644 index b26bd7649630..000000000000 --- a/src/relay/transforms/pattern_utils.h +++ /dev/null @@ -1,887 +0,0 @@ -/* - * 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 tvm/relay/transforms/pattern_utils.h - * \brief Header of internal operator functions - * These can be used for writing passes. - */ -#ifndef TVM_RELAY_TRANSFORMS_PATTERN_UTILS_H_ -#define TVM_RELAY_TRANSFORMS_PATTERN_UTILS_H_ - -#include -#include -#include -#include -#include -#include -#include -#include -#include -#include -#include - -#include -#include -#include -#include -#include - -#include "../backend/utils.h" -#include "../op/make_op.h" - -namespace tvm { -namespace relay { - -/*! - * \brief Dispatch DataType to the C++ data type - * during runtime. - */ -#define TVM_DTYPE_DISPATCH(type, DType, ...) \ - if (type == DataType::Float(64)) { \ - typedef double DType; \ - { __VA_ARGS__ } \ - } else if (type == DataType::Float(32)) { \ - typedef float DType; \ - { __VA_ARGS__ } \ - } else if (type == DataType::Float(16)) { \ - typedef uint16_t DType; \ - { __VA_ARGS__ } \ - } else if (type == DataType::BFloat(16)) { \ - typedef uint16_t DType; \ - { __VA_ARGS__ } \ - } else if (type == DataType::Int(64)) { \ - typedef int64_t DType; \ - { __VA_ARGS__ } \ - } else if (type == DataType::Int(32)) { \ - typedef int32_t DType; \ - { __VA_ARGS__ } \ - } else if (type == DataType::Int(16)) { \ - typedef int16_t DType; \ - { __VA_ARGS__ } \ - } else if (type == DataType::Int(8)) { \ - typedef int8_t DType; \ - { __VA_ARGS__ } \ - } else if (type == DataType::UInt(64)) { \ - typedef uint64_t DType; \ - { __VA_ARGS__ } \ - } else if (type == DataType::UInt(32)) { \ - typedef uint32_t DType; \ - { __VA_ARGS__ } \ - } else if (type == DataType::UInt(16)) { \ - typedef uint16_t DType; \ - { __VA_ARGS__ } \ - } else if (type == DataType::UInt(8)) { \ - typedef uint8_t DType; \ - { __VA_ARGS__ } \ - } else if (type == DataType::Bool()) { \ - typedef bool DType; \ - { __VA_ARGS__ } \ - } else if ((*tvm::runtime::Registry::Get("runtime._datatype_get_type_registered"))( \ - static_cast(type.code()))) { \ - typedef double DType; \ - { __VA_ARGS__ } \ - } else { \ - LOG(FATAL) << "unknown data type " << type; \ - } - -/*! - * \brief Try to do the type inference over expr: - * - * Do the infer_type over each node in expr - * - * \param expr The IR expression - * \return infered expr if succeed. - */ -inline Expr InferType(const Expr& expr) { - auto mod = IRModule::FromExpr(expr); - mod = transform::InferType()(mod); - if (expr.as()) { - return mod->Lookup("main"); - } else { - return mod->Lookup("main").as()->body; - } -} - -/*! - * \brief Try to match lhs and rhs via broadcasting rule, such that: - * - * rhs matches the dimension of lhs specified by lhs_axes - * rhs's value equals 1 on rest of dimensions. - * - * \param tlhs The type of left operand (data) - * \param trhs The type right operand (bias) - * \param lhs_axes The axes on lhs to match. - * \param rhs_value A squeezed version of rhs which only contains matched dimension. - * \return Whether match is successful. - */ -inline bool MatchBroadcastToLeftAxes(const TensorTypeNode* tlhs, const TensorTypeNode* trhs, - const Array& lhs_axes, Expr* rhs_value = nullptr) { - if (tlhs->shape.size() < trhs->shape.size()) return false; - StructuralEqual equal; - size_t base = tlhs->shape.size() - trhs->shape.size(); - size_t j = 0; - - // handle case trhs is simple constant - if (trhs->shape.size() == 0 && rhs_value != nullptr && lhs_axes.size() > 0) { - *rhs_value = MakeExpandDims(*rhs_value, 0, lhs_axes.size()); - for (size_t i = 0; i < lhs_axes.size(); i++) { - int repeat_value = - tlhs->shape[static_cast(lhs_axes[j]->value)].as()->value; - *rhs_value = MakeRepeat(*rhs_value, repeat_value, i); - } - return true; - } - - ObjectPtr squeeze_attrs; - if (rhs_value != nullptr) { - squeeze_attrs = make_object(); - } - - for (size_t i = 0; i < tlhs->shape.size(); ++i) { - if (j < lhs_axes.size() && i == static_cast(lhs_axes[j]->value)) { - if (i < base || !equal(tlhs->shape[i], trhs->shape[i - base])) { - return false; - } - ++j; - } else if (i >= base) { - if (!tir::is_const_int(trhs->shape[i - base], 1)) { - return false; - } - if (rhs_value != nullptr) { - squeeze_attrs->axis.push_back(static_cast(i - base)); - } - } - } - if (rhs_value != nullptr && squeeze_attrs->axis.size() != 0) { - static const Op& squeeze_op = Op::Get("squeeze"); - *rhs_value = Call(squeeze_op, {rhs_value[0]}, Attrs(squeeze_attrs), {}); - } - return true; -} - -/*! - * \brief Expand 1D Tensor to match axis. - * - * The result bias can be used to add or multiply to - * the target Tensor on the specified axis via broadcasting rule. - * - * \param bias The bias. - * \param target_ndim Target dimension. - * \param axes The axis on the output we want to match on. - */ -inline Expr ExpandBiasToMatchAxis(Expr bias, int target_ndim, const Array& axes) { - static const Op& expand_dims = Op::Get("expand_dims"); - for (size_t i = axes.size(); i != 0; --i) { - if (i == axes.size()) { - int64_t num_pad_axis = target_ndim - axes[i - 1]->value - 1; - if (num_pad_axis > 0) { - auto attrs = make_object(); - attrs->axis = i; - attrs->num_newaxis = static_cast(num_pad_axis); - bias = Call(expand_dims, {bias}, Attrs(attrs), {}); - } - } else { - int64_t diff = axes[i]->value - axes[i - 1]->value; - ICHECK_GE(diff, 0L); - if (diff > 0) { - auto attrs = make_object(); - attrs->axis = i; - attrs->num_newaxis = static_cast(diff); - bias = Call(expand_dims, {bias}, Attrs(attrs), {}); - } - } - } - return bias; -} - -/*! - * \brief Check if the call is depthwise conv3d. - * - * \param call The conv call. - * \param param The conv attributes. - * \return Whether it is depthwise_conv3d. - */ -template -inline bool IsDepthwiseConv(const Call& call, ATTRS param, const Layout& kernel_layout) { - static const Layout kOIXX = - backend::IsOp(call.as(), "nn.conv2d") ? Layout("OIHW") : Layout("OIDHW"); - const auto bilayout = tir::BijectiveLayout(kernel_layout, kOIXX); - auto wshape = bilayout.ForwardShape(call->args[1]->type_as()->shape); - return tir::is_const_int(wshape[0], param->groups) && tir::is_const_int(wshape[1], 1); -} - -/*! - * \brief Get super-dimension of output channels of conv2d - * \param call The conv2d call. - * \return Super-dimension size of output channels of conv2d. - */ -inline int64_t GetConv2DSuperChannelsDim(const CallNode* call) { - auto param = call->attrs.as(); - auto tweight = call->args[1]->type_as(); - auto index = param->kernel_layout.operator std::string().find('O'); - ICHECK_NE(index, std::string::npos); - auto channels = tir::as_const_int(tweight->shape[index]); - return *channels; -} - -/*! - * \brief Is single value tensor (scalar). - * \param expr The expr. - * \return True if single value tensor. - */ -inline bool IsScalar(const Expr& expr) { - if (auto tensor_type = expr->checked_type().as()) { - for (auto dim_index_expr : tensor_type->shape) { - if (auto dim_index = dim_index_expr.as()) { - if (dim_index->value != 1) { - return false; - } - } else { - return false; - } - } - } else { - return false; - } - return true; -} - -/*! - * \brief Check if expr is a const scalar. - * \param expr The expr. - * \return True if const scalar. - */ -inline bool IsConstScalar(const Expr& expr) { - const auto* const_expr = expr.as(); - if (const_expr) { - return const_expr->is_scalar(); - } - return false; -} - -/*! - * \brief Create a Constant with a scalar - * - * \param dtype The data type. - * \param value The value of the scalar. - * \return A Constant. - */ -template -inline Constant MakeConstantScalar(DataType dtype, T value) { - runtime::NDArray arr = runtime::NDArray::Empty({}, dtype, {kDLCPU, 0}); - TVM_DTYPE_DISPATCH(dtype, DType, { - if (dtype == DataType::Float(16)) { - // convert to float16 - // storage is uint16_t - *static_cast(arr->data) = - __truncXfYf2__(static_cast(value)); - } else if (dtype == DataType::BFloat(16)) { - // convert to bfloat16 - // storage is uint16_t - *static_cast(arr->data) = - __truncXfYf2__(static_cast(value)); - } else { - *static_cast(arr->data) = value; - } - }) - return Constant(arr); -} - -/*! - * \brief Create a Constant with a tensor. - * - * \param dtype The data type. - * \param value The vector of the tensor values. - * \return A Constant. - */ -template -static inline Constant MakeConstantTensor(DataType dtype, std::vector shape, - std::vector value) { - runtime::NDArray arr = runtime::NDArray::Empty(shape, dtype, {kDLCPU, 0}); - TVM_DTYPE_DISPATCH(dtype, DType, { - for (size_t i = 0; i < value.size(); i++) { - if (dtype == DataType::Float(16)) { - // convert to float16 - // storage is uint16_t - // Similar handling as that in MakeConstantScalar - *(static_cast(arr->data) + i) = - __truncXfYf2__( - static_cast(value[i])); - } else if (dtype == DataType::BFloat(16)) { - // convert to bfloat16 - // storage is uint16_t - *(static_cast(arr->data) + i) = - __truncXfYf2__( - static_cast(value[i])); - } else { - *(static_cast(arr->data) + i) = value[i]; - } - } - }) - return Constant(arr); -} - -/*! - * \brief Create a Constant with a tensor. - * - * \param dtype The data type. - * \param value The array of the tensor values. - * \return A Constant. - */ -template -static inline Constant MakeConstantTensor(DataType dtype, std::vector shape, - Array value) { - runtime::NDArray arr = runtime::NDArray::Empty(shape, dtype, {kDLCPU, 0}); - TVM_DTYPE_DISPATCH(dtype, DType, { - for (size_t i = 0; i < value.size(); i++) { - if (dtype == DataType::Float(16)) { - // convert to float16 - // storage is uint16_t - // Similar handling as that in MakeConstantScalar - *(static_cast(arr->data) + i) = - __truncXfYf2__( - static_cast(value[i])); - } else if (dtype == DataType::BFloat(16)) { - // convert to bfloat16 - // storage is uint16_t - *(static_cast(arr->data) + i) = - __truncXfYf2__( - static_cast(value[i])); - } else { - *(static_cast(arr->data) + i) = value[i]; - } - } - }) - return Constant(arr); -} - -/*! - * \brief Create a Constant tensor of zeros. - * - * \param dtype The data type. - * \param shape The shape of the output constant tensor. - * \return A Constant. - */ -static inline Constant MakeConstantZeros(DataType dtype, std::vector shape) { - runtime::NDArray arr = runtime::NDArray::Empty(shape, dtype, {kDLCPU, 0}); - int64_t data_size = 1; - for (int64_t dim : shape) { - data_size *= dim; - } - TVM_DTYPE_DISPATCH(dtype, DType, { - for (int64_t i = 0; i < data_size; i++) { - if (dtype == DataType::Float(16)) { - // convert to float16 - // storage is uint16_t - // Similar handling as that in MakeConstantScalar - *(static_cast(arr->data) + i) = - __truncXfYf2__(static_cast(0)); - } else if (dtype == DataType::BFloat(16)) { - // convert to bfloat16 - // storage is uint16_t - *(static_cast(arr->data) + i) = - __truncXfYf2__(static_cast(0)); - } else { - *(static_cast(arr->data) + i) = 0; - } - } - }) - return Constant(arr); -} - -/*! - * \brief Check whether a shape is static and create corresponding Constant. - Eventually this will be removed and replaced with CheckConstantShapeArrayInteger - * - * \param shape The Array of the shape values. - * \return A Constant. - */ -static inline Constant CheckConstantShape(const Array& shape) { - auto shape_array = - runtime::NDArray::Empty({int64_t(shape.size())}, DataType::Int(64), {kDLCPU, 0}); - auto* shape_data = static_cast(shape_array->data); - for (size_t i = 0; i < shape.size(); ++i) { - const auto& dim_val = shape[i].as(); - ICHECK(dim_val) << "Do not support symbolic shape for " - "Array format. Pass shape as Expr instead."; - shape_data[i] = dim_val->value; - } - return Constant(shape_array); -} - -/*! - * \brief Check whether a shape is static and create corresponding Array. Will replace - * CheckConstantShape after dynamic refactorization is complete - * - * \param shape The Array of the shape values. - * \return A Constant. - */ -static inline Array CheckConstantShapeArrayInteger(const Array& shape) { - Array constShape; - - for (size_t i = 0; i < shape.size(); ++i) { - const auto& dim_val = shape[i].as(); - ICHECK(dim_val) << "Do not support symbolic shape for " - "Array format. Pass shape as Expr instead."; - - constShape.push_back(dim_val->value); - } - return constShape; -} - -/*! - * \brief Check if two expressions are equal scalars. - * \param a The expression to be checked. - * \param b The expression to be checked - * \return Whether two expressions are equal scalars. - */ -inline bool IsEqualScalar(const Expr& a, const Expr& b) { - const auto* constant_a = a.as(); - const auto* constant_b = b.as(); - if (!constant_a || !constant_b || !constant_a->is_scalar() || !constant_b->is_scalar()) { - return false; - } - return tvm::StructuralEqual()(a, b); -} - -/*! - * \brief Convert an element of a NDArray with type int or float to scalar. - * \param array Input NDArray - * \param i element index - * \return Converted scalar value, or None if conversion failed - */ -template -static inline std::optional TryToScalar(const runtime::NDArray& array, size_t i = 0) { - if (array->dtype.code == kDLInt) { - if (array->dtype.bits == 8) { - return std::optional(reinterpret_cast(array->data)[i]); - } else if (array->dtype.bits == 16) { - return std::optional(reinterpret_cast(array->data)[i]); - } else if (array->dtype.bits == 32) { - return std::optional(reinterpret_cast(array->data)[i]); - } else if (array->dtype.bits == 64) { - return std::optional(reinterpret_cast(array->data)[i]); - } - } else if (array->dtype.code == kDLUInt) { - if (array->dtype.bits == 1) { // bool - return std::optional(reinterpret_cast(array->data)[i]); - } else if (array->dtype.bits == 8) { - return std::optional(reinterpret_cast(array->data)[i]); - } else if (array->dtype.bits == 16) { - return std::optional(reinterpret_cast(array->data)[i]); - } else if (array->dtype.bits == 32) { - return std::optional(reinterpret_cast(array->data)[i]); - } else if (array->dtype.bits == 64) { - return std::optional(reinterpret_cast(array->data)[i]); - } - } else if (array->dtype.code == kDLFloat) { - if (array->dtype.bits == 16) { - return std::optional(__extendXfYf2__( - reinterpret_cast(array->data)[i])); - } - if (array->dtype.bits == 32) { - return std::optional(reinterpret_cast(array->data)[i]); - } else if (array->dtype.bits == 64) { - return std::optional(reinterpret_cast(array->data)[i]); - } - } else if (array->dtype.code == kDLBfloat) { - if (array->dtype.bits == 16) { - return std::optional(__extendXfYf2__( - reinterpret_cast(array->data)[i])); - } - } - return std::nullopt; -} - -/*! - * \brief Convert an element of a NDArray with type int or float to scalar. - * \param array Input NDArray - * \param i element index - * \return Converted scalar value - */ -template -static inline T ToScalar(const runtime::NDArray& array, size_t i = 0) { - auto try_value = TryToScalar(array, i); - ICHECK(try_value) << "Unknown data type: " << tvm::runtime::DLDataType2String(array->dtype); - return try_value.value(); -} - -static inline long double ToScalar(const runtime::NDArray& array, size_t i = 0) { - auto try_value = TryToScalar(array, i); - ICHECK(try_value) << "Unknown data type: " << tvm::runtime::DLDataType2String(array->dtype); - return try_value.value(); -} - -/*! - * \brief Convert a NDArray with type int or float to Array. - * \param array Input NDArray - * \return Converted Array. - */ -static inline Array ToVector(const runtime::NDArray& array) { - size_t ndim = array.Shape().size(); - ICHECK_EQ(ndim, 1) << "This function should only be used for 1D NDArrays"; - size_t len = array.Shape().front(); - Array out; - for (size_t i = 0; i < len; ++i) { - uint64_t elem_val = ToScalar(array, i); - out.push_back(Integer(IntImm(DataType::Int(32), static_cast(elem_val)))); - } - return out; -} - -/*! - * \brief Convert a NDArray with type int or float to Array. - * \param array Input NDArray - * \return Converted Array. - */ -static inline Array ToFloatVector(const runtime::NDArray& array) { - size_t ndim = array.Shape().size(); - ICHECK_EQ(ndim, 1) << "This function should only be used for 1D NDArrays"; - size_t len = array.Shape().front(); - Array out; - for (size_t i = 0; i < len; ++i) { - long double elem_val = ToScalar(array, i); - out.push_back(FloatImm(DataType::Float(32), static_cast(elem_val))); - } - return out; -} - -/*! - * \brief Convert a NDArray with type int or float to Array>. - * \param array Input NDArray - * \return Converted Array. - */ -static inline Array> ToMatrix(const runtime::NDArray& array) { - size_t ndim = array.Shape().size(); - ICHECK_EQ(ndim, 2) << "This function should only used for 2D NDArrays"; - size_t dim1 = array.Shape().at(0); - size_t dim2 = array.Shape().at(1); - - Array> out; - - for (size_t i = 0; i < dim1; ++i) { - Array inner_out; - for (size_t j = 0; j < dim2; ++j) { - double elem_val = ToScalar(array, i * dim2 + j); - inner_out.push_back(Integer(static_cast(elem_val))); - } - out.push_back(inner_out); - } - return out; -} - -inline Expr GetField(Expr t, size_t i) { return TupleGetItem(t, i); } - -inline Expr Pair(Expr l, Expr r) { return Tuple({l, r}); } - -inline Expr Exp(Expr e) { - static const Op& op = Op::Get("exp"); - return Call(op, {e}); -} - -inline Expr Erf(Expr e) { - static const Op& op = Op::Get("erf"); - return Call(op, {e}); -} - -inline Expr FastExp(Expr e) { - static const Op& op = Op::Get("fast_exp"); - return Call(op, {e}); -} - -inline Expr FastErf(Expr e) { - static const Op& op = Op::Get("fast_erf"); - return Call(op, {e}); -} - -inline Expr FastTanh(Expr e) { - static const Op& op = Op::Get("fast_tanh"); - return Call(op, {e}); -} - -inline Expr FastSoftmax(Expr e, tvm::Attrs attr) { - static const Op& op = Op::Get("nn.fast_softmax"); - return Call(op, {e}, attr); -} - -inline Expr Log(Expr e) { - static const Op& op = Op::Get("log"); - return Call(op, {e}); -} - -inline Expr Tanh(Expr e) { - static const Op& op = Op::Get("tanh"); - return Call(op, {e}); -} - -inline Expr Abs(Expr e) { - static const Op& op = Op::Get("abs"); - return Call(op, {e}); -} -/*! - * \brief Get an immediate scalar from a Constant expr. - * - * \param expr The Constant expr. - * \return A scalar with type T. - */ -template -T GetScalarFromConstant(Expr expr) { - const auto* n = expr.as(); - ICHECK(n) << "Expr must be a constant expr - " << AsText(expr, false); - ICHECK(n->is_scalar()); - return static_cast(n->data->data)[0]; -} - -inline Expr Cast(Expr x, DataType dtype) { return MakeCast(x, dtype); } - -inline Expr Negative(Expr x) { - static const Op& op = Op::Get("negative"); - return Call(op, {x}, Attrs(), {}); -} - -inline Expr Sqrt(Expr x) { - static const Op& op = Op::Get("sqrt"); - return Call(op, {x}, Attrs(), {}); -} - -inline Expr Sigmoid(Expr x) { - static const Op& op = Op::Get("sigmoid"); - return Call(op, {x}, Attrs(), {}); -} - -inline Expr Rsqrt(Expr x) { - static const Op& op = Op::Get("rsqrt"); - return Call(op, {x}, Attrs(), {}); -} - -inline Expr Relu(Expr x) { - static const Op& op = Op::Get("nn.relu"); - return Call(op, {x}, Attrs(), {}); -} - -inline Expr Round(Expr x) { - static const Op& op = Op::Get("round"); - return Call(op, {x}, Attrs(), {}); -} - -inline Expr Floor(Expr x) { - static const Op& op = Op::Get("floor"); - return Call(op, {x}, Attrs(), {}); -} - -inline Expr Clip(Expr x, double a_min, double a_max) { return MakeClip(x, a_min, a_max); } - -inline Expr FixedPointMultiply(Expr x, int32_t multiplier, int32_t shift) { - static const Op& op = Op::Get("fixed_point_multiply"); - auto attrs = make_object(); - attrs->multiplier = multiplier; - attrs->shift = shift; - return Call(op, {x}, Attrs(attrs), {}); -} - -inline Expr FixedPointMultiplyPerAxis(Expr x, Expr m, Expr lshift, Expr rshift, - bool is_lshift_required, bool is_rshift_required, - Array axes) { - return MakeFixedPointMultiplyPerAxis(x, m, lshift, rshift, is_lshift_required, is_rshift_required, - axes); -} - -inline Expr Add(Expr lhs, Expr rhs) { - static const Op& op = Op::Get("add"); - return Call(op, {lhs, rhs}, Attrs(), {}); -} - -inline Expr Subtract(Expr lhs, Expr rhs) { - static const Op& op = Op::Get("subtract"); - return Call(op, {lhs, rhs}, Attrs(), {}); -} - -inline Expr Multiply(Expr lhs, Expr rhs) { - static const Op& op = Op::Get("multiply"); - return Call(op, {lhs, rhs}, Attrs(), {}); -} - -inline Expr Divide(Expr lhs, Expr rhs) { - static const Op& op = Op::Get("divide"); - return Call(op, {lhs, rhs}, Attrs(), {}); -} - -inline Expr Maximum(Expr lhs, Expr rhs) { - static const Op& op = Op::Get("maximum"); - return Call(op, {lhs, rhs}, Attrs(), {}); -} - -inline Expr ZerosLike(Expr e) { - static const Op& op = Op::Get("zeros_like"); - return Call(op, {e}); -} - -inline Expr Zeros(Array shape, DataType dtype) { - return MakeZeros(CheckConstantShapeArrayInteger(shape), dtype); -} - -inline Expr OnesLike(Expr e) { - static const Op& op = Op::Get("ones_like"); - return Call(op, {e}); -} - -inline Expr Ones(Array shape, DataType dtype) { - return MakeOnes(CheckConstantShapeArrayInteger(shape), dtype); -} - -inline Expr CollapseSumLike(Expr e) { - static const Op& op = Op::Get("collapse_sum_like"); - return Call(op, {e}); -} - -inline Expr Power(Expr lhs, Expr rhs) { - static const Op& op = Op::Get("power"); - return Call(op, {lhs, rhs}, Attrs(), {}); -} - -inline Expr RightShift(Expr x, Expr nbit) { - static const Op& op = Op::Get("right_shift"); - return Call(op, {x, nbit}, Attrs(), {}); -} - -inline Expr LeftShift(Expr x, Expr nbit) { - static const Op& op = Op::Get("left_shift"); - return Call(op, {x, nbit}, Attrs(), {}); -} - -inline Expr ReshapeLike(Expr lhs, Expr rhs, int lhs_begin, Integer lhs_end, int rhs_begin, - Integer rhs_end) { - return MakeReshapeLike(lhs, rhs, lhs_begin, lhs_end, rhs_begin, rhs_end); -} - -inline Expr Copy(Expr data) { - static const Op& op = Op::Get("copy"); - return Call(op, {data}, Attrs(), {}); -} - -inline Expr Max(Expr data, Array axis, bool keepdims, bool exclude) { - return MakeReduce(data, axis, keepdims, exclude, "max"); -} - -inline Expr Mean(Expr data, Array axis, bool keepdims, bool exclude) { - return MakeReduce(data, axis, keepdims, exclude, "mean"); -} - -inline Expr Variance(Expr data, Expr mean, Array axis, bool keepdims, bool exclude, - bool unbiased = false) { - return MakeVariance(data, mean, axis, keepdims, exclude, unbiased); -} - -static inline Expr Where(const Expr& condition, const Expr& x, const Expr& y) { - static const Op& op = Op::Get("where"); - return Call(op, {condition, x, y}); -} - -static inline Expr LogicalOr(const Expr& lhs, const Expr& rhs) { - static const Op& op = Op::Get("logical_or"); - return Call(op, {lhs, rhs}, Attrs(), {}); -} - -static inline Expr GreaterEqual(const Expr& lhs, const Expr& rhs) { - static const Op& op = Op::Get("greater_equal"); - return Call(op, {lhs, rhs}, Attrs(), {}); -} - -static inline Expr Equal(const Expr& lhs, const Expr& rhs) { - static const Op& op = Op::Get("equal"); - return Call(op, {lhs, rhs}, Attrs(), {}); -} - -static inline Expr Less(const Expr& lhs, const Expr& rhs) { - static const Op& op = Op::Get("less"); - return Call(op, {lhs, rhs}, Attrs(), {}); -} - -static inline Expr IsFinite(const Expr x) { - static const Op& op = Op::Get("isfinite"); - return Call(op, {x}, Attrs(), {}); -} - -static inline Expr Full(Expr fill_value, Array shape, DataType dtype) { - return MakeFull(fill_value, CheckConstantShapeArrayInteger(shape), dtype); -} - -static inline Expr Conv2D(Expr data, Expr weight, Array strides, - Array padding, Array dilation, int groups, - IndexExpr channels, Array kernel_size, std::string data_layout, - std::string kernel_layout, std::string out_layout, DataType out_dtype) { - return MakeConv(data, weight, strides, padding, dilation, groups, channels, - kernel_size, data_layout, kernel_layout, out_layout, out_dtype, - "nn.conv2d"); -} - -static inline Expr Dense(Expr data, Expr weight, IndexExpr units, DataType out_dtype) { - return MakeDense(data, weight, units, out_dtype); -} - -static inline Expr Sum(Expr data, Array axis, bool keepdims, bool exclude) { - return MakeReduce(data, axis, keepdims, exclude, "sum"); -} - -static inline Expr Prod(Expr data, Array axis, bool keepdims, bool exclude) { - return MakeReduce(data, axis, keepdims, exclude, "prod"); -} - -static inline Expr Reshape(Expr data, Array newshape) { - return MakeReshape(data, newshape); -} - -static inline Expr AvgPool2D(Expr data, Array pool_size, Array strides, - Array dilation, Array padding, - std::string layout, std::string out_layout, bool ceil_mode, - bool count_include_pad) { - return MakeAvgPool(data, pool_size, strides, dilation, padding, layout, - out_layout, ceil_mode, count_include_pad, "nn.avg_pool2d"); -} - -static inline Expr Pad(Expr data, Array> pad_width, Expr pad_value, - std::string pad_mode) { - Array> pad_width_int; - for (size_t i = 0; i < pad_width.size(); ++i) { - pad_width_int.push_back(CheckConstantShapeArrayInteger(pad_width[i])); - } - return MakePad(data, pad_width_int, pad_value, pad_mode); -} - -static inline Expr Tile(Expr data, Array reps) { return MakeTile(data, reps); } - -static inline Expr BroadCastTo(Expr data, Array shape) { - return MakeBroadCastTo(data, CheckConstantShapeArrayInteger(shape)); -} - -inline Expr Hardswish(Expr x) { - auto three = MakeConstantScalar(DataType::Float(32), 3.0); - auto six = MakeConstantScalar(DataType::Float(32), 6.0); - auto x2 = Add(x, three); - x2 = Clip(x2, 0.0, 6.0); - x2 = Multiply(x, x2); - x2 = Divide(x2, six); - return x2; -} - -} // namespace relay -} // namespace tvm -#endif // TVM_RELAY_TRANSFORMS_PATTERN_UTILS_H_ diff --git a/src/relay/transforms/remove_standalone_reshapes.cc b/src/relay/transforms/remove_standalone_reshapes.cc deleted file mode 100644 index 063060b3ebf9..000000000000 --- a/src/relay/transforms/remove_standalone_reshapes.cc +++ /dev/null @@ -1,118 +0,0 @@ -/* - * 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 src/relay/transforms/remove_standalone_reshapes.cc - * \brief This file contains the Relay pass for removing unfused reshapes from lowered graph. - */ - -#include -#include - -#include "../op/call/call.h" -#include "../op/memory/on_device.h" - -namespace tvm { -namespace relay { - -TVM_REGISTER_PASS_CONFIG_OPTION("relay.remove_standalone_reshapes.enable", Bool); -/*! Removes reshapes right after LowerTE. Removes preceding on_device calls - * while removing reshapes. - */ -class RemoveStandaloneReshapesMutator : public MixedModeMutator { - public: - explicit RemoveStandaloneReshapesMutator(IRModule& mod) {} // NOLINT(runtime/references) - - using MixedModeMutator::VisitExpr_; - - /*! * \brief Generated map of let variables to preceding CallLowered */ - Expr VisitExpr_(const LetNode* let) final { - Let ret_let; - Var var = Downcast(this->Mutate(let->var)); - auto value = this->Mutate(let->value); - if (auto* on_device_call = value.as()) { - OnDeviceProps on_device_props = GetOnDeviceProps(on_device_call); - if (on_device_props.body.defined() && on_device_props.body->IsInstance()) { - const Call call_lowered = Downcast(on_device_props.body); - if (call_lowered.defined() && call_lowered->op.same_as(CallLoweredOp())) { - let_var_to_call_lowered_.Set(var, call_lowered); - } - } - } - auto body = this->Mutate(let->body); - return WithFields(GetRef(let), var, value, body); - } - - /*! * \brief Returns preceding CallLowered when call is a CallLowered(Reshape) */ - Expr Rewrite_(const CallNode* call, const Expr& post) final { - /* - %1 = call_lowered(@tvmgen_default_non_reshape_function, %input, ...); - let %x: = on_device(%1, ...); - %2 = (%x,); - %3 = call_lowered(@tvmgen_default_fused_reshape, %2, ..., - "relay_attrs"=__dict__="relay.reshape_only"=1, ...); - */ - const CallNode* post_call = post.as(); - CallLoweredProps call_lowered_props = GetCallLoweredProps(post_call); - if (call_lowered_props.lowered_func.defined() && IsReshapeOnly(call_lowered_props)) { - if (!call_lowered_props.arguments.empty() && - call_lowered_props.arguments[0]->IsInstance()) { - Var var = Downcast(call_lowered_props.arguments[0]); - if (var.defined() && let_var_to_call_lowered_.find(var) != let_var_to_call_lowered_.end()) { - return let_var_to_call_lowered_[var]; - } - } - } - - return post; - } - - private: - /*! \brief Map of LetNode's var to previous call_lowered. */ - Map let_var_to_call_lowered_; -}; - -namespace transform { - -Pass RemoveStandaloneReshapes() { - auto pass_func = [=](IRModule mod, const PassContext& pass_ctx) { - VLOG(1) << "RemoveStandaloneReshapes before:" << std::endl << PrettyPrint(mod); - RemoveStandaloneReshapesMutator remove_reshapes_mutator(mod); - Function main_func = Downcast(mod->Lookup("main")); - Expr new_main_body = remove_reshapes_mutator.VisitExpr(main_func->body); - if (!new_main_body.same_as(main_func->body)) { - auto main_var = mod->GetGlobalVar("main"); - auto new_main_func = Function(main_func->params, new_main_body, main_func->ret_type, - main_func->type_params, main_func->attrs); - mod->Update(main_var, new_main_func); - } - Array entry_functions{"main"}; - mod = RemoveUnusedFunctions(entry_functions)(mod); - - VLOG(1) << "RemoveStandaloneReshapes after:" << std::endl << PrettyPrint(mod); - return mod; - }; - return tvm::transform::CreateModulePass(pass_func, 0, "RemoveStandaloneReshapes", {}); -} - -TVM_REGISTER_GLOBAL("relay._transform.RemoveStandaloneReshapes") - .set_body_typed(RemoveStandaloneReshapes); - -} // namespace transform -} // namespace relay -} // namespace tvm diff --git a/src/relay/transforms/simplify_expr.cc b/src/relay/transforms/simplify_expr.cc deleted file mode 100644 index 8036d301e191..000000000000 --- a/src/relay/transforms/simplify_expr.cc +++ /dev/null @@ -1,1165 +0,0 @@ -/* - * 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 src/relay/transforms/simplify_expr.cc - * \brief A pass for simplifying the Relay expression. - */ - -#include "simplify_expr.h" - -#include -#include -#include -#include -#include - -#include -#include -#include -#include -#include - -#include "../op/tensor/transform.h" -#include "fold_constant.h" -#include "pattern_utils.h" - -namespace tvm { -namespace relay { - -/*! - * \brief SimplifyReshape matches the pattern of consecutive reshape or reverse_reshape ops, - * and merges into one reshape op. - */ -class SimplifyReshape : public DFPatternRewrite { - public: - SimplifyReshape() { - x_ = IsWildcard(); - auto reshape1 = IsOp("reshape") || IsOp("contrib_reverse_reshape"); - auto reshape2 = IsOp("reshape") || IsOp("contrib_reverse_reshape"); - pattern_ = reshape1({reshape2({x_})}); - } - - Expr Callback(const Expr& pre, const Expr& post, - const Map>& node_map) const override { - auto x = node_map[x_][0]; - bool const_shape = true; - Array newshape; - for (auto dim : Downcast(pre->checked_type())->shape) { - if (dim.as() == nullptr) { - const_shape = false; - break; - } - newshape.push_back(Downcast(dim)); - } - if (const_shape) { - return MakeReshape(x, newshape); - } - return post; - } - - private: - /*! \brief Pattern input */ - DFPattern x_; -}; - -/*! - * \brief SimplifySameCast matches the pattern of cast data to the same dtype. - */ -class SimplifySameCast : public DFPatternRewrite { - public: - SimplifySameCast() { - data_pat_ = IsWildcard(); - like_pat_ = IsWildcard(); - pattern_ = IsOp("cast_like")({data_pat_, like_pat_}) || IsOp("cast")({data_pat_}); - } - - Expr Callback(const Expr& pre, const Expr& post, - const Map>& node_map) const override { - const CallNode* call = pre.as(); - const TensorTypeNode* data_ty = call->args[0]->checked_type().as(); - const TensorTypeNode* like_ty = pre->checked_type().as(); - if (like_ty->dtype == data_ty->dtype) { - return node_map[data_pat_][0]; - } - return post; - } - - protected: - DFPattern data_pat_; - DFPattern like_pat_; -}; - -/*! - * \brief SimplifyConsecutiveCast matches the pattern of consecutive cast/cast_like ops - */ -class SimplifyConsecutiveCast : public DFPatternRewrite { - public: - SimplifyConsecutiveCast() { - data_ = IsWildcard(); - cast1_ = IsOp("cast_like")({data_, IsWildcard()}) || IsOp("cast")({data_}); - pattern_ = IsOp("cast_like")({cast1_, IsWildcard()}) || IsOp("cast")({cast1_}); - } - - Expr Callback(const Expr& pre, const Expr& post, - const Map>& node_map) const override { - auto data = node_map[data_][0]; - auto cast1 = Downcast(node_map[cast1_][0]); - auto data_type = Downcast(data->checked_type()); - DataType cast1_dtype = Downcast(cast1->checked_type())->dtype; - - if (!IsWidenCast(data_type->dtype, cast1_dtype)) { - // Cannot remove the narrow cast - return post; - } - - const CallNode* cast2 = post.as(); - DataType cast2_dtype = Downcast(cast2->checked_type())->dtype; - auto expr = MakeCast(data, cast2_dtype); - - // We need to set the checked type as it may be needed in the next callback - expr->checked_type_ = TensorType(data_type->shape, cast2_dtype); - return expr; - } - - bool IsWidenCast(DataType origin, DataType cast) const { - /* Return whether casting from origin to cast results in more or the same precision.*/ - if (origin.code() == cast.code() && origin.bits() <= cast.bits()) { - return true; - } - if (origin.code() == DataType::kBFloat || cast.code() == DataType::kBFloat) { - // BFloat cast cannot be omitted - return false; - } - if (origin.code() < cast.code() && origin.bits() <= cast.bits()) { - // Loosely have a hiearchy to datatypes - // e.g. int --> uint --> float has increasing range of numbers they can represent - return true; - } - return false; - } - - protected: - DFPattern data_; - DFPattern cast1_; -}; - -/*! If mode == 0, return true if the interval [min_value, max_value] contains the range of dtype, - * and return false otherwise. If mode == 1, return true if the interval [min_value, max_value] is - * contained by the range of dtype, and return false otherwise.*/ -bool CheckDataTypeMaxMinValue(DataType dtype, double min_value, double max_value, int mode = 0) { - double lbound{}, ubound{}; - if (dtype.is_int() || dtype.is_uint()) { - ubound = static_cast(Downcast(tvm::max_value(dtype))->value); - lbound = static_cast(Downcast(tvm::min_value(dtype))->value); - } else if (dtype.is_float() || dtype.is_bfloat16()) { - ubound = Downcast(tvm::max_value(dtype))->value; - lbound = Downcast(tvm::min_value(dtype))->value; - } - if (mode == 0) { - return max_value >= ubound && min_value <= lbound; - } else if (mode == 1) { - return max_value <= ubound && min_value >= lbound; - } else { - LOG(FATAL) << "invalid mode " << mode << " in CheckDataTypeMaxMinValue"; - return false; - } -} - -/*! - * \brief SimplifyClipAndConsecutiveCast matches the pattern clip->cast->...->cast and remove - * redundant casts. Analysis of "redundancy" is done based on clip min/max values and min/max values - * of casted data type. - * - * Example: - * %0 == [type=int32] - * %1 = clip(%0, a_min=0f, a_max=255f) [type=int32] - * %2 = cast(%1, dtype="uint8") [type=uint8] - * %3 = cast(%2, dtype="int32") [type=int32] - * - * Optimized to (both casts can be removed): - * %1 = clip(%0, a_min=0f, a_max=255f) [type=int32] - */ -class SimplifyClipAndConsecutiveCast : public DFPatternRewrite { - public: - SimplifyClipAndConsecutiveCast() { - clip_ = IsOp("clip")({IsWildcard()}); - ObjectPtr pattern_ptr = make_object(); - pattern_ptr->op = IsOp("cast"); - pattern_ptr->args.clear(); - pattern_ = CallPattern(pattern_ptr); - AltPattern or_pattern{pattern_, clip_}; - pattern_ptr->args.push_back(or_pattern); - } - - Expr Callback(const Expr& pre, const Expr& post, - const Map>& node_map) const override { - auto clip = Downcast(node_map[clip_][0]); - const CallNode* clip_node = clip.as(); - const ClipAttrs* clip_attrs = clip_node->attrs.as(); - - std::vector remaining_casts{}; - Expr cast_expr{post}; - while (cast_expr != clip) { - DataType cast_dtype = Downcast(cast_expr->checked_type())->dtype; - if (!CheckDataTypeMaxMinValue(cast_dtype, clip_attrs->a_min, clip_attrs->a_max, 1)) { - remaining_casts.push_back(cast_expr); - } - cast_expr = cast_expr.as()->args[0]; - } - - Expr last_op = (remaining_casts.size() == 0) ? clip : remaining_casts[0]; - DataType last_op_dtype = Downcast(last_op->checked_type())->dtype; - bool need_additional_cast{false}; - if (last_op_dtype != Downcast(post->checked_type())->dtype) { - need_additional_cast = true; - } - - Expr res{clip}; - for (size_t i = remaining_casts.size(); i > 0; --i) { - auto attrs = make_object(); - attrs->dtype = remaining_casts[i - 1].as()->attrs.as()->dtype; - res = Call(Op::Get("cast"), {res}, Attrs(attrs), {}); - } - if (need_additional_cast) { - auto attrs = make_object(); - attrs->dtype = Downcast(post->checked_type())->dtype; - res = Call(Op::Get("cast"), {res}, Attrs(attrs), {}); - } - return res; - } - - protected: - DFPattern clip_; -}; - -/*! - * \brief SimplifyClip removes redundant Clip based on its a_min/a_max values and the min/max values - * of the data type. - * - * Example: - * %1 = cast(%0, dtype="uint8") [type=uint8] - * %2 = clip(%1, a_min=0f, a_max=255f) [type=int8] - * - * Optimized to (remove Clip): - * %1 = cast(%0, dtype="uint8") [type=uint8] - */ -class SimplifyClip : public DFPatternRewrite { - public: - SimplifyClip() { - x_ = IsWildcard(); - pattern_ = IsOp("clip")({x_}); - } - - Expr Callback(const Expr& pre, const Expr& post, - const Map>& node_map) const override { - DataType cast_dtype = Downcast(pre->checked_type())->dtype; - - const CallNode* clip_node = post.as(); - const ClipAttrs* clip_attrs = clip_node->attrs.as(); - - // TODO(kfeng123): For now, the arg of "clip" is forced to not be "qnn.requantize" and - // "qnn.add". This is to avoid destroying the structure required by LegalizeQnnOpForDnnl - auto child{post.as()->args[0].as()}; - if (child && child->op.as()) { - String op_name{child->op.as()->name}; - if (op_name == "qnn.requantize" || op_name == "qnn.add") { - return post; - } - } - - if (CheckDataTypeMaxMinValue(cast_dtype, clip_attrs->a_min, clip_attrs->a_max)) { - return node_map[x_][0]; - } - return post; - } - - protected: - DFPattern x_; -}; - -/*! - * \brief Return the axis order for layout transform and transpose - * ops. - */ -static std::vector GetTransposeAxisOrder(const Call& call, int ndim) { - std::vector attr_axes; - if (auto attr = call->attrs.as()) { - if (attr->axes.defined()) { - for (int i = 0; i < ndim; ++i) { - int64_t axis = attr->axes[i].IntValue(); - axis += (axis < 0) ? ndim : 0; - attr_axes.push_back(axis); - } - } else { - // Empty axes means reverse - for (int i = ndim - 1; i >= 0; --i) { - attr_axes.push_back(i); - } - } - } else if (auto attr = call->attrs.as()) { - Layout src_layout(attr->src_layout); - Layout dst_layout(attr->dst_layout); - for (int i = 0; i < ndim; ++i) { - attr_axes.push_back(src_layout.IndexOf(dst_layout[i])); - } - } else { - CHECK(false) << "Expected transpose or layout_transform, but got " - << Downcast(call->op)->name; - } - return std::move(attr_axes); -} - -/*! - * \brief SimplifyTranspose matches the pattern of consecutive transpose op, - * and merges or cancels them. - */ -class SimplifyTranspose : public DFPatternRewrite { - public: - SimplifyTranspose() { - x_ = IsWildcard(); - auto trans1 = IsOp("transpose") || IsOp("layout_transform"); - auto trans2 = IsOp("transpose") || IsOp("layout_transform"); - pattern_ = trans1({trans2({x_})}); - } - - Expr Callback(const Expr& pre, const Expr& post, - const Map>& node_map) const override { - auto x = node_map[x_][0]; - - Call trans_call = Downcast(post); - - // Try to fuse any rank changing layout transformations - if (auto layout_trans = FoldRankChangingLayoutTrans(x, trans_call)) { - if (auto attr = layout_trans.value()->attrs.as()) { - // Prune any trivial layout transformation - if (attr->src_layout == attr->dst_layout) { - return x; - } - } - return layout_trans.value(); - } - - // Initialize axes - int ndim = Downcast(pre->checked_type())->shape.size(); - Array axes; - for (int i = 0; i < ndim; ++i) { - axes.push_back(i); - } - - // Collect axes changes from the matched pattern, including two consecutive transposes. - std::vector> interm_axes; - interm_axes.push_back(GetTransposeAxisOrder(trans_call, ndim)); - trans_call = Downcast(trans_call->args[0]); - interm_axes.push_back(GetTransposeAxisOrder(trans_call, ndim)); - - // Calculate the final axes in reverse order (from root to output) - auto it = interm_axes.rbegin(); - while (it != interm_axes.rend()) { - auto interm = *it; - - Array new_axes; - for (int i = 0; i < ndim; ++i) { - new_axes.push_back(axes[interm[i]]); - } - axes = new_axes; - it++; - } - - return MakeTranspose(x, axes); - } - - String PermuteLayout(const String& layout, std::vector axes_order) const { - std::string new_layout{}; - std::string old_layout{layout}; - ICHECK_EQ(axes_order.size(), layout.size()) - << "Number of axes must match the number of named axes in the layout to permute: length(" - << old_layout << ") != " << axes_order.size(); - std::stringstream order; - for (auto axis : axes_order) { - new_layout += old_layout[axis]; - order << axis << ", "; - } - DLOG(INFO) << "Using transpose axes order {" << order.str() - << "} to permute layout: " << old_layout << " to " << new_layout; - return new_layout; - } - - struct RankChangingLayoutDescriptor { - Layout src_layout; - Layout dst_layout; - // Either a rank changing layout transform or a transpose - Call other_transform; - }; - - std::unique_ptr GetRankChangeDescriptor(const Call& call) const { - std::unique_ptr desc{nullptr}; - if (auto attr = call->attrs.as()) { - if (attr->src_layout.length() != attr->dst_layout.length()) { - desc = std::make_unique(); - desc->src_layout = Layout(attr->src_layout); - desc->dst_layout = Layout(attr->dst_layout); - desc->other_transform = Downcast(call->args[0]); - } - } - if (auto attr = Downcast(call->args[0])->attrs.as()) { - if (attr->src_layout.length() != attr->dst_layout.length()) { - if (!desc) { - desc = std::make_unique(); - desc->src_layout = Layout(attr->src_layout); - desc->dst_layout = Layout(attr->dst_layout); - desc->other_transform = call; - } else { - ICHECK(desc->src_layout->name == attr->dst_layout) - << "Back-to-back layout transforms must have the same intermediate layout: " - << desc->src_layout->name << " != " << attr->dst_layout; - desc->src_layout = Layout(attr->src_layout); - } - } - } - return desc; - } - - /* - * \brief Fuse call and it's argument into a single layout_transform operator - * when either call or it's argument is a rang changing layout_transform, e.g., - * - * Simplify - * - * [N, H, W, C] -> Transpose -> [N, C, H, W] -> LayoutTrans -> [N, C, H, W, 4c] - * - * to, - * - * [N, H, W, C] -> LayoutTrans -> [N, C, H, W, 4c]. - * - * \param The input expression to the matched pattern - * \param The pattern root; the second of two consecutive Transpose/LayoutTransform ops - */ - Optional FoldRankChangingLayoutTrans(const Expr& data, const Call& call) const { - // Check to see if either the first or second call in matched pattern - // is a rank changing layout transform. If so, return a descriptor containing - // the layouts and any additional transpose or layout transform op. - auto desc = GetRankChangeDescriptor(call); - if (desc == nullptr) { - // No rank changing layout transform - return Optional{nullptr}; - } - - Optional output_layout_trans; - // Fuse a rank increasing layout transform and a preceeding transpose - if (desc->src_layout->axes.size() < desc->dst_layout->axes.size()) { - auto axes = GetTransposeAxisOrder(desc->other_transform, desc->src_layout->axes.size()); - // Calculate the reverse axis order and apply to the source layout - std::vector inverse(axes.size()); - for (size_t i = 0; i < axes.size(); i++) { - inverse[axes[i]] = i; - } - String new_layout = PermuteLayout(desc->src_layout->name, inverse); - output_layout_trans = MakeLayoutTransform(data, new_layout, desc->dst_layout->name); - // Fuse a rank descreasing layout transform followed by a transpose - } else if (desc->src_layout->axes.size() > desc->dst_layout->axes.size()) { - auto axes = GetTransposeAxisOrder(desc->other_transform, desc->dst_layout->axes.size()); - String new_layout = PermuteLayout(desc->dst_layout->name, axes); - output_layout_trans = MakeLayoutTransform(data, desc->src_layout->name, new_layout); - // Fuse two back-to-back layout transformations which change rank - } else if (desc->other_transform->attrs.as()) { - output_layout_trans = - MakeLayoutTransform(data, desc->src_layout->name, desc->dst_layout->name); - } - return Downcast(output_layout_trans); - } - - private: - /*! \brief Pattern input */ - DFPattern x_; -}; - -/*! - * \brief SimplifyNoOpTranspose matches the pattern of transpose or - * layout transform ops which do not change the layout or rank and - * removes the op. - */ -class SimplifyNoOpTranspose : public DFPatternRewrite { - public: - SimplifyNoOpTranspose() { - x_ = IsWildcard(); - auto trans1 = IsOp("transpose") || IsOp("layout_transform"); - pattern_ = trans1({x_}); - } - - Expr Callback(const Expr& pre, const Expr& post, - const Map>& node_map) const override { - auto x = node_map[x_][0]; - Call trans_call = Downcast(post); - - // Do not remove ops which change rank - if (auto attr = trans_call->attrs.as()) { - if (attr->src_layout != attr->dst_layout) { - return post; - } - } - - int ndim = Downcast(pre->checked_type())->shape.size(); - auto axes = GetTransposeAxisOrder(trans_call, ndim); - - bool need_transpose = false; - for (int i = 0; i < ndim; ++i) { - if (axes[i] != i) { - need_transpose = true; - break; - } - } - - if (!need_transpose) return x; - - return post; - } - - private: - /*! \brief Pattern input */ - DFPattern x_; -}; - -/*! - * \brief FullElementwise finds full like ops followed by broadcasting ops, and eliminates - * the full op by directly passing the fill value into the broadcasting op. - */ -class FullElementwise : public DFPatternRewrite { - public: - FullElementwise() { - x_ = IsWildcard(); - data_ = IsWildcard(); - value_ = IsConstant(); - - full_ = IsOp("full")({value_}) || IsOp("full_like")({data_, value_}); - ones_ = IsOp("ones")({}) || IsOp("ones_like")({data_}); - zeros_ = IsOp("zeros")({}) || IsOp("zeros_like")({data_}); - - Map attrs; - attrs.Set("TOpPattern", Integer(static_cast(kBroadcast))); - DFPattern op = IsWildcard().HasAttr(attrs); - DFPattern full = full_ || ones_ || zeros_; - pattern_ = op({full, x_}) || op({x_, full}); - } - - Expr Callback(const Expr& pre, const Expr& post, - const Map>& node_map) const override { - const CallNode* call = pre.as(); - ICHECK(call); - Type pre_type = pre->checked_type_; - ICHECK(pre_type.as()); - auto dtype = pre_type.as()->dtype; - auto x = node_map[x_][0]; - bool is_left = post.as()->args[1] == x; - Type x_type; - if (is_left) { - x_type = call->args[1]->checked_type_; - } else { - x_type = call->args[0]->checked_type_; - } - - if (StructuralEqual()(x_type, pre_type)) { - Expr value; - if (node_map.count(full_)) { - value = node_map[value_][0]; - ICHECK(IsConstScalar(value)); - } else if (node_map.count(ones_)) { - value = MakeConstantScalar(dtype, 1); - } else if (node_map.count(zeros_)) { - value = MakeConstantScalar(dtype, 0); - } else { - ICHECK(false) << "Didn't find a full op while matching full + elementwise"; - } - if (is_left) { - return Call(call->op, {value, x}, call->attrs, call->type_args, call->span); - } else { - return Call(call->op, {x, value}, call->attrs, call->type_args, call->span); - } - } - return post; - } - - private: - /*! \brief binary argument */ - DFPattern x_; - /*! \brief data ops get shape from */ - DFPattern data_; - /*! \brief constant input */ - DFPattern value_; - /*! \brief full op */ - DFPattern full_; - /*! \brief ones op */ - DFPattern ones_; - /*! \brief zeros op */ - DFPattern zeros_; -}; - -/*! - * \brief Converts `*_like` operators to their explicit shape equivalent (e.g. `zeros_like(x, y)` to - * `zeros(x, y.shape)`), when the target shape is concrete. This removes unnecessary dependencies - * and can enable more opportunities for operator fusion. - */ -class ConcretizeLikeRewrite : public DFPatternRewrite { - public: - explicit ConcretizeLikeRewrite(const Op& op) { - ICHECK(op->num_inputs == 1 || op->num_inputs == 2) - << "ConcretizeLike does not handle operators that aren't unary or binary, got: " << op; - like_pat_ = IsWildcard(); - data_pat_ = IsWildcard(); - if (op->num_inputs == 1) { - pattern_ = IsExpr(op)({like_pat_}); - } else { - pattern_ = IsExpr(op)({data_pat_, like_pat_}); - } - } - - virtual bool Check(const Expr& pre, const Expr& post, - const Map>& node_map) const { - const CallNode* call_node = pre.as(); - ICHECK(call_node); - - if (!call_node->checked_type().as()) { - return false; - } - - return true; - } - - virtual Expr Concretize(const Map>& node_map, Array shape, - DataType dtype) const = 0; - - Expr Callback(const Expr& pre, const Expr& post, - const Map>& node_map) const override { - if (!Check(pre, post, node_map)) { - return post; - } - - const TensorTypeNode* like_ty = pre->checked_type().as(); - Array cshape; - for (const auto& dim : like_ty->shape) { - if (auto imm = dim.as()) { - cshape.push_back(Integer(imm.value())); - } else { - // shape is not static, don't concretize - return post; - } - } - - return Concretize(node_map, cshape, like_ty->dtype); - } - - protected: - DFPattern data_pat_; - DFPattern like_pat_; -}; - -class ConcretizeZerosLikeRewrite : public ConcretizeLikeRewrite { - public: - ConcretizeZerosLikeRewrite() : ConcretizeLikeRewrite(Op::Get("zeros_like")) {} - - Expr Concretize(const Map>& node_map, Array shape, - DataType dtype) const override { - return MakeZeros(shape, dtype); - } -}; - -class ConcretizeOnesLikeRewrite : public ConcretizeLikeRewrite { - public: - ConcretizeOnesLikeRewrite() : ConcretizeLikeRewrite(Op::Get("ones_like")) {} - - Expr Concretize(const Map>& node_map, Array shape, - DataType dtype) const override { - return MakeOnes(shape, dtype); - } -}; - -class ConcretizeFullLikeRewrite : public ConcretizeLikeRewrite { - public: - ConcretizeFullLikeRewrite() : ConcretizeLikeRewrite(Op::Get("full_like")) {} - - Expr Concretize(const Map>& node_map, Array shape, - DataType dtype) const override { - // `like_pat_` here is `fill_value` - return MakeFull(node_map[like_pat_][0], shape, dtype); - } -}; - -class ConcretizeReshapeLikeRewrite : public ConcretizeLikeRewrite { - public: - ConcretizeReshapeLikeRewrite() : ConcretizeLikeRewrite(Op::Get("reshape_like")) {} - - Expr Concretize(const Map>& node_map, Array shape, - DataType dtype) const override { - return MakeReshape(node_map[data_pat_][0], shape); - } -}; - -class ConcretizeCollapseSumLikeRewrite : public ConcretizeLikeRewrite { - public: - ConcretizeCollapseSumLikeRewrite() : ConcretizeLikeRewrite(Op::Get("collapse_sum_like")) {} - - Expr Concretize(const Map>& node_map, Array shape, - DataType dtype) const override { - ICHECK_LE(shape.size(), std::numeric_limits::max()); - static const Op& op = Op::Get("collapse_sum_to"); - auto attrs = make_object(); - attrs->shape = shape; - std::vector s; - std::transform(shape.begin(), shape.end(), std::back_inserter(s), - [](Integer i) { return i.IntValue(); }); - auto cshape = MakeConstantTensor(DataType::Int(32), {static_cast(shape.size())}, s); - return Call(op, {node_map[data_pat_][0], cshape}, Attrs(attrs)); - } -}; - -class ConcretizeBroadcastToLikeRewrite : public ConcretizeLikeRewrite { - public: - ConcretizeBroadcastToLikeRewrite() : ConcretizeLikeRewrite(Op::Get("broadcast_to_like")) {} - - Expr Concretize(const Map>& node_map, Array shape, - DataType dtype) const override { - return MakeBroadCastTo(node_map[data_pat_][0], shape); - } -}; - -/*! - * \brief Converts cast_like operator to cast. Not inheriting from ConcretizeLikeRewrite - * because even if shape is not static, still can concretize. - */ -class ConcretizeCastLikeRewrite : public DFPatternRewrite { - public: - ConcretizeCastLikeRewrite() { - data_pat_ = IsWildcard(); - like_pat_ = IsWildcard(); - pattern_ = IsOp("cast_like")({data_pat_, like_pat_}); - } - - Expr Callback(const Expr& pre, const Expr& post, - const Map>& node_map) const override { - const CallNode* call_node = pre.as(); - ICHECK(call_node); - - if (!call_node->checked_type().as()) { - return post; - } - - const TensorTypeNode* like_ty = pre->checked_type().as(); - return MakeCast(node_map[data_pat_][0], like_ty->dtype); - } - - protected: - DFPattern data_pat_; - DFPattern like_pat_; -}; - -/*! \brief Eliminates expressions that are equivalent to identity. */ -class EliminateIdentityRewrite : public DFPatternRewrite { - public: - EliminateIdentityRewrite() { - x_ = IsWildcard(); - const_ = IsConstant(); - - DFPattern add_op = IsOp("add"); - DFPattern mul_op = IsOp("multiply"); - DFPattern zeros_expr = IsOp("zeros")({}) || IsOp("zeros_like")({IsWildcard()}) || const_; - DFPattern ones_expr = IsOp("ones")({}) || IsOp("ones_like")({IsWildcard()}) || const_; - - // add and multiply are commutative so we don't need another pattern for reversed args - DFPattern add_id = add_op({x_, zeros_expr}); - DFPattern mul_id = mul_op({x_, ones_expr}); - - DFPattern sub_id = IsOp("subtract")({x_, zeros_expr}); - DFPattern div_id = IsOp("divide")({x_, ones_expr}); - - pattern_ = add_id || mul_id || sub_id || div_id; - } - - bool CheckConstant(const OpNode* op, const ConstantNode* constant) const { - if (!IsScalar(GetRef(constant))) { - return false; - } - auto value = TryToScalar(constant->data, 0); - if (!value) { - // unsupported dtype - return false; - } - if (op->name == "add" || op->name == "subtract") { - return value.value() == 0.0; - } else if (op->name == "multiply" || op->name == "divide") { - return value.value() == 1.0; - } - return false; - } - - Expr Callback(const Expr& pre, const Expr& post, - const Map>& node_map) const override { - const CallNode* call = pre.as(); - ICHECK(call); - Type pre_type = pre->checked_type_; - ICHECK(pre_type.as()); - auto x = node_map[x_][0]; - bool is_left = post.as()->args[1] == x; - Type x_type; - if (is_left) { - x_type = call->args[1]->checked_type_; - } else { - x_type = call->args[0]->checked_type_; - } - - if (node_map.count(const_)) { - // the other argument is a Constant in this case - const ConstantNode* constant = node_map[const_][0].as(); - const OpNode* op = call->op.as(); - ICHECK(constant); - ICHECK(op); - if (!CheckConstant(op, constant)) { - return post; - } - } - - if (StructuralEqual()(x_type, pre_type)) { - return x; - } - - return post; - } - - private: - DFPattern x_; - DFPattern const_; -}; - -/*! \brief Switch adjacent add-mul with constants to mul-add. - * As mul-add pattern is more friendly to FoldScaleAxis. - */ -class SwitchAddMultiply : public DFPatternRewrite { - public: - SwitchAddMultiply() { - x_ = IsWildcard(); - c1_ = IsConstant(); - c2_ = IsConstant(); - pattern_ = (x_ + c1_) * c2_; - } - - Expr Callback(const Expr& pre, const Expr& post, - const Map>& node_map) const override { - auto x = node_map[x_][0]; - auto c1 = node_map[c1_][0]; - auto c2 = node_map[c2_][0]; - - if (x.as()) { - return post; - } - - Expr const_expr = Call(Op::Get("multiply"), {c1, c2}); - Expr const_val = transform::FoldConstantExpr(const_expr); - - return Call(Op::Get("add"), {Call(Op::Get("multiply"), {x, c2}), const_val}); - } - - private: - DFPattern x_; - DFPattern c1_; - DFPattern c2_; -}; - -/*! \brief Simplify two adjacent multiply or add with constants for further constant folding. - * The pattern matching supports commutative property. - */ -class SimplifyAdjacentMultiplyOrAdd : public DFPatternRewrite { - public: - SimplifyAdjacentMultiplyOrAdd() { - x_ = IsWildcard(); - c1_ = IsConstant(); - c2_ = IsConstant(); - pattern_ = (x_ * c1_ * c2_) || (x_ + c1_ + c2_); - } - - Expr Callback(const Expr& pre, const Expr& post, - const Map>& node_map) const override { - const CallNode* call = pre.as(); - auto x = node_map[x_][0]; - auto c1 = node_map[c1_][0]; - auto c2 = node_map[c2_][0]; - - if (x.as()) { - return post; - } - - Expr const_expr = Call(call->op, {c1, c2}); - Expr const_val = transform::FoldConstantExpr(const_expr); - - return Call(call->op, {x, const_val}); - } - - private: - DFPattern x_; - DFPattern c1_; - DFPattern c2_; -}; - -/*! \brief Simplifying x+x to x*2 */ -class SimplifyAdd : public DFPatternRewrite { - public: - SimplifyAdd() { - x_ = IsWildcard(); - y_ = IsWildcard(); - pattern_ = IsOp("add")({x_, y_}); - } - - Expr Callback(const Expr& pre, const Expr& post, - const Map>& node_map) const override { - Type pre_type = pre->checked_type_; - auto dtype = pre_type.as()->dtype; - auto x = node_map[x_][0]; - auto y = node_map[y_][0]; - auto data_type = Downcast(x->checked_type()); - - if (x == y) { - Expr value; - value = MakeConstantScalar(dtype, 2); - return InferType(Call(Op::Get("multiply"), {x, value})); - } - return post; - } - - private: - /*! \brief Pattern input */ - DFPattern x_; - DFPattern y_; -}; - -/*! \brief Simplifying a * x * x + b * x * y + c * y * y to a * (x + p * y) * (x + q * y) */ -class SimplifyBinomial : public DFPatternRewrite { - public: - SimplifyBinomial() { - x_ = IsWildcard(); - y_ = IsWildcard(); - a_ = IsConstant(); - b_ = IsConstant(); - c_ = IsConstant(); - DFPattern add = IsOp("add"); - DFPattern mul = IsOp("multiply"); - DFPattern x_sq = mul({a_, mul({x_, x_})}) || mul({x_, mul({a_, x_})}) || mul({x_, x_}); - DFPattern xy = mul({b_, mul({x_, y_})}) || mul({x_, mul({b_, y_})}) || - mul({y_, mul({b_, x_})}) || mul({x_, y_}); - DFPattern y_sq = mul({c_, mul({y_, y_})}) || mul({y_, mul({c_, y_})}) || mul({y_, y_}); - - pattern_ = add({add({xy, x_sq}), y_sq}) || add({add({xy, y_sq}), x_sq}) || - add({add({x_sq, y_sq}), xy}); - } - - Expr Callback(const Expr& pre, const Expr& post, - const Map>& node_map) const override { - Type pre_type = pre->checked_type_; - auto dtype = pre_type.as()->dtype; - auto x = node_map[x_][0]; - auto y = node_map[y_][0]; - double a_val = 1; - double b_val = 1; - double c_val = 1; - double* vals[] = {&a_val, &b_val, &c_val}; - DFPattern nodes[] = {a_, b_, c_}; - for (int i = 0; i < 3; i++) { - if (node_map.count(nodes[i]) > 0) { - if (dtype == DataType::Int(32, 1)) - *vals[i] = static_cast( - transform::FoldConstantExpr(node_map[nodes[i]][0]).as()->data->data)[0]; - else if (dtype == DataType::Float(32, 1)) - *vals[i] = static_cast( - transform::FoldConstantExpr(node_map[nodes[i]][0]).as()->data->data)[0]; - else if (dtype == DataType::Float(64, 1)) - *vals[i] = static_cast( - transform::FoldConstantExpr(node_map[nodes[i]][0]).as()->data->data)[0]; - } - } - if (c_val == 1 && a_val > 1) { - auto temp_exp = x; - x = y; - y = temp_exp; - float temp_val = a_val; - a_val = c_val; - c_val = temp_val; - } - - double sub_value = b_val * b_val - 4 * a_val * c_val; - if (sub_value < 0) return pre; - bool same_multiplicands = sub_value < 10e-5; - - double discriminant = std::sqrt(sub_value); - Expr first_val = MakeConstantScalar(dtype, (b_val + discriminant) / (2 * a_val)); - Expr second_val = same_multiplicands - ? first_val - : MakeConstantScalar(dtype, (b_val - discriminant) / (2 * a_val)); - - Expr first_multiplicand = Call(Op::Get("add"), {x, Call(Op::Get("multiply"), {y, first_val})}); - Expr second_multiplicand = - same_multiplicands ? first_multiplicand - : Call(Op::Get("add"), {x, Call(Op::Get("multiply"), {y, second_val})}); - Expr a = MakeConstantScalar(dtype, a_val); - return Call(Op::Get("multiply"), - {a, Call(Op::Get("multiply"), {first_multiplicand, second_multiplicand})}); - } - - private: - /*! \brief Pattern input */ - DFPattern a_; - DFPattern b_; - DFPattern c_; - DFPattern x_; - DFPattern y_; -}; - -/*! \brief Simplifying x/sqrt to x*rsqrt */ -class SimplifyRSqrt : public DFPatternRewrite { - public: - SimplifyRSqrt() { - x_ = IsWildcard(); - numerator_ = IsWildcard(); - auto sqrt = IsOp("sqrt"); - pattern_ = IsOp("divide")({numerator_, sqrt({x_})}); - } - - Expr Callback(const Expr& pre, const Expr& post, - const Map>& node_map) const override { - static const Op& op = Op::Get("rsqrt"); - auto x = node_map[x_][0]; - auto numerator = node_map[numerator_][0]; - return Call(Op::Get("multiply"), {numerator, Call(op, {x})}); - } - - private: - /*! \brief Pattern input */ - DFPattern x_; - DFPattern numerator_; -}; - -/*! \brief Base class for simplifying dequantize followed by arg ops */ -class SimplifyDQArgFunc : public DFPatternRewrite { - public: - explicit SimplifyDQArgFunc(std::string op) : op_(op) { - x_ = IsWildcard(); - dq_ = IsOp("qnn.dequantize")({x_, IsWildcard(), IsWildcard()}); - pattern_ = IsOp(op_)({dq_}); - } - - Expr Callback(const Expr& pre, const Expr& post, - const Map>& node_map) const override { - const CallNode* call = pre.as(); - ICHECK(call); - auto x = node_map[x_][0]; - return Call(Op::Get(op_), {x}, call->attrs); - } - - protected: - /*! \brief Pattern input */ - DFPattern x_; - /*! \brief dequantize op */ - DFPattern dq_; - /*! \brief Name of op to simplify */ - String op_; -}; - -/*! \brief Simplify dequantize follwed by argmax */ -class SimplifyDQArgMax : public SimplifyDQArgFunc { - public: - SimplifyDQArgMax() : SimplifyDQArgFunc("argmax") {} -}; - -/*! \brief Simplify dequantize follwed by argmin */ -class SimplifyDQArgMin : public SimplifyDQArgFunc { - public: - SimplifyDQArgMin() : SimplifyDQArgFunc("argmin") {} -}; - -/*! \brief Simplify dequantize follwed by argsort */ -class SimplifyDQArgSort : public SimplifyDQArgFunc { - public: - SimplifyDQArgSort() : SimplifyDQArgFunc("argsort") {} -}; - -Expr SimplifyExpr(const Expr& expr, const IRModule& mod) { - // the rewrites will be applied in the given order, and repeated until fixed point - DFPatternRewriteComposer composer; - composer.AddRewrite(); - composer.AddRewrite(); - composer.AddRewrite(); - composer.AddRewrite(); - composer.AddRewrite(); - composer.AddRewrite(); - composer.AddRewrite(); - composer.AddRewrite(); - composer.AddRewrite(); - composer.AddRewrite(); - composer.AddRewrite(); - composer.AddRewrite(); - composer.AddRewrite(); - composer.AddRewrite(); - composer.AddRewrite(); - composer.AddRewrite(); - composer.AddRewrite(); - composer.AddRewrite(); - composer.AddRewrite(); - composer.AddRewrite(); - composer.AddRewrite(); - composer.AddRewrite(); - composer.AddRewrite(); - composer.AddRewrite(); - return RewritePatterns(composer.MakeCallbacks(), expr, mod); -} - -Expr SimplifyExprPostAlterOp(const Expr& expr, const IRModule& mod) { - // stripped-down version of AlterOp that cleans up some patterns - // often left by the AlterOpLayout pass. - DFPatternRewriteComposer composer; - composer.AddRewrite(); - composer.AddRewrite(); - composer.AddRewrite(); - composer.AddRewrite(); - composer.AddRewrite(); - composer.AddRewrite(); - return RewritePatterns(composer.MakeCallbacks(), expr, mod); -} - -namespace transform { - -Pass SimplifyExpr() { - runtime::TypedPackedFunc pass_func = - [=](Function f, IRModule m, PassContext pc) { - return Downcast(SimplifyExpr(f, m)); - }; - return CreateFunctionPass(pass_func, 0, "SimplifyExpr", {"InferType"}); -} - -Pass SimplifyExprPostAlterOp() { - runtime::TypedPackedFunc pass_func = - [=](Function f, IRModule m, PassContext pc) { - return Downcast(SimplifyExprPostAlterOp(f, m)); - }; - return CreateFunctionPass(pass_func, 0, "SimplifyExprPostAlterOp", {"InferType"}); -} - -TVM_REGISTER_GLOBAL("relay._transform.SimplifyExpr").set_body_typed(SimplifyExpr); -TVM_REGISTER_GLOBAL("relay._transform.SimplifyExprPostAlterOp") - .set_body_typed(SimplifyExprPostAlterOp); - -} // namespace transform - -} // namespace relay -} // namespace tvm diff --git a/src/relay/transforms/simplify_expr.h b/src/relay/transforms/simplify_expr.h deleted file mode 100644 index bdd3f2ca6e6f..000000000000 --- a/src/relay/transforms/simplify_expr.h +++ /dev/null @@ -1,93 +0,0 @@ -/* - * 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 src/relay/transforms/simplify_expr.h - * \brief Utility data structures for simplifying Relay expressions. - */ -#ifndef TVM_RELAY_TRANSFORMS_SIMPLIFY_EXPR_H_ -#define TVM_RELAY_TRANSFORMS_SIMPLIFY_EXPR_H_ - -#include -#include - -#include -#include - -namespace tvm { -namespace relay { - -/*! \brief A wrapper class defining a rewrite matching a specific pattern. */ -class DFPatternRewrite { - public: - /*! \brief Returns the rewritten expression. */ - virtual Expr Callback(const Expr& pre, const Expr& post, - const Map>& node_map) const = 0; - - virtual ~DFPatternRewrite() = default; - - /*! \brief Returns the pattern to be used for matching and rewriting. */ - inline DFPattern Pattern() const { return pattern_; } - - inline bool RequireType() const { return require_type_; } - - inline DFPatternCallback MakeCallback() const { - auto func = [this](TVMArgs args, TVMRetValue* rv) { - Expr pre = args[0]; - Expr post = args[1]; - Map> node_map = args[2]; - *rv = this->Callback(pre, post, node_map); - }; - return DFPatternCallback(pattern_, PackedFunc(func), require_type_, rewrite_once_); - } - - protected: - /*! \brief The pattern for matching and rewriting. */ - DFPattern pattern_; - /*! \brief Whether or not the rewrite requires types to be inferred. */ - bool require_type_ = true; - /*! \brief Whether or not run the callback only once */ - bool rewrite_once_ = false; -}; - -/*! \brief Helper class for composing rewrites and getting callbacks. */ -class DFPatternRewriteComposer { - public: - template - inline void AddRewrite(Args... args) { - rewrites_.push_back(std::make_shared(args...)); - } - - inline Array MakeCallbacks() const { - Array callbacks; - for (const auto& rewrite : rewrites_) { - callbacks.push_back(rewrite->MakeCallback()); - } - return callbacks; - } - - private: - /*! \brief the rewrites to be composed. */ - std::vector> rewrites_; -}; - -} // namespace relay -} // namespace tvm - -#endif // TVM_RELAY_TRANSFORMS_SIMPLIFY_EXPR_H_ diff --git a/src/relay/transforms/simplify_fc_transpose.cc b/src/relay/transforms/simplify_fc_transpose.cc deleted file mode 100644 index ad38ea6cb8df..000000000000 --- a/src/relay/transforms/simplify_fc_transpose.cc +++ /dev/null @@ -1,146 +0,0 @@ -/* - * 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 simplify_fc_transpose.cc - * - * \brief Mutate ```y = nn.dense(x, tranpose(w, [1, 0]))``` to - * ```y = nn.dense(x, wt)``` - */ -#include -#include -#include -#include -#include -#include -#include - -#include -#include - -namespace tvm { -namespace relay { - -// Find name of weight in ```y = nn.dense(x, tranpose(w, [1, 0]))``` -class FCTransposeVisitor : private ExprVisitor { - public: - FCTransposeVisitor() : dense_op_(Op::Get("nn.dense")), transpose_op_(Op::Get("transpose")) {} - - Array Search(const Expr& expr) { - VisitExpr(expr); - return memo_; - } - - private: - void VisitExpr_(const CallNode* n) final { - if (n->op == dense_op_) { - const auto weight = n->args[1].as(); - if (weight) { - if (weight->op == transpose_op_) { - if (weight->args[0].as()) { - const auto arg = weight->args[0].as(); - memo_.push_back(arg->name_hint()); - } - } - } - } - for (const auto& arg : n->args) { - VisitExpr(arg); - } - } - - const Op& dense_op_; - const Op& transpose_op_; - Array memo_; -}; // SearchDenseOpWeight - -Array SearchFCTranspose(const Expr& e) { return FCTransposeVisitor().Search(e); } - -TVM_REGISTER_GLOBAL("relay.analysis.search_fc_transpose").set_body_typed(SearchFCTranspose); - -// Mutate ```y = nn.dense(x, tranpose(w, [1, 0]))``` to ```y = nn.dense(x, wt)``` -class FCTransposeMutator : public ExprRewriter { - public: - explicit FCTransposeMutator(const Array& target_weights) - : dense_op_(Op::Get("nn.dense")), transpose_op_(Op::Get("transpose")) { - for (size_t i = 0; i < target_weights.size(); ++i) { - ICHECK(target_weights[i]->IsInstance()); - std::string k = target_weights[i].as()->data; - target_weights_.emplace(k); - } - } - - Expr Rewrite_(const CallNode* pre, const Expr& post) override { - if (pre->op == dense_op_) { - const auto data = post.as()->args[0]; - const auto weight = pre->args[1].as(); - if (weight) { - if (weight->op == transpose_op_) { - const auto arg = weight->args[0]; - if (arg.as()) { - const auto& arg_node = arg.as(); - ICHECK_GT(target_weights_.count(arg_node->name_hint()), 0); - const auto& tt = arg_node->type_annotation.as(); - auto wt_type = TensorType({tt->shape[1], tt->shape[0]}, tt->dtype); - Var wt(arg_node->name_hint() + ".T", wt_type); - return Call(dense_op_, {data, wt}, pre->attrs, pre->type_args); - } - } - } - } - return post; - } - - private: - // Cached op - const Op& dense_op_; - const Op& transpose_op_; - std::unordered_set target_weights_; -}; // class DenseToSparseDenseAlter - -Expr SimplifyFCTranspose(const Expr& e, const Array& target_weights) { - auto rewriter = FCTransposeMutator(target_weights); - return PostOrderRewrite(e, &rewriter); -} - -namespace transform { - -Pass SimplifyFCTranspose(const Array& target_weights) { - runtime::TypedPackedFunc pass_func = - [=](Function f, IRModule m, PassContext pc) { - // Remove FreeVar warning - auto f0 = Downcast(SimplifyFCTranspose(f, target_weights)); - Array wt_params = FreeVars(f0); - auto f1 = WithFields(f0, wt_params); - Array params = FreeVars(f1); - for (const auto& var : wt_params) { - params.push_back(var); - } - return WithFields(f1, params); - }; - return CreateFunctionPass(pass_func, 4, "SimplifyFCTranspose", {"DeadCodeElimination"}); -} - -TVM_REGISTER_GLOBAL("relay._transform.SimplifyFCTranspose").set_body_typed(SimplifyFCTranspose); - -} // namespace transform - -} // namespace relay -} // namespace tvm diff --git a/src/relay/transforms/simplify_inference.cc b/src/relay/transforms/simplify_inference.cc deleted file mode 100644 index e7eef41e41c4..000000000000 --- a/src/relay/transforms/simplify_inference.cc +++ /dev/null @@ -1,259 +0,0 @@ -/* - * 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 simplify_inference.cc - */ -#include -#include -#include -#include -#include - -#include "pattern_utils.h" - -namespace tvm { -namespace relay { - -Expr BatchNormToInferUnpack(const Attrs attrs, Expr data, Expr gamma, Expr beta, Expr moving_mean, - Expr moving_var, Type tdata) { - auto ttype = tdata.as(); - ICHECK(ttype); - const auto param = attrs.as(); - Expr epsilon = MakeConstantScalar(ttype->dtype, static_cast(param->epsilon)); - Expr var_add_eps = Add(moving_var, epsilon); - Expr sqrt_var = Sqrt(var_add_eps); - Expr scale = Divide(MakeConstantScalar(ttype->dtype, 1.0f), sqrt_var); - - if (param->scale) { - scale = Multiply(scale, gamma); - } - Expr neg_mean = Negative(moving_mean); - Expr shift = Multiply(neg_mean, scale); - if (param->center) { - shift = Add(shift, beta); - } - - auto ndim = ttype->shape.size(); - int axis = (param->axis < 0) ? param->axis + ndim : param->axis; - scale = ExpandBiasToMatchAxis(scale, ndim, {axis}); - shift = ExpandBiasToMatchAxis(shift, ndim, {axis}); - - Expr out = Multiply(data, scale); - out = Add(out, shift); - return out; -} - -Expr GroupNormToInferUnpack(const Attrs attrs, Expr data, Expr gamma, Expr beta, Type tdata) { - auto ttype = tdata.as(); - ICHECK(ttype); - const auto param = attrs.as(); - ICHECK(param); - - int ndim = ttype->shape.size(); - int axis = (param->axis < 0) ? param->axis + ndim : param->axis; - Array reduced_axes; - Array new_shape; - Array old_shape; - - int num_groups = param->num_groups; - int channel = ttype->shape[axis].as()->value; - - // old_shape = N, C, H, W - // new shape = N, num_groups, C/num_groups, H, W - // reduce_axes = axis of (C/num_groups, H, W) - for (int i = 0; i < ndim; ++i) { - auto val = ttype->shape[i].as()->value; - - // Save the old shape to reshape later - old_shape.push_back(val); - if (i == axis) { - new_shape.push_back(num_groups); - new_shape.push_back(channel / num_groups); - reduced_axes.push_back(i + 1); - continue; - } - if (i >= axis) { - reduced_axes.push_back(i + 1); - } - new_shape.push_back(val); - } - - data = Reshape(data, new_shape); - - Expr epsilon = MakeConstantScalar(ttype->dtype, static_cast(param->epsilon)); - Expr mean = Mean(data, {reduced_axes}, true, false); - Expr var = Variance(data, mean, {reduced_axes}, true, false); - Expr denom = Sqrt(Add(var, epsilon)); - Expr out = Divide(Subtract(data, mean), denom); - - out = Reshape(out, old_shape); - - if (param->scale) { - out = Multiply(out, ExpandBiasToMatchAxis(gamma, ndim, {axis})); - } - if (param->center) { - out = Add(out, ExpandBiasToMatchAxis(beta, ndim, {axis})); - } - - return out; -} - -Expr LayerNormToInferUnpack(const Attrs attrs, Expr data, Expr gamma, Expr beta, Type tdata) { - auto ttype = tdata.as(); - ICHECK(ttype); - const auto param = attrs.as(); - ICHECK(param); - - Expr epsilon = MakeConstantScalar(ttype->dtype, static_cast(param->epsilon)); - Expr mean = Mean(data, {param->axis}, true, false); - Expr var = Variance(data, mean, {param->axis}, true, false); - Expr denom = Sqrt(Add(var, epsilon)); - Expr out = Divide(Subtract(data, mean), denom); - - size_t ndim = ttype->shape.size(); - int axis = (param->axis < 0) ? param->axis + ndim : param->axis; - if (param->scale) { - out = Multiply(out, ExpandBiasToMatchAxis(gamma, ndim, {axis})); - } - if (param->center) { - out = Add(out, ExpandBiasToMatchAxis(beta, ndim, {axis})); - } - return out; -} - -Expr InstanceNormToInferUnpack(const Attrs attrs, Expr data, Expr gamma, Expr beta, Type tdata) { - auto ttype = tdata.as(); - ICHECK(ttype); - const auto param = attrs.as(); - ICHECK(param); - - int ndim = ttype->shape.size(); - int axis = (param->axis < 0) ? param->axis + ndim : param->axis; - Array reduced_axes; - for (int i = 1; i < ndim; ++i) { - if (i != axis) reduced_axes.push_back(i); - } - - Expr epsilon = MakeConstantScalar(ttype->dtype, static_cast(param->epsilon)); - Expr mean = Mean(data, reduced_axes, true, false); - Expr var = Variance(data, mean, reduced_axes, true, false); - Expr denom = Sqrt(Add(var, epsilon)); - Expr out = Divide(Subtract(data, mean), denom); - - if (param->scale) { - out = Multiply(out, ExpandBiasToMatchAxis(gamma, ndim, {axis})); - } - if (param->center) { - out = Add(out, ExpandBiasToMatchAxis(beta, ndim, {axis})); - } - return out; -} - -Expr L2NormToInferUnpack(const Attrs attrs, Expr data) { - const auto param = attrs.as(); - ICHECK(param); - - Expr epsilon = MakeConstantScalar(DataType::Float(32), static_cast(param->eps)); - - Expr sqr = Multiply(data, data); - Expr sum = Maximum(Sum(sqr, param->axis, true, false), epsilon); - Expr sqrt = Sqrt(sum); - return Divide(data, sqrt); -} - -class InferenceSimplifier : public MixedModeMutator { - public: - InferenceSimplifier() - : batch_norm_op_(Op::Get("nn.batch_norm")), - dropout_op_(Op::Get("nn.dropout")), - instance_norm_op_(Op::Get("nn.instance_norm")), - layer_norm_op_(Op::Get("nn.layer_norm")), - group_norm_op_(Op::Get("nn.group_norm")), - l2_norm_op_(Op::Get("nn.l2_normalize")) {} - - Expr Rewrite_(const TupleGetItemNode* n, const Expr& new_e) final { - const auto* new_n = new_e.as(); - if (new_n->index != 0) { - return new_e; - } - if (const auto* call = new_n->tuple.as()) { - if (call->op == batch_norm_op_) { - return BatchNormToInferUnpack(call->attrs, call->args[0], call->args[1], call->args[2], - call->args[3], call->args[4], ty_map_.at(call->args[0])); - } else if (call->op == dropout_op_) { - return call->args[0]; - } - } - return new_e; - } - - Expr Rewrite_(const CallNode* n, const Expr& new_n) { - if (n->op == batch_norm_op_) { - ty_map_[new_n.as()->args[0]] = n->args[0]->checked_type(); - } else if (n->op == layer_norm_op_) { - const auto* call = new_n.as(); - return LayerNormToInferUnpack(call->attrs, call->args[0], call->args[1], call->args[2], - n->args[0]->checked_type()); - } else if (n->op == group_norm_op_) { - const auto* call = new_n.as(); - return GroupNormToInferUnpack(call->attrs, call->args[0], call->args[1], call->args[2], - n->args[0]->checked_type()); - } else if (n->op == instance_norm_op_) { - const auto* call = new_n.as(); - return InstanceNormToInferUnpack(call->attrs, call->args[0], call->args[1], call->args[2], - n->args[0]->checked_type()); - } else if (n->op == l2_norm_op_) { - const auto* call = new_n.as(); - return L2NormToInferUnpack(call->attrs, call->args[0]); - } - return new_n; - } - - private: - // Cache the following ops. They will be used in the passes repeatedly for - // operator equivalence checking so that the registry lookup overhead can be - // reduced. - const Op& batch_norm_op_; - const Op& dropout_op_; - const Op& instance_norm_op_; - const Op& layer_norm_op_; - const Op& group_norm_op_; - const Op& l2_norm_op_; - std::unordered_map ty_map_; -}; - -Expr SimplifyInference(const Expr& e) { return InferenceSimplifier().Mutate(e); } - -namespace transform { - -Pass SimplifyInference() { - runtime::TypedPackedFunc pass_func = - [=](Function f, IRModule m, PassContext pc) { - return Downcast(SimplifyInference(f)); - }; - return CreateFunctionPass(pass_func, 0, "SimplifyInference", {"InferType"}); -} - -TVM_REGISTER_GLOBAL("relay._transform.SimplifyInference").set_body_typed(SimplifyInference); - -} // namespace transform - -} // namespace relay -} // namespace tvm diff --git a/src/relay/transforms/split_args.cc b/src/relay/transforms/split_args.cc deleted file mode 100644 index 423adff9a4cb..000000000000 --- a/src/relay/transforms/split_args.cc +++ /dev/null @@ -1,143 +0,0 @@ -/* - * 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 split_args.cc - */ -#include -#include - -#include "../op/annotation/annotation.h" -#include "./pattern_utils.h" - -namespace tvm { -namespace relay { - -class ArgumentSplitter : public ExprRewriter { - public: - explicit ArgumentSplitter(size_t max_function_args) - : max_function_args_(max_function_args), concat_op_(Op::Get("concatenate")) {} - - Expr ConcatSplitter(const TupleNode* tuple_node, const tvm::Array& args, int axis, - size_t limit) { - tvm::Array new_args; - size_t added_args = 0; - for (const auto& it : args) { - size_t curr_args = 1; - if (const auto* ttype = it->checked_type().as()) { - ICHECK(additional_args_cache_.count(ttype)); - curr_args += additional_args_cache_[ttype]; - } - if (added_args + curr_args > limit) { - Tuple new_tuple = WithFields(GetRef(tuple_node), new_args); - Expr stop = StopFusion(new_tuple); - Expr lastExpr = MakeConcatenate(stop, axis); - new_args.clear(); - new_args.push_back(lastExpr); - added_args = curr_args; - } - added_args += curr_args; - new_args.push_back(it); - } - Tuple new_tuple = WithFields(GetRef(tuple_node), new_args); - Expr stop = StopFusion(new_tuple); - Expr lastExpr = MakeConcatenate(stop, axis); - return lastExpr; - } - - // In the case of dynamic shape in tensor, the sizes of any_dims and strides are passed as - // function args - size_t CalculateNumberOfAdditionalArgs_(const TensorTypeNode* arg, bool isOutput = false) { - size_t num = 0; - for (const auto& dim : arg->shape) { - if (dim.as()) { - num++; - } - } - // In the case of dynamic shape, strides are also passed to a function as arguments. The number - // of strides equals the rank of the tensor. - if (num > 0 && isOutput) - return arg->shape.size(); - else if (num > 0) - num += arg->shape.size(); - return num; - } - - Expr Rewrite_(const CallNode* call, const Expr& post) final { - if (max_function_args_ == 0) return post; - if (call->op == concat_op_) { - auto tuple_node = call->args[0].as(); - if (tuple_node == nullptr) return post; - const auto param = call->attrs.as(); - size_t outputsNum = 1; - if (const auto* tuple_type = call->checked_type().as()) { - outputsNum = tuple_type->fields.size(); - for (const auto& it : tuple_type->fields) { - if (const auto* ttype = it.as()) { - outputsNum += CalculateNumberOfAdditionalArgs_(ttype, true); - } - } - } else if (const auto* ttype = call->checked_type().as()) { - outputsNum += CalculateNumberOfAdditionalArgs_(ttype, true); - } - CHECK_GT(max_function_args_, outputsNum); - size_t limit = max_function_args_ - outputsNum; - - size_t argsNum = tuple_node->fields.size(); - for (const auto& it : tuple_node->fields) { - if (const auto* ttype = it->checked_type().as()) { - size_t any_dims = CalculateNumberOfAdditionalArgs_(ttype); - argsNum += any_dims; - additional_args_cache_[ttype] = any_dims; - } - } - if (argsNum < limit) return post; - return ConcatSplitter(tuple_node, tuple_node->fields, param->axis, limit); - } - return post; - } - - private: - const size_t max_function_args_; - const Op& concat_op_; - std::unordered_map additional_args_cache_; -}; - -Expr SplitArgs(const Expr& expr, size_t max_function_args) { - auto rewriter = ArgumentSplitter(max_function_args); - return PostOrderRewrite(expr, &rewriter); -} - -namespace transform { - -Pass SplitArgs(uint64_t max_function_args) { - runtime::TypedPackedFunc pass_func = - [=](Function f, IRModule m, PassContext pc) { - auto r = Downcast(SplitArgs(f, max_function_args)); - return m->attrs.defined() ? WithAttrs(r, {m->attrs->dict}) : r; - }; - return CreateFunctionPass(pass_func, 1, "SplitArgs", {"InferType"}); -} - -TVM_REGISTER_GLOBAL("relay._transform.SplitArgs").set_body_typed(SplitArgs); - -} // namespace transform - -} // namespace relay -} // namespace tvm diff --git a/src/relay/transforms/target_hooks.cc b/src/relay/transforms/target_hooks.cc deleted file mode 100644 index f52e95b2adbf..000000000000 --- a/src/relay/transforms/target_hooks.cc +++ /dev/null @@ -1,179 +0,0 @@ -/* - * 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 target_hooks.cc - * \brief Relay passes for processing Target Hooks which have been registered on functions within - * the IRModule - */ - -#include -#include - -namespace tvm { -namespace relay { -namespace transform { - -namespace { - -/*! - * \brief A pass extracted from a target kind's "RelayToTIR" attribute, along with any - * 'external codegen' Target instance with matching kind name which should be current when - * the pass is applied. - */ -struct CustomPass { - std::string target_kind_name; - Pass pass; - Optional opt_target; - - CustomPass(std::string target_kind_name, Pass pass, Optional opt_target) - : target_kind_name(std::move(target_kind_name)), - pass(std::move(pass)), - opt_target(std::move(opt_target)) {} -}; - -/*! - * \brief Collect all the \p CustomPasses needed according to the "Compiler" attributes on - * inlined or global functions. - */ -class TargetHookVisitor : public MixedModeVisitor { - public: - TargetHookVisitor(IRModule mod, CompilationConfig config) - : mod_(std::move(mod)), - config_(std::move(config)), - target_attr_map_(tvm::TargetKind::GetAttrMap(tvm::attr::kRelayToTIR)) {} - - std::vector Visit() { - ICHECK(custom_passes_.empty()); - // To ensure the passes are run in a deterministic order we'll search for functions in - // lexicographic order. - std::vector> functions; - for (const auto& kv : mod_->functions) { - functions.emplace_back(kv.first->name_hint, kv.second); - } - std::sort(functions.begin(), functions.end()); - for (const auto& kv : functions) { - if (const auto* function_node = kv.second.as()) { - // May be a top-level function with a "Compiler" attribute. - MaybeAddPassForFunction(function_node); - } - if (const auto* function_node = AsOptimizableFunctionNode(kv.second)) { - // May have calls to inlined "Compiler" functions in body. - VisitExpr(GetRef(function_node)); - } - } - return std::move(custom_passes_); - } - - private: - using tvm::relay::MixedModeVisitor::VisitExpr_; - - void VisitExpr_(const LetNode* let_node) final { - auto pre_visit = [this](const LetNode* inner_let_node) { - this->VisitExpr(inner_let_node->var); - this->VisitExpr(inner_let_node->value); - }; - auto post_visit = [this](const LetNode* inner_let_node) { - this->VisitExpr(inner_let_node->body); - this->visit_counter_[inner_let_node] += 1; - }; - ExpandANormalForm(let_node, pre_visit, post_visit); - } - - void VisitExpr_(const FunctionNode* function_node) override { - ExprVisitor::VisitExpr_(function_node); - MaybeAddPassForFunction(function_node); - } - - /*! - * \brief If \p function_node has a "Compiler" attribute, checks if we should include a - * matching custom pass. Otherwise no-op. - */ - void MaybeAddPassForFunction(const FunctionNode* function_node) { - Optional opt_compiler = function_node->GetAttr(attr::kCompiler); - if (!opt_compiler) { - // No external codegen required. - return; - } - // First cross-over: use "Compiler" attribute name as target kind. - std::string kind_name = opt_compiler.value(); - Optional opt_target_kind = tvm::TargetKind::Get(kind_name); - if (!opt_target_kind || !target_attr_map_.count(opt_target_kind.value())) { - // Target kind does not exist or have the "RelayToTIR" attribute, no custom pass to consider. - return; - } - if (!seen_kinds_.emplace(kind_name).second) { - // Already accounted for custom pass. - return; - } - // Second (optional) cross-over: find unique Target instance in overall available targets with - // the same kind so that it can be made available when custom pass is invoked. - Optional opt_target = config_->FindPrimitiveTargetForKind(opt_compiler.value()); - Pass custom_target_pass = target_attr_map_[opt_target_kind.value()]; - custom_passes_.emplace_back(std::move(kind_name), std::move(custom_target_pass), - std::move(opt_target)); - } - - /*! \brief IRModule we are visiting. */ - IRModule mod_; - /*! \brief All available targets. */ - CompilationConfig config_; - /*! \brief Cached attribute map for all registered targets */ - TargetKindAttrMap target_attr_map_; - /*! \brief Which target kind names have already contributed to the custom passes list. */ - std::unordered_set seen_kinds_; - /*! - * \brief All the custom passes to run, paired with their corresponding target instances, if any. - */ - std::vector custom_passes_; -}; - -} // namespace - -Pass RelayToTIRTargetHook(CompilationConfig config) { - auto pass_func = [config = std::move(config)](IRModule mod, const PassContext& pass_ctx) { - VLOG(1) << "RelayToTIRTargetHook before:" << std::endl << PrettyPrint(mod); - TargetHookVisitor target_hook_visitor(mod, config); - std::vector custom_passes = target_hook_visitor.Visit(); - for (const auto& custom_pass : custom_passes) { - if (custom_pass.opt_target.defined()) { - VLOG(0) << "Invoking custom pass for target " - << custom_pass.opt_target.value()->ToDebugString(); - // Push the target on the stack. - With with_target(custom_pass.opt_target.value()); - // Invoke the pass with target in scope. - mod = custom_pass.pass(mod); - } else { - // Invoke the pass. - // Note that there may be a non-external codegen target in scope. Each custom pass - // must be prepared to handle this, eg by creating a default target instance if the - // current target is either null or of a generic kind such as 'cuda' or 'llvm'. - VLOG(0) << "Invoking custom pass for target kind '" << custom_pass.target_kind_name << "'"; - mod = custom_pass.pass(mod); - } - } - VLOG(1) << "RelayToTIRTargetHook after:" << std::endl << PrettyPrint(mod); - return mod; - }; - return tvm::transform::CreateModulePass(pass_func, 0, "RelayToTIRTargetHook", {}); -} - -} // namespace transform -} // namespace relay -} // namespace tvm diff --git a/src/relay/transforms/to_a_normal_form.cc b/src/relay/transforms/to_a_normal_form.cc deleted file mode 100644 index 8319726b79c5..000000000000 --- a/src/relay/transforms/to_a_normal_form.cc +++ /dev/null @@ -1,472 +0,0 @@ -/* - * 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 to_a_normal_form.cc - * - * \brief Turn implicit sharing into observable sharing. - */ -#include -#include -#include -#include - -#include "../../support/arena.h" -#include "../analysis/dependency_graph.h" -#include "../op/annotation/annotation.h" -#include "./device_aware_visitors.h" -#include "./let_list.h" -#include "./pass_utils.h" - -namespace tvm { -namespace relay { - -Scope LCA(Scope lhs, Scope rhs) { - while (lhs != rhs) { - if (lhs->level > rhs->level) { - lhs = lhs->parent; - } else if (lhs->level < rhs->level) { - rhs = rhs->parent; - } else { - lhs = lhs->parent; - rhs = rhs->parent; - } - } - return lhs; -} - -std::pair CalcScope(const DependencyGraph& dg) { - NodeScopeMap expr_scope; - ExprSet lifted_exprs; - std::unordered_map node_to_expr; - for (auto expr_node : dg.expr_node) { - node_to_expr[expr_node.second] = expr_node.first; - } - bool global_scope_used = false; - Scope global_scope = std::make_shared(); - - for (auto it = dg.post_dfs_order.rbegin(); it != dg.post_dfs_order.rend(); ++it) { - DependencyGraph::Node* n = *it; - auto iit = n->parents.head; - Scope s; - if (iit == nullptr) { - ICHECK(!global_scope_used); - s = global_scope; - global_scope_used = true; - } else { - s = expr_scope.at(iit->value); - const auto original_s = s; - iit = iit->next; - for (; iit != nullptr; iit = iit->next) { - s = LCA(s, expr_scope.at(iit->value)); - } - if (s != original_s && node_to_expr.find(n) != node_to_expr.end()) { - // filter out exprs whose scope do not matter - Expr expr = node_to_expr[n]; - if (!expr.as()) { - lifted_exprs.insert(expr); - } - } - } - if (n->new_scope) { - auto child_scope = std::make_shared(s); - expr_scope.insert({n, child_scope}); - } else { - expr_scope.insert({n, s}); - } - } - ICHECK(global_scope_used); - return std::make_pair(expr_scope, lifted_exprs); -} - -namespace { - -/* Special care is needed to handle local recursion. - * Fill additionally take a (possibly null) Var argument, - * If it is not null, Fill is required to bind the transformed result to that var. - * - * ToANormalForm and PlanDevices - * ----------------------------- - * If PlanDevices has run this transform must respect the lexical scoping rules for the residual - * "on_device" calls. Eg: - * \code - * on_device(add(subtract(x, y), add(y, z)), device_type=2, is_fixed=true) - * ==> - * let %x0 = on_device(subtract(x, y), device_type=2, is_fixed=true) - * let %x1 = on_device(add(y, z), device_type=2, is_fixed=true) - * let %x2 = on_device(add(%x0, %x1), device_type=2, is_fixed=true) - * %x2 - * \endcode - * - * In addition to conversion to ANF this pass is also handling hoisting implicitly shared - * sub-expressions to the inner-most scope common to all their uses: - * \code - * on_device( - * if y { - * on_device(%0, device_type=2, is_fixed=true) - * } else { - * on_device(subtract(%0, b), device_type=2, is_fixed=true) - * }, - * device_type=1, is_fixed=true) - * (where %0 = add(a, b)) - * ==> - * let %x0 = on_device(add(a, b), device_type=2, is_fixed=true); - * on_device( - * if y { - * on_device(%x0, device_type=2, is_fixed=true) - * } else { - * let %x1 = on_device(subtract(%x0, b), device_type=2, is_fixed=true); - * %x1 - * }, - * device_type=1, is_fixed=true) - * \endcode - * Though the PlanDevices has already avoided inserting "on_device" calls where they are redundant - * due to lexical scope, it's fiddly to do the same in this pass since the notion of 'scope' is - * now determined by the scope map. So we'll just insert them mechanically on every let-binding. - * - * TODO(mbs): Rewrite to derive from DeviceAwareExprMutator and not track device types - * explicitly. It's easy to get rid of the need for the extra var argument on VisitExpr by shifting - * the recursion a '1/2 step' to return a possibly compound expression who's inner expressions are - * all atomic. However the use of the scope map is currently subtle enough I want to leave it - * alone for now. - */ -class Fill : ExprFunctor, private transform::LexicalOnDeviceMixin { - public: - static Expr ToANormalForm(const Expr& e, const DependencyGraph& dg, NodeScopeMap* node_scope) { - Fill fi(dg, node_scope, nullptr); - return fi.GetScope(e)->let_list->Get(fi.VisitExpr(e)); - } - - // For basic block normal form, bind expressions only if the original expression's scope - // should be lifted - static Expr ToBasicBlockNormalForm(const Expr& e, const DependencyGraph& dg, - NodeScopeMap* node_scope, ExprSet* lifted) { - Fill fi(dg, node_scope, lifted); - return fi.GetScope(e)->let_list->Get(fi.VisitExpr(e)); - } - - private: - // Note: Conversion to ANF needn't care about the devices for global vars since all that can - // happen with them is to go from: - // ...@g... - // to: - // let %x = @g; - // ... - // ...%x... - // In that case the code will ask for the device for @g, get kInvalidDeviceType, then - // MaybeOnDevice @g, which is always a no-op. - Fill(const DependencyGraph& dg, NodeScopeMap* node_scope, ExprSet* include_set) - : transform::LexicalOnDeviceMixin(Optional()), - dg_(dg), - node_scope_(node_scope), - include_set_(include_set) {} - - Scope GetScope(const Expr& e) { return node_scope_->at(dg_.expr_node.at(e)); } - - Scope GetSubScope(const Expr& e, size_t i) { - DependencyGraph::Node* n = dg_.expr_node.at(e); - auto h = n->children.head; - while (i != 0) { - ICHECK(h); - --i; - h = h->next; - } - ICHECK(h); - return node_scope_->at(h->value); - } - - Expr VisitExpr(const Expr& e) { return this->VisitExpr(e, Var()); } - - Expr VisitExpr(const Expr& e, const Var& v) final { - if (memo.count(e) == 0) { - memo.insert({e, ExprFunctor::VisitExpr(e, v)}); - } else if (v.defined()) { - GetScope(e)->let_list->Push(v, memo.at(e)); - } - auto ret = memo.at(e); - // if no include_set is specified, every expression should be atomic. - // TODO(mbs): Note that Constants must be let-bound even though they are considered 'atomic' - // by this test. - if (include_set_ == nullptr && function_nesting() > 0) { - ICHECK(IsAtomic(ret)) << "expression:" << std::endl << PrettyPrint(ret); - } - return ret; - } - - Expr Atomic(const Expr& e, const Var& v) { - Expr annotated_expr = MaybeOnDeviceFixed(e, GetVirtualDevice(e)); - return v.defined() ? GetScope(e)->let_list->Push(v, annotated_expr) : annotated_expr; - } - - // Bind expression `now` to var `v` if the original expression is in the include set, or if - // v is already defined (e.g. coming from a Let expression). Otherwise return `now` directly - Expr Compound(const Expr& orig, const Expr& now, const Var& v) { - Expr annotated_expr = MaybeOnDeviceFixed(now, GetVirtualDevice(orig)); - Var var = v.defined() ? v : Var::GenSym(); - bool not_included = include_set_ && include_set_->find(orig) == include_set_->end(); - if (!v.defined() && not_included) { - return annotated_expr; - } else if (const LetNode* let = AsIgnoringOnDevice(now)) { - // Instead of making a nested binding "let var = (let x = ...; bindings...; body)", we push - // the inner bindings into the outer scope and bind body to var, giving - // "let x = ...; bindings...; let var = body;" as the resulting bindings. - Expr e = GetRef(let); - while (const LetNode* inner_let = AsIgnoringOnDevice(e)) { - GetScope(orig)->let_list->Push(inner_let->var, inner_let->value); - e = inner_let->body; - } - Expr annotated_body = MaybeOnDeviceFixed(e, GetVirtualDevice(orig)); - return GetScope(orig)->let_list->Push(var, annotated_body); - } else { - return GetScope(orig)->let_list->Push(var, annotated_expr); - } - } - - Expr VisitExpr_(const CallNode* c, const Var& v) final { - OnDeviceProps props = GetOnDeviceProps(c); - if (props.body.defined() && props.is_fixed()) { - // Keep track of expression device type for lexically enclosing sub-expressions. - PushVirtualDevice(props.virtual_device); - Expr body = VisitExpr(props.body, v); - // We are done with this sub-expression. - PopVirtualDevice(); - // Preserve the "on_device" annotations. - return OnDeviceWithProps(body, props); - } - - Expr e = GetRef(c); - std::vector args; - for (const auto& a : c->args) { - args.push_back(VisitExpr(a)); - } - return Compound(e, Call(VisitExpr(c->op), args, c->attrs, c->type_args), v); - } - - Expr VisitExpr_(const TupleNode* tuple_node, const Var& v) final { - Expr e = GetRef(tuple_node); - Array fields; - fields.reserve(tuple_node->fields.size()); - for (const auto& a : tuple_node->fields) { - fields.push_back(VisitExpr(a)); - } - return Compound(e, WithFields(GetRef(tuple_node), fields), v); - } - - Expr VisitExpr_(const TupleGetItemNode* t, const Var& v) final { - Expr e = GetRef(t); - return Compound(e, TupleGetItem(VisitExpr(t->tuple), t->index), v); - } - - Expr VisitExpr_(const RefCreateNode* r, const Var& v) final { - Expr e = GetRef(r); - return Compound(e, RefCreate(VisitExpr(r->value)), v); - } - - Expr VisitExpr_(const RefReadNode* r, const Var& v) final { - Expr e = GetRef(r); - return Compound(e, RefRead(VisitExpr(r->ref)), v); - } - - Expr VisitExpr_(const RefWriteNode* r, const Var& v) final { - Expr e = GetRef(r); - return Compound(e, RefWrite(VisitExpr(r->ref), VisitExpr(r->value)), v); - } - - Expr VisitExpr_(const IfNode* i, const Var& v) final { - Expr e = GetRef(i); - Expr ret = If(VisitExpr(i->cond), GetSubScope(e, 1)->let_list->Get(VisitExpr(i->true_branch)), - GetSubScope(e, 2)->let_list->Get(VisitExpr(i->false_branch))); - return Compound(e, ret, v); - } - - Expr VisitExpr_(const FunctionNode* f, const Var& v) final { - Expr e = GetRef(f); - Expr ret; - if (f->HasNonzeroAttr(attr::kPrimitive)) { - ret = e; - } else { - // Keep track of expression and bound variable device types for lexically enclosing - // sub-expressions. - PushVirtualDevice(f->virtual_device()); - for (auto param : f->params) { - PushBoundVar(param, param->virtual_device()); - } - EnterFunctionBody(); - ret = WithFields(GetRef(f), f->params, - GetSubScope(e, 0)->let_list->Get(VisitExpr(f->body))); - // We are done with this function. - ExitFunctionBody(); - for (size_t i = 0; i < f->params.size(); ++i) { - PopBoundVar(f->params[i]); - } - PopVirtualDevice(); - } - if (function_nesting() == 0) { - ICHECK(!v.defined()); - // This is a global function which can be bound directly in the module. - return ret; - } else { - // This is a local function which must be let-bound. - return Compound(e, ret, v); - } - } - - Expr VisitExpr_(const LetNode* l, const Var& v) final { - Expr e = GetRef(l); - // Keep track of bound variable device types for lexically enclosing sub-expressions. - PushBoundVar(l->var, GetVirtualDevice(l->value)); - VisitExpr(l->value, l->var); - Expr ret = GetSubScope(e, 0)->let_list->Get(VisitExpr(l->body)); - // We are done with these sub-expressions. - PopBoundVar(l->var); - return Compound(e, ret, v); - } - - Expr VisitExpr_(const ConstantNode* c, const Var& v) final { - Expr e = GetRef(c); - return Compound(e, e, v); - } - - Expr VisitExpr_(const VarNode* vn, const Var& v) final { - Expr e = GetRef(vn); - return Atomic(e, v); - } - - Expr VisitExpr_(const GlobalVarNode* gvn, const Var& v) final { - GlobalVar gv = GetRef(gvn); - return Atomic(gv, v); - } - - Expr VisitExpr_(const OpNode* op, const Var& v) final { - Expr e = GetRef(op); - return Atomic(e, v); - } - - Expr VisitExpr_(const ConstructorNode* c, const Var& v) final { - Expr e = GetRef(c); - return Atomic(e, v); - } - - Expr VisitExpr_(const MatchNode* m, const Var& v) final { - Expr e = GetRef(m); - Expr data = VisitExpr(m->data); - std::vector clauses; - for (const Clause& c : m->clauses) { - clauses.emplace_back(c->lhs, - GetSubScope(e, 1 + clauses.size())->let_list->Get(VisitExpr(c->rhs))); - } - return Compound(e, Match(data, clauses, m->complete), v); - } - - const DependencyGraph& dg_; - NodeScopeMap* node_scope_ = nullptr; - std::unordered_map memo; - // a set of Expressions to include for let bindings. If set to nullptr - // all Exprs will be pushed to the let list. - ExprSet* include_set_ = nullptr; -}; - -IRModule ModuleToANormalForm(const IRModule& mod) { - tvm::Map updates; - auto funcs = mod->functions; - for (const auto& it : funcs) { - ICHECK_EQ(FreeVars(it.second).size(), 0); - if (const auto* n = it.second.as()) { - if (n->GetAttr(attr::kCompiler).defined()) continue; - Function func = GetRef(n); - Function ret = Downcast(transform::ToANormalForm(func)); - ICHECK_EQ(FreeVars(ret).size(), 0) << "rewritten:" << std::endl - << PrettyPrint(ret) << std::endl - << "should not have free vars: " << FreeVars(ret); - VLOG(1) << "rewritten:" << std::endl - << PrettyPrint(func) << std::endl - << "to ANF:" << std::endl - << PrettyPrint(ret); - updates.Set(it.first, ret); - } - } - - for (auto pair : updates) { - mod->Add(pair.first, pair.second, true); - } - - return mod; -} - -} // namespace - -Expr ToBasicBlockNormalFormAux(const Expr& e) { - // calculate all the dependency between nodes. - support::Arena arena; - DependencyGraph dg = DependencyGraph::Create(&arena, e); - /* The scope of the whole expr is global. - * The scope of any subexpr, is the lowest common ancestor of all incoming edge. - * We also record the set of expressions whose scope is lifted. - */ - std::pair scopes = CalcScope(dg); - return Fill::ToBasicBlockNormalForm(e, dg, &scopes.first, &scopes.second); -} - -namespace transform { - -Expr ToANormalForm(const Expr& e) { - /* When you lift a lambda, what is inside is also being lift. - * - * So we must determine the scope of the lambda before determining the scope of it's body. - * - * To make this more principled, - * we always determine the scope of parent before determining the scope of children. - * - * So we calculate all the dependency between nodes. - */ - support::Arena arena; - DependencyGraph dg = DependencyGraph::Create(&arena, e); - /* In order to model new subscopes created by lambda, if else and pattern matching, - * we also assign scope to edge as well. - * The scope of an edge is either the parent's scope, or a new subscope of the parent's scope. - * - * So, the scope of the whole expr is global. - * The scope of any subexpr, is the lowest common ancestor of all incoming edge. - * - * Every scope additionally contain a LetList which collect all value of that scope. - * We do an additional pass to fill all the LetList and we are done. - */ - std::pair scopes = CalcScope(dg); - return Fill::ToANormalForm(e, dg, &scopes.first); -} - -Pass ToANormalForm() { - runtime::TypedPackedFunc pass_func = - [=](IRModule m, PassContext pc) { return ModuleToANormalForm(m); }; - return CreateModulePass(pass_func, 1, "ToANormalForm", {}); -} - -TVM_REGISTER_GLOBAL("relay._transform.ToANormalForm").set_body_typed([]() { - return ToANormalForm(); -}); - -TVM_REGISTER_GLOBAL("relay._transform.ToANormalFormExpr").set_body_typed([](const Expr& e) { - return ToANormalForm(e); -}); - -} // namespace transform - -} // namespace relay -} // namespace tvm diff --git a/src/relay/transforms/to_basic_block_normal_form.cc b/src/relay/transforms/to_basic_block_normal_form.cc deleted file mode 100644 index 931543d2640c..000000000000 --- a/src/relay/transforms/to_basic_block_normal_form.cc +++ /dev/null @@ -1,95 +0,0 @@ -/* - * 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 to_basic_block_normal_form.cc - * - * \brief Turn an expression to the basic normal form. - */ -#include -#include -#include -#include - -#include "../../support/arena.h" -#include "../analysis/dependency_graph.h" -#include "./pass_utils.h" - -namespace tvm { -namespace relay { - -IRModule ToBasicBlockNormalForm(const IRModule& mod) { - // Create a new module by shallow copy. - IRModule new_mod = mod->ShallowCopy(); - - tvm::Map updates; - auto funcs = new_mod->functions; - for (const auto& it : funcs) { - ICHECK_EQ(FreeVars(it.second).size(), 0) << "Expected no free variables"; - if (const auto* n = it.second.as()) { - if (n->GetAttr(attr::kCompiler).defined()) continue; - Function func = GetRef(n); - Function ret = Downcast(ToBasicBlockNormalFormAux(func)); - VLOG(1) << "rewritten:" << std::endl - << PrettyPrint(func) << std::endl - << "to BasicBlockANF:" << std::endl - << PrettyPrint(ret); - updates.Set(it.first, Downcast(ret)); - } - } - - for (auto pair : updates) { - new_mod->Add(pair.first, pair.second, true); - } - - return new_mod; -} - -bool BasicBlockNormalFormCheck(const Expr& e) { - // calculate all the dependency between nodes. - support::Arena arena; - DependencyGraph dg = DependencyGraph::Create(&arena, e); - std::pair scopes = CalcScope(dg); - for (auto expr : scopes.second) { - LOG(FATAL) << "The expression below violates the basic block normal form in that " - << "its scope should be lifted:\n" - << expr; - } - return scopes.second.size() == 0; -} - -TVM_REGISTER_GLOBAL("relay.analysis.check_basic_block_normal_form") - .set_body_typed(BasicBlockNormalFormCheck); - -namespace transform { - -Pass ToBasicBlockNormalForm() { - runtime::TypedPackedFunc pass_func = - [=](IRModule m, PassContext pc) { return relay::ToBasicBlockNormalForm(m); }; - return CreateModulePass(pass_func, 1, "ToBasicBlockNormalForm", {}); -} - -TVM_REGISTER_GLOBAL("relay._transform.ToBasicBlockNormalForm") - .set_body_typed(ToBasicBlockNormalForm); - -} // namespace transform - -} // namespace relay -} // namespace tvm diff --git a/src/relay/transforms/to_cps.cc b/src/relay/transforms/to_cps.cc deleted file mode 100644 index 05d49cf5047c..000000000000 --- a/src/relay/transforms/to_cps.cc +++ /dev/null @@ -1,371 +0,0 @@ -/* - * 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 to_cps.cc - * - * \brief Turn a program to continuation passing style. - * - * Given a fresh type variable 'answer', - * continuation passing style(CPS) convert every function of a -> b to a -> (b -> anwer) -> answer. - * - * That is, instead of returning the result directly, - * function will now call another function (called the continuation) - * and return that value as a result instead. - * - * Continuation passing style turn all function call into tail call, - * which bound the stack size, prevent stack from overflowing during recursion, - * and allow tail call optimization. - * - * In relay, as tensor operation is the bottleneck, - * CPS is currently intended to transform the program before partial eval (PE), - * as it reify the control flow and enable PE to handle control flow join more aggressively. - * - * For example, given 'let a = if b then c else d in e', it will transform the code into - * 'let f a = e in if b then f c else f d'. - * This allow f to be optimized individually in both branch. - * - * We implement CPS conversion by higher order transform - * (see http://matt.might.net/articles/cps-conversion/). - * The basic idea is that we will recursively traverse the AST. - * During the traversal, there is an extra parameter, mcont, of expr -> expr. - * It is basically a continuation at the metalevel. - * All cases in the transform must return via the mcont, - * wheter directly invoking it, or indirectly by recursion. - */ -#include -#include -#include -#include -#include - -#include "let_list.h" -#include "pass_utils.h" - -namespace tvm { -namespace relay { - -// we assume the data type has no closure - no idea how to look into datatype right now. - -Type Arrow(const Type& l, const Type& r) { return FuncType({l}, r, {}, {}); } - -Type CPSType(const Type& t, const TypeVar& answer); - -FuncType CPSFuncType(const FuncType& f, const TypeVar& answer) { - tvm::Array new_arg_types; - for (const Type& t : f->arg_types) { - new_arg_types.push_back(CPSType(t, answer)); - } - new_arg_types.push_back(Arrow(CPSType(f->ret_type, answer), answer)); - return FuncType(new_arg_types, answer, f->type_params, f->type_constraints); -} - -Type CPSType(const Type& t, const TypeVar& answer) { - struct CPSTypeMutator : TypeMutator { - explicit CPSTypeMutator(const TypeVar& answer) : answer(answer) {} - TypeVar answer; - Type VisitType_(const FuncTypeNode* t) final { - return CPSFuncType(GetRef(t), answer); - } - } mut(answer); - return mut(t); -} - -// transform global functions into cps form. -using CPSMap = std::unordered_map; - -// transform vars from the original program into new vars, so their type will be correct. -using VarMap = std::unordered_map; - -/* - * The meta continuation. - * There is 3 rules on the metacontinuation: - * 0: It can only use the argument once. - * The argument is code, and using it twice will duplicate code. - * Bound the argument via let instead. - * 1: If the size of the metacontinuation is unbounded, it can only be called once. - * It contain code, so calling it twice duplicate code. - * Reify the continuation and bound it instead. - * See the function 'reify' and the if case for more detail. - * 2: The argument must be effect free. - * It might reorder or drop the argument. - * Again, bound the argument via let instead. - * See the call case for more detail. - */ -using MCont = std::function; - -Function ToCPS(const Function& f, const IRModule& m, CPSMap* cm); - -Function ToCPS(const Function& f, const IRModule& m, CPSMap* cm, VarMap* vm, - const TypeVar& answer) { - std::function remap = [&](const Var& v) { return vm->count(v) == 0 ? v : vm->at(v); }; - auto function_type = Downcast(f->checked_type()); - // Each MCont can be used at most once. - struct CPSFunctor : ExprFunctor, PatternMutator { - CPSFunctor(const std::function& remap, const TypeVar& answer, const IRModule& m, - VarMap* vm, CPSMap* cm) - : remap(remap), answer(answer), m(m), vm(vm), cm(cm) {} - const std::function& remap; - TypeVar answer; - IRModule m; - VarMap* vm; - CPSMap* cm; - - Expr VisitExpr_(const LetNode* op, const MCont& k) final { - return VisitExpr( - op->value, [&](const Expr& v) { return Let(remap(op->var), v, VisitExpr(op->body, k)); }); - } - - Expr VisitExpr_(const FunctionNode* op, const MCont& k) final { - ICHECK(!op->HasNonzeroAttr(attr::kPrimitive)) << "primitive func not supported yet."; - return k(ToCPS(GetRef(op), m, cm, vm, answer)); - } - - Expr VisitExpr_(const ConstantNode* op, const MCont& k) final { - return k(GetRef(op)); - } - - Expr VisitExpr_(const VarNode* op, const MCont& k) final { return k(remap(GetRef(op))); } - - Pattern VisitPattern_(const PatternVarNode* op) final { return PatternVar(remap(op->var)); } - - Expr VisitExpr_(const GlobalVarNode* op, const MCont& k) final { - auto gv = GetRef(op); - if (cm->count(gv) == 0) { - // only look unfold non-external calls. - BaseFunc base_func = m->Lookup(gv); - if (auto* n = base_func.as()) { - auto cps_gv = GlobalVar(std::string(gv->name_hint) + "_cps"); - cm->insert({gv, cps_gv}); - m->Add(cps_gv, ToCPS(GetRef(n), m, cm)); - } else { - // return the original global var if it is - // an external call to non-relay function. - return GetRef(op); - } - } - return k(cm->at(gv)); - } - - Expr VisitExpr_(const RefCreateNode* op, const MCont& k) final { - return VisitExpr(op->value, [&](const Expr& v) { return k(RefCreate(v)); }); - } - - Expr reify(const MCont& k) { - Var arg = Var("arg", Type()); - return Function({arg}, k(arg), Type(), {}); - } - - Expr reify(const MCont& k, const std::function& cont) { - return LetList::LetBind(reify(k), [&](const Var& f) { - return cont([&](const Expr& e) { return Call(f, {e}); }); - }); - } - - Expr VisitExpr_(const IfNode* op, const MCont& k) final { - return reify(k, [&](const MCont& kf) { - return VisitExpr(op->cond, [&](const Expr& v) { - return If(v, VisitExpr(op->true_branch, kf), VisitExpr(op->false_branch, kf)); - }); - }); - } - - Expr VisitExpr_(const MatchNode* op, const MCont& k) final { - return reify(k, [&](const MCont& kf) { - return VisitExpr(op->data, [&](const Expr& v) { - tvm::Array clauses; - for (const auto& c : op->clauses) { - clauses.push_back(Clause(VisitPattern(c->lhs), VisitExpr(c->rhs, kf))); - } - return Match(v, clauses, op->complete); - }); - }); - } - - Expr VisitExpr_(const RefReadNode* op, const MCont& k) final { - return VisitExpr(op->ref, [&](const Expr& r) { return LetList::LetBind(RefRead(r), k); }); - } - - Expr VisitExpr_(const RefWriteNode* op, const MCont& k) final { - return VisitExpr(op->ref, [&](const Expr& r) { - return VisitExpr(op->value, - [&](const Expr& v) { return LetList::LetBind(RefWrite(r, v), k); }); - }); - } - - Expr VisitExpr_(const TupleNode* tuple_node, const MCont& k) final { - tvm::Array fields; - fields.reserve(tuple_node->fields.size()); - std::function next; - next = [&]() { - return (fields.size() == tuple_node->fields.size()) - ? k(WithFields(GetRef(tuple_node), fields)) - : VisitExpr(tuple_node->fields[fields.size()], [&](const Expr& v) { - fields.push_back(v); - return next(); - }); - }; - return next(); - } - - Expr VisitExpr_(const TupleGetItemNode* op, const MCont& k) final { - return VisitExpr(op->tuple, [&](const Expr& v) { return k(TupleGetItem(v, op->index)); }); - } - - Expr VisitExpr_(const CallNode* op, const MCont& k) final { - if (op->op.as() || op->op.as()) { - tvm::Array args; - std::function next; - next = [&]() { - if (args.size() == op->args.size()) { - return LetList::LetBind(Call(op->op, args, op->attrs, op->type_args), k); - } else { - return VisitExpr(op->args[args.size()], [&](const Expr& v) { - args.push_back(v); - return next(); - }); - } - }; - return next(); - } else { - Expr f; - tvm::Array args; - std::function next; - next = [&]() { - if (args.size() == op->args.size()) { - args.push_back(reify(k)); - return Expr(Call(f, args, op->attrs, op->type_args)); - } else { - return VisitExpr(op->args[args.size()], [&](const Expr& v) { - args.push_back(v); - return next(); - }); - } - }; - return VisitExpr(op->op, [&](const Expr& v) { - f = v; - return next(); - }); - } - } - } mut(remap, answer, m, vm, cm); - Var k = Var("k", Arrow(CPSType(function_type->ret_type, answer), answer)); - tvm::Array new_params; - for (const Var& v : f->params) { - new_params.push_back(remap(v)); - } - new_params.push_back(k); - return WithFields(f, new_params, - mut.VisitExpr(f->body, [&](const Expr& e) { return Call(k, {e}); }), answer); -} - -Function ToCPS(const Function& f, const IRModule& m, CPSMap* cm) { - TypeVar answer = TypeVar("answer", kType); - VarMap var; - struct Remapper : ExprVisitor, PatternVisitor { - Remapper(const TypeVar& answer, VarMap* vm) : answer(answer), vm(vm) {} - TypeVar answer; - VarMap* vm; - void VisitExpr_(const VarNode* vn) final { - Var v = GetRef(vn); - if (vm->count(v) == 0) { - auto ret = Var(v->name_hint(), CPSType(v->checked_type(), answer)); - vm->insert({v, ret}); - } - } - - void VisitPattern(const Pattern& p) final { PatternVisitor::VisitPattern(p); } - - void VisitPattern_(const PatternVarNode* op) final { VisitExpr(op->var); } - } remap(answer, &var); - remap.VisitExpr(f); - Function ret = ToCPS(f, m, cm, &var, answer); - auto new_type_params = ret->type_params; - new_type_params.push_back(answer); - return WithFields(ret, ret->params, ret->body, ret->ret_type, new_type_params); -} - -Function ToCPS(const Function& f, const IRModule& m) { - CheckFeature(f, m, FeatureSet::All() - fGraph); - CPSMap cps; - return ToCPS(f, m, &cps); -} - -Function UnCPS(const Function& f) { - CheckFeature(f, FeatureSet::All() - fGraph); - ICHECK_GT(f->params.size(), 0); - Array new_params; - for (const auto& p : f->params) { - new_params.push_back(Var(p->name_hint(), p->checked_type())); - } - auto cont_type = Downcast(new_params.back()->type_annotation); - new_params.pop_back(); - ICHECK_EQ(cont_type->arg_types.size(), 1); - auto new_ret_type = Type(cont_type->arg_types[0]); - Array new_type_params; - for (const auto& tp : f->type_params) { - new_type_params.push_back(TypeVar(tp->name_hint, tp->kind)); - } - auto answer_type = new_type_params.back(); - new_type_params.pop_back(); - // TODO(@M.K.): make alphaequal work on free term - // ICHECK(tvm::StructuralEqual()(cont_type, Arrow(new_ret_type, answer_type))); - auto x = Var("x", new_ret_type); - auto cont = Function({x}, x, new_ret_type, {}); - tvm::Array args; - for (const auto& p : new_params) { - args.push_back(p); - } - args.push_back(cont); - tvm::Array type_args; - for (const auto& tp : new_type_params) { - type_args.push_back(tp); - } - type_args.push_back(new_ret_type); - return WithFields(f, new_params, Call(f, args, {}, type_args), new_ret_type, new_type_params); -} - -TVM_REGISTER_GLOBAL("relay._transform.to_cps") - .set_body_typed(static_cast(ToCPS)); - -TVM_REGISTER_GLOBAL("relay._transform.un_cps").set_body_typed(UnCPS); - -namespace transform { - -Pass ToCPS() { - runtime::TypedPackedFunc pass_func = - [=](Function f, IRModule m, PassContext pc) { return Function(ToCPS(f, m)); }; - return CreateFunctionPass(pass_func, 1, "ToCPS", {}); -} - -TVM_REGISTER_GLOBAL("relay._transform.ToCPS").set_body_typed(ToCPS); - -Pass UnCPS() { - runtime::TypedPackedFunc pass_func = - [=](Function f, IRModule m, PassContext pc) { return Function(UnCPS(f)); }; - return CreateFunctionPass(pass_func, 1, "UnCPS", {}); -} - -TVM_REGISTER_GLOBAL("relay._transform.UnCPS").set_body_typed(UnCPS); - -} // namespace transform - -} // namespace relay -} // namespace tvm diff --git a/src/relay/transforms/to_graph_normal_form.cc b/src/relay/transforms/to_graph_normal_form.cc deleted file mode 100644 index ff5ff568b048..000000000000 --- a/src/relay/transforms/to_graph_normal_form.cc +++ /dev/null @@ -1,89 +0,0 @@ -/* - * 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 to_gnf.cc - * - * \brief Turn A normal form into graph normal form. - */ -#include -#include -#include - -#include "let_list.h" - -namespace tvm { -namespace relay { - -class UseVarVisitor : public ExprVisitor { - public: - explicit UseVarVisitor(const Var& v) : v(v) {} - - static bool UseVar(const Var& v, const Expr& e) { - UseVarVisitor uv(v); - uv(e); - return uv.use_var; - } - - private: - bool use_var = false; - Var v; - - void VisitExpr_(const VarNode* vn) override { use_var = use_var || (v == GetRef(vn)); } -}; - -class GNF : public ExprMutator { - private: - std::unordered_map var_map_; - Expr VisitExpr_(const VarNode* vn) override { - Var v = GetRef(vn); - return var_map_.count(v) == 0 ? v : var_map_.at(v); - } - - static bool UseVar(const Var& v, const Expr& e) { return UseVarVisitor::UseVar(v, e); } - - static Expr WrapRec(const Var& var, const Expr& val) { - return UseVar(var, val) ? Let(var, val, var) : val; - } - - Expr VisitExpr_(const LetNode* ln) override { - var_map_.insert(std::pair(ln->var, WrapRec(ln->var, VisitExpr(ln->value)))); - return VisitExpr(ln->body); - } -}; - -Expr ToGraphNormalForm(const Expr& e) { return GNF()(e); } - -namespace transform { - -Pass ToGraphNormalForm() { - runtime::TypedPackedFunc pass_func = - [=](Function f, IRModule m, PassContext pc) { - return Downcast(ToGraphNormalForm(f)); - }; - return CreateFunctionPass(pass_func, 1, "ToGraphNormalForm", {}); -} - -TVM_REGISTER_GLOBAL("relay._transform.ToGraphNormalForm").set_body_typed(ToGraphNormalForm); - -} // namespace transform - -} // namespace relay -} // namespace tvm diff --git a/src/relay/transforms/to_mixed_precision.cc b/src/relay/transforms/to_mixed_precision.cc deleted file mode 100644 index 1112755b76a0..000000000000 --- a/src/relay/transforms/to_mixed_precision.cc +++ /dev/null @@ -1,560 +0,0 @@ -/* - * 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 to_mixed_precision.cc - * \brief Automatic mixed floating point precision for relay graphs. i.e. turn a graph into fp16. - * - */ - -#include -#include -#include -#include - -#include - -#include "../../support/scalars.h" -#include "pattern_utils.h" - -namespace tvm { -namespace relay { - -TVM_REGISTER_PASS_CONFIG_OPTION("relay.ToMixedPrecision.keep_orig_output_dtype", Bool); -// A callable which hashes std::pair -struct pair_hash { - template - std::size_t operator()(const std::pair& pair) const { - auto h1 = std::hash()(pair.first); - auto h2 = std::hash()(pair.second); - - // Use boost's combine_hash strategy - return h1 ^ (h1 + 0x9e3779b9 + (h2 << 6) + (h2 >> 2)); - } -}; - -// MIXED_PRECISION_ALWAYS ops should always be done in lower precision due to the speed and memory -// savings. MIXED_PRECISION_FOLLOW ops can be done in lower precision but don't have speedups to -// justify a cast. MIXED_PRECISION_NEVER colored ops should not be done in lower precision due to -// numerical reasons. -enum MixedTypeConversionCategory : int { - MIXED_PRECISION_ALWAYS = 0, - MIXED_PRECISION_FOLLOW = 1, - MIXED_PRECISION_NEVER = 2 -}; - -// A map of a parent node and a wanted dtype to existing nodes casted to the wanted dtype -using CachedCastNodes = std::unordered_map, Expr, pair_hash>; - -// Return array is of type : [MixedTypeConversionCategory (int), String, String] -// The fields are : [ConversionCategory, accumulation_datatype, output_datatype] -// Call is a call node, DataType is the mixed precision type -using FTVMMixedPrecisionConversionType = runtime::TypedPackedFunc>( - const Call& call_node, const std::string& target_dtype_str)>; - -/*! \brief This class transforms the given relay module into a version where - * as many operations as possible operate in the target mixed precision dtype. - * - * Input : A Relay module with operations registered with FTVMMixedPrecisionConversionType - * functions. These describe when and how the operations will be transformed - * into the target precision dtype. - * - * Output : A Relay module with some operations transformed according to the below - * methodology. - * - * Methodology : - * 1) Each relay Op is either of conversion category ALWAYS, FOLLOW, NEVER - * defined by the associated FTVMMixedPrecisionConversionType function. - * If an operation is not registered, it by default is assumed to be - * FOLLOW. - * 2) ALWAYS operations always convert the input floating point args into - * the target mixed precision dtype. FOLLOW Ops will convert the input - * floating point args back into FP32 unless all floating point args - * are in the target mixed precision dtypes. NEVER ops will always cast - * inputs back into FP32. - * 3) Each ALWAYS Op, and FOLLOW Op with mixed precision dtype arguments - * also have an associated accumulation_dtype and output_dtype which - * describe whether a larger dtype is used to accumulate the results - * of the operation. The output_dtype meanwhile describes the dtype - * most Ops should use from this accumulator. - */ -class MixedPrecisionPass : public MixedModeMutator { - private: - /*! \brief A cache of nodes + target dtype to a cast version of the node with target dtype. */ - CachedCastNodes cast_nodes_cache_; - - /*! \brief The target datatype we want to convert to e.g. FP16 */ - const DataType mixed_precision_type_; - - /*! \brief Map of Ops with no associated FTVMMixedPrecisionConversionType to the times they were - * encountered. Used for emitting warnings on missing ops in the pass. - */ - std::unordered_map missing_ops_; - const RelayExprNode* root_; - std::vector original_dtype_; - bool keep_orig_output_dtype_; - - /*! \brief If some of the constant attributes are out of mixed_precision_type_ bounds, then - * computation cannot be performed in mixed precision. */ - bool IsMixedPrecisionApplicableToAttrs(const Attrs& attrs) const { - if (attrs.get() != nullptr) { - double min_bound; - double max_bound; - if (mixed_precision_type_.is_float16()) { - min_bound = -support::kMaxFloat16; - max_bound = support::kMaxFloat16; - } else if (mixed_precision_type_.is_bfloat16()) { - min_bound = -support::kMaxBFloat16; - max_bound = support::kMaxBFloat16; - } else if (mixed_precision_type_.is_float8()) { - double bound = (mixed_precision_type_.code() == DataType::kE4M3Float) ? support::kMaxE4M3 - : support::kMaxE5M2; - min_bound = -bound; - max_bound = bound; - } else if (mixed_precision_type_.is_float()) { - min_bound = std::numeric_limits::lowest(); - max_bound = std::numeric_limits::max(); - } else { - return true; - } - - if (auto cur_attrs = attrs.as()) { - if (cur_attrs->a_min < min_bound || cur_attrs->a_max > max_bound) { - return false; - } - } - } - return true; - } - - Attrs GetNewAttrs(const CallNode* call, const DataType& accumulation_dtype) const { - /* If the accumulation dtype is in the attributes make a copy and mutate the field. */ - Attrs cur_attrs = call->attrs; - if (cur_attrs.get() != nullptr) { - // TODO(AndrewZhaoLuo): Figure out a better way to do this - // modify output_dtype attributes (accumulation dtypes for ops) - if (auto attrs = cur_attrs.as()) { - return ModifyAttrsOutputDType(attrs, accumulation_dtype); - } else if (auto attrs = cur_attrs.as()) { - return ModifyAttrsOutputDType(attrs, accumulation_dtype); - } else if (auto attrs = cur_attrs.as()) { - return ModifyAttrsOutputDType(attrs, accumulation_dtype); - } else if (auto attrs = cur_attrs.as()) { - return ModifyAttrsOutputDType(attrs, accumulation_dtype); - } else if (auto attrs = cur_attrs.as()) { - return ModifyAttrsOutputDType(attrs, accumulation_dtype); - } else if (auto attrs = cur_attrs.as()) { - return ModifyAttrsOutputDType(attrs, accumulation_dtype); - } else if (auto attrs = cur_attrs.as()) { - return ModifyAttrsOutputDType(attrs, accumulation_dtype); - } else if (auto attrs = cur_attrs.as()) { - return ModifyAttrsOutputDType(attrs, accumulation_dtype); - } else if (auto attrs = cur_attrs.as()) { - return ModifyAttrsOutputDType(attrs, accumulation_dtype); - } else if (auto attrs = cur_attrs.as()) { - return ModifyAttrsOutputDType(attrs, accumulation_dtype); - } else if (auto attrs = cur_attrs.as()) { - return ModifyAttrsOutputDType(attrs, accumulation_dtype); - } else if (auto attrs = cur_attrs.as()) { - return ModifyAttrsOutputDType(attrs, accumulation_dtype); - } - - // modify dtype attributes (creating new tensors of type dtype) - if (auto attrs = cur_attrs.as()) { - return ModifyAttrsDType(attrs, accumulation_dtype); - } - } - - return cur_attrs; - } - - template - Attrs ModifyAttrsOutputDType(const T* attrs, const DataType& accumulation_dtype) const { - /* - Helper template to modify relevant attributes with out_dtype type. - These represent accumulation dtypes for some operations e.g. - conv2d might take in fp16 and give a fp32 result. - Attrs is const because we get it as a const. - */ - DataType cur_type = (attrs->out_dtype); - ObjectPtr new_attrs = make_object(*attrs); - if (cur_type.is_float() || cur_type.is_bfloat16() || cur_type.is_void()) { - new_attrs->out_dtype = accumulation_dtype; - } - return Attrs(new_attrs); - } - - template - Attrs ModifyAttrsDType(const T* attrs, const DataType& accumulation_dtype) const { - /* - Helper template to modify relevant attributes with dtype type. - This determines the output dtype for some ops. For example - zeros creates a tensor of zeros of the specified dtype. - Attrs is const because we get it as a const. - */ - DataType cur_type = (attrs->dtype); - ObjectPtr new_attrs = make_object(*attrs); - if (cur_type.is_float() || cur_type.is_bfloat16() || cur_type.is_void()) { - new_attrs->dtype = accumulation_dtype; - } - return Attrs(new_attrs); - } - - Type GetType(const Expr& expr) const { - // The expression has not been changed AND it's existing type - // is known to still be valid. (See special handling for tuples etc - // below for where we null out checked_type_ when we can not - // sure it is still valid. - Type checked_type = expr->checked_type_; - if (checked_type.defined()) { - return checked_type; - } - - // This also populates the checked_type_ field for expr - return transform::InferTypeLocal(expr); - } - - bool IsMixedPrecisionType(const Type& t, bool ignore_non_float = false) const { - /* Returns whether t is a type with only target mixed precision type elements. - If ignore_non_float, then ignore non-floating types. - */ - if (const TensorTypeNode* tensor_type = t.as()) { - bool is_supported_floating_point_type = - (tensor_type->dtype).is_float() || (tensor_type->dtype).is_bfloat16(); - return (ignore_non_float && !is_supported_floating_point_type) || - tensor_type->dtype == mixed_precision_type_; - } else if (const TupleTypeNode* tuple_type = t.as()) { - for (Type t : tuple_type->fields) { - if (!IsMixedPrecisionType(t, ignore_non_float)) return false; - } - return true; - } else { - LOG(FATAL) << "Unsupported type " << t << " we don't know how to handle"; - } - } - - Expr CachedCast(const Expr& expr, const DataType& expr_dtype, const DataType& wanted_dtype) { - /* Cast tensor to the wanted datatype, returning a cached version if it's already been done. */ - - // If this is not a floating point type, do not cast. E.g. it might be an integer - if (!(expr_dtype.is_float() || expr_dtype.is_bfloat16())) { - return expr; - } - - if (expr_dtype == wanted_dtype) { - return expr; - } - - const ExprNode* expr_node = expr.as(); - CHECK(expr_node) << "Non-expression node found in cast: " << expr; - - // Use cached result if possible. - auto search = cast_nodes_cache_.find({expr_node, wanted_dtype}); - if (search != cast_nodes_cache_.end()) { - return search->second; - } - - Expr result = Cast(expr, wanted_dtype); - cast_nodes_cache_[{expr_node, wanted_dtype}] = result; - - // Reverse the cache result, e.g. if we want to reverse the cast simply point to original node - const ExprNode* new_expr_node = result.as(); - cast_nodes_cache_[{new_expr_node, expr_dtype}] = expr; - return result; - } - - Expr CastArg(const Expr& expr, const Type& expr_type, const DataType& wanted_dtype) { - /* Helper for casting arguments to call_nodes handling all relevant cases. */ - if (const TensorTypeNode* tensor_type = expr_type.as()) { - return CachedCast(expr, tensor_type->dtype, wanted_dtype); - } else if (const TupleTypeNode* tuple_type = expr_type.as()) { - Array new_expr; - bool all_same = true; - for (size_t i = 0; i < (tuple_type->fields).size(); i++) { - Expr tuple_element = GetField(expr, i); - Type tuple_element_dtype = (tuple_type->fields)[i]; - Expr casted_element = CastArg(tuple_element, tuple_element_dtype, wanted_dtype); - new_expr.push_back(casted_element); - all_same &= casted_element.same_as(tuple_element); - } - return all_same ? expr : Tuple(new_expr); - } - CHECK(0) << "Unsupported type " << expr_type << " we don't know how to cast for arguments!"; - return expr; - } - - std::pair, Array> CastAllArgs(const Array& cur_args, - const Array& cur_arg_types, - const DataType& wanted_dtype) { - Array new_args; - Array new_arg_types; - for (size_t i = 0; i < cur_args.size(); i++) { - Expr cur_arg = cur_args[i]; - Type cur_arg_type = cur_arg_types[i]; - Expr new_arg = CastArg(cur_arg, cur_arg_type, wanted_dtype); - Type new_arg_type = GetType(new_arg); - new_args.push_back(new_arg); - new_arg_types.push_back(new_arg_type); - } - return {new_args, new_arg_types}; - } - - public: - using MixedModeMutator::VisitExpr_; - - explicit MixedPrecisionPass(Expr base, bool keep_orig_output_dtype, - DataType mixed_precision_type = DataType::Float(16)) - : MixedModeMutator(), - mixed_precision_type_(mixed_precision_type), - root_(Downcast(base)->body.get()), - keep_orig_output_dtype_(keep_orig_output_dtype) { - if (keep_orig_output_dtype_) { - if (root_->IsInstance()) { - const TupleTypeNode* tuple_type = (root_->checked_type_).as(); - for (Type t : tuple_type->fields) { - const TensorTypeNode* tensor_type = t.as(); - original_dtype_.push_back(tensor_type->dtype); - } - } else if (root_->IsInstance()) { - original_dtype_.push_back((root_->checked_type_).as()->dtype); - } - } - if (!(mixed_precision_type_.is_float() || mixed_precision_type_.is_bfloat16())) { - LOG(FATAL) << "Only support IEEE floating point mixed precision types and bfloat16, but got " - << mixed_precision_type_; - } - } - - Expr Rewrite_(const CallNode* pre_call_node, const Expr& post) final { - const CallNode* post_call_node = post.as(); - CHECK(post_call_node) << "Expected a CallNode, but got " << post; - - Expr cur_op = post_call_node->op; - - // TODO(AndrewZhaoLuo): Support ADTs - // Relay's algebraic data types are not supported yet. - bool isADT = (cur_op.as() // used to declare functions for recursion - || cur_op.as() // constructing ADT types - || cur_op.as() // used for binding lambdas - || cur_op.as()); // used for calling recursive functions - if (isADT) return post; - - // Get info on the operation being called: - // conversion category (int), accumulation dtype (str), output dtype (str) - MixedTypeConversionCategory initial_category; - DataType accumulation_dtype, output_dtype; - if (cur_op.as()) { - // Avoid messing with functions to avoid changing signature - initial_category = MIXED_PRECISION_NEVER; - accumulation_dtype = DataType::Float(32); - output_dtype = DataType::Float(32); - } else if (cur_op.as()) { - static auto attr_map = - Op::GetAttrMap("FTVMMixedPrecisionConversionType"); - Op op = Downcast(cur_op); - if (attr_map.count(op)) { - // Calculate the conversion category and dtypes from registered attribute. - FTVMMixedPrecisionConversionType func = attr_map[op]; - Array> op_descriptor = - func(GetRef(pre_call_node), DLDataType2String(mixed_precision_type_)); - ICHECK(op_descriptor.size() == 3) - << "got the wrong number of returned arguments (expected 3 got " << op_descriptor.size() - << ") from FTVMMixedPrecisionConversionType for " << AsText(op, false); - - int64_t op_conversion_type = Downcast(op_descriptor[0])->value; - initial_category = static_cast(op_conversion_type); - accumulation_dtype = DataType(String2DLDataType(Downcast(op_descriptor[1]))); - output_dtype = DataType(String2DLDataType(Downcast(op_descriptor[2]))); - } else { - missing_ops_[op->name] += 1; - - // If not registered, by default assume is a generic FOLLOW operation. - initial_category = MIXED_PRECISION_FOLLOW; - accumulation_dtype = mixed_precision_type_; - output_dtype = mixed_precision_type_; - } - } else { - LOG(FATAL) << "Unsupported op type in CallNode: " << pre_call_node->op; - } - - // First check if all the new mutated args are in lower precision form - Array cur_arg_types; - bool all_args_mixed_type_compatible = true; - for (Expr arg : post_call_node->args) { - Type cur_arg_type = GetType(arg); - cur_arg_types.push_back(cur_arg_type); - - if (initial_category == MIXED_PRECISION_FOLLOW && all_args_mixed_type_compatible) { - // We can cast Vars and Constants to the right types so don't care about the types. - bool is_mixed_type_compatible = IsMixedPrecisionType(cur_arg_type, true) || - arg->IsInstance() || - arg->IsInstance(); - all_args_mixed_type_compatible &= is_mixed_type_compatible; - } - } - - // Determine the final category we want for conversion - MixedTypeConversionCategory final_category = initial_category; - if (initial_category == MIXED_PRECISION_FOLLOW) { - final_category = - all_args_mixed_type_compatible ? MIXED_PRECISION_ALWAYS : MIXED_PRECISION_NEVER; - } - - bool is_mixed_precision_applicable = - static_cast(final_category == MIXED_PRECISION_ALWAYS && - IsMixedPrecisionApplicableToAttrs(pre_call_node->attrs)); - // Create the new arguments to the call. - DataType wanted_arg_dtypes = - is_mixed_precision_applicable ? mixed_precision_type_ : DataType::Float(32); - auto call_args_and_types = CastAllArgs(post_call_node->args, cur_arg_types, wanted_arg_dtypes); - Array new_args = call_args_and_types.first; - Array new_arg_types; - - if (pre_call_node->op.as()) { - // Function Nodes don't store type info in the Call, it should be a [] - new_arg_types = pre_call_node->type_args; - } else { - new_arg_types = call_args_and_types.second; - } - - // Finally create the new attributes. - if (is_mixed_precision_applicable) { - Attrs new_attrs = GetNewAttrs(pre_call_node, accumulation_dtype); - Expr output = Call(cur_op, new_args, new_attrs, new_arg_types, pre_call_node->span); - if (accumulation_dtype != output_dtype) { - output = CastArg(output, GetType(output), output_dtype); - } - if (pre_call_node == root_ && keep_orig_output_dtype_) { - if (original_dtype_[0] != output_dtype) { - output = CastArg(output, GetType(output), original_dtype_[0]); - } - } - return output; - } - - return Call(cur_op, new_args, pre_call_node->attrs, new_arg_types, pre_call_node->span); - } - - Expr Rewrite_(const TupleGetItemNode* pre, const Expr& post) { - // The old checked type in the expression may not be valid so clear it - post->checked_type_ = Type(nullptr); - return post; - } - - Expr Rewrite_(const TupleNode* pre, const Expr& post) { - // The old checked type in the expression may not be valid so clear it - post->checked_type_ = Type(nullptr); - if (pre == root_ && keep_orig_output_dtype_) { - Array new_expr; - bool all_same = true; - for (size_t i = 0; i < original_dtype_.size(); i++) { - Expr output_element = GetField(post, i); - Expr casted_element; - auto output_element_type = transform::InferTypeLocal(output_element); - casted_element = CastArg(output_element, output_element_type, original_dtype_[i]); - new_expr.push_back(casted_element); - all_same &= casted_element.same_as(output_element); - } - if (!all_same) { - return Tuple(new_expr); - } - } - return post; - } - - Expr VisitExpr_(const FunctionNode* func) final { - // Erase the ret_type annotation and let the normal pass recalculate - const_cast(func)->ret_type = Type(nullptr); - return ExprMutator::VisitExpr_(func); - } - - Expr VisitExpr_(const LetNode* op) final { - // First convert as much of the bound computation to lower precision as possible - Expr value = this->Mutate(op->value); - - // Then rewrite the var type and associated expression - Var var = Downcast(this->Mutate(op->var)); - VarNode* mutable_var = const_cast((op->var).as()); - mutable_var->type_annotation = GetType(value); - mutable_var->checked_type_ = mutable_var->type_annotation; - - // Mutate body last as it may depend on previous results - Expr body = this->Mutate(op->body); - return Let(var, value, body, op->span); - } - - // To access map of ops not registered for error reporting - friend Expr ToMixedPrecision(const Expr& expr, bool keep_orig_output_dtype, - const DataType& mixed_precision_type, int missing_op_mode); -}; - -Expr ToMixedPrecision(const Expr& expr, bool keep_orig_output_dtype, - const DataType& mixed_precision_type, int missing_op_mode) { - /* - missing_op_mode: - - 0: Does not allow any missing ops. Will throw errors and terminate the pass when encountering any. - 1: Allow missing ops but throw warnings. - 2: Allow missing ops and silently ignore them. - */ - ICHECK(missing_op_mode >= 0 && missing_op_mode <= 2) - << " missing_op_mode must be either 0, 1, or 2 got " << missing_op_mode; - - MixedPrecisionPass converter = - MixedPrecisionPass(expr, keep_orig_output_dtype, mixed_precision_type); - auto result = converter.Mutate(expr); - - for (auto it = converter.missing_ops_.begin(); - missing_op_mode != 2 && it != converter.missing_ops_.end(); it++) { - std::string op_name = it->first; - int appear_count = it->second; - - LOG(WARNING) << "Op \"" << op_name << "\" not registered " - << "FTVMMixedPrecisionConversionType appears " << appear_count - << " times in graph."; - } - - if (converter.missing_ops_.size() != 0 && missing_op_mode == 0) { - CHECK(0) << "Missing ops were found!"; - } - return result; -} - -namespace transform { - -Pass ToMixedPrecision(DataType mixed_precision_type, int missing_op_mode) { - runtime::TypedPackedFunc pass_func = - [=](Function f, IRModule m, PassContext pc) { - bool keep_orig_output_dtype = false; - keep_orig_output_dtype = pc->GetConfig("relay.ToMixedPrecision.keep_orig_output_dtype", - Bool(keep_orig_output_dtype)) - .value(); - return Downcast( - ToMixedPrecision(f, keep_orig_output_dtype, mixed_precision_type, missing_op_mode)); - }; - return CreateFunctionPass(pass_func, 0, "ToMixedPrecision", {}); -} - -TVM_REGISTER_GLOBAL("relay._transform.ToMixedPrecision").set_body_typed(ToMixedPrecision); - -} // namespace transform - -} // namespace relay -} // namespace tvm diff --git a/src/relay/transforms/transform_layout.h b/src/relay/transforms/transform_layout.h deleted file mode 100644 index 117096e1334a..000000000000 --- a/src/relay/transforms/transform_layout.h +++ /dev/null @@ -1,463 +0,0 @@ -/* - * 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 transform_layout.h - * \brief Common infrastructure for transforming the layouts. This is used for AlterOpLayout and - * ConvertLayout pass. */ - -#ifndef TVM_RELAY_TRANSFORMS_TRANSFORM_LAYOUT_H_ -#define TVM_RELAY_TRANSFORMS_TRANSFORM_LAYOUT_H_ - -#include -#include - -#include -#include -#include -#include -#include - -#include "infer_layout_utils.h" -#include "pattern_utils.h" - -namespace tvm { -namespace relay { - -/*! - * \brief Memorizes layout transformations to reuse. - */ -class TransformMemorizerNode : public Object { - public: - /*! \brief The key for the memorizer map is (Expr, src_layout, dst_layout). */ - using TransformKey = std::tuple; - - struct key_hash : public std::function { - std::size_t operator()(const TransformKey& k) const { - return dmlc::HashCombine( - dmlc::HashCombine(std::hash()(std::get<0>(k)), - std::get<1>(k)), - (std::get<2>(k))); - } - }; - - /*! - * \brief Defines the call transformation for derived passes. The new layouts are defined by - * used for different targets using a packed func. - * \param ref_call The original call. - * \param new_attrs Updated attributes consistent with new layouts. - * \param new_args The traversed/recursed args to the call. - * \return The new Call after calling the packed func. - */ - virtual Call CallWithNewLayouts(const Call& ref_call, Attrs new_attrs, - const std::vector& new_args) = 0; - - virtual Call CallWithNewLayouts(const Call& ref_call, const std::vector& new_args) { - return CallWithNewLayouts(ref_call, ref_call->attrs, new_args); - } - - /*! \brief The memorizer map. */ - std::unordered_map memo; - - static constexpr const char* _type_key = "relay.alter_op_layout.TransformMemorizerNode"; - TVM_DECLARE_FINAL_OBJECT_INFO(TransformMemorizerNode, Object); -}; - -/*! - * \brief Container that transforms the layouts and memorizes them. - */ -class TransformMemorizer : public ObjectRef { - public: - TransformMemorizer() = default; - explicit TransformMemorizer(ObjectPtr n) : ObjectRef(n) {} - - TransformMemorizerNode* operator->() { - return static_cast(get_mutable()); - } - - /* - * \brief Memorizes and transforms the layout. - * \param expr The initial expr. - * \param src_layout The source layout. - * \param dst_layout The dest layout. - * \return The new expr with the dst layout. - */ - Expr Transform(Expr raw, const Layout& src_layout, const Layout& dst_layout) { - if (src_layout.Equals(dst_layout)) { - return raw; - } - - std::tuple key = - std::make_tuple<>(raw.get(), src_layout.name(), dst_layout.name()); - auto& memo = operator->()->memo; - - auto iter = memo.find(key); - if (iter != memo.end()) { - return iter->second; - } else { - Expr transform = TransformHelper(raw, src_layout, dst_layout); - memo[key] = transform; - return transform; - } - } - - /* - * \brief Helper to transform the layouts. - * \param expr The initial expr. - * \param src_layout The source layout. - * \param dst_layout The dest layout. - * \return The new expr with the dst layout. - * \note It performs following 2 operations - * 1) If src_layout ndim is smaller then dst_layout, expand_dim is inserted to match the dim - * size. For example, src_layout = C, dst_layout = NCHW16c. The src is expanded to NHWC. - * 2) Call layout transform with new src layout. - */ - Expr TransformHelper(Expr raw, Layout src_layout, Layout dst_layout) { - if (src_layout.Equals(dst_layout)) { - return raw; - } - - // 1) Check if the shape lengths are different. If yes, expand dims. - Expr input_expr = raw; - Layout new_src_layout = src_layout; - if (src_layout.ndim_primal() < dst_layout.ndim_primal()) { - // If scalar, then no need of layout transformation as scalar can be broadcasted easily even - // if the other operand has a transformed layout. - if (input_expr->checked_type_.defined() && IsScalar(input_expr)) { - return raw; - } - int num_new_axis = dst_layout.ndim_primal() - src_layout.ndim_primal(); - new_src_layout = src_layout.ExpandPrimal(dst_layout); - input_expr = MakeExpandDims(input_expr, 0, num_new_axis); - if (new_src_layout.Equals(dst_layout)) { - return input_expr; - } - } - - // 2) Insert layout transform on the transformed src. - ICHECK(new_src_layout.defined() && dst_layout.defined()) - << "Cannot insert layout transform because there are undefined layouts"; - ICHECK(tir::BijectiveLayout(new_src_layout, dst_layout).defined()) - << "Cannot insert layout transform because there are inconvertible layouts: " - << new_src_layout << " v.s. " << dst_layout; - return MakeLayoutTransform(input_expr, new_src_layout.name(), dst_layout.name()); - } - - using ContainerType = TransformMemorizerNode; -}; - -/* - * \brief TempExprNode during layout transform. Instance of this expr will be Realized to normal - * expr ultimately. - * \tparam TransformMemorizerT The derived TransformMemorizer type. - */ -template -class LayoutAlternatedExprNode : public TempExprNode { - public: - Expr value; - Layout old_layout; - Layout new_layout; - TransformMemorizerT memorizer; - - Expr Realize() const final { - // NOTE: use a copy to discard the "const" qualifier - TransformMemorizerT tmp_memorizer = memorizer; - // fallback to old layout - return tmp_memorizer.Transform(value, new_layout, old_layout); - } - - void VisitAttrs(AttrVisitor* v) { - v->Visit("value", &value); - v->Visit("old_layout", &old_layout); - v->Visit("new_layout", &new_layout); - } - - static constexpr const char* _type_key = "relay.alter_op_layout.LayoutAlternatedExprNode"; - TVM_DECLARE_FINAL_OBJECT_INFO(LayoutAlternatedExprNode, TempExprNode); -}; - -/*! - * \brief Container for the layout alternated expr. - * \tparam TransformMemorizerT The derived TransformMemorizer type. - */ -template -class LayoutAlternatedExpr : public ObjectRef { - public: - LayoutAlternatedExpr() {} - explicit LayoutAlternatedExpr(ObjectPtr n) : ObjectRef(n) {} - - LayoutAlternatedExprNode* operator->() { - return static_cast*>(get_mutable()); - } - - using ContainerType = LayoutAlternatedExprNode; -}; - -/*! - * Call registered FInferCorrectLayout of an op. - * Parameters are the same as the parameters for FInferCorrectLayout - * Returns inferred_input_layout, inferred_output_layout, updated attributes, and a flag - * indicating whether or not layout conversion is successful. - */ -static inline std::tuple InferCorrectLayouts( - const Call& call, const Array& new_in_layouts, const Array& old_in_layouts, - const Array& old_in_types) { - static auto finfer_layout = Op::GetAttrMap("FInferCorrectLayout"); - auto null_res = std::make_tuple( - InferCorrectLayoutOutput(Array(nullptr), Array(nullptr), Attrs(nullptr)), - false); - if (!call->op.as()) { - return null_res; - } - - Op op = Downcast(call->op); - if (finfer_layout.count(op)) { - auto out = finfer_layout[op](call->attrs, new_in_layouts, old_in_layouts, old_in_types); - for (auto inferred_layouts : {out->input_layouts, out->output_layouts}) { - for (auto layout : inferred_layouts) { - if (!layout.defined()) { // inference fails - return null_res; - } - } - } - return std::make_tuple(out, true); - } else { - return null_res; - } -} - -/* - * \brief Used with ForwardRewrite to transform the expr. The input args are same as - * FForwardRewrite. - * \param ref_call The reference old call type to be rewritten. - * We can make use of the op and type information. - * \param new_args The new arguments (some of them could be TempExpr). - * \param ctx Optional context information about ref_call. - * \tparam TransformMemorizerT The derived TransformMemorizer type. - * \return The rewriten result call, can also return nullptr, - * which indicate the rewriter should use the default fallback - * rule that realizes all its input and compose the call. - * - * \note The ctx can be used to provide extra information during transformation. The ctx is - * templated to reuse across AlterOpLayout and ConvertLayout pass. The steps are - * - Extract the original layouts. - * - Use ctx transformation to get a Call with new layouts - CallWithNewLayouts. - * - Extract the new layouts from the returned Call. - * - Transform the original call to reuse the new layouts using TransformMemorizer. - */ -template -Expr LayoutRewriter(const Call& ref_call, const Array& new_args, const ObjectRef& ctx) { - std::vector> inputs; - std::vector normal_new_args; - - // NOTE: discard the "const" qualifier - // TransformMemorizer memorizer = Downcast(ctx); - // TransformMemorizerT* ctx_transformer = - // static_cast(memorizer.operator->()); - TransformMemorizerT memorizer = Downcast(ctx); - - // fill incomplete state and flatten tuple - auto push_back_one_arg = [&inputs, memorizer](Expr arg) { - // We always expect LayoutAlternatedExpr. - // This is used to convert the normal Expr to LayoutAlternatedExpr. - if (const LayoutAlternatedExprNode* inp = - arg.as>()) { - inputs.push_back(GetRef>(inp)); - return inp->value; - } else { - auto inode = make_object>(); - inode->value = arg; - inode->memorizer = memorizer; - inputs.push_back(LayoutAlternatedExpr(inode)); - return arg; - } - }; - - for (auto new_arg : new_args) { - // NOTE: do not support nested tuple - if (new_arg->IsInstance()) { - Tuple tuple_new_arg = Downcast(new_arg); - Array fields; - fields.reserve(tuple_new_arg->fields.size()); - for (auto x : tuple_new_arg->fields) { - Expr tmp = push_back_one_arg(x); - fields.push_back(tmp); - } - normal_new_args.push_back(WithFields(tuple_new_arg, fields)); - } else { - Expr tmp = push_back_one_arg(new_arg); - normal_new_args.push_back(tmp); - } - } - - // If there is no FInferCorrectLayout for the type, then we just assume the layout is correct. - static auto finfer_layout = Op::GetAttrMap("FInferCorrectLayout"); - if (Op::HasAttrMap("FTVMAlterOpLayout")) { - static auto falter_layout = Op::GetAttrMap("FTVMAlterOpLayout"); - if (ref_call->op.as()) { - Op op = Downcast(ref_call->op); - if (falter_layout.count(op) && !finfer_layout.count(op)) { - return memorizer->CallWithNewLayouts(ref_call, normal_new_args); - } - } - } - - // old_prd, new_prd = state[inputs] - // different ops can view a tensor with different layouts, e.g. conv_1->transpose(H, W)->conv_2 - // transpose view its output having NCWH layout, but conv_2 still views it as NCHW to operate - // old_prd, new_prd: the input layouts from the perspective of the producer (transpose) - // old_cur, new_cur: the input layouts from the perspective of the current node (conv_2) - // old_prd->new_prd tells how producer changed the layout - // old_cur->new_cur tells what change the current node wants to see - // No layout transforms are needed when they mean the same (NCHW->NCHW4c == NCWH->NCWH4c) - - // The workflow: - // 1. Run InferCorrectLayouts(NULL, old_prd) to get old_cur - // 2. Run InferCorrectLayouts(new_prd, old_prd) to get new_cur and rewrite the current op - - Array old_prd, old_cur, old_out, new_prd, new_out, new_cur; - for (auto inp : inputs) { - old_prd.push_back(inp->old_layout); - new_prd.push_back(inp->new_layout); - } - - // Collect input types to pass on to Infer Correct Layout. - tvm::Array types; - for (auto arg : ref_call->args) { - types.push_back(arg->checked_type()); - } - - bool success = false; - InferCorrectLayoutOutput infer_out; - std::tie(infer_out, success) = - InferCorrectLayouts(ref_call, Array(nullptr), old_prd, types); - old_cur = infer_out->input_layouts; - old_out = infer_out->output_layouts; - if (!success) { - return Expr(nullptr); - } - ICHECK_EQ(old_cur.size(), new_prd.size()); - - // for backward compatibility of InferCorrectLayouts - Array new_prd_inferred = new_prd; - // if new_prd_inferred == 'undef': new_prd_inferred = old_cur - for (size_t i = 0; i < new_prd_inferred.size(); ++i) { - if (!new_prd_inferred[i].defined()) { - new_prd_inferred.Set(i, old_cur[i]); - } - } - Array old_prd_inferred = old_prd; - // if old_prd_inferred == 'undef': old_prd_inferred = old_cur - for (size_t i = 0; i < old_prd_inferred.size(); ++i) { - if (!old_prd_inferred[i].defined()) { - old_prd_inferred.Set(i, old_cur[i]); - } - } - - // new_op = alter(op) - Call new_call = memorizer->CallWithNewLayouts(ref_call, infer_out->new_attrs, normal_new_args); - - // new_cur, new_out = op.infer(new_prd) - if (new_call->op->IsInstance()) { - success = false; - std::tie(infer_out, success) = - InferCorrectLayouts(new_call, new_prd_inferred, old_prd_inferred, types); - new_cur = infer_out->input_layouts; - new_out = infer_out->output_layouts; - if (!success) { - return Expr(nullptr); - } - } else { - return Expr(nullptr); - } - - ICHECK_EQ(new_out.size(), old_out.size()) - << "The number of output nodes should keep the same during alter_op_layout"; - ICHECK_EQ(new_prd.size(), new_cur.size()) - << "The number of input nodes should keep the same during alter_op_layout"; - - auto transform_layout = [&memorizer](Expr arg_item, const Layout& old_prd, const Layout& old_cur, - const Layout& new_prd, const Layout& new_cur) { - if (old_cur.Equals(old_prd)) { // the two transforms can be fused to one - arg_item = memorizer.Transform(arg_item, new_prd, new_cur); - } else { - if (old_prd.defined()) arg_item = memorizer.Transform(arg_item, new_prd, old_prd); - arg_item = memorizer.Transform(arg_item, old_cur, new_cur); - } - return arg_item; - }; - - DLOG(INFO) << "Transforming layout for `" << ref_call->op << "`"; - DLOG(INFO) << " old_prd=" << old_prd; - DLOG(INFO) << " new_prd=" << new_prd; - DLOG(INFO) << " old_cur=" << old_cur; - DLOG(INFO) << " new_cur=" << new_cur; - - // if (new_prd != new_cur): insert transform (new_prd -> new_cur) - Array transformed_args; - size_t pt = 0; - for (auto arg : new_call->args) { - if (arg->IsInstance()) { // unflatten tuple - Tuple tuple_arg = Downcast(arg); - Array transformed_tuple_arg; - transformed_tuple_arg.reserve(tuple_arg->fields.size()); - for (auto arg_item : tuple_arg->fields) { - transformed_tuple_arg.push_back( - transform_layout(arg_item, old_prd[pt], old_cur[pt], new_prd[pt], new_cur[pt])); - pt++; - } - transformed_args.push_back(WithFields(tuple_arg, transformed_tuple_arg)); - } else { - transformed_args.push_back( - transform_layout(arg, old_prd[pt], old_cur[pt], new_prd[pt], new_cur[pt])); - pt++; - } - } - ICHECK_EQ(pt, inputs.size()); - - // state[node] = (old_out, new_out) - // (handle tuple output) - if (ref_call->checked_type()->IsInstance()) { - Expr tuple_output = Call(new_call->op, transformed_args, infer_out->new_attrs); - Array fields; - for (size_t i = 0; i < new_out.size(); ++i) { - auto rnode = make_object>(); - rnode->value = TupleGetItem(tuple_output, i); - rnode->old_layout = old_out[i]; - rnode->new_layout = new_out[i]; - rnode->memorizer = memorizer; - fields.push_back(Expr(rnode)); - } - return Tuple(fields); - } else { - auto rnode = make_object>(); - ICHECK_EQ(new_out.size(), 1); - rnode->value = Call(new_call->op, transformed_args, infer_out->new_attrs, {}, ref_call->span); - rnode->old_layout = old_out[0]; - rnode->new_layout = new_out[0]; - rnode->memorizer = memorizer; - return Expr(rnode); - } -} - -} // namespace relay -} // namespace tvm - -#endif // TVM_RELAY_TRANSFORMS_TRANSFORM_LAYOUT_H_ diff --git a/src/relay/transforms/type_infer.cc b/src/relay/transforms/type_infer.cc deleted file mode 100644 index 4f4c998eee58..000000000000 --- a/src/relay/transforms/type_infer.cc +++ /dev/null @@ -1,1015 +0,0 @@ -/* - * 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 type_infer.cc - * \brief Relay type inference and checking. - * - * This file implements one of the most important passes to the - * Relay IR. In order to do many transformations and generate the - * most efficient code we need to obtain type information for the - * IR. - * - * Similar to previous computation graph based IRs, the Relay IR leaves - * type information implicit and computes types by performing program - * analysis. - * - * Given an expression `e` this pass infers a type `t` for - * the expression as well as simultaneously checking the property `e : t` - * (i.e., we can show e has type t). - * - * If we can not infer a type or there is a conflicting - * constraint it will emit errors. - */ - -#include -#include -#include -#include -#include -#include -#include - -#include "../analysis/type_solver.h" -#include "pass_utils.h" - -namespace tvm { -namespace relay { - -// Necessary deferred relation for TupleGetItem -struct TupleGetItemAttrs : public tvm::AttrsNode { - int index; - - TVM_DECLARE_ATTRS(TupleGetItemAttrs, "relay.attrs.TupleGetItemAttrs") { TVM_ATTR_FIELD(index); } -}; - -bool TupleGetItemRel(const Array& types, int num_inputs, const Attrs& attrs, - const TypeReporter& reporter) { - ICHECK_EQ(types.size(), 2); - if (types[0].as()) return false; - const auto* data = types[0].as(); - ICHECK(data != nullptr) << "TupleGetItem expect input type to be TupleType " - << " get " << types[0] << " instead"; - const auto* param = attrs.as(); - ICHECK(param != nullptr); - ICHECK_GE(param->index, 0); - ICHECK_LT(param->index, data->fields.size()); - reporter->Assign(types[1], data->fields[param->index]); - return true; -} - -TVM_REGISTER_NODE_TYPE(TupleGetItemAttrs); -TVM_REGISTER_GLOBAL("tvm.relay.type_relation.TupleGetItem").set_body_typed(TupleGetItemRel); - -struct ResolvedTypeInfo { - explicit ResolvedTypeInfo(Type checked_type, Array type_args) - : checked_type(checked_type), type_args(type_args) {} - ResolvedTypeInfo() {} - - Type checked_type; - // Only allocated when the expression is a call. - - Array type_args = Array(ObjectPtr(nullptr)); -}; - -// -// The inference algorithm can roughly be divided into three stages: -// - Populate the constraints by visiting the expression (TypeInferencer.GetType) -// - solver.AddConstraint and solver.Unify are called to populate the necessary constraints -// - Solve the constraints (solver_.Solve) -// - Recreate expression with the resolved checked_type (Resolver.VisitExpr) -// -class TypeInferencer : private ExprFunctor, - private PatternFunctor { - public: - // constructors - - explicit TypeInferencer(IRModule mod, DiagnosticContext diag_ctx) - : mod_(mod), diag_ctx(diag_ctx), solver_(GlobalVar(), diag_ctx) { - ICHECK(mod.defined()) << "Module must not be null in the type inferencer."; - } - - // Infer the types inside of a function. - Expr Infer(GlobalVar var, Function expr); - - private: - // type resolver that maps back to type - class Resolver; - // internal environment - IRModule mod_; - - // The current function being type checked. - GlobalVar current_func_; - - /*! \brief The diagnostic context. */ - DiagnosticContext diag_ctx; - - // map from expression to checked type - // type inferencer will populate it up - std::unordered_map type_map_; - - // The solver used by the inferencer. - TypeSolver solver_; - // relation function - TypeRelationFn tuple_getitem_rel_; - TypeRelationFn make_tuple_rel_; - - /*! \brief Internal map used for memoization. */ - std::unordered_map memo_; - - void VisitLeaf(const Expr& expr) { - if (!memo_.count(expr)) { - Type ret = this->DispatchVisitExpr(expr); - memo_[expr] = ret; - } - } - - bool CheckVisited(const Expr& expr) { - if (memo_.count(expr)) { - return true; - } else { - return false; - } - } - - Type DispatchVisitExpr(const Expr& expr) { return ExprFunctor::VisitExpr(expr); } - - Type VisitExpr(const Expr& expr) final { - auto fcheck_visited = [this](const Expr& expr) { return this->CheckVisited(expr); }; - auto fvisit_leaf = [this](const Expr& expr) { return this->VisitLeaf(expr); }; - if (memo_.count(expr)) { - return memo_[expr]; - } else { - ExpandDataflow(expr, fcheck_visited, fvisit_leaf); - return memo_[expr]; - } - } - - // Perform unification on two types and report the error at the expression - // or the span of the expression. - Type Unify(const Type& t1, const Type& t2, const Span& span, bool assign_lhs = true, - bool assign_rhs = true) { - try { - return solver_.Unify(t1, t2, span, assign_lhs, assign_rhs); - } catch (const Error& e) { - this->EmitFatal(Diagnostic::Error(span) - << "Error unifying `" << t1 << "` and `" << t2 << "`: " << e.what()); - return Type(); - } - } - - // Lazily get type for expr - // expression, we will populate it now, and return the result. - Type GetType(const Expr& expr) { - auto it = type_map_.find(expr); - if (it != type_map_.end() && it->second.checked_type.defined()) { - return it->second.checked_type; - } - Type ret = this->VisitExpr(expr); - ICHECK(ret.defined()) << "expression:" << std::endl << PrettyPrint(expr); - KindCheck(ret, mod_, this->diag_ctx); - ResolvedTypeInfo& rti = type_map_[expr]; - rti.checked_type = ret; - return ret; - } - - void EmitFatal(const Diagnostic& diag) { this->diag_ctx.EmitFatal(diag); } - - // Visitor Logic - Type VisitExpr_(const VarNode* op) final { - if (op->type_annotation.defined()) { - return op->type_annotation; - } else { - return IncompleteType(Kind::kType); - } - } - - Type VisitExpr_(const GlobalVarNode* op) final { - GlobalVar var = GetRef(op); - if (!mod_.defined()) { - this->EmitFatal(Diagnostic::Error(op->span) << "Cannot do type inference on global variables " - << "without a module"); - } - if (mod_->ContainGlobalVar(var->name_hint)) { - BaseFunc func = mod_->Lookup(var->name_hint); - - if (const auto* function_node = func.as()) { - VLOG(1) << "global var '" << op->name_hint << "' bound to Function"; - return function_node->checked_type(); - } else { - VLOG(1) << "global var '" << op->name_hint << "' bound to PrimFunc"; - return op->checked_type_; - } - } else { - // TODO(mbs): extern function cleanup - // Assume the function is extern thus no longer in the IRModule. - VLOG(1) << "global var '" << op->name_hint << "' not in module"; - return op->checked_type_; - } - } - - Type VisitExpr_(const ConstantNode* op) final { return op->tensor_type(); } - - Type VisitExpr_(const TupleNode* op) final { - Array types; - for (Expr field : op->fields) { - types.push_back(GetType(field)); - } - return TupleType(types); - } - - Type VisitExpr_(const TupleGetItemNode* op) final { - if (!tuple_getitem_rel_.defined()) { - tuple_getitem_rel_ = - Downcast(EnvFunc::Get("tvm.relay.type_relation.TupleGetItem")); - } - Type tuple_type = GetType(op->tuple); - Type rtype = IncompleteType(Kind::kType); - auto attrs = make_object(); - attrs->index = op->index; - solver_.AddConstraint(TypeRelation(tuple_getitem_rel_, {tuple_type, rtype}, 1, Attrs(attrs)), - op->span); - return rtype; - } - - void VisitPattern_(const PatternConstructorNode* con, const Type& t) { - ICHECK(mod_.defined()) << "Cannot do type inference without a environment:" - << con->constructor->name_hint; - TypeData td = mod_->type_definitions.at(con->constructor->belong_to); - auto pc = GetRef(con); - - // we can expect a certain number of arguments - Array unknown_args; - for (size_t i = 0; i < td->type_vars.size(); i++) { - unknown_args.push_back(IncompleteType(Kind::kType)); - } - - Type expected = TypeCall(con->constructor->belong_to, unknown_args); - Type unified = Unify(t, expected, pc->span); - - auto* tc = unified.as(); - if (!tc) { - this->EmitFatal(Diagnostic::Error(pc->span) << "Expected a type call, got " << unified); - } - - if (td->header != tc->func) { - this->EmitFatal(Diagnostic::Error(pc->span) << "ADT headers must match, but we have " - << td->header << " and " << tc->func); - } - - if (td->type_vars.size() != tc->args.size()) { - this->EmitFatal(Diagnostic::Error(pc->span) - << "The number of type args must match" - << "the number of type vars in the type data: " << td->type_vars.size() - << " != " << tc->args.size()); - } - std::unordered_map type_var_map_; - for (size_t i = 0; i < td->type_vars.size(); ++i) { - type_var_map_[td->type_vars[i]] = tc->args[i]; - } - - if (con->constructor->inputs.size() != con->patterns.size()) { - this->EmitFatal(Diagnostic::Error(pc->span) << "Not enough inputs for the constructor; " - << "expected " << con->constructor->inputs.size() - << ", got " << con->patterns.size()); - } - - for (size_t i = 0; i < con->constructor->inputs.size(); ++i) { - VisitPattern(con->patterns[i], Bind(con->constructor->inputs[i], type_var_map_)); - } - } - - void VisitPattern_(const PatternTupleNode* tup, const Type& t) { - auto pt = GetRef(tup); - - // we can expect a certain number of arguments - Array unknown_args; - for (size_t i = 0; i < tup->patterns.size(); i++) { - unknown_args.push_back(IncompleteType(Kind::kType)); - } - - Type expected = TupleType(unknown_args); - Type unified = Unify(t, expected, tup->span); - - auto* tt = unified.as(); - if (!tt) { - this->EmitFatal(Diagnostic::Error(pt->span) << "Expected a tuple type, got " << unified); - } - ICHECK(tup->patterns.size() == tt->fields.size()) << "not enough pattern"; - for (size_t i = 0; i < tup->patterns.size(); ++i) { - VisitPattern(tup->patterns[i], tt->fields[i]); - } - } - - void VisitPattern_(const PatternVarNode* pv, const Type& t) { - Type vt = GetType(pv->var); - Unify(vt, t, pv->span); - } - - void VisitPattern_(const PatternWildcardNode* wc, const Type& t) {} - - Type VisitExpr_(const MatchNode* op) final { - Type dtype = GetType(op->data); - for (const auto& c : op->clauses) { - VisitPattern(c->lhs, dtype); - } - Type rtype = IncompleteType(Kind::kType); - for (const auto& c : op->clauses) { - rtype = this->Unify(rtype, GetType(c->rhs), op->span); - } - - if (op->complete) { - // check completness - Match match = GetRef(op); - Array unmatched_cases = UnmatchedCases(match, this->mod_); - if (unmatched_cases.size() != 0) { - ErrorBuilder ss; - auto err = Diagnostic::Error(match->span); - err << "match expression does not handle the following cases: "; - int i = 0; - for (auto cs : unmatched_cases) { - err << "case " << i++ << ": \n" << PrettyPrint(cs); - } - this->EmitFatal(err); - } - } - - return rtype; - } - - Type VisitExpr_(const OpNode* op) final { return op->op_type; } - - Type VisitExpr_(const LetNode* let) final { - auto pre_visit = [this](const LetNode* op) { - // if the definition is a function literal, permit recursion - bool is_functional_literal = op->value.as() != nullptr; - Type let_type = IncompleteType(Kind::kType); - - if (is_functional_literal) { - let_type = this->GetType(op->var); - this->type_map_[op->var].checked_type = let_type; - } - - if (op->var->type_annotation.defined()) { - let_type = this->Unify(let_type, op->var->type_annotation, op->span); - } - - Type vtype = this->GetType(op->value); - let_type = this->Unify(let_type, vtype, op->span); - - ICHECK(is_functional_literal || !this->type_map_.count(op->var)); - // NOTE: no scoping is necessary because var are unique in program - this->type_map_[op->var].checked_type = let_type; - }; - auto post_visit = [this](const LetNode* op) { - Expr expr = GetRef(op); - this->memo_[expr] = this->GetType(op->body); - this->type_map_[expr].checked_type = this->memo_[expr]; - }; - ExpandANormalForm(let, pre_visit, post_visit); - return memo_[GetRef(let)]; - } - - Type VisitExpr_(const IfNode* ite) final { - // Ensure the type of the guard is of Tensor[Bool, ()], - // that is a rank-0 boolean tensor. - Type cond_type = this->GetType(ite->cond); - this->Unify(cond_type, TensorType::Scalar(tvm::DataType::Bool()), ite->cond->span); - Type checked_true = this->GetType(ite->true_branch); - Type checked_false = this->GetType(ite->false_branch); - return this->Unify(checked_true, checked_false, ite->span); - } - - // This code is special-cased for primitive operators, - // which are registered in the style defined in src/relay/op/*. - // - // The result will be the return type of the operator. - Type PrimitiveCall(const FuncTypeNode* op, Array arg_types, const Attrs& attrs, - const Span& span) { - if (op->type_params.size() != arg_types.size() + 1) return Type(); - if (op->type_constraints.size() != 1) return Type(); - const TypeRelationNode* rel = op->type_constraints[0].as(); - if (rel == nullptr) return Type(); - // validate if the type parameter matches up - for (size_t i = 0; i < op->type_params.size(); ++i) { - if (!op->type_params[i].same_as(rel->args[i])) return Type(); - } - Type rtype = IncompleteType(Kind::kType); - arg_types.push_back(rtype); - // we can do simple replacement here - solver_.AddConstraint(TypeRelation(rel->func, arg_types, arg_types.size() - 1, attrs), span); - return rtype; - } - - // substitute the type args in the function type - FuncType InstantiateFuncType(const FuncTypeNode* fn_ty, const Array& ty_args) { - tvm::Map subst_map; - - // Build a subsitituion map up from the function type and type arguments. - // Eventually allow the type vars to be passed in. - ICHECK(fn_ty->type_params.size() == ty_args.size()) - << "number of type parameters does not match expected"; - for (size_t i = 0; i < ty_args.size(); ++i) { - subst_map.Set(fn_ty->type_params[i], ty_args[i]); - } - - Type ret_type = fn_ty->ret_type; - - // If the function type is incomplete, place a new IncompleteType - // This relax the fn_ty to inputs -> Any - // The type checking can still pass when there are additional constraints on the type - // This is a temporary work around to check recursive functions whose - // return type is not yet known. - if (!ret_type.defined()) { - ret_type = IncompleteType(Kind::kType); - } - - Type inst_ty = FuncType(fn_ty->arg_types, ret_type, {}, fn_ty->type_constraints); - inst_ty = Bind(inst_ty, subst_map); - return Downcast(inst_ty); - } - - // instantiates starting from incompletes - FuncType InstantiateFuncType(const FuncTypeNode* fn_ty) { - if (fn_ty->type_params.size() == 0) { - return GetRef(fn_ty); - } - - Array type_args; - for (size_t i = 0; i < fn_ty->type_params.size(); i++) { - type_args.push_back(IncompleteType(Kind::kType)); - } - return InstantiateFuncType(fn_ty, type_args); - } - - void AddTypeArgs(const Expr& expr, Array type_args) { - auto type_info = type_map_.find(expr); - if (type_info == type_map_.end()) { - type_map_.insert({expr, ResolvedTypeInfo(Type(), type_args)}); - } else { - ICHECK(!type_info->second.type_args.defined()); - type_info->second.type_args = type_args; - } - } - - // Handle general call node. - Type GeneralCall(const CallNode* call, Array arg_types) { - Type ftype = GetType(call->op); - auto* fn_ty_node = ftype.as(); - auto* inc_ty_node = ftype.as(); - - if (fn_ty_node == nullptr && inc_ty_node == nullptr) { - this->EmitFatal(Diagnostic::Error(call->span) - << "only expressions with function types can be called, found " << ftype); - } - - // incomplete type => it must be a function taking the arg types - // with an unknown return type - if (inc_ty_node != nullptr) { - Type ret_type = IncompleteType(Kind::kType); - Type func_type = FuncType(arg_types, ret_type, {}, {}); - Type unified = this->Unify(ftype, func_type, call->op->span); - fn_ty_node = unified.as(); - } - - Array type_args = call->type_args; - if (type_args.size() > fn_ty_node->type_params.size()) { - this->EmitFatal(Diagnostic::Error(call->span) - << "Incorrect number of type args in " << call->span << ": " - << "Expected " << fn_ty_node->type_params.size() << " but got " - << type_args.size() << " for call:\n" - << PrettyPrint(GetRef(call))); - } - for (size_t i = type_args.size(); i < fn_ty_node->type_params.size(); i++) { - type_args.push_back(IncompleteType(TypeKind::kType)); - } - - FuncType fn_ty = InstantiateFuncType(fn_ty_node, type_args); - - AddTypeArgs(GetRef(call), type_args); - - size_t type_arity = fn_ty->arg_types.size(); - size_t number_of_args = arg_types.size(); - bool is_variable = false; - - if (const OpNode* opnode = call->op.as()) { - if (opnode->num_inputs == -1) { - is_variable = true; - } - } - - if ((type_arity < number_of_args) && !is_variable) { - this->EmitFatal(Diagnostic::Error(call->span) - << "the function is provided too many arguments " - << "expected " << type_arity << ", found " << number_of_args); - } else if (type_arity > number_of_args) { - this->EmitFatal(Diagnostic::Error(call->span) - << "the function is provided too few arguments " - << "expected " << type_arity << ", found " << number_of_args); - } - - Array unified_arg_types; - if (!is_variable) { - for (size_t i = 0; i < fn_ty->arg_types.size(); i++) { - this->Unify(fn_ty->arg_types[i], arg_types[i], call->span, true, false); - } - } else { - for (size_t i = 0; i < number_of_args; i++) { - if (i < fn_ty->arg_types.size()) { - unified_arg_types.push_back( - this->Unify(fn_ty->arg_types[i], arg_types[i], call->span, false, false)); - } else { - unified_arg_types.push_back(arg_types[i]); - } - } - unified_arg_types.push_back(fn_ty->ret_type); - } - for (auto cs : fn_ty->type_constraints) { - if (const auto* tr = cs.as()) { - if (!is_variable) { - solver_.AddConstraint(TypeRelation(tr->func, tr->args, tr->num_inputs, call->attrs), - call->span); - } else { - solver_.AddConstraint( - TypeRelation(tr->func, unified_arg_types, number_of_args, call->attrs), call->span); - } - } else { - solver_.AddConstraint(cs, call->span); - } - } - - return fn_ty->ret_type; - } - - Type VisitExpr_(const CallNode* call) final { - Array arg_types; - for (Expr arg : call->args) { - arg_types.push_back(GetType(arg)); - } - - if (const OpNode* opnode = call->op.as()) { - Type rtype = - PrimitiveCall(opnode->op_type.as(), arg_types, call->attrs, call->span); - - if (rtype.defined()) { - AddTypeArgs(GetRef(call), arg_types); - return rtype; - } - } - - solver_.Solve(); - return GeneralCall(call, arg_types); - } - - Type VisitExpr_(const FunctionNode* f) final { - solver_.Solve(); - Array arg_types; - for (auto param : f->params) { - arg_types.push_back(GetType(param)); - } - Type rtype = GetType(f->body); - if (auto* ft = rtype.as()) { - rtype = InstantiateFuncType(ft); - } - if (f->ret_type.defined()) { - rtype = this->Unify(f->ret_type, rtype, GetRef(f)->span); - } - ICHECK(rtype.defined()); - auto ret = FuncType(arg_types, rtype, f->type_params, {}); - return solver_.Resolve(ret); - } - - Type VisitExpr_(const RefCreateNode* op) final { return RelayRefType(GetType(op->value)); } - - Type VisitExpr_(const RefReadNode* op) final { - Type it = IncompleteType(Kind::kType); - this->Unify(GetType(op->ref), RelayRefType(it), op->span); - return it; - } - - Type VisitExpr_(const RefWriteNode* op) final { - Type it = IncompleteType(Kind::kType); - this->Unify(GetType(op->ref), RelayRefType(it), op->span); - this->Unify(GetType(op->value), it, op->span); - return TupleType::Empty(); - } - - Type VisitExpr_(const ConstructorNode* c) final { - ICHECK(mod_.defined()) << "Cannot do type inference without a environment:" << c->name_hint; - TypeData td = mod_->LookupTypeDef(c->belong_to); - std::vector types; - for (const auto& t : td->type_vars) { - types.push_back(t); - } - return FuncType(c->inputs, TypeCall(c->belong_to, types), td->type_vars, {}); - } - - void Solve() { solver_.Solve(); } -}; - -class TypeInferencer::Resolver : public MixedModeMutator, PatternMutator { - public: - Resolver(const std::unordered_map& tmap, - TypeSolver* solver) - : tmap_(tmap), solver_(solver) {} - - using MixedModeMutator::VisitExpr_; - - Expr VisitExpr_(const VarNode* op) final { return VisitVar(GetRef(op)); } - - Expr VisitExpr_(const ConstantNode* op) final { return AttachCheckedType(op); } - - Expr VisitExpr_(const GlobalVarNode* op) final { return GetRef(op); } - - Expr VisitExpr_(const OpNode* op) final { return ExprMutator::VisitExpr_(op); } - - Expr Rewrite_(const TupleNode* op, const Expr& post) final { return AttachCheckedType(op, post); } - - Expr Rewrite_(const TupleGetItemNode* op, const Expr& post) final { - return AttachCheckedType(op, post); - } - - Expr VisitExpr_(const FunctionNode* op) final { return AttachCheckedType(op); } - - Expr Rewrite_(const CallNode* op, const Expr& post) final { return AttachCheckedType(op, post); } - - Expr VisitExpr_(const LetNode* op) final { - auto pre_visit = [this](const LetNode* op) { - this->VisitExpr(op->var); - this->VisitExpr(op->value); - }; - auto post_visit = [this](const LetNode* op) { - Expr expr = GetRef(op); - Var var = Downcast(this->VisitExpr(op->var)); - Expr value = this->VisitExpr(op->value); - Expr body = this->VisitExpr(op->body); - this->memo_[expr] = this->AttachCheckedType(op, Let(var, value, body)); - }; - ExpandANormalForm(op, pre_visit, post_visit); - return memo_[GetRef(op)]; - } - - Expr VisitExpr_(const IfNode* op) final { return AttachCheckedType(op); } - - Expr VisitExpr_(const RefCreateNode* op) final { return AttachCheckedType(op); } - - Expr VisitExpr_(const RefReadNode* op) final { return AttachCheckedType(op); } - - Expr VisitExpr_(const RefWriteNode* op) final { return AttachCheckedType(op); } - - Expr VisitExpr_(const ConstructorNode* op) final { return AttachCheckedType(op); } - - Expr VisitExpr_(const MatchNode* op) final { return AttachCheckedType(op); } - - Pattern VisitPattern(const Pattern& p) final { return PatternMutator::VisitPattern(p); } - - Var VisitVar(const Var& v) final { - if (vmap_.count(v) == 0) { - vmap_[v] = Downcast(AttachCheckedType(v.as())); - } - return vmap_.at(v); - } - - // attach checked type to the mutated node. - template - Expr AttachCheckedType(const T* op, const Expr& post = Expr()) { - auto it = tmap_.find(GetRef(op)); - ICHECK(it != tmap_.end()); - Type checked_type = solver_->Resolve(it->second.checked_type); - - if (checked_type.as() != nullptr) { - this->solver_->Emit( - Diagnostic::Error(op->span) - << "The type inference pass was unable to infer a type for this expression.\n" - << "This usually occurs when an operator call is under constrained in some way," - << " check other reported errors for hints of what may of happened."); - } - - Expr new_e = post.defined() ? post : ExprMutator::VisitExpr_(op); - // new_call and new_var's code is only going to be valid for VarNode/CallNode. - // Compiler optimization will likely fold these away for other nodes. - CallNode* new_call = (std::is_base_of::value - ? const_cast(static_cast(new_e.get())) - : nullptr); - VarNode* new_var = (std::is_base_of::value - ? const_cast(static_cast(new_e.get())) - : nullptr); - FunctionNode* new_fn = - (std::is_base_of::value - ? const_cast(static_cast(new_e.get())) - : nullptr); - - // check if we need update the new_e - bool need_update_type = !checked_type.same_as(new_e->checked_type_); - bool need_update_call = - (std::is_base_of::value && it->second.type_args.defined() && - !it->second.type_args.same_as(new_call->type_args)); - bool need_update_var = (std::is_base_of::value && update_missing_type_annotation_ && - !new_var->type_annotation.defined()); - - bool need_update_fn = (std::is_base_of::value && - update_missing_type_annotation_ && !new_fn->ret_type.defined()); - - if (!need_update_type && !need_update_var && !need_update_call && !need_update_fn) { - return new_e; - } - - if (!new_e.unique()) { - // Copy on write optimization - // If new_e is an old expression, - // we make a copy mutating an existing reference. - ObjectPtr ptr = make_object(*new_e.as()); - new_e = Expr(ptr); - new_call = - (std::is_base_of::value ? static_cast(ptr.get()) : nullptr); - new_var = (std::is_base_of::value ? static_cast(ptr.get()) : nullptr); - new_fn = (std::is_base_of::value ? static_cast(ptr.get()) - : nullptr); - } - - // attach the information. - if (need_update_type) { - new_e->checked_type_ = checked_type; - } - - if (need_update_call) { - new_call->type_args = it->second.type_args; - for (size_t i = 0; i < new_call->type_args.size(); i++) { - new_call->type_args.Set(i, solver_->Resolve(new_call->type_args[i])); - } - } - if (need_update_var) { - new_var->type_annotation = checked_type; - } - if (need_update_fn) { - auto* fn_type = checked_type.as(); - ICHECK(fn_type != nullptr); - new_fn->ret_type = fn_type->ret_type; - } - return new_e; - } - - Type VisitType(const Type& t) final { return solver_->Resolve(t); } - - private: - std::unordered_map vmap_; - const std::unordered_map& tmap_; - TypeSolver* solver_; - // whether attach the checked type as type_annotation - // if original type anntation is missing. - bool update_missing_type_annotation_{true}; -}; - -Expr TypeInferencer::Infer(GlobalVar var, Function function) { - // Set the current function being type checked. - this->current_func_ = var; - - // Step 1: Populate the constraints. - GetType(function); - - // Step 2: Solve the constraints. - Solve(); - - // Step 3: Attach resolved types to checked_type field. - auto resolved_expr = Resolver(type_map_, &solver_).VisitExpr(function); - - if (!WellFormed(resolved_expr, this->diag_ctx)) { - this->diag_ctx.Emit(Diagnostic::Bug(function->span) - << "the type checked function is malformed, please report this"); - } - - return resolved_expr; -} - -struct AllCheckTypePopulated : MixedModeVisitor { - using MixedModeVisitor::VisitExpr_; - void DispatchExprVisit(const Expr& e) { - if (e.as()) { - return; - } - if (e.as()) { - return; - } - if (e.as()) { - return; - } - ICHECK(e->checked_type_.defined()) << "Expression: " << e; - return ExprVisitor::VisitExpr(e); - } - void VisitExpr_(const LetNode* op) final { - auto pre_visit = [this](const LetNode* op) { - this->VisitExpr(op->var); - this->VisitExpr(op->value); - }; - auto post_visit = [this](const LetNode* op) { - this->VisitExpr(op->body); - this->visit_counter_[op] += 1; - }; - ExpandANormalForm(op, pre_visit, post_visit); - } -}; - -void EnsureCheckedType(const Expr& e) { AllCheckTypePopulated().VisitExpr(e); } - -// TODO(@jroesch): Can we optimize this? -void AddGlobalTypes(IRModule mod) { - std::vector> updates; - for (const auto& it : mod->functions) { - // Currently we don't type check TIR. - // The inferencer will only check Relay functions - // the future plan is to have a unified type checker - // that works on TIR and Relay at the same time. - if (auto* func_node = it.second.as()) { - Function func = Function(make_object(*func_node)); - func->checked_type_ = func->func_type_annotation(); - updates.push_back({it.first, Downcast(func)}); - } - } - - for (const auto& pair : updates) { - mod->Add(pair.first, pair.second, true); - } -} - -/*! - * \brief Returns a possibly much smaller subgraph whose inner nodes have the same type. - * - * Returns the largest sub-graph who's inner nodes need types and leaves are vars standing in - * for already typed sub-expressions. This creates a graph whose inner nodes have the same - * type as the original graph and when running type inference, we can avoid copying and - * recursing through most of the expression graph when running type inference. Note, this assumes - * that current populated type information is correct! - * - * ExprMutator is sufficient over MixedModemutator since we will not recurse much. - */ -class SameTypedSubgraphExtractor : public ExprMutator { - Expr VisitExpr_(const VarNode* op) { return Var(op->vid, op->type_annotation, op->span); } - Expr VisitExpr_(const ConstantNode* op) { return Constant(op->data, op->span); } - Expr VisitExpr_(const GlobalVarNode* op) { return GlobalVar(op->name_hint); } - Expr VisitExpr_(const OpNode* op) { return Op(GetRef(op)); } - Expr VisitExpr_(const TupleNode* op) { - return Tuple(GetAnalogousExpression(op->fields), op->span); - } - Expr VisitExpr_(const FunctionNode* op) { - // Unfortunately our strategy of inserting variables as dummies would change the signature of - // existing function nodes so we have to copy all used functions always :/ - return Function(op->params, op->body, op->ret_type, op->type_params, op->attrs, op->span); - } - Expr VisitExpr_(const CallNode* op) { - return Call(op->op, GetAnalogousExpression(op->args), op->attrs, op->type_args, op->span); - } - Expr VisitExpr_(const LetNode* op) { - return Let(op->var, GetAnalogousExpression(op->value), GetAnalogousExpression(op->body), - op->span); - } - Expr VisitExpr_(const IfNode* op) { - return If(GetAnalogousExpression(op->cond), GetAnalogousExpression(op->true_branch), - GetAnalogousExpression(op->false_branch), op->span); - } - Expr VisitExpr_(const TupleGetItemNode* op) { - return TupleGetItem(GetAnalogousExpression(op->tuple), op->index, op->span); - } - Expr VisitExpr_(const RefCreateNode* op) { - return RefCreate(GetAnalogousExpression(op->value), op->span); - } - Expr VisitExpr_(const RefReadNode* op) { - return RefRead(GetAnalogousExpression(op->ref), op->span); - } - Expr VisitExpr_(const RefWriteNode* op) { - return RefWrite(GetAnalogousExpression(op->ref), GetAnalogousExpression(op->value), op->span); - } - Expr VisitExpr_(const ConstructorNode* op) { - return Constructor(op->name_hint, op->inputs, op->belong_to); - } - Expr VisitExpr_(const MatchNode* op) { - return Match(GetAnalogousExpression(op->data), op->clauses, op->complete, op->span); - } - - private: - Expr GetAnalogousExpression(const Expr& expr) { - // Replace the expression with a potentially simpler expression of the same type - if (expr->checked_type_.defined()) { - // Since the expression already has a checked_type which we assume is correct we don't need - // full type inference to enter it. So stub it out with a dummy var of the same type. - return Var("dummy_var", expr->checked_type(), expr->span); - } - - return VisitExpr(expr); - } - Array GetAnalogousExpression(const Array& fields) { - Array new_fields; - for (Expr expr : fields) { - new_fields.push_back(GetAnalogousExpression(expr)); - } - return new_fields; - } -}; - -namespace transform { - -Type InferTypeLocal(const Expr& expr) { - /* - This type inference differs from InferType in that it uses existing type information - to avoid recursing over much of the graph, and it only examines the type of the input - node. This makes it faster if you need to run type inference iteratively throughout - a pass for example. - - However, it assumes any existing populated type inference is correct! If some populated - type inference is incorrect, an incorrect type may be returned or a type error will be - raised. If you know not all populated type fields are correct with the current graph, - you should use InferType() instead. - */ - SameTypedSubgraphExtractor subgraph_extractor; - Expr sub_graph = subgraph_extractor(expr); - - Type result_type; - result_type = relay::InferType(sub_graph)->checked_type(); - - expr->checked_type_ = result_type; - return result_type; -} - -TVM_REGISTER_GLOBAL("relay._transform.InferTypeLocal").set_body_typed([](const Expr& expr) { - return InferTypeLocal(expr); -}); - -Pass InferType() { - auto pass_info = PassInfo(0, "InferType", {}, /* trace */ false); - return tvm::transform::CreateModulePass( - [=](IRModule mod, const PassContext& pass_ctx) { - // Execute the pass function and return a new module. - IRModule updated_mod = mod->ShallowCopy(); - - pass_ctx->diag_ctx = DiagnosticContext::Default(updated_mod); - - // Add all the type annotations to the functions in the model. - AddGlobalTypes(mod); - - std::vector> updates; - for (const auto& it : updated_mod->functions) { - // Currently we don't type check TIR. - // - // The inferencer will only check Relay functions. - - // In the future we plan a unified type checker - // that works on TIR and Relay at the same time. - if (auto func = it.second.as()) { - // // If a function already has type information we can skip checking it. - // if (func->checked_type_.defined()) { - // continue; - // } - - // TODO(@jroesch): we should be able to move the type inferencer outside - // of this function but it seems to be more stateful then I expect. - auto inferencer = TypeInferencer(mod, pass_ctx->diag_ctx.value()); - auto updated_func = inferencer.Infer(it.first, func.value()); - - pass_ctx->diag_ctx.value().Render(); - - // After we are done checking write the global type back - // into the global var. - it.first->checked_type_ = updated_func->checked_type(); - - if (!WellFormed(updated_func, pass_ctx->diag_ctx)) { - LOG(FATAL) << "The type checked intermediate representation is malformed"; - } - - auto free_tvars = FreeTypeVars(updated_func, mod); - ICHECK(free_tvars.size() == 0) - << "Found unbound type variables in " << updated_func << ": " << free_tvars; - EnsureCheckedType(updated_func); - updates.push_back({it.first, Downcast(updated_func)}); - } - } - - for (const auto& pair : updates) { - updated_mod->Add(pair.first, pair.second, true); - } - - return updated_mod; - }, - 0, "InferType", {}); -} - -TVM_REGISTER_GLOBAL("relay._transform.InferType").set_body_typed([]() { return InferType(); }); - -} // namespace transform - -} // namespace relay -} // namespace tvm diff --git a/src/runtime/aot_executor/aot_executor.cc b/src/runtime/aot_executor/aot_executor.cc deleted file mode 100644 index 955f97adf8fc..000000000000 --- a/src/runtime/aot_executor/aot_executor.cc +++ /dev/null @@ -1,232 +0,0 @@ -/* - * 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. - */ - -/*! - * \brief Defines an implementation of Module-based Model Runtime Interface that works with - * Ahead-of-Time compilation. - * \file aot_executor.cc - */ - -#include "aot_executor.h" - -#include -#include -#include - -#include -#include - -#include "../meta_data.h" - -namespace tvm { -namespace runtime { - -AotExecutor::AotExecutor(tvm::runtime::Module module, const std::vector& devs) - : module_{module}, devices_{devs} { - auto fmetadata = module->GetFunction("get_metadata"); - CHECK(fmetadata != nullptr) << "Expected a module with PackedFunc get_metadata"; - auto ret_value = fmetadata(); - metadata_ = ret_value.AsObjectRef(); - - ICHECK_EQ(devices_.size(), 1) << "Expect exactly 1 device passed."; - DLDevice expected_device{kDLCPU, 0}; - ICHECK_EQ(devices_[0].device_id, expected_device.device_id) - << "At this time, AOTExecutor supports only execution on kDLCPU 0"; - // TODO(tvm-team): Temporary hack since Hexagon is defined different than kDLCPU. - bool is_valid_device = - (devices_[0].device_type == kDLHexagon) || (devices_[0].device_type == kDLCPU); - CHECK(is_valid_device) - << "At this time, AOTExecutor supports only execution on kDLCPU 0 or kDLHexagon 0"; - - for (auto input : metadata_->inputs()) { - // TODO(areusch): Encode device information in Metadata. - args_.emplace_back(NDArray::Empty(ShapeTuple(input->shape().begin(), input->shape().end()), - input->dtype(), devices_[0])); - } - - for (auto output : metadata_->outputs()) { - args_.emplace_back(NDArray::Empty(ShapeTuple(output->shape().begin(), output->shape().end()), - output->dtype(), devices_[0])); - } - - // USMP is used - if (metadata_->num_workspace_pools()) { - // merge all constants into one ndarray - int64_t blob_len = 0; - for (const auto& c : metadata_->constant_pools()) { - auto data = c->data(); - int64_t byte_size = GetDataSize(*data.operator->()) + c->byte_offset(); - blob_len = blob_len > byte_size ? blob_len : byte_size; - } - ICHECK(blob_len < std::numeric_limits::max()); - NDArray ci = NDArray::Empty({blob_len}, DataType::UInt(8), devices_[0]); - for (const auto& c : metadata_->constant_pools()) { - auto data = c->data(); - data.CopyToBytes(static_cast(ci->data) + c->byte_offset(), - GetDataSize(*data.operator->())); - } - // Emplace constant node pool only if workspace pools supplied - args_.emplace_back(ci); - - int32_t pool_len = 0; - for (auto pool : metadata_->workspace_pools()) { - pool_len = - GetDataSize(*NDArray::Empty({pool->shape()}, pool->dtype(), devices_[0]).operator->()); - args_.emplace_back(NDArray::Empty({pool_len}, DataType::UInt(8), devices_[0])); - } - } -} - -PackedFunc AotExecutor::GetFunction(const String& name, const ObjectPtr& sptr_to_self) { - // Return member functions during query. - if (name == "set_input") { - return PackedFunc([sptr_to_self, this](TVMArgs args, TVMRetValue* rv) { - if (String::CanConvertFrom(args[0])) { - int in_idx = this->GetInputIndex(tvm::runtime::SanitizeName(args[0].operator String())); - if (in_idx >= 0) this->SetInput(in_idx, args[1]); - } else { - this->SetInput(args[0], args[1]); - } - }); - } else if (name == "set_input_zero_copy") { - return PackedFunc([sptr_to_self, this](TVMArgs args, TVMRetValue* rv) { - if (String::CanConvertFrom(args[0])) { - int in_idx = this->GetInputIndex(tvm::runtime::SanitizeName(args[0].operator String())); - if (in_idx >= 0) this->SetInputZeroCopy(in_idx, args[1]); - } else { - this->SetInputZeroCopy(args[0], args[1]); - } - }); - } else if (name == "set_output_zero_copy") { - return PackedFunc([sptr_to_self, this](TVMArgs args, TVMRetValue* rv) { - if (String::CanConvertFrom(args[0])) { - int out_idx = this->GetOutputIndex(tvm::runtime::SanitizeName(args[0].operator String())); - if (out_idx >= 0) this->SetOutputZeroCopy(out_idx, args[1]); - } else { - this->SetOutputZeroCopy(args[0], args[1]); - } - }); - } else if (name == "get_output") { - return PackedFunc([sptr_to_self, this](TVMArgs args, TVMRetValue* rv) { - if (args.num_args == 2) { - this->CopyOutputTo(args[0], args[1]); - } else { - *rv = this->GetOutput(args[0]); - } - }); - } else if (name == "get_input") { - return PackedFunc([sptr_to_self, this](TVMArgs args, TVMRetValue* rv) { - int in_idx = 0; - if (String::CanConvertFrom(args[0])) { - in_idx = this->GetInputIndex(tvm::runtime::SanitizeName(args[0].operator String())); - } else { - in_idx = args[0]; - } - if (in_idx >= 0) { - *rv = this->GetInput(in_idx); - } - }); - } else if (name == "get_num_outputs") { - return PackedFunc( - [sptr_to_self, this](TVMArgs args, TVMRetValue* rv) { *rv = this->NumOutputs(); }); - } else if (name == "get_num_inputs") { - return PackedFunc( - [sptr_to_self, this](TVMArgs args, TVMRetValue* rv) { *rv = this->NumInputs(); }); - } else if (name == "run") { - return PackedFunc([sptr_to_self, this](TVMArgs args, TVMRetValue* rv) { this->Run(); }); - } else if (name == "get_input_index") { - return PackedFunc([sptr_to_self, this](TVMArgs args, TVMRetValue* rv) { - CHECK(String::CanConvertFrom(args[0])) << "Input key is not a string"; - *rv = this->GetInputIndex(tvm::runtime::SanitizeName(args[0].operator String())); - }); - } else if (name == "get_input_name") { - return PackedFunc( - [sptr_to_self, this](TVMArgs args, TVMRetValue* rv) { *rv = this->GetInputName(args[0]); }); - } else { - return PackedFunc(); - } -} - -void AotExecutor::Run() { - auto pf = module_.GetFunction( - get_name_mangled(metadata_->mod_name(), ::tvm::runtime::symbol::tvm_module_main), - true /* query_imports */); - ICHECK(pf != nullptr) << "Module entrypoint is not defined"; - - const int num_args = args_.size(); - auto call_values = ::std::make_unique(num_args); - auto call_type_codes = ::std::make_unique(num_args); - for (int i = 0; i < num_args; ++i) { - auto managed = args_[i].ToDLPack(); - call_values.get()[i].v_handle = &managed->dl_tensor; - call_type_codes.get()[i] = kTVMDLTensorHandle; - } - - TVMArgs args{call_values.get(), call_type_codes.get(), num_args}; - TVMRetValue rv; - pf.CallPacked(args, &rv); -} - -int AotExecutor::GetInputIndex(const std::string& name) { - auto inputs = metadata_->inputs(); - for (unsigned int i = 0; i < inputs.size(); i++) { - if (inputs[i]->name() == name) { - return i; - } - } - ICHECK(false) << "Invalid input name."; -} - -std::string AotExecutor::GetInputName(int index) { - auto inputs = metadata_->inputs(); - return inputs[index]->name(); -} - -int AotExecutor::GetOutputIndex(const std::string& name) { - auto outputs = metadata_->outputs(); - for (unsigned int i = 0; i < outputs.size(); i++) { - if (outputs[i]->name() == name) { - return i; - } - } - return -1; -} - -void AotExecutor::SetInput(int index, DLTensor* data_ref) { args_[index].CopyFrom(data_ref); } - -void AotExecutor::SetInputZeroCopy(int index, DLTensor* data_ref) { - ICHECK(false) << "not implemented"; -} - -void AotExecutor::SetOutputZeroCopy(int index, DLTensor* data_ref) { - ICHECK(false) << "not implemented"; -} - -int AotExecutor::NumOutputs() const { return metadata_->num_outputs(); } - -int AotExecutor::NumInputs() const { return metadata_->num_inputs(); } - -NDArray AotExecutor::GetInput(int index) const { return args_[index]; } - -NDArray AotExecutor::GetOutput(int index) const { return args_[metadata_->num_inputs() + index]; } - -void AotExecutor::CopyOutputTo(int index, DLTensor* data_out) { GetOutput(index).CopyTo(data_out); } - -} // namespace runtime -} // namespace tvm diff --git a/src/runtime/aot_executor/aot_executor.h b/src/runtime/aot_executor/aot_executor.h deleted file mode 100644 index 164deb507830..000000000000 --- a/src/runtime/aot_executor/aot_executor.h +++ /dev/null @@ -1,157 +0,0 @@ -/* - * 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. - */ - -/*! - * \brief Defines an implementation of Module-based Model Runtime Interface that works with - * Ahead-of-Time compilation. - * \file aot_executor.h - */ -#ifndef TVM_RUNTIME_AOT_EXECUTOR_AOT_EXECUTOR_H_ -#define TVM_RUNTIME_AOT_EXECUTOR_AOT_EXECUTOR_H_ - -#include -#include -#include -#include - -#include -#include - -namespace tvm { -namespace runtime { - -class TVM_DLL AotExecutor : public ModuleNode { - public: - /*! - * \brief Implements member function lookup for this Module for the frontend. - * \param name The name of the function. - * \param sptr_to_self The pointer to the module node. - * \return The corresponding member function. - */ - PackedFunc GetFunction(const String& name, const ObjectPtr& sptr_to_self) override; - - /*! - * \return The type key of the executor. - */ - const char* type_key() const final { return "AotExecutor"; } - - /*! \brief Get the property of the runtime module .*/ - int GetPropertyMask() const final { return ModulePropertyMask::kRunnable; } - - void Run(); - - /*! - * \brief Initialize the AOT executor with metadata, runtime::Module, and device. - * \param module The module containing the compiled functions for the host - * processor. - * \param devs A 1-element vector. The Device which AOT compute will run on. Currently, only - * Device(kDLCPU, 0) is supported. - */ - AotExecutor(tvm::runtime::Module module, const std::vector& devs); - - /*! - * \brief Get the input index given the name of input. - * \param name The name of the input. - * \return The index of input. - */ - int GetInputIndex(const std::string& name); - - /*! - * \brief Get the input name given the index of input. - * \param index The index of the input. - * \return The name of input. - */ - std::string GetInputName(int index); - - /*! - * \brief Get the output index given the name of output. - * \param name The name of the output. - * \return The index of output. - */ - int GetOutputIndex(const std::string& name); - - /*! - * \brief set index-th input to the graph. - * \param index The input index. - * \param data_in The input data. - */ - void SetInput(int index, DLTensor* data_in); - /*! - * \brief set index-th input to the graph without copying the data - * \param index The input index. - * \param data_ref The input data that is referred. - */ - void SetInputZeroCopy(int index, DLTensor* data_ref); - /*! - * \brief set index-th output to the graph without copying the data. - * \param index The output index. - * \param data_ref The output data that is referred. - */ - void SetOutputZeroCopy(int index, DLTensor* data_ref); - /*! - * \brief Get the number of outputs - * - * \return The number of outputs from graph. - */ - int NumOutputs() const; - /*! - * \brief Get the number of inputs - * - * \return The number of inputs to the graph. - */ - int NumInputs() const; - /*! - * \brief Return NDArray for given input index. - * \param index The input index. - * - * \return NDArray corresponding to given input node index. - */ - NDArray GetInput(int index) const; - /*! - * \brief Return NDArray for given output index. - * \param index The output index. - * - * \return NDArray corresponding to given output node index. - */ - NDArray GetOutput(int index) const; - /*! - * \brief Copy index-th output to data_out. - * \param index The output index. - * \param data_out the output data. - */ - void CopyOutputTo(int index, DLTensor* data_out); - - private: - /*! \brief Metadata provided to the runtime from the compiler. */ - metadata::Metadata metadata_; - - /*! \brief Runtime module which contains the AOT top-level function. */ - Module module_; - - /*! \brief The devices which should be used to execute the computations. */ - std::vector devices_; - - /*! \brief Holds one NDArray per function argument in the same order. */ - std::vector args_; -}; - -} // namespace runtime -} // namespace tvm - -#endif // TVM_RUNTIME_AOT_EXECUTOR_AOT_EXECUTOR_H_ diff --git a/src/runtime/aot_executor/aot_executor_factory.cc b/src/runtime/aot_executor/aot_executor_factory.cc deleted file mode 100644 index 011e0824fbad..000000000000 --- a/src/runtime/aot_executor/aot_executor_factory.cc +++ /dev/null @@ -1,137 +0,0 @@ -/* - * 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 aot_executor_factory.cc - * \brief AOT executor factory implementations - */ - -#include "./aot_executor_factory.h" - -#include -#include -#include - -#include -#include - -namespace tvm { -namespace runtime { - -AotExecutorFactory::AotExecutorFactory( - const std::unordered_map& params, - const std::string& module_name) { - params_ = params; - module_name_ = module_name; -} - -PackedFunc AotExecutorFactory::GetFunction( - const String& name, const tvm::runtime::ObjectPtr& sptr_to_self) { - if (name == module_name_) { - return PackedFunc([sptr_to_self, this](TVMArgs args, TVMRetValue* rv) { - ICHECK_GT(args.num_args, 0) << "Must supply at least one device argument"; - std::vector devices; - for (int i = 0; i < args.num_args; ++i) { - devices.emplace_back(args[i].operator Device()); - } - *rv = this->ExecutorCreate(devices); - }); - } else if (name == "list_module_names") { - return PackedFunc([sptr_to_self, this](TVMArgs args, TVMRetValue* rv) { - Array names = {module_name_}; - *rv = names; - }); - } else if (name == "remove_params") { - return PackedFunc([sptr_to_self, this](TVMArgs args, TVMRetValue* rv) { - std::unordered_map empty_params{}; - auto exec = make_object(empty_params, this->module_name_); - exec->Import(this->imports_[0]); - *rv = Module(exec); - }); - } else { - return PackedFunc(); - } -} - -void AotExecutorFactory::SaveToBinary(dmlc::Stream* stream) { - std::vector names; - std::vector arrays; - for (const auto& v : params_) { - names.emplace_back(v.first); - arrays.emplace_back(const_cast(v.second.operator->())); - } - uint64_t sz = arrays.size(); - ICHECK(sz == names.size()); - stream->Write(sz); - stream->Write(names); - for (size_t i = 0; i < sz; ++i) { - tvm::runtime::SaveDLTensor(stream, arrays[i]); - } - stream->Write(module_name_); -} - -Module AotExecutorFactory::ExecutorCreate(const std::vector& devs) { - auto exec = make_object(this->imports_[0], devs); - // set params - SetParams(exec.get(), this->params_); - return Module(exec); -} - -Module AotExecutorFactoryModuleLoadBinary(void* strm) { - dmlc::Stream* stream = static_cast(strm); - std::unordered_map params; - std::string module_name; - uint64_t sz; - ICHECK(stream->Read(&sz)); - std::vector names; - ICHECK(stream->Read(&names)); - ICHECK(sz == names.size()); - for (size_t i = 0; i < sz; ++i) { - tvm::runtime::NDArray temp; - temp.Load(stream); - params[names[i]] = temp; - } - ICHECK(stream->Read(&module_name)); - auto exec = make_object(params, module_name); - return Module(exec); -} - -TVM_REGISTER_GLOBAL("tvm.aot_executor_factory.create").set_body([](TVMArgs args, TVMRetValue* rv) { - ICHECK_GE(args.num_args, 2) << "The expected number of arguments for " - "aot_executor_factory.create needs at least 2, " - "but it has " - << args.num_args; - // The argument order is module, module_name, param0_name, param0_tensor, - // [param1_name, param1_tensor], ... - ICHECK_EQ((args.size() - 2) % 2, 0); - std::unordered_map params; - for (size_t i = 2; i < static_cast(args.size()); i += 2) { - std::string name = args[i].operator String(); - params[name] = args[i + 1].operator tvm::runtime::NDArray(); - } - auto exec = make_object(params, args[1]); - exec->Import(args[0]); - *rv = Module(exec); -}); - -TVM_REGISTER_GLOBAL("runtime.module.loadbinary_AotExecutorFactory") - .set_body_typed(AotExecutorFactoryModuleLoadBinary); - -} // namespace runtime -} // namespace tvm diff --git a/src/runtime/aot_executor/aot_executor_factory.h b/src/runtime/aot_executor/aot_executor_factory.h deleted file mode 100644 index 15ac6f5e7f23..000000000000 --- a/src/runtime/aot_executor/aot_executor_factory.h +++ /dev/null @@ -1,122 +0,0 @@ -/* - * 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 tvm/runtime/aot_executor/aot_executor_factory.h - * \brief Aot executor factory creating aot executor. - */ - -#ifndef TVM_RUNTIME_AOT_EXECUTOR_AOT_EXECUTOR_FACTORY_H_ -#define TVM_RUNTIME_AOT_EXECUTOR_AOT_EXECUTOR_FACTORY_H_ - -#include -#include -#include -#include - -#include -#include -#include -#include -#include -#include - -#include "./aot_executor.h" - -namespace tvm { -namespace runtime { - -class TVM_DLL AotExecutorFactory : public runtime::ModuleNode { - public: - /*! - * \brief Construct the AotExecutorFactory. - * \param params The params of aot. - * \param module_name The module name of aot. - */ - AotExecutorFactory(const std::unordered_map& params, - const std::string& module_name); - - /*! - * \brief Get member function to front-end - * \param name The name of the function. - * \param sptr_to_self The pointer to the module node. - * \return The corresponding member function. - */ - PackedFunc GetFunction(const String& name, const ObjectPtr& sptr_to_self) final; - - /*! - * \return The type key of the executor. - */ - const char* type_key() const final { return "AotExecutorFactory"; } - - /*! \brief Get the property of the runtime module .*/ - int GetPropertyMask() const final { return ModulePropertyMask::kBinarySerializable; } - - /*! - * \brief Save the module to binary stream. - * \param stream The binary stream to save to. - */ - void SaveToBinary(dmlc::Stream* stream) override; - - /*! - * \brief Create a specific executor module - * \param devs The device of the host and devices where the model will be - * executed. - * \return created executor module - */ - Module ExecutorCreate(const std::vector& devs); - - /*! - * \brief Set params. - * \param aot_executor The aot executor we want to set the params into. - * \param params The aot params value we want to set. - */ - void SetParams(AotExecutor* aot_executor, - const std::unordered_map& params) const { - std::unordered_map value = params; - // upload big arrays first to avoid memory issue in rpc mode - std::vector keys; - for (const auto& p : value) { - keys.emplace_back(p.first); - } - std::sort(std::begin(keys), std::end(keys), - [&](const std::string& lhs, const std::string& rhs) -> bool { - auto lhs_size = GetDataSize(*value[lhs].operator->()); - auto rhs_size = GetDataSize(*value[rhs].operator->()); - return lhs_size > rhs_size; - }); - for (const auto& key : keys) { - int in_idx = aot_executor->GetInputIndex(key); - if (in_idx >= 0) { - aot_executor->SetInput(in_idx, const_cast(value[key].operator->())); - } - } - } - - protected: - /*! \brief The params. */ - std::unordered_map params_; - /*! \brief module name */ - std::string module_name_; -}; - -} // namespace runtime -} // namespace tvm - -#endif // TVM_RUNTIME_AOT_EXECUTOR_AOT_EXECUTOR_FACTORY_H_ diff --git a/src/runtime/const_loader_module.cc b/src/runtime/const_loader_module.cc index 35da78a83eea..ecedaf6321c9 100644 --- a/src/runtime/const_loader_module.cc +++ b/src/runtime/const_loader_module.cc @@ -249,8 +249,6 @@ Module ConstLoaderModuleCreate( return Module(n); } -TVM_REGISTER_GLOBAL("runtime.module.loadbinary_metadata") - .set_body_typed(ConstLoaderModuleNode::LoadFromBinary); TVM_REGISTER_GLOBAL("runtime.module.loadbinary_const_loader") .set_body_typed(ConstLoaderModuleNode::LoadFromBinary); diff --git a/src/runtime/contrib/amx/amx_config.cc b/src/runtime/contrib/amx/amx_config.cc index 2e034bd478b5..8aae82fbdfa0 100644 --- a/src/runtime/contrib/amx/amx_config.cc +++ b/src/runtime/contrib/amx/amx_config.cc @@ -28,13 +28,13 @@ namespace tvm { namespace runtime { #ifdef __linux__ -#include #include #include #include #include #include #include +#include #include #define XFEATURE_XTILECFG 17 diff --git a/src/runtime/contrib/json/json_node.h b/src/runtime/contrib/json/json_node.h index dd16c606815a..6681e15975f0 100644 --- a/src/runtime/contrib/json/json_node.h +++ b/src/runtime/contrib/json/json_node.h @@ -26,6 +26,7 @@ #define TVM_RUNTIME_CONTRIB_JSON_JSON_NODE_H_ #include +#include #include #include #include @@ -333,14 +334,38 @@ inline bool SameType(const dmlc::any& data) { return std::type_index(data.type()) == std::type_index(typeid(T)); } +template <> +struct Handler> { + inline static void Write(dmlc::JSONWriter* writer, + const std::shared_ptr& data) { + data->Save(writer); + } + + inline static void Read(dmlc::JSONReader* reader, + std::shared_ptr* data) { + (*data)->Load(reader); + } +}; + template <> struct Handler> { inline static void Write(dmlc::JSONWriter* writer, const std::unordered_map& data) { + writer->BeginObject(); for (const auto& kv : data) { auto k = kv.first; const dmlc::any& v = kv.second; - if (SameType>(v)) { + if (SameType(v)) { + writer->WriteObjectKeyValue(k, dmlc::get(v)); + } else if (SameType(v)) { + writer->WriteObjectKeyValue(k, dmlc::get(v)); + } else if (SameType>(v)) { + writer->WriteObjectKeyValue(k, dmlc::get>(v)); + } else if (SameType>>(v)) { + writer->WriteObjectKeyValue(k, dmlc::get>>(v)); + } else if (SameType>(v)) { + writer->WriteObjectKeyValue(k, dmlc::get>(v)); + } else if (SameType>(v)) { writer->WriteObjectKeyValue(k, dmlc::get>(v)); } else { LOG(FATAL) << "Not supported"; @@ -350,20 +375,33 @@ struct Handler> { } inline static void Read(dmlc::JSONReader* reader, std::unordered_map* data) { - LOG(FATAL) << "Not implemented"; + LOG(FATAL) << "Not implemented."; } }; template <> -struct Handler> { - inline static void Write(dmlc::JSONWriter* writer, - const std::shared_ptr& data) { - data->Save(writer); +struct Handler> { + inline static void Write(dmlc::JSONWriter* writer, const std::vector& data) { + writer->BeginArray(); + for (const auto& v : data) { + if (SameType(v)) { + writer->WriteArrayItem(dmlc::get(v)); + } else if (SameType(v)) { + writer->WriteArrayItem(dmlc::get(v)); + } else if (SameType>(v)) { + writer->WriteArrayItem(dmlc::get>(v)); + } else if (SameType>>(v)) { + writer->WriteArrayItem(dmlc::get>>(v)); + } else if (SameType>(v)) { + writer->WriteArrayItem(dmlc::get>(v)); + } else { + LOG(FATAL) << "Not supported"; + } + } + writer->EndArray(); } - - inline static void Read(dmlc::JSONReader* reader, - std::shared_ptr* data) { - (*data)->Load(reader); + inline static void Read(dmlc::JSONReader* reader, std::vector* data) { + LOG(FATAL) << "Not implemented."; } }; } // namespace json diff --git a/src/runtime/graph_executor/cuda_graph/graph_runtime_cuda_graph.cc b/src/runtime/graph_executor/cuda_graph/graph_runtime_cuda_graph.cc deleted file mode 100644 index 5cd331807da7..000000000000 --- a/src/runtime/graph_executor/cuda_graph/graph_runtime_cuda_graph.cc +++ /dev/null @@ -1,136 +0,0 @@ -/* - * 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 graph_executor_cuda_graph.cc - */ - -#include - -#include "../../cuda/cuda_common.h" -#include "../graph_executor.h" - -namespace tvm { -namespace runtime { - -/*! - * \brief Graph executor with CUDA Graph Support. - * - * This is the extension of GraphExecutor class used for CUDA graph launch - * instead of CUDA kernel launch. CUDA graph launch requires CUDA 10.0 or - * above, currently there are two ways of constructing CUDA graphs: - * (1) Using CUDA stream capture API to capture a series of operations on - * CUDA stream, and automatically generates a graph (2) Building a graph - * using CUDA graph API manually. This implementation uses stream capture. - */ -class GraphExecutorCudaGraph : public GraphExecutor { - public: - /*! - * \brief Begin CUDA graph capture on stream, the stream enters capture mode. - */ - void StartCapture() { - const Device& dev = data_entry_[entry_id(0, 0)]->device; - - TVMStreamCreate(dev.device_type, dev.device_id, &capture_stream_); - TVMSetStream(dev.device_type, dev.device_id, capture_stream_); - - CUDA_CALL(cudaStreamBeginCapture(static_cast(capture_stream_), - cudaStreamCaptureModeGlobal)); - } - - /*! - * \brief Launch the instantiated graph on stream - */ - void RunCudaGraph() { - cudaStream_t cuStream = static_cast(capture_stream_); - CUDA_CALL(cudaGraphLaunch(cuda_graph_exec_, cuStream)); - CUDA_CALL(cudaStreamSynchronize(cuStream)); - } - - /*! - * \brief End CUDA graph capture on stream, a graph will be created and - * instantiated. - */ - void EndCapture() { - cudaGraph_t graph; - CUDA_CALL(cudaStreamEndCapture(static_cast(capture_stream_), &graph)); - - cudaGraphNode_t* nodes = NULL; - size_t numNodes = 0; - CUDA_CALL(cudaGraphGetNodes(graph, nodes, &numNodes)); - LOG(INFO) << "Num of nodes in the cuda graph created using stream capture API = " << numNodes; - - CUDA_CALL(cudaGraphInstantiate(&cuda_graph_exec_, graph, NULL, NULL, 0)); - } - - /*! - * \brief GetFunction Get the function based on input. - * \param name The function which needs to be invoked. - * \param sptr_to_self Packed function pointer. - */ - PackedFunc GetFunction(const String& name, const ObjectPtr& sptr_to_self); - - private: - /*! \brief The Cuda stream on which to capture a CUDA graph. */ - TVMStreamHandle capture_stream_; - /*! \brief The captured CUDA graph will be instantiated to this. */ - cudaGraphExec_t cuda_graph_exec_; -}; - -PackedFunc GraphExecutorCudaGraph::GetFunction(const String& name, - const ObjectPtr& sptr_to_self) { - if (name == "run_cuda_graph") { - return PackedFunc( - [sptr_to_self, this](TVMArgs args, TVMRetValue* rv) { this->RunCudaGraph(); }); - } else if (name == "start_capture") { - return PackedFunc( - [sptr_to_self, this](TVMArgs args, TVMRetValue* rv) { this->StartCapture(); }); - } else if (name == "end_capture") { - return PackedFunc([sptr_to_self, this](TVMArgs args, TVMRetValue* rv) { this->EndCapture(); }); - } else { - return GraphExecutor::GetFunction(name, sptr_to_self); - } -} - -Module GraphExecutorCudaGraphCreate(const std::string& sym_json, const tvm::runtime::Module& m, - const std::vector& devs, - PackedFunc lookup_linked_param_func) { - auto exec = make_object(); - exec->Init(sym_json, m, devs, lookup_linked_param_func); - return Module(exec); -} - -TVM_REGISTER_GLOBAL("tvm.graph_executor_cuda_graph.create") - .set_body([](TVMArgs args, TVMRetValue* rv) { - ICHECK_GE(args.num_args, 4) - << "The expected number of arguments for graph_executor.create is " - "at least 4, but it has " - << args.num_args; - PackedFunc lookup_linked_param_func; - int dev_start_arg = 2; - if (args[2].type_code() == kTVMPackedFuncHandle) { - lookup_linked_param_func = args[2]; - dev_start_arg++; - } - - *rv = GraphExecutorCudaGraphCreate(args[0], args[1], GetAllDevice(args, dev_start_arg), - lookup_linked_param_func); - }); -} // namespace runtime -} // namespace tvm diff --git a/src/runtime/graph_executor/debug/graph_executor_debug.cc b/src/runtime/graph_executor/debug/graph_executor_debug.cc deleted file mode 100644 index a9cd4d544d3b..000000000000 --- a/src/runtime/graph_executor/debug/graph_executor_debug.cc +++ /dev/null @@ -1,461 +0,0 @@ -/* - * 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 graph_executor_debug.cc - */ -#include "./graph_executor_debug.h" - -#include -#include -#include -#include -#include - -#include -#include -#include -#include - -#include "../../rpc/rpc_session.h" - -namespace tvm { -namespace runtime { -std::string GraphExecutorDebug::RunIndividual(int number, int repeat, int min_repeat_ms, - int limit_zero_time_iterations, - int cooldown_interval_ms, int repeats_to_cooldown) { - // warmup run - GraphExecutor::Run(); - std::string tkey = module_->type_key(); - std::vector> time_sec_per_op(op_execs_.size()); - if (tkey == "rpc") { - // RPC modules rely on remote timing which implements the logic from the else branch. - for (size_t index = 0; index < op_execs_.size(); ++index) { - time_sec_per_op[index] = - RunOpRPC(index, number, repeat, min_repeat_ms, limit_zero_time_iterations, - cooldown_interval_ms, repeats_to_cooldown); - } - } else { - int op = 0; - for (size_t index = 0; index < op_execs_.size(); ++index) { - std::string result_str = - RunIndividualNode(index, number, repeat, min_repeat_ms, limit_zero_time_iterations, - cooldown_interval_ms, repeats_to_cooldown); - const double* blob_ptr = reinterpret_cast(result_str.data()); - for (int i = 0; i < repeat; ++i, ++blob_ptr) { - time_sec_per_op[index].push_back(*blob_ptr); - } - if (op_execs_[index]) { - LOG(INFO) << "Op #" << op << " " << GetNodeName(index) << ":"; - for (size_t cur_repeat = 0; cur_repeat < time_sec_per_op[index].size(); cur_repeat++) { - const auto& data = time_sec_per_op[index][cur_repeat]; - LOG(INFO) << "Iteration: " << cur_repeat << ": " << (data * 1e6) << " us/iter"; - } - ++op; - } - } - } - - std::ostringstream os; - int64_t size = time_sec_per_op.size(); - os.write(reinterpret_cast(&size), sizeof(int64_t)); - for (size_t index = 0; index < time_sec_per_op.size(); ++index) { - for (auto& repeat_data : time_sec_per_op[index]) { - // To have good behavior when calculating total time, etc. - double data = std::isnan(repeat_data) ? 0 : repeat_data; - os.write(reinterpret_cast(&data), sizeof(double)); - } - } - return os.str(); -} - -std::string GraphExecutorDebug::RunIndividualNode(int node_index, int number, int repeat, - int min_repeat_ms, int limit_zero_time_iterations, - int cooldown_interval_ms, - int repeats_to_cooldown) { - std::string tkey = module_->type_key(); - - if (tkey == "rpc") { - LOG(FATAL) << "RPC measurements should not use RunIndividualNode!"; - } - - if (!op_execs_[node_index]) { - // don't return anything... - std::ostringstream os; - double zero = 0; - for (int i = 0; i < repeat; ++i) { - os.write(reinterpret_cast(&zero), sizeof(double)); - } - return os.str(); - } - - // assume host runs things which is first device - Device& d = devices_[0]; - PackedFunc time_evaluator = profiling::WrapTimeEvaluator( - TypedPackedFunc([this, node_index]() { this->RunOpHost(node_index); }), d, number, - repeat, min_repeat_ms, limit_zero_time_iterations, cooldown_interval_ms, repeats_to_cooldown); - return time_evaluator(); -} - -std::vector GraphExecutorDebug::RunOpRPC(int index, int number, int repeat, - int min_repeat_ms, int limit_zero_time_iterations, - int cooldown_interval_ms, - int repeats_to_cooldown) { - std::vector results(repeat, 0); - // Right now we expect either "tvm_op" for nodes which run PackedFunc or "null" for nodes - // which represent inputs/parameters to the graph. Other types may be supported in the - // future, but consideration would be needed as to how to do that over RPC before we support - // it here. - if (nodes_[index].op_type != "tvm_op") { - CHECK_EQ(nodes_[index].op_type, "null") - << "Don't know how to run op type " << nodes_[index].op_type - << " remotely over RPC right now"; - - // NOTE: GraphExecutorDebug expects graph nodes to have an "op" attribute of "tvm_op" or - // "null" and "null" is a placeholder node for a parameter or input. - return results; - } - - const Device& dev = data_entry_[entry_id(index, 0)]->device; - TVMOpParam param = nodes_[index].param; - std::string name = param.func_name; - uint32_t num_inputs = param.num_inputs; - uint32_t num_outputs = param.num_outputs; - - PackedFunc time_eval = - runtime::Registry::Get("runtime.RPCTimeEvaluator") - -> - operator()(module_, name, static_cast(dev.device_type), dev.device_id, number, - repeat, min_repeat_ms, limit_zero_time_iterations, cooldown_interval_ms, - repeats_to_cooldown, /*cache_flush_bytes=*/0, ""); - - int num_flat_args = num_inputs + num_outputs; - auto values = std::make_unique(num_flat_args); - auto type_codes = std::make_unique(num_flat_args); - TVMArgsSetter setter(values.get(), type_codes.get()); - int offs = 0; - const auto& inode = nodes_[index]; - for (const auto& e : inode.inputs) { - uint32_t eid = this->entry_id(e); - DLTensor* arg = const_cast(data_entry_[eid].operator->()); - setter(offs, arg); - offs++; - } - for (uint32_t i = 0; i < num_outputs; ++i) { - uint32_t eid = this->entry_id(index, i); - DLTensor* arg = const_cast(data_entry_[eid].operator->()); - setter(offs, arg); - offs++; - } - TVMRetValue rv; - time_eval.CallPacked(TVMArgs(values.get(), type_codes.get(), num_flat_args), &rv); - std::string results_str = rv.operator std::string(); - const double* blob_ptr = reinterpret_cast(results_str.data()); - for (int i = 0; i < repeat; ++i, ++blob_ptr) { - results[i] = *blob_ptr; - } - - std::ostringstream os; - for (auto& repeat_data : results) { - os << std::to_string(repeat_data) << ", "; - } - LOG(INFO) << "Got op timing: " << os.str(); - return results; -} - -Timer GraphExecutorDebug::RunOpHost(int index) { - const Device& dev = data_entry_[entry_id(index, 0)]->device; - Timer t = Timer::Start(dev); - op_execs_[index](); - t->Stop(); - return t; -} - -/*! - * \brief GetFunction Get the function based on input. - * \param name The function which needs to be invoked. - * \param sptr_to_self Packed function pointer. - */ -PackedFunc GraphExecutorDebug::GetFunction(const String& name, - const ObjectPtr& sptr_to_self) { - // return member functions during query. - if (name == "debug_get_output") { - return PackedFunc([sptr_to_self, this](TVMArgs args, TVMRetValue* rv) { - int args0 = -1; - if (String::CanConvertFrom(args[0])) { - args0 = this->GetNodeIndex(args[0]); - } else { - args0 = args[0]; - } - - if (args.num_args == 2) { - this->DebugGetNodeOutput(args0, args[1]); - } else { - *rv = this->DebugGetNodeOutput(args0); - } - }); - } else if (name == "execute_node") { - return PackedFunc( - [sptr_to_self, this](TVMArgs args, TVMRetValue* rv) { this->ExecuteNode(args[0]); }); - } else if (name == "debug_run_ext_compiler") { - return PackedFunc( - [sptr_to_self, this](TVMArgs args, TVMRetValue* rv) { *rv = this->DebugRunExtCompiler(); }); - } else if (name == "get_node_output") { - return PackedFunc([sptr_to_self, this](TVMArgs args, TVMRetValue* rv) { - *rv = this->GetNodeOutput(args[0], args[1]); - }); - } else if (name == "run_individual") { - return PackedFunc([sptr_to_self, this](TVMArgs args, TVMRetValue* rv) { - int number = args[0]; - int repeat = args[1]; - int min_repeat_ms = args[2]; - int limit_zero_time_iterations = args[3]; - int cooldown_interval_ms = args[4]; - int repeats_to_cooldown = args[5]; - ICHECK_GT(number, 0); - ICHECK_GT(repeat, 0); - ICHECK_GE(min_repeat_ms, 0); - ICHECK_GE(limit_zero_time_iterations, 0); - ICHECK_GE(cooldown_interval_ms, 0); - ICHECK_GT(repeats_to_cooldown, 0); - std::string blob = - this->RunIndividual(number, repeat, min_repeat_ms, limit_zero_time_iterations, - cooldown_interval_ms, repeats_to_cooldown); - TVMByteArray arr; - arr.size = blob.length(); - arr.data = blob.data(); - *rv = arr; - }); - } else if (name == "run_individual_node") { - return PackedFunc([sptr_to_self, this](TVMArgs args, TVMRetValue* rv) { - int node_index = args[0]; - int number = args[1]; - int repeat = args[2]; - int min_repeat_ms = args[3]; - int limit_zero_time_iterations = args[4]; - int cooldown_interval_ms = args[5]; - int repeats_to_cooldown = args[6]; - ICHECK_GE(node_index, 0); - ICHECK_LT(node_index, nodes_.size()); - ICHECK_GT(number, 0); - ICHECK_GT(repeat, 0); - ICHECK_GE(min_repeat_ms, 0); - ICHECK_GE(limit_zero_time_iterations, 0); - ICHECK_GE(cooldown_interval_ms, 0); - ICHECK_GT(repeats_to_cooldown, 0); - std::string blob = this->RunIndividualNode(node_index, number, repeat, min_repeat_ms, - limit_zero_time_iterations, cooldown_interval_ms, - repeats_to_cooldown); - TVMByteArray arr; - arr.size = blob.length(); - arr.data = blob.data(); - *rv = arr; - }); - } else if (name == "profile") { - return TypedPackedFunc)>( - [sptr_to_self, this](Array collectors) { - // We cannot send Arrays over rpc, so in order to support profiling - // on remotes, we accept a nullptr for collectors. - if (collectors.defined()) { - return this->Profile(collectors); - } else { - return this->Profile({}); - } - }); - } else if (name == "profile_rpc") { - // We cannot return a Report over RPC because TMV RPC mechanism only - // supports a subset of Object classes. Instead we serialize it on the - // remote (here) and deserialize it on the other end. - return TypedPackedFunc([sptr_to_self, this]() { - PackedFunc profile = GetFunction("profile", sptr_to_self); - profiling::Report report = profile(Array()); - return report->AsJSON(); - }); - } else { - return GraphExecutor::GetFunction(name, sptr_to_self); - } -} - -int GraphExecutorDebug::GetNodeIndex(const std::string& name) const { - for (size_t nid = 0; nid < GetNumOfNodes(); ++nid) { - if (GetNodeName(nid) == name) { - return static_cast(nid); - } - } - LOG(FATAL) << "cannot find " << name << " among nodex"; - return -1; -} - -void GraphExecutorDebug::ExecuteNode(int node) { - ICHECK_LT(static_cast(node), op_execs_.size()); - - int start_ind; - int end_ind; - if (node < last_executed_node_) { - start_ind = 0; - end_ind = node; - } else if (node > last_executed_node_) { - start_ind = last_executed_node_ + 1; - end_ind = node; - } else { - return; - } - - for (int i = start_ind; i <= end_ind; i++) { - if (op_execs_[i]) op_execs_[i](); - } - last_executed_node_ = end_ind; -} - -std::string GraphExecutorDebug::DebugRunExtCompiler(void) { - std::ostringstream os; - dmlc::JSONWriter writer(&os); - writer.BeginArray(); - for (size_t i = 0; i < op_execs_.size(); ++i) { - if (!nodes_[i].param.compiler.empty() && op_profile_execs_[i]) { - TVMRetValue rv; - rv = String("debug_dump"); - this->op_profile_execs_[i](&rv); - std::string debug_ret = rv; - - writer.BeginObject(); - writer.WriteObjectKeyValue("compiler", nodes_[i].param.compiler); - writer.WriteObjectKeyValue("op", nodes_[i].param.func_name); - writer.WriteObjectKeyValue("dump", debug_ret); - writer.EndObject(); - } else { - if (op_execs_[i]) op_execs_[i](); - } - } - writer.EndArray(); - - return os.str(); -} - -void GraphExecutorDebug::DebugGetNodeOutput(int index, DLTensor* data_out) { - ICHECK_LT(static_cast(index), op_execs_.size()); - uint32_t eid = index; - - for (size_t i = 0; i < op_execs_.size(); ++i) { - if (op_execs_[i]) op_execs_[i](); - if (static_cast(i) == index) break; - } - - data_entry_[eid].CopyTo(data_out); -} - -NDArray GraphExecutorDebug::DebugGetNodeOutput(int index) { - ICHECK_LT(static_cast(index), op_execs_.size()); - uint32_t eid = index; - - for (size_t i = 0; i < op_execs_.size(); ++i) { - if (op_execs_[i]) op_execs_[i](); - if (static_cast(i) == index) break; - } - - return data_entry_[eid]; -} - -NDArray GraphExecutorDebug::GetNodeOutput(int node, int out_ind) { - ICHECK_EQ(node, last_executed_node_); - ICHECK_LT(entry_id(node, out_ind), data_entry_.size()); - return data_entry_[entry_id(node, out_ind)].CopyTo({kDLCPU, 0}); -} - -profiling::Report GraphExecutorDebug::Profile(Array collectors) { - std::vector cs(collectors.begin(), collectors.end()); - profiling::Profiler prof(devices_, cs, {{String("Executor"), String("Graph")}}); - - // warm up. 1 iteration does not seem enough. - for (int i = 0; i < 3; i++) { - GraphExecutor::Run(); - } - - prof.Start(); - for (size_t i = 0; i < op_execs_.size(); ++i) { - if (op_execs_[i]) { - // get argument shapes - std::vector shapes; - for (const auto& e : nodes_[i].inputs) { - uint32_t eid = entry_id(e); - shapes.push_back(data_entry_[eid]); - } - for (uint32_t j = 0; j < nodes_[i].param.num_outputs; ++j) { - uint32_t eid = entry_id(i, j); - shapes.push_back(data_entry_[eid]); - } - - uint32_t eid = entry_id(i, 0); - const Device& dev = data_entry_[eid]->device; - - std::unordered_map metrics; - for (auto p : nodes_[i].param.attrs) { - if (std::string(p.first).find("layout") != std::string::npos) { - metrics[p.first] = p.second; - } - } - if (nodes_[i].param.attrs.find("hash") != nodes_[i].param.attrs.end()) { - metrics["Hash"] = Downcast(nodes_[i].param.attrs.at("hash")); - } - metrics["Argument Shapes"] = profiling::ShapeString(shapes); - if (!nodes_[i].param.compiler.empty() && op_profile_execs_[i]) { - TVMRetValue rv; - rv = static_cast(&prof); - this->op_profile_execs_[i](&rv); - } else { - prof.StartCall(nodes_[i].param.func_name, dev, metrics); - op_execs_[i](); - prof.StopCall(); - } - } - } - prof.Stop(); - return prof.Report(); -} - -/*! - * \brief GraphExecutorDebugCreate Get the function based on input. - * \param sym_json The graph symbol in json format. - * \param m Compiled module which will be loaded. - * \param devs All devices. - */ -Module GraphExecutorDebugCreate(const std::string& sym_json, const tvm::runtime::Module& m, - const std::vector& devs, - PackedFunc lookup_linked_param_func) { - auto exec = make_object(); - exec->Init(sym_json, m, devs, lookup_linked_param_func); - return Module(exec); -} - -TVM_REGISTER_GLOBAL("tvm.graph_executor_debug.create").set_body([](TVMArgs args, TVMRetValue* rv) { - ICHECK_GE(args.num_args, 4) << "The expected number of arguments for graph_executor.create is " - "at least 4, but it has " - << args.num_args; - PackedFunc lookup_linked_param_func; - int dev_start_arg = 2; - if (args[2].type_code() == kTVMPackedFuncHandle) { - lookup_linked_param_func = args[2]; - dev_start_arg++; - } - - *rv = GraphExecutorDebugCreate(args[0], args[1], GetAllDevice(args, dev_start_arg), - lookup_linked_param_func); -}); -} // namespace runtime -} // namespace tvm diff --git a/src/runtime/graph_executor/debug/graph_executor_debug.h b/src/runtime/graph_executor/debug/graph_executor_debug.h deleted file mode 100644 index 8ede2a3a5f84..000000000000 --- a/src/runtime/graph_executor/debug/graph_executor_debug.h +++ /dev/null @@ -1,168 +0,0 @@ -/* - * 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. - */ - -#ifndef TVM_RUNTIME_GRAPH_EXECUTOR_DEBUG_GRAPH_EXECUTOR_DEBUG_H_ -#define TVM_RUNTIME_GRAPH_EXECUTOR_DEBUG_GRAPH_EXECUTOR_DEBUG_H_ - -#include - -#include -#include - -#include "../graph_executor.h" - -namespace tvm { -namespace runtime { - -/*! - * \brief Graph executor with debug . - * - * This is the extension of GraphExecutor class used for debugging - * TVM runtime PackedFunc API. - */ -class GraphExecutorDebug : public GraphExecutor { - public: - /*! - * \brief Run each operation in the graph and get the time per op for all ops. - * \param number The number of times to run this function for taking average. - * \param repeat The number of times to repeat the measurement. - * In total, the function will be invoked (1 + number x repeat) times, - * where the first one is warmed up and will be discarded in case - * there is lazy initialization. - * \param min_repeat_ms The minimum duration of one `repeat` in milliseconds. - * By default, one `repeat` contains `number` runs. If this parameter is set, - * the parameters `number` will be dynamically adjusted to meet the - * minimum duration requirement of one `repeat`. - * \param limit_zero_time_iterations The maximum number of repeats when - * measured time is equal to 0. It helps to avoid hanging during - * measurements. - * \param cooldown_interval_ms The cooldown interval in milliseconds between the number of repeats - * defined by `repeats_to_cooldown`. - * \param repeats_to_cooldown The number of repeats before the - * cooldown is activated. - * \return Returns a string with an encoded byte array. Where the first 8 bytes are int64_t - * representing the number of layers. Next the encoded real numbers are float32_t in the number of - * repeat multiplied by the number of layers. - */ - std::string RunIndividual(int number, int repeat, int min_repeat_ms, - int limit_zero_time_iterations, int cooldown_interval_ms, - int repeats_to_cooldown); - - std::string RunIndividualNode(int node_index, int number, int repeat, int min_repeat_ms, - int limit_zero_time_iterations, int cooldown_interval_ms, - int repeats_to_cooldown); - - std::vector RunOpRPC(int index, int number, int repeat, int min_repeat_ms, - int limit_zero_time_iterations, int cooldown_interval_ms, - int repeats_to_cooldown); - - Timer RunOpHost(int index); - - /*! - * \brief GetFunction Get the function based on input. - * \param name The function which needs to be invoked. - * \param sptr_to_self Packed function pointer. - */ - PackedFunc GetFunction(const String& name, const ObjectPtr& sptr_to_self); - - /*! - * \brief Get the node index given the name of node. - * \param name The name of the node. - * \return The index of node. - */ - int GetNodeIndex(const std::string& name) const; - - /*! - * \brief Execute index-th node in the network. - * - * This method will do a partial run of the graph - * up to index-th node. - * - * \param node: The index of the node. - */ - void ExecuteNode(int node); - - /*! - * \brief debug external comilers if supported. - * - * This method invokes the external compilers to generate any debug trace info. - * - * \return Returns serialized debug trace information to the caller - */ - std::string DebugRunExtCompiler(void); - - /*! - * \brief Returns index-th output of node. - * - * This method will return index-th out_ind output - * of index-th node in the network. - * - * \param node: The index of the node. - * \param out_ind: The index of the output. - * \return Output array. - */ - NDArray GetNodeOutput(int node, int out_ind); - - /*! - * \brief Copy index-th node to data_out. - * - * This method will do a partial run of the graph - * from begining upto the index-th node and return output of index-th node. - * This is costly operation and suggest to use only for debug porpose. - * - * \param index: The index of the node. - * \param data_out the node data. - */ - void DebugGetNodeOutput(int index, DLTensor* data_out); - - /*! - * \brief return output of index-th node. - * - * This method will do a partial run of the graph - * from begining up to the index-th node and return output of index-th node. - * This is costly operation and suggest to use only for debug porpose. - * - * \param index: The index of the node. - * - */ - NDArray DebugGetNodeOutput(int index); - - /*! - * \brief Profile execution time of the module. - * - * We run the entire module while recording overall and per-op timing - * information. The module may be run multiple times to ensure everything is - * warmed up. This function is a more correct reflection of actual runtime of - * the module compared to GraphRuntimeDebug::RunIndividual as it runs the - * entire graph in order. - * - * \param collectors Optional user defined `MetricCollector`s to use with this profiling run. - * - * \returns A table of per-op runtimes and total times. - */ - profiling::Report Profile(Array collectors); - - private: - int last_executed_node_ = -1; -}; - -} // namespace runtime -} // namespace tvm - -#endif // TVM_RUNTIME_GRAPH_EXECUTOR_DEBUG_GRAPH_EXECUTOR_DEBUG_H_ diff --git a/src/runtime/graph_executor/graph_executor.cc b/src/runtime/graph_executor/graph_executor.cc deleted file mode 100644 index 3cc3ea396e17..000000000000 --- a/src/runtime/graph_executor/graph_executor.cc +++ /dev/null @@ -1,832 +0,0 @@ -/* - * 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 graph_executor.cc - */ -#include "graph_executor.h" - -#include -#include -#include -#include -#include -#include -#include -#include -#include -#include - -#include -#include -#include -#include -#include -#include -#include -#include - -#include "../file_utils.h" -#include "../texture.h" - -namespace tvm { -namespace runtime { -namespace details { -inline size_t GetDataAlignment(const DLTensor& arr) { - size_t align = (arr.dtype.bits / 8) * arr.dtype.lanes; - if (align < kAllocAlignment) return kAllocAlignment; - return align; -} -constexpr auto Is2DStorage = IsTextureStorage; -} // namespace details - -/*! - * \brief Run all the operations one by one. - */ -void GraphExecutor::Run() { - // setup the array and requirements. - for (size_t i = 0; i < op_execs_.size(); ++i) { - if (op_execs_[i]) op_execs_[i](); - } -} - -/*! - * \brief Initialize the graph executor with graph and device. - * \param graph_json The execution graph. - * \param module The module containing the compiled functions for the host - * processor. - * \param devs The devices of the host and devices where graph nodes will be - * executed on. - * \param lookup_linked_param_func Linked parameter lookup function. Default is nullptr. - */ -void GraphExecutor::Init(const std::string& graph_json, tvm::runtime::Module module, - const std::vector& devs, - const PackedFunc lookup_linked_param_func) { - std::istringstream is(graph_json); - dmlc::JSONReader reader(&is); - this->Load(&reader); - module_ = module; - devices_ = devs; - lookup_linked_param_ = lookup_linked_param_func; - if (lookup_linked_param_ == nullptr) { - lookup_linked_param_ = PackedFunc( - [this](TVMArgs args, TVMRetValue* rv) { this->DefaultLookupLinkedParam(args, rv); }); - } - this->SetupStorage(); - this->SetupOpExecs(); - for (size_t i = 0; i < input_nodes_.size(); i++) { - const uint32_t nid = input_nodes_[i]; - std::string& name = nodes_[nid].name; - input_map_[name] = i; - } - for (size_t i = 0; i < outputs_.size(); i++) { - const uint32_t nid = outputs_[i].node_id; - std::string& name = nodes_[nid].name; - std::stringstream ss; - ss << name << ":" << i; - output_map_[ss.str()] = i; - } -} - -/*! - * \brief Get the input index given the name of input. - * \param name The name of the input. - * \return The index of input. - */ -int GraphExecutor::GetInputIndex(const std::string& name) { - auto it = input_map_.find(name); - if (it != input_map_.end()) { - return it->second; - } - return -1; -} - -/*! - * \brief Get the input info of Graph by parsing the input nodes. - * \return The shape and dtype tuple. - */ -std::tuple GraphExecutor::GetInputInfo() const { - GraphExecutor::ShapeInfo shape_dict; - GraphExecutor::DtypeInfo dtype_dict; - for (uint32_t nid : input_nodes_) { - CHECK_LE(nid, nodes_.size()); - std::string name = nodes_[nid].name; - if (param_names_.find(name) == param_names_.end()) { - CHECK_LE(nid, attrs_.shape.size()); - auto shape = attrs_.shape[nid]; - shape_dict.Set(name, ShapeTuple(shape)); - CHECK_LE(nid, attrs_.dltype.size()); - auto dtype = attrs_.dltype[nid]; - dtype_dict.Set(name, String(dtype)); - } - } - return std::make_tuple(shape_dict, dtype_dict); -} - -/*! - * \brief Get the output info of Graph by parsing the output nodes. - * \return The shape and dtype tuple. - */ -std::tuple GraphExecutor::GetOutputInfo() - const { - GraphExecutor::ShapeInfo shape_dict; - GraphExecutor::DtypeInfo dtype_dict; - for (auto out : outputs_) { - uint32_t nid = out.node_id; - CHECK_LE(nid, nodes_.size()); - std::string name = nodes_[nid].name; - CHECK_LE(nid, attrs_.shape.size()); - auto shape = attrs_.shape[nid]; - shape_dict.Set(name, ShapeTuple(shape)); - CHECK_LE(nid, attrs_.dltype.size()); - auto dtype = attrs_.dltype[nid]; - dtype_dict.Set(name, String(dtype)); - } - return std::make_tuple(shape_dict, dtype_dict); -} - -/*! - * \brief Get the output index given the name of output. - * \param name The name of the output. - * \return The index of output. - */ -int GraphExecutor::GetOutputIndex(const std::string& name) { - auto it = output_map_.find(name); - if (it != output_map_.end()) { - return it->second; - } - return -1; -} -/*! - * \brief set index-th input to the graph. - * \param index The input index. - * \param data_in The input data. - */ -void GraphExecutor::SetInput(int index, DLTensor* data_in) { - ICHECK_LT(static_cast(index), input_nodes_.size()); - uint32_t eid = this->entry_id(input_nodes_[index], 0); - data_entry_[eid].CopyFrom(data_in); -} -/*! - * \brief Check the legality of external DLTensor*. - * \param external The external DLTensor*. - * \param eid The data_enrty_ index. - */ -void GraphExecutor::CheckExternalDLTensor(const DLTensor* external, uint32_t eid) const { - const DLTensor* internal = data_entry_[eid].operator->(); - - ICHECK_EQ(data_alignment_[eid], details::GetDataAlignment(*external)); - ICHECK_EQ(reinterpret_cast(static_cast(external->data) + external->byte_offset) % - kAllocAlignment, - 0); - ICHECK_EQ(internal->ndim, static_cast(external->ndim)); - ICHECK_EQ(internal->device.device_type, external->device.device_type); - ICHECK_EQ(internal->device.device_id, external->device.device_id); - for (auto i = 0; i < external->ndim; ++i) { - ICHECK_EQ(internal->shape[i], external->shape[i]); - } -} -/*! - * \brief set index-th input to the graph without copying the data. - * \param index The input index. - * \param data_ref The input data that is referred. - */ -void GraphExecutor::SetInputZeroCopy(int index, DLTensor* data_ref) { - ICHECK_LT(static_cast(index), input_nodes_.size()); - uint32_t eid = this->entry_id(input_nodes_[index], 0); - // check the consistency of input - CheckExternalDLTensor(data_ref, eid); - // Update the data pointer for each argument of each op - for (DLTensor* t : input_dltensors_[eid]) { - t->data = static_cast(data_ref->data) + data_ref->byte_offset; - } -} -/*! - * \brief set index-th output to the graph without copying the data. - * \param index The output index. - * \param data_ref The output data that is referred. - */ -void GraphExecutor::SetOutputZeroCopy(int index, DLTensor* data_ref) { - ICHECK_LT(static_cast(index), outputs_.size()); - ICHECK_LT(static_cast(index), output_dltensors_.size()); - const NodeEntry& output_node = outputs_[index]; - uint32_t output_node_eid = this->entry_id(output_node); - - // check the consistency of output - CheckExternalDLTensor(data_ref, output_node_eid); - - if (nodes_[output_node.node_id].op_type == "tvm_op" && - nodes_[output_node.node_id].param.func_name == "__nop") { - const NodeEntry& input_node = nodes_[output_node.node_id].inputs[0]; - output_node_eid = this->entry_id(input_node); - ICHECK_NE(node_output_dltensors_[output_node_eid].size(), 0); - for (DLTensor* t : node_output_dltensors_[output_node_eid]) { - t->data = static_cast(data_ref->data) + data_ref->byte_offset; - } - } - - // Update the data pointer for output op - for (DLTensor* t : output_dltensors_[output_node_eid]) { - t->data = static_cast(data_ref->data) + data_ref->byte_offset; - } - - // Update the input of the op connected to the output - for (DLTensor* t : both_output_opinput_dltensors_[output_node_eid]) { - t->data = static_cast(data_ref->data) + data_ref->byte_offset; - } -} -/*! - * \brief Get the number of outputs - * - * \return The number of outputs from graph. - */ -int GraphExecutor::NumOutputs() const { return outputs_.size(); } -/*! - * \brief Get the number of inputs - * - * \return The number of inputs to the graph. - */ -int GraphExecutor::NumInputs() const { return input_nodes_.size(); } -/*! - * \brief Return NDArray for given input index. - * \param index The input index. - * - * \return NDArray corresponding to given input node index. - */ -NDArray GraphExecutor::GetInput(int index) const { - ICHECK_LT(static_cast(index), input_nodes_.size()); - uint32_t eid = this->entry_id(input_nodes_[index], 0); - return data_entry_[eid]; -} -/*! - * \brief Return NDArray for given output index. - * \param index The output index. - * - * \return NDArray corresponding to given output node index. - */ -NDArray GraphExecutor::GetOutput(int index) const { - ICHECK_LT(static_cast(index), outputs_.size()); - uint32_t eid = this->entry_id(outputs_[index]); - return data_entry_[eid]; -} -/*! - * \brief Copy index-th output to data_out. - * \param index The output index. - * \param data_out the output data. - */ -void GraphExecutor::CopyOutputTo(int index, DLTensor* data_out) { - ICHECK_LT(static_cast(index), outputs_.size()); - uint32_t eid = this->entry_id(outputs_[index]); - - // Check the shapes to avoid receiving in different dimension but same size. - const NDArray& data = data_entry_[eid]; - ICHECK_EQ(data->ndim, data_out->ndim); - for (int32_t j = 0; j < data->ndim; ++j) { - ICHECK_EQ(data->shape[j], data_out->shape[j]); - } - - data_entry_[eid].CopyTo(data_out); -} - -/*! - * \brief Load parameters from parameter blob. - * \param param_blob A binary blob of parameter. - */ -void GraphExecutor::LoadParams(const std::string& param_blob) { - dmlc::MemoryStringStream strm(const_cast(¶m_blob)); - this->LoadParams(&strm); -} - -void GraphExecutor::LoadParams(dmlc::Stream* strm) { - Map params = ::tvm::runtime::LoadParams(strm); - for (auto& p : params) { - param_names_.insert(p.first); - int in_idx = GetInputIndex(p.first); - if (in_idx < 0) continue; - uint32_t eid = this->entry_id(input_nodes_[in_idx], 0); - data_entry_[eid].CopyFrom(p.second); - } -} - -void GraphExecutor::ShareParams(const GraphExecutor& other, dmlc::Stream* strm) { - uint64_t header, reserved; - ICHECK(strm->Read(&header)) << "Invalid parameters file format"; - ICHECK(header == kTVMNDArrayListMagic) << "Invalid parameters file format"; - ICHECK(strm->Read(&reserved)) << "Invalid parameters file format"; - std::vector names; - ICHECK(strm->Read(&names)) << "Invalid parameters file format"; - uint64_t sz; - strm->Read(&sz); - size_t size = static_cast(sz); - ICHECK(size == names.size()) << "Invalid parameters file format"; - for (size_t i = 0; i < size; ++i) { - int in_idx = GetInputIndex(names[i]); - if (in_idx < 0) continue; - uint32_t eid = this->entry_id(input_nodes_[in_idx], 0); - ICHECK_LT(eid, data_entry_.size()); - ICHECK_EQ(data_entry_[eid].use_count(), 1); - data_entry_[eid] = other.GetInput(GetInputIndex(names[i])); - ICHECK_GT(data_entry_[eid].use_count(), 1); - const DLTensor* tmp = data_entry_[eid].operator->(); - data_alignment_[eid] = details::GetDataAlignment(*tmp); - } - this->SetupOpExecs(); -} - -void GraphExecutor::LinkedNDArrayDeleter(Object* container) { - // container is the NDArray::Container which needs to get deleted. - // The data member points to global const memory, so it does not need deleting. - delete static_cast(container); -} - -void GraphExecutor::DefaultLookupLinkedParam(TVMArgs args, TVMRetValue* rv) { - Module mod = args[0]; - int64_t storage_id = args[1]; - DLTensor* template_tensor = args[2]; - Device dev = args[3]; - // Get pre-linked parameter lookup function, if it was generated. When pf == nullptr, no linked - // params are present. - if (!module_lookup_linked_param_valid_) { - module_lookup_linked_param_ = - mod.GetFunction(::tvm::runtime::symbol::tvm_lookup_linked_param, true); - } - if (module_lookup_linked_param_ == nullptr) { - *rv = nullptr; - return; - } - - TVMRetValue opaque_handle = module_lookup_linked_param_(storage_id); - if (opaque_handle.type_code() == kTVMNullptr) { - *rv = nullptr; - return; - } - - std::vector shape_vec{template_tensor->shape, - template_tensor->shape + template_tensor->ndim}; - - auto* container = new NDArray::Container(static_cast(opaque_handle), shape_vec, - template_tensor->dtype, dev); - container->SetDeleter(GraphExecutor::LinkedNDArrayDeleter); - *rv = NDArray(GetObjectPtr(container)); -} - -void GraphExecutor::SetupStorage() { - // Grab saved optimization plan from graph. - std::vector vtype; - for (const std::string& s_type : attrs_.dltype) { - vtype.push_back(tvm::runtime::String2DLDataType(s_type)); - } - - // Size and device type of each storage pool entry. - std::vector pool_entry; - // Find the maximum space size. - for (size_t i = 0; i < attrs_.shape.size(); ++i) { - int storage_id = attrs_.storage_id[i]; - std::string storage_scope = attrs_.storage_scope.empty() ? "" : attrs_.storage_scope[i]; - // Use the fallback device if no device index is available. - int device_type = static_cast(devices_[0].device_type); - if (!attrs_.device_index.empty()) { - device_type = attrs_.device_index[i]; - } - - uint32_t sid = static_cast(storage_id); - if (sid >= pool_entry.size()) { - pool_entry.resize(sid + 1, {-1, {0}, {}}); - } else { - ICHECK(pool_entry[sid].device_type == -1 || pool_entry[sid].device_type == device_type) - << "The same pool entry cannot be assigned to multiple devices"; - } - TVMRetValue lookup_rv; - { - std::vector shape_vec{attrs_.shape[i].begin(), attrs_.shape[i].end()}; - DLTensor template_tensor{nullptr, Device{kDLCPU, 0}, static_cast(shape_vec.size()), - vtype[i], shape_vec.data(), nullptr, - 0}; - lookup_rv = lookup_linked_param_(module_, sid, &template_tensor, devices_[0]); - } - if (lookup_rv.type_code() != kTVMNullptr) { - pool_entry[sid].linked_param = lookup_rv; - } - pool_entry[sid].param_data_entry = i; - pool_entry[sid].device_type = device_type; - - DLDataType t = vtype[i]; - - auto dev_type = pool_entry[sid].device_type; - const auto& cit = std::find_if(devices_.begin(), devices_.end(), [&dev_type](const Device& d) { - return dev_type == static_cast(d.device_type); - }); - Device dev = cit == devices_.end() ? devices_[0] : *cit; - - DLTensor temp; - temp.data = nullptr; - temp.device = dev; - temp.ndim = attrs_.shape[i].size(); - temp.dtype = t; - temp.shape = static_cast(attrs_.shape[i].data()); - temp.strides = nullptr; - temp.byte_offset = 0; - - int64_t alloc_size = DeviceAPI::Get(dev)->GetDataSize(temp, String(storage_scope)); - - if (pool_entry[sid].alloc_size < alloc_size) { - pool_entry[sid].dtype = t; - pool_entry[sid].shape = attrs_.shape[i]; - pool_entry[sid].alloc_size = alloc_size; - pool_entry[sid].scope = storage_scope; - } - } - - // Allocate the space. - for (const auto& pit : pool_entry) { - // This for loop is very fast since there are usually only a couple of - // devices available on the same hardware. - const auto& cit = std::find_if(devices_.begin(), devices_.end(), [&pit](const Device& d) { - return pit.device_type == static_cast(d.device_type); - }); - Device dev = cit == devices_.end() ? devices_[0] : *cit; - if (pit.linked_param.defined()) { - ndarray_pool_.push_back(pit.linked_param); - } else { - std::vector shape = pit.shape; - String mem_scope = pit.scope.empty() ? "global" : String(pit.scope); - auto allocator = MemoryManager::GetOrCreateAllocator(dev, AllocatorType::kPooled); - auto buffer = allocator->Alloc(dev, pit.alloc_size, kAllocAlignment, pit.dtype); - auto stor = Storage(buffer, allocator); - storage_pool_.push_back(stor); - } - } - - // Assign the pooled entries. A unified memory pool is used to simplifiy - // memory assignment for each node entry. The allocated memory on each device - // is mapped to this pool. - data_entry_.resize(num_node_entries()); - data_alignment_.resize(num_node_entries()); - // sid_to_eid has a size of storage_id's size, which is the size of pool_entry. - sid_to_eid_.resize(pool_entry.size()); - for (size_t i = 0, j = 0; i < data_entry_.size(); ++i) { - int storage_id = attrs_.storage_id[i]; - // Update "storage_id -> entry_id" pair. - sid_to_eid_[storage_id].push_back(i); - - ICHECK_LT(static_cast(storage_id), pool_entry.size()); - - if (pool_entry[storage_id].linked_param.defined()) { - data_entry_[i] = ndarray_pool_[j++]; - } else { - std::string storage_scope = attrs_.storage_scope.empty() ? "global" : attrs_.storage_scope[i]; - data_entry_[i] = storage_pool_[storage_id]->AllocNDArrayScoped(0, ShapeTuple(attrs_.shape[i]), - vtype[i], storage_scope); - } - const DLTensor* tmp = data_entry_[i].operator->(); - data_alignment_[i] = details::GetDataAlignment(*tmp); - } -} - -void GraphExecutor::SetupOpExecs() { - op_execs_.resize(this->GetNumOfNodes()); - op_profile_execs_.resize(this->GetNumOfNodes()); - input_dltensors_.resize(num_node_entries()); - output_dltensors_.resize(num_node_entries()); - both_output_opinput_dltensors_.resize(num_node_entries()); - std::unordered_set input_node_eids; - for (size_t i = 0; i < input_nodes_.size(); i++) { - uint32_t nid = input_nodes_[i]; - input_node_eids.insert(entry_id(nid, 0)); - } - std::unordered_set output_node_eids; - for (size_t i = 0; i < outputs_.size(); i++) { - output_node_eids.insert(entry_id(outputs_[i])); - } - - // setup the array and requirements. - for (uint32_t nid = 0; nid < this->GetNumOfNodes(); ++nid) { - const auto& inode = nodes_[nid]; - if (inode.op_type == "null") continue; - std::vector args; - for (const auto& e : inode.inputs) { - uint32_t eid = this->entry_id(e); - args.push_back(const_cast(data_entry_[eid].operator->())); - } - for (uint32_t index = 0; index < inode.param.num_outputs; ++index) { - uint32_t eid = this->entry_id(nid, index); - args.push_back(const_cast(data_entry_[eid].operator->())); - } - ICHECK(inode.op_type == "tvm_op") << "Can only take tvm_op as op"; - - std::shared_ptr op_args = nullptr; - std::tie(op_execs_[nid], op_profile_execs_[nid], op_args) = CreateTVMOp(inode.param, args); - - for (size_t i = 0; i < inode.inputs.size(); i++) { - uint32_t input_eid = this->entry_id(inode.inputs[i]); - // check if op input is model input - if (input_node_eids.count(input_eid) > 0) { - input_dltensors_[input_eid].push_back( - static_cast(op_args->arg_values[i].v_handle)); - - // Data entry who has the same storage_id should also be pushed into "input_dltensors" and - // being able to be updated by "SetInputZeroCopy()". This is to handle the situation that a - // "relay.reshape" follows immediately after input and input dltensor and reshape's output - // dltensor point to the same data_entry. - auto storage_id = attrs_.storage_id[input_eid]; - for (auto eid : sid_to_eid_[storage_id]) { - input_dltensors_[input_eid].push_back( - const_cast(data_entry_[eid].operator->())); - } - } else { - const auto& arg_node = nodes_[inode.inputs[i].node_id]; - if (arg_node.op_type == "tvm_op" && arg_node.param.func_name == "__nop") { - uint32_t arg_input_eid = this->entry_id(arg_node.inputs[0]); - input_dltensors_[arg_input_eid].push_back( - static_cast(op_args->arg_values[i].v_handle)); - } - } - // check if any model output is the input of the op - if (output_node_eids.count(input_eid) > 0) { - both_output_opinput_dltensors_[input_eid].push_back( - static_cast(op_args->arg_values[i].v_handle)); - } - } - - for (uint32_t i = inode.inputs.size(); i < inode.inputs.size() + inode.param.num_outputs; ++i) { - uint32_t output_eid = this->entry_id(nid, i - inode.inputs.size()); - // check if op output is model output - if (output_node_eids.count(output_eid) > 0) { - output_dltensors_[output_eid].push_back( - static_cast(op_args->arg_values[i].v_handle)); - } else { - // If the node is not an output, keep its output for record and support set_output_zero_copy - // of reshape __nop nodes. - node_output_dltensors_[output_eid].push_back( - static_cast(op_args->arg_values[i].v_handle)); - } - } - } -} - -std::tuple, std::function, - std::shared_ptr> -GraphExecutor::CreateTVMOp(const TVMOpParam& param, const std::vector& args) { - std::shared_ptr arg_ptr = std::make_shared(); - // setup address. - arg_ptr->args = args; - if (param.flatten_data) { - arg_ptr->shape_data.resize(arg_ptr->args.size()); - } - for (size_t i = 0; i < arg_ptr->args.size(); ++i) { - TVMValue v; - DLTensor* t = arg_ptr->args[i]; - v.v_handle = t; - arg_ptr->arg_values.push_back(v); - arg_ptr->arg_tcodes.push_back(kTVMDLTensorHandle); - if (param.flatten_data) { - arg_ptr->shape_data[i] = - std::accumulate(t->shape, t->shape + t->ndim, 1, std::multiplies()); - t->ndim = 1; - t->shape = &(arg_ptr->shape_data[i]); - } - } - - if (param.func_name == "__nop") { - return {[]() {}, [](TVMRetValue* rv) {}, arg_ptr}; - } else if (param.func_name == "__copy") { - // Perform cross device data copy. - // Directly copy data from the input to the output. - // TODO(mbs): device_copy cleanup. - auto fexec = [arg_ptr]() { - DLTensor* from = static_cast(arg_ptr->arg_values[0].v_handle); - DLTensor* to = static_cast(arg_ptr->arg_values[1].v_handle); - TVM_CCALL(TVMArrayCopyFromTo(from, to, nullptr)); - }; - return {fexec, [](TVMRetValue* rv) {}, arg_ptr}; - } - - // Get compiled function from the module that contains both host and device - // code. - tvm::runtime::PackedFunc pf = module_.GetFunction(param.func_name, true); - ICHECK(pf != nullptr) << "no such function in module: " << param.func_name; - auto fexec = [arg_ptr, pf]() { - TVMRetValue rv; - TVMArgs targs(arg_ptr->arg_values.data(), arg_ptr->arg_tcodes.data(), - static_cast(arg_ptr->arg_values.size())); - pf.CallPacked(targs, &rv); - }; - - pf = module_.GetFunction(param.func_name + "_debug", true); - std::function fexec_profile = nullptr; - if (pf != nullptr) { - fexec_profile = [arg_ptr, pf](TVMRetValue* rv) { - TVMArgs targs(arg_ptr->arg_values.data(), arg_ptr->arg_tcodes.data(), - static_cast(arg_ptr->arg_values.size())); - pf.CallPacked(targs, rv); - }; - } - - return {fexec, fexec_profile, arg_ptr}; -} - -PackedFunc GraphExecutor::GetFunction(const String& name, const ObjectPtr& sptr_to_self) { - // Return member functions during query. - if (name == "set_input") { - return PackedFunc([sptr_to_self, this](TVMArgs args, TVMRetValue* rv) { - if (String::CanConvertFrom(args[0])) { - int in_idx = this->GetInputIndex(args[0].operator String()); - if (in_idx >= 0) this->SetInput(in_idx, args[1]); - } else { - this->SetInput(args[0], args[1]); - } - }); - } else if (name == "set_input_zero_copy") { - return PackedFunc([sptr_to_self, this](TVMArgs args, TVMRetValue* rv) { - if (String::CanConvertFrom(args[0])) { - int in_idx = this->GetInputIndex(args[0].operator String()); - if (in_idx >= 0) this->SetInputZeroCopy(in_idx, args[1]); - } else { - this->SetInputZeroCopy(args[0], args[1]); - } - }); - } else if (name == "set_output_zero_copy") { - return PackedFunc([sptr_to_self, this](TVMArgs args, TVMRetValue* rv) { - if (String::CanConvertFrom(args[0])) { - int out_idx = this->GetOutputIndex(args[0].operator String()); - if (out_idx >= 0) this->SetOutputZeroCopy(out_idx, args[1]); - } else { - this->SetOutputZeroCopy(args[0], args[1]); - } - }); - } else if (name == "get_output") { - return PackedFunc([sptr_to_self, this](TVMArgs args, TVMRetValue* rv) { - if (args.num_args == 2) { - this->CopyOutputTo(args[0], args[1]); - } else { - int out_idx = -1; - if (String::CanConvertFrom(args[0])) { - for (size_t i = 0; i < outputs_.size(); i++) { - std::string& name = nodes_[outputs_[i].node_id].name; - if (args[0].operator String() == name) { - out_idx = i; - } - } - CHECK(out_idx != -1) << "Invalid output node:" << args[0].operator String(); - } else { - out_idx = args[0]; - } - *rv = this->GetOutput(out_idx); - } - }); - } else if (name == "get_input") { - return PackedFunc([sptr_to_self, this](TVMArgs args, TVMRetValue* rv) { - int in_idx = 0; - if (String::CanConvertFrom(args[0])) { - in_idx = this->GetInputIndex(args[0].operator String()); - } else { - in_idx = args[0]; - } - if (in_idx >= 0) { - *rv = this->GetInput(in_idx); - } - }); - } else if (name == "get_num_outputs") { - return PackedFunc( - [sptr_to_self, this](TVMArgs args, TVMRetValue* rv) { *rv = this->NumOutputs(); }); - } else if (name == "get_num_inputs") { - return PackedFunc( - [sptr_to_self, this](TVMArgs args, TVMRetValue* rv) { *rv = this->NumInputs(); }); - } else if (name == "run") { - return PackedFunc([sptr_to_self, this](TVMArgs args, TVMRetValue* rv) { this->Run(); }); - } else if (name == "run_from_inputs") { - return PackedFunc( - [sptr_to_self, this](TVMArgs args, TVMRetValue* rv) { - CHECK(args.size() % 2 == 0) - << "Number of arguments to run_from_inputs must be an even number of key-value pairs"; - Device host{static_cast(args[0].operator int()), args[1].operator int()}; - for (int i = 2; i < args.size(); i += 2) { - if (String::CanConvertFrom(args[i])) { - int in_idx = this->GetInputIndex(args[i].operator String()); - if (in_idx >= 0) { - this->SetInput(in_idx, args[i + 1]); - } else { - LOG(FATAL) << args[i].operator String() << " is not a valid input name"; - } - } else { - this->SetInput(args[i], args[i + 1]); - } - } - this->Run(); - Array outputs; - for (int i = 0; i < this->NumOutputs(); i++) { - NDArray out = this->GetOutput(i); - NDArray a = NDArray::Empty(out.Shape(), out.DataType(), host); - a.CopyFrom(out); - outputs.push_back(a); - } - *rv = outputs; - }); - } else if (name == "load_params") { - return PackedFunc([sptr_to_self, this](TVMArgs args, TVMRetValue* rv) { - this->LoadParams(args[0].operator std::string()); - }); - } else if (name == "share_params") { - return PackedFunc([sptr_to_self, this](TVMArgs args, TVMRetValue* rv) { - const auto& module = args[0].operator Module(); - ICHECK_EQ(module.operator->()->type_key(), std::string("GraphExecutor")); - const auto& param_blob = args[1].operator std::string(); - dmlc::MemoryStringStream strm(const_cast(¶m_blob)); - this->ShareParams(dynamic_cast(*module.operator->()), &strm); - }); - } else if (name == "get_input_index") { - return PackedFunc([sptr_to_self, this](TVMArgs args, TVMRetValue* rv) { - CHECK(String::CanConvertFrom(args[0])) << "Input key is not a string"; - *rv = this->GetInputIndex(args[0].operator String()); - }); - } else if (name == "get_output_index") { - return PackedFunc([sptr_to_self, this](TVMArgs args, TVMRetValue* rv) { - CHECK(String::CanConvertFrom(args[0])) << "Output key is not a string"; - int out_idx = -1; - for (size_t i = 0; i < outputs_.size(); i++) { - std::string& name = nodes_[outputs_[i].node_id].name; - if (args[0].operator String() == name) { - out_idx = i; - } - } - *rv = out_idx; - }); - } else if (name == "get_input_info") { - return PackedFunc([sptr_to_self, this](TVMArgs args, TVMRetValue* rv) { - auto [shape_info, dtype_info] = this->GetInputInfo(); - Map input_info; - input_info.Set("shape", shape_info); - input_info.Set("dtype", dtype_info); - *rv = input_info; - }); - } else if (name == "get_output_info") { - return PackedFunc([sptr_to_self, this](TVMArgs args, TVMRetValue* rv) { - auto [shape_info, dtype_info] = this->GetOutputInfo(); - Map input_info; - input_info.Set("shape", shape_info); - input_info.Set("dtype", dtype_info); - *rv = input_info; - }); - } else { - return PackedFunc(); - } -} - -Module GraphExecutorCreate(const std::string& sym_json, const tvm::runtime::Module& m, - const std::vector& devs, - const PackedFunc lookup_linked_param_func) { - auto exec = make_object(); - exec->Init(sym_json, m, devs, lookup_linked_param_func); - return Module(exec); -} - -// Get all devices for the host and other runtime devices. -std::vector GetAllDevice(const TVMArgs& args, int dev_start_arg) { - // Reserve the first item as the fallback device. - std::vector ret; - Device dev; - for (int i = dev_start_arg; i < args.num_args; i += 2) { - int dev_type = args[i]; - dev.device_type = static_cast(dev_type); - dev.device_id = args[i + 1]; - ret.push_back(dev); - } - return ret; -} - -// 4-argument version is currently reserved to keep support of calling -// from tvm4j and javascript, since they don't have heterogeneous -// execution support yet. For heterogenenous execution, at least 5 arguments will -// be passed in. The third one is the number of devices. -// Eventually, we will only probably pass Device for all the languages. -TVM_REGISTER_GLOBAL("tvm.graph_executor.create").set_body([](TVMArgs args, TVMRetValue* rv) { - ICHECK_GE(args.num_args, 4) << "The expected number of arguments for graph_executor.create is " - "at least 4, but it has " - << args.num_args; - PackedFunc lookup_linked_param_func; - int dev_start_arg = 2; - if (args[2].type_code() == kTVMPackedFuncHandle) { - lookup_linked_param_func = args[2]; - dev_start_arg++; - } - const auto& devices = GetAllDevice(args, dev_start_arg); - *rv = GraphExecutorCreate(args[0], args[1], devices, lookup_linked_param_func); -}); -} // namespace runtime -} // namespace tvm diff --git a/src/runtime/graph_executor/graph_executor.h b/src/runtime/graph_executor/graph_executor.h deleted file mode 100644 index e1c61001f1d9..000000000000 --- a/src/runtime/graph_executor/graph_executor.h +++ /dev/null @@ -1,514 +0,0 @@ -/* - * 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. - */ - -/*! - * \brief Tiny graph executor that can run graph - * containing only tvm PackedFunc. - * \file graph_executor.h - */ -#ifndef TVM_RUNTIME_GRAPH_EXECUTOR_GRAPH_EXECUTOR_H_ -#define TVM_RUNTIME_GRAPH_EXECUTOR_GRAPH_EXECUTOR_H_ - -#include -#include -#include -#include -#include -#include - -#include -#include -#include -#include -#include -#include -#include - -namespace tvm { -namespace runtime { - -using memory::AllocatorType; -using memory::MemoryManager; -using tvm::runtime::memory::Storage; - -/*! \brief macro to do C API call */ -#define TVM_CCALL(func) \ - { \ - int ret = (func); \ - ICHECK_EQ(ret, 0) << TVMGetLastError(); \ - } - -/*! \brief operator attributes about tvm op */ -struct TVMOpParam { - std::string func_name; - std::string compiler; - std::unordered_map attrs; - uint32_t num_inputs; - uint32_t num_outputs; - uint32_t flatten_data; -}; - -/*! - * \brief Tiny graph executor. - * - * This runtime can be accessible in various languages via - * TVM runtime PackedFunc API. - */ -class TVM_DLL GraphExecutor : public ModuleNode { - struct OpArgs { - std::vector args; - std::vector arg_values; - std::vector arg_tcodes; - std::vector shape_data; - }; - - public: - using ShapeInfo = Map; - using DtypeInfo = Map; - /*! - * \brief Get member function to front-end - * \param name The name of the function. - * \param sptr_to_self The pointer to the module node. - * \return The corresponding member function. - */ - virtual PackedFunc GetFunction(const String& name, const ObjectPtr& sptr_to_self); - - /*! - * \return The type key of the executor. - */ - const char* type_key() const final { return "GraphExecutor"; } - void Run(); - - /*! \brief Get the property of the runtime module .*/ - int GetPropertyMask() const final { return ModulePropertyMask::kRunnable; } - - /*! - * \brief Initialize the graph executor with graph and device. - * \param graph_json The execution graph. - * \param module The module containing the compiled functions for the host - * processor. - * \param devs The device of the host and devices where graph nodes will be - * executed on. - * \param lookup_linked_param_func If given, a PackedFunc invoked to lookup linked parameters - * by storage_id. If not given, linked parameters are looked-up using an internal implementation, - * which is not compatible with RPCModules. Default is nullptr. - */ - - void Init(const std::string& graph_json, tvm::runtime::Module module, - const std::vector& devs, const PackedFunc lookup_linked_param_func = nullptr); - - /*! - * \brief Get the input index given the name of input. - * \param name The name of the input. - * \return The index of input. - */ - int GetInputIndex(const std::string& name); - - /*! - * \brief Get the input info of Graph by parsing the input nodes. - * \return The shape and dtype tuple. - */ - std::tuple GetInputInfo() const; - - /*! - * \brief Get the output info of Graph by parsing the output nodes. - * \return The shape and dtype tuple. - */ - std::tuple GetOutputInfo() const; - - /*! - * \brief Get the output index given the name of output. - * \param name The name of the output. - * \return The index of output. - */ - int GetOutputIndex(const std::string& name); - - /*! - * \brief set index-th input to the graph. - * \param index The input index. - * \param data_in The input data. - */ - void SetInput(int index, DLTensor* data_in); - /*! - * \brief set index-th input to the graph without copying the data - * \param index The input index. - * \param data_ref The input data that is referred. - */ - void SetInputZeroCopy(int index, DLTensor* data_ref); - /*! - * \brief set index-th output to the graph without copying the data. - * \param index The output index. - * \param data_ref The output data that is referred. - */ - void SetOutputZeroCopy(int index, DLTensor* data_ref); - /*! - * \brief Get the number of outputs - * - * \return The number of outputs from graph. - */ - int NumOutputs() const; - /*! - * \brief Get the number of inputs - * - * \return The number of inputs to the graph. - */ - int NumInputs() const; - /*! - * \brief Return NDArray for given input index. - * \param index The input index. - * - * \return NDArray corresponding to given input node index. - */ - NDArray GetInput(int index) const; - /*! - * \brief Return NDArray for given output index. - * \param index The output index. - * - * \return NDArray corresponding to given output node index. - */ - NDArray GetOutput(int index) const; - /*! - * \brief Copy index-th output to data_out. - * \param index The output index. - * \param data_out the output data. - */ - void CopyOutputTo(int index, DLTensor* data_out); - /*! - * \brief Load parameters from binary stream - * \param strm The input stream. - */ - void LoadParams(dmlc::Stream* strm); - /*! - * \brief Load parameters from parameter blob. - * \param param_blob A binary blob of parameter. - */ - void LoadParams(const std::string& param_blob); - - /*! - * \brief Share parameters from pre-existing GraphExecutor instance. - * \param other A GraphExecutor instance, previously with |LoadParams| called with the - * identical input |param_blob|. - * \param strm The input stream. - */ - void ShareParams(const GraphExecutor& other, dmlc::Stream* strm); - - /*! - * \brief Get total number of nodes. - * \return Total number of nodes. - */ - uint32_t GetNumOfNodes() const { return static_cast(nodes_.size()); } - - std::string GetNodeName(uint32_t nid) const { return nodes_[nid].name; } - - protected: - // Memory pool entry. - struct PoolEntry { - int device_type; - std::vector shape; - DLDataType dtype; - int param_data_entry; - NDArray linked_param; - std::string scope; - int64_t alloc_size{-1}; - // PoolEntry(int s, int dev_type, void* pre_linked_param) : - // size(s), device_type(dev_type), pre_linked_param(std::move(pre_linked_param)) {} - }; - // Node entry - struct NodeEntry { - uint32_t node_id; - uint32_t index; - uint32_t version; - inline bool operator==(const NodeEntry& other) const { - return node_id == other.node_id && index == other.index && version == other.version; - } - // JSON Loader - void Load(dmlc::JSONReader* reader) { - reader->BeginArray(); - ICHECK(reader->NextArrayItem()) << "invalid json format"; - reader->Read(&node_id); - ICHECK(reader->NextArrayItem()) << "invalid json format"; - reader->Read(&index); - if (reader->NextArrayItem()) { - reader->Read(&version); - ICHECK(!reader->NextArrayItem()) << "invalid json format"; - } else { - version = 0; - } - } - }; - // Node - struct Node { - // operator type in string - std::string op_type; - // name of the op - std::string name; - // parameters - TVMOpParam param; - // inputs - std::vector inputs; - // control deps - std::vector control_deps; - - // JSON Loader - void LoadAttrs(dmlc::JSONReader* reader, TVMOpParam* param) { - int bitmask = 0; - std::string key, value; - reader->BeginObject(); - while (reader->NextObjectItem(&key)) { - reader->Read(&value); - if (key == "func_name") { - param->func_name = value; - bitmask |= 1; - } - if (key == "Compiler") { - param->compiler = value; - } else if (key == "num_inputs") { - param->num_inputs = strtoul(value.c_str(), nullptr, 10); - bitmask |= 2; - } else if (key == "num_outputs") { - param->num_outputs = strtoul(value.c_str(), nullptr, 10); - bitmask |= 4; - } else if (key == "flatten_data") { - param->flatten_data = strtoul(value.c_str(), nullptr, 10); - bitmask |= 8; - } else { - param->attrs[key] = String(value); - } - } - ICHECK_EQ(bitmask, 1 | 2 | 4 | 8) << "invalid format"; - } - // JSON Loader - void Load(dmlc::JSONReader* reader) { - reader->BeginObject(); - int bitmask = 0; - std::string key; - while (reader->NextObjectItem(&key)) { - if (key == "op") { - reader->Read(&op_type); - bitmask |= 1; - } else if (key == "name") { - reader->Read(&name); - bitmask |= 2; - } else if (key == "inputs") { - reader->Read(&inputs); - bitmask |= 4; - } else if (key == "attr" || key == "attrs") { - this->LoadAttrs(reader, ¶m); - } else if (key == "control_deps") { - reader->Read(&control_deps); - } else { - LOG(FATAL) << "do not support key " << key; - } - } - ICHECK_EQ(bitmask, 1 | 2 | 4) << "invalid format"; - } - }; - struct GraphAttr { - size_t storage_num_not_alloctaed{0}; - std::vector storage_id; - std::vector device_index; - std::vector dltype; - std::vector storage_scope; - std::vector> shape; - // The graph attribute fields. - void Load(dmlc::JSONReader* reader) { - reader->BeginObject(); - int bitmask = 0; - std::string key, type; - while (reader->NextObjectItem(&key)) { - if (key == "dltype") { - reader->BeginArray(); - ICHECK(reader->NextArrayItem()); - reader->Read(&type); - ICHECK_EQ(type, "list_str"); - ICHECK(reader->NextArrayItem()); - reader->Read(&dltype); - ICHECK(!reader->NextArrayItem()); - bitmask |= 1; - } else if (key == "storage_id") { - reader->BeginArray(); - ICHECK(reader->NextArrayItem()); - reader->Read(&type); - ICHECK_EQ(type, "list_int"); - ICHECK(reader->NextArrayItem()); - reader->Read(&storage_id); - ICHECK(!reader->NextArrayItem()); - bitmask |= 2; - } else if (key == "storage_scope") { - reader->BeginArray(); - ICHECK(reader->NextArrayItem()); - reader->Read(&type); - ICHECK_EQ(type, "list_str"); - ICHECK(reader->NextArrayItem()); - reader->Read(&storage_scope); - ICHECK(!reader->NextArrayItem()); - bitmask |= 1; - } else if (key == "shape") { - reader->BeginArray(); - ICHECK(reader->NextArrayItem()); - reader->Read(&type); - ICHECK_EQ(type, "list_shape"); - ICHECK(reader->NextArrayItem()); - reader->Read(&shape); - ICHECK(!reader->NextArrayItem()); - bitmask |= 4; - } else if (key == "device_index") { - reader->BeginArray(); - ICHECK(reader->NextArrayItem()); - reader->Read(&type); - ICHECK_EQ(type, "list_int"); - ICHECK(reader->NextArrayItem()); - reader->Read(&device_index); - ICHECK(!reader->NextArrayItem()); - } else { - reader->BeginArray(); - ICHECK(reader->NextArrayItem()); - reader->Read(&type); - if (type == "list_int") { - ICHECK(reader->NextArrayItem()); - std::vector temp; - reader->Read(&temp); - } else if (type == "size_t") { - ICHECK(reader->NextArrayItem()); - size_t temp; - reader->Read(&temp); - } else { - LOG(FATAL) << "cannot skip graph attr " << key; - } - ICHECK(!reader->NextArrayItem()); - } - } - ICHECK_EQ(bitmask, 1 | 2 | 4) << "invalid format"; - } - }; - // The graph attribute fields. - void Load(dmlc::JSONReader* reader) { - reader->BeginObject(); - int bitmask = 0; - std::string key; - while (reader->NextObjectItem(&key)) { - if (key == "nodes") { - reader->Read(&nodes_); - bitmask |= 1; - } else if (key == "arg_nodes") { - reader->Read(&input_nodes_); - bitmask |= 2; - } else if (key == "node_row_ptr") { - reader->Read(&node_row_ptr_); - bitmask |= 4; - } else if (key == "heads") { - reader->Read(&outputs_); - bitmask |= 8; - } else if (key == "attrs") { - reader->Read(&attrs_); - bitmask |= 16; - } else if (key == "metadata") { - break; - } else { - LOG(FATAL) << "key " << key << " is not supported"; - } - } - ICHECK_EQ(bitmask, 1 | 2 | 4 | 8 | 16) << "invalid format"; - } - /*! \brief PackedFunc to lookup a linked paramter from a local Module. */ - void DefaultLookupLinkedParam(TVMArgs args, TVMRetValue* rv); - /*! \brief Delete NDArray::Container with linked (i.e. static) data. */ - static void LinkedNDArrayDeleter(Object* container); - /*! \brief Setup the temporal storage */ - void SetupStorage(); - /*! \brief Setup the executors. */ - void SetupOpExecs(); - /*! - * \brief Check the legality of external DLTensor*. - * \param external The external DLTensor*. - * \param eid The data_enrty_ index. - */ - void CheckExternalDLTensor(const DLTensor* external, uint32_t eid) const; - /*! - * \brief Create an execution function given input. - * \param attrs The node attributes. - * \param args The arguments to the functor, including inputs and outputs. - * \return The created executor. - */ - std::tuple, std::function, std::shared_ptr> - CreateTVMOp(const TVMOpParam& attrs, const std::vector& args); - // Get node entry index. - uint32_t entry_id(uint32_t nid, uint32_t index) const { return node_row_ptr_[nid] + index; } - // Get node entry index. - uint32_t entry_id(const NodeEntry& e) const { return entry_id(e.node_id, e.index); } - // Number of node entries. - uint32_t num_node_entries() const { return node_row_ptr_.back(); } - /*! \brief The graph nodes. */ - std::vector nodes_; - /*! \brief The argument nodes. */ - std::vector input_nodes_; - /*! \brief The parameter names. */ - std::unordered_set param_names_; - /*! \brief Map of input names to input indices. */ - std::unordered_map input_map_; - /*! \brief Map of output names to output indices. */ - std::unordered_map output_map_; - /*! \brief Used for quick node input DLTensor* lookup given an input eid. */ - std::vector> input_dltensors_; - /*! \brief Used for quick node output DLTensor* lookup given an output eid. */ - std::vector> output_dltensors_; - /*! \brief Used for quick node(both model output and op input) DLTensor* lookup given an eid. */ - std::vector> both_output_opinput_dltensors_; - /*! \brief Used for quick node output DLTensor* lookup given a nop's input eid. */ - std::unordered_map> node_output_dltensors_; - /*! \brief Used for quick entry_id lookup given an storage_id. */ - std::vector> sid_to_eid_; - /*! \brief Used for quick entry indexing. */ - std::vector node_row_ptr_; - /*! \brief Output entries. */ - std::vector outputs_; - /*! \brief Additional graph attributes. */ - GraphAttr attrs_; - /*! \brief The code module that contains both host and device code. */ - tvm::runtime::Module module_; - /*! \brief Execution context of all devices including the host. */ - std::vector devices_; - /*! \brief Common storage pool for all devices. */ - std::vector storage_pool_; - /*! \brief Common NDArray pool for all devices. */ - std::vector ndarray_pool_; - /*! \brief Data entry of each node. */ - std::vector data_entry_; - /*! \brief Data alignment of each node. */ - std::vector data_alignment_; - /*! \brief Operator on each node. */ - std::vector> op_execs_; - /*! \brief Profilable Operator on each node. */ - std::vector> op_profile_execs_; - /*! \brief Linked parameter lookup function. */ - PackedFunc lookup_linked_param_; - /*! \brief Module's _lookup_linked_param function, used by DefaultLookupLinkedParam. */ - PackedFunc module_lookup_linked_param_; - /*! - * \brief True when module_lookup_linked_param_ is valid. - * When the module does not include linked parmeters, module_lookup_linked_param_ will be nullptr. - */ - bool module_lookup_linked_param_valid_; -}; - -std::vector GetAllDevice(const TVMArgs& args, int dev_start_arg); -} // namespace runtime -} // namespace tvm - -#endif // TVM_RUNTIME_GRAPH_EXECUTOR_GRAPH_EXECUTOR_H_ diff --git a/src/runtime/graph_executor/graph_executor_factory.cc b/src/runtime/graph_executor/graph_executor_factory.cc deleted file mode 100644 index 90c42ca74be1..000000000000 --- a/src/runtime/graph_executor/graph_executor_factory.cc +++ /dev/null @@ -1,234 +0,0 @@ -/* - * 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 graph_executor_factory.cc - * \brief Graph executor factory implementations - */ - -#include "./graph_executor_factory.h" - -#include -#include -#include -#include - -#include -#include - -namespace tvm { -namespace runtime { - -GraphExecutorFactory::GraphExecutorFactory( - const std::string& graph_json, - const std::unordered_map& params, - const std::string& module_name) { - graph_json_ = graph_json; - params_ = params; - module_name_ = module_name; -} - -PackedFunc GraphExecutorFactory::GetFunction( - const String& name, const tvm::runtime::ObjectPtr& sptr_to_self) { - if (name == module_name_) { - return PackedFunc([sptr_to_self, this](TVMArgs args, TVMRetValue* rv) { - std::vector devices; - for (int i = 0; i < args.num_args; ++i) { - devices.emplace_back(args[i].operator Device()); - } - *rv = this->ExecutorCreate(devices); - }); - } else if (name == "get_graph_json") { - return PackedFunc( - [sptr_to_self, this](TVMArgs args, TVMRetValue* rv) { *rv = this->graph_json_; }); - - } else if (name == "get_graph_params") { - return PackedFunc([sptr_to_self, this](TVMArgs args, TVMRetValue* rv) { - Map params; - for (const auto& kv : params_) { - params.Set(kv.first, kv.second); - } - *rv = params; - }); - } else if (name == "debug_create") { - return PackedFunc([sptr_to_self, this](TVMArgs args, TVMRetValue* rv) { - ICHECK_GE(args.size(), 2); - std::string module_name = args[0].operator String(); - ICHECK(module_name == module_name_) << "Currently we only support single model for now."; - std::vector devices; - for (int i = 1; i < args.num_args; ++i) { - devices.emplace_back(args[i].operator Device()); - } - *rv = this->DebugExecutorCreate(devices); - }); - } else if (name == "remove_params") { - return PackedFunc([sptr_to_self, this](TVMArgs args, TVMRetValue* rv) { - std::unordered_map empty_params{}; - auto exec = - make_object(this->graph_json_, empty_params, this->module_name_); - exec->Import(this->imports_[0]); - *rv = Module(exec); - }); - } else if (name == "cuda_graph_create") { - return PackedFunc([sptr_to_self, this](TVMArgs args, TVMRetValue* rv) { - std::vector devices; - for (int i = 0; i < args.num_args; ++i) { - devices.emplace_back(args[i].operator Device()); - } - *rv = this->CudaGraphExecutorCreate(devices); - }); - } else { - return PackedFunc(); - } -} - -void GraphExecutorFactory::SaveToBinary(dmlc::Stream* stream) { - stream->Write(graph_json_); - std::vector names; - std::vector arrays; - for (const auto& v : params_) { - names.emplace_back(v.first); - arrays.emplace_back(const_cast(v.second.operator->())); - } - uint64_t sz = arrays.size(); - ICHECK(sz == names.size()); - stream->Write(sz); - stream->Write(names); - for (size_t i = 0; i < sz; ++i) { - tvm::runtime::SaveDLTensor(stream, arrays[i]); - } - stream->Write(module_name_); -} - -Module GraphExecutorFactory::ExecutorCreate(const std::vector& devs) { - auto exec = make_object(); - exec->Init(this->graph_json_, this->imports_[0], devs, PackedFunc()); - // set params - SetParams(exec.get(), this->params_); - return Module(exec); -} - -Module GraphExecutorFactory::DebugExecutorCreate(const std::vector& devs) { - const PackedFunc* pf = tvm::runtime::Registry::Get("tvm.graph_executor_debug.create"); - ICHECK(pf != nullptr) << "Cannot find function tvm.graph_executor_debug.create in registry. " - "Do you enable debug graph executor build?"; - // Debug executor create packed function will call GetAllContexs, so we unpack the devs. - std::vector unpacked_devs; - for (const auto& dev : devs) { - unpacked_devs.emplace_back(dev.device_type); - unpacked_devs.emplace_back(dev.device_id); - } - size_t args_size = unpacked_devs.size() + 2; - std::vector values(args_size); - std::vector codes(args_size); - runtime::TVMArgsSetter setter(values.data(), codes.data()); - setter(0, this->graph_json_); - setter(1, this->imports_[0]); - for (size_t i = 0; i < unpacked_devs.size(); ++i) { - setter(i + 2, unpacked_devs[i]); - } - TVMRetValue rv; - pf->CallPacked(TVMArgs(values.data(), codes.data(), args_size), &rv); - Module mod = rv.operator Module(); - // debug graph executor is one child class of graph executor. - SetParams(const_cast(mod.as()), this->params_); - return mod; -} - -Module GraphExecutorFactory::CudaGraphExecutorCreate(const std::vector& devs) { - const PackedFunc* pf = tvm::runtime::Registry::Get("tvm.graph_executor_cuda_graph.create"); - ICHECK(pf != nullptr) << "Cannot find function tvm.graph_executor_cuda_graph.create in registry. " - "Did you set(USE_GRAPH_EXECUTOR_CUGRAPH=ON)?"; - std::vector unpacked_devs; - for (const auto& dev : devs) { - unpacked_devs.emplace_back(dev.device_type); - unpacked_devs.emplace_back(dev.device_id); - } - size_t args_size = unpacked_devs.size() + 2; - std::vector values(args_size); - std::vector codes(args_size); - runtime::TVMArgsSetter setter(values.data(), codes.data()); - setter(0, this->graph_json_); - setter(1, this->imports_[0]); - for (size_t i = 0; i < unpacked_devs.size(); ++i) { - setter(i + 2, unpacked_devs[i]); - } - TVMRetValue rv; - pf->CallPacked(TVMArgs(values.data(), codes.data(), args_size), &rv); - Module mod = rv.operator Module(); - SetParams(const_cast(mod.as()), this->params_); - return mod; -} - -Module GraphExecutorFactoryModuleLoadBinary(void* strm) { - dmlc::Stream* stream = static_cast(strm); - std::string graph_json; - std::unordered_map params; - std::string module_name; - ICHECK(stream->Read(&graph_json)); - uint64_t sz; - ICHECK(stream->Read(&sz)); - std::vector names; - ICHECK(stream->Read(&names)); - ICHECK(sz == names.size()); - for (size_t i = 0; i < sz; ++i) { - tvm::runtime::NDArray temp; - temp.Load(stream); - params[names[i]] = temp; - } - ICHECK(stream->Read(&module_name)); - auto exec = make_object(graph_json, params, module_name); - return Module(exec); -} - -TVM_REGISTER_GLOBAL("tvm.graph_executor_factory.create") - .set_body([](TVMArgs args, TVMRetValue* rv) { - ICHECK_GE(args.num_args, 3) << "The expected number of arguments for " - "graph_executor_factory.create needs at least 3, " - "but it has " - << args.num_args; - // The argument order is graph_json, module, module_name, param0_name, param0_tensor, - // [param1_name, param1_tensor], ... - ICHECK_EQ((args.size() - 3) % 2, 0); - std::unordered_map params; - for (size_t i = 3; i < static_cast(args.size()); i += 2) { - std::string name = args[i].operator String(); - params[name] = args[i + 1].operator tvm::runtime::NDArray(); - } - auto exec = make_object(args[0], params, args[2]); - exec->Import(args[1]); - *rv = Module(exec); - }); - -TVM_REGISTER_GLOBAL("runtime.module.loadbinary_GraphExecutorFactory") - .set_body_typed(GraphExecutorFactoryModuleLoadBinary); - -Module GraphRuntimeFactoryModuleLoadBinary(void* strm) { - LOG(WARNING) << "You are loading a module which was built with GraphRuntimeFactory. " - << "GraphRuntime has been renamed to GraphExecutor, and support for loading " - << "GraphRuntimeFactory modules will be removed after the next TVM release. " - << "Please rebuild the module before then to avoid breakage."; - return GraphExecutorFactoryModuleLoadBinary(strm); -} - -TVM_REGISTER_GLOBAL("runtime.module.loadbinary_GraphRuntimeFactory") - .set_body_typed(GraphRuntimeFactoryModuleLoadBinary); - -} // namespace runtime -} // namespace tvm diff --git a/src/runtime/graph_executor/graph_executor_factory.h b/src/runtime/graph_executor/graph_executor_factory.h deleted file mode 100644 index 2f41bb4e2eb2..000000000000 --- a/src/runtime/graph_executor/graph_executor_factory.h +++ /dev/null @@ -1,142 +0,0 @@ -/* - * 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 tvm/runtime/graph_executor/graph_executor_factory.h - * \brief Graph executor factory creating graph executor. - */ - -#ifndef TVM_RUNTIME_GRAPH_EXECUTOR_GRAPH_EXECUTOR_FACTORY_H_ -#define TVM_RUNTIME_GRAPH_EXECUTOR_GRAPH_EXECUTOR_FACTORY_H_ - -#include -#include -#include -#include - -#include -#include -#include -#include -#include -#include - -#include "./graph_executor.h" - -namespace tvm { -namespace runtime { - -class TVM_DLL GraphExecutorFactory : public runtime::ModuleNode { - public: - /*! - * \brief Construct the GraphExecutorFactory. - * \param graph_json The execution graph. - * \param params The params of graph. - * \param module_name The module name of graph. - */ - GraphExecutorFactory(const std::string& graph_json, - const std::unordered_map& params, - const std::string& module_name = "default"); - - /*! - * \brief Get member function to front-end - * \param name The name of the function. - * \param sptr_to_self The pointer to the module node. - * \return The corresponding member function. - */ - PackedFunc GetFunction(const String& name, const ObjectPtr& sptr_to_self) final; - - /*! - * \return The type key of the executor. - */ - const char* type_key() const final { return "GraphExecutorFactory"; } - - /*! \brief Get the property of the runtime module .*/ - int GetPropertyMask() const final { return ModulePropertyMask::kBinarySerializable; } - - /*! - * \brief Save the module to binary stream. - * \param stream The binary stream to save to. - */ - void SaveToBinary(dmlc::Stream* stream) override; - - /*! - * \brief Create a specific executor module - * \param devs The device of the host and devices where graph nodes will be - * executed on. - * \return created executor module - */ - Module ExecutorCreate(const std::vector& devs); - - /*! - * \brief Create a specific debug executor module - * \param devs The device of the host and devices where graph nodes will be - * executed on. - * \return created debug executor module - */ - Module DebugExecutorCreate(const std::vector& devs); - - /*! - * \brief Create a specific cuda graph executor module - * \param devs The device of the host and devices where graph nodes will be - * executed on. - * \return created cuda graph executor module - */ - Module CudaGraphExecutorCreate(const std::vector& devs); - - /*! - * \brief Set params. - * \param graph_executor The graph executor we want to set the params into. - * \param params The graph params value we want to set. - */ - void SetParams(GraphExecutor* graph_executor, - const std::unordered_map& params) const { - std::unordered_map value = params; - // upload big arrays first to avoid memory issue in rpc mode - std::vector keys; - for (const auto& p : value) { - keys.emplace_back(p.first); - } - std::sort(std::begin(keys), std::end(keys), - [&](const std::string& lhs, const std::string& rhs) -> bool { - auto lhs_size = GetDataSize(*value[lhs].operator->()); - auto rhs_size = GetDataSize(*value[rhs].operator->()); - return lhs_size > rhs_size; - }); - for (const auto& key : keys) { - int in_idx = graph_executor->GetInputIndex(key); - if (in_idx >= 0) { - graph_executor->SetInput(in_idx, const_cast(value[key].operator->())); - } - } - } - - protected: - /*! \brief The execution graph. */ - std::string graph_json_; - /*! \brief The params. */ - std::unordered_map params_; - /*! \brief module name */ - std::string module_name_; -}; - -} // namespace runtime -} // namespace tvm - -#endif // TVM_RUNTIME_GRAPH_EXECUTOR_GRAPH_EXECUTOR_FACTORY_H_ diff --git a/src/runtime/logging.cc b/src/runtime/logging.cc index 844a8bcf1cc2..2d4164ce4425 100644 --- a/src/runtime/logging.cc +++ b/src/runtime/logging.cc @@ -270,7 +270,7 @@ std::string Backtrace() { #else -#include +#include namespace tvm { namespace runtime { diff --git a/src/runtime/meta_data.h b/src/runtime/meta_data.h index 766b93261ac0..257c931df9bf 100644 --- a/src/runtime/meta_data.h +++ b/src/runtime/meta_data.h @@ -27,7 +27,6 @@ #include #include #include -#include #include #include #include @@ -50,15 +49,6 @@ inline String get_name_mangled(const String& module_name, const String& name) { return ss.str(); } -/*! - * \brief Create a metadata module object. - * - * \param metadata Exported metadata structure. - * - * \return The created metadata module. - */ -Module MetadataModuleCreate(metadata::Metadata metadata); - namespace launch_param { /*! \brief A tag to specify whether or not dynamic shared memory is used */ diff --git a/src/runtime/metadata.cc b/src/runtime/metadata.cc deleted file mode 100644 index 40f91d16e1ed..000000000000 --- a/src/runtime/metadata.cc +++ /dev/null @@ -1,143 +0,0 @@ -/* - * 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 src/runtime/metadata.cc - * \brief Defines implementations of TVM metadata which can exist in the runtime. - */ - -#include -#include -#include -#include - -#include - -namespace tvm { -namespace runtime { -namespace metadata { - -TVM_REGISTER_OBJECT_TYPE(MetadataBaseNode); - -ArrayAccessor MetadataNode::inputs() { - return ArrayAccessor(data_->inputs, data_->num_inputs); -} -ArrayAccessor MetadataNode::outputs() { - return ArrayAccessor(data_->outputs, data_->num_outputs); -} -ArrayAccessor MetadataNode::workspace_pools() { - return ArrayAccessor(data_->workspace_pools, - data_->num_workspace_pools); -} -ArrayAccessor MetadataNode::constant_pools() { - return ArrayAccessor(data_->constant_pools, - data_->num_constant_pools); -} - -TVM_REGISTER_OBJECT_TYPE(MetadataBaseNode); - -MetadataArray::MetadataArray(Array array, MetadataKind kind, const char* struct_name) - : MetadataBase{make_object(array, kind, struct_name)} {} - -const char* MetadataArrayNode::get_c_struct_name() const { - ICHECK(false) << "MetadataArrayNode get_c_struct_name is unimplemented"; - return nullptr; -} -TVM_REGISTER_OBJECT_TYPE(MetadataArrayNode); - -Metadata::Metadata(const struct ::TVMMetadata* data) - : MetadataBase{make_object(data)} {} -TVM_REGISTER_OBJECT_TYPE(MetadataNode); - -const char* MetadataNode::get_c_struct_name() const { return "TVMMetadata"; } - -TensorInfo::TensorInfo(const struct ::TVMTensorInfo* data) - : MetadataBase{make_object(data)} {} -TVM_REGISTER_OBJECT_TYPE(TensorInfoNode); - -const char* TensorInfoNode::get_c_struct_name() const { return "TVMTensorInfo"; } - -ConstantInfoMetadata::ConstantInfoMetadata(const struct ::TVMConstantInfo* data) - : MetadataBase{make_object(data)} {} -TVM_REGISTER_OBJECT_TYPE(ConstantInfoMetadataNode); - -const char* ConstantInfoMetadataNode::get_c_struct_name() const { return "TVMConstantInfo"; } - -} // namespace metadata - -class MetadataModuleNode : public ::tvm::runtime::ModuleNode { - public: - explicit MetadataModuleNode(runtime::metadata::Metadata metadata) - : metadata_{::std::move(metadata)} {} - - const char* type_key() const final { return "metadata_module"; } - - /*! \brief Get the property of the runtime module .*/ - int GetPropertyMask() const final { return ModulePropertyMask::kBinarySerializable; } - - static Module LoadFromBinary() { - return Module(make_object(runtime::metadata::Metadata())); - } - - void SaveToBinary(dmlc::Stream* stream) final {} - - PackedFunc GetFunction(const String& name, const ObjectPtr& sptr_to_self) { - if (name == "get_metadata") { - return PackedFunc([this, sptr_to_self](TVMArgs args, TVMRetValue* rv) { - if (!metadata_.defined()) { - TVMFunctionHandle f_handle; - int32_t ret_code = TVMBackendGetFuncFromEnv(this, symbol::tvm_get_c_metadata, &f_handle); - ICHECK_EQ(ret_code, 0) << "Unable to locate " << symbol::tvm_get_c_metadata - << " PackedFunc"; - - TVMValue ret_value; - int ret_type_code; - ret_code = TVMFuncCall(f_handle, nullptr, nullptr, 0, &ret_value, &ret_type_code); - ICHECK_EQ(ret_code, 0) << "Invoking " << symbol::tvm_get_c_metadata - << ": TVMFuncCall returned " << ret_code; - - ICHECK_EQ(ret_type_code, kTVMOpaqueHandle) - << "Expected kOpaqueHandle returned; got " << ret_type_code; - ICHECK(ret_value.v_handle != nullptr) - << symbol::tvm_get_c_metadata << " returned nullptr"; - - metadata_ = runtime::metadata::Metadata( - static_cast(ret_value.v_handle)); - } - - *rv = metadata_; - }); - } - - return PackedFunc(); - } - - private: - runtime::metadata::Metadata metadata_; -}; - -Module MetadataModuleCreate(metadata::Metadata metadata) { - return Module(make_object(metadata)); -} - -TVM_REGISTER_GLOBAL("runtime.module.loadbinary_metadata_module") - .set_body([](TVMArgs args, TVMRetValue* rv) { *rv = MetadataModuleNode::LoadFromBinary(); }); - -} // namespace runtime -} // namespace tvm diff --git a/src/runtime/pipeline/pipeline_executor.cc b/src/runtime/pipeline/pipeline_executor.cc deleted file mode 100644 index a0013742932f..000000000000 --- a/src/runtime/pipeline/pipeline_executor.cc +++ /dev/null @@ -1,290 +0,0 @@ -/* - * 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 pipeline_executor.cc - */ -#include "pipeline_executor.h" -namespace tvm { -namespace runtime { -/*! - * \brief Give frontends an access to packed functions. - * \param name The name of the function. - * \param sptr_to_self The pointer to the module node. - * \return The corresponding packed function. - */ -PackedFunc PipelineExecutor::GetFunction(const String& name, - const ObjectPtr& sptr_to_self) { - if (name == "get_num_outputs") { - return PackedFunc( - [sptr_to_self, this](TVMArgs args, TVMRetValue* rv) { *rv = this->NumOutputs(); }); - } else if (name == "get_num_inputs") { - return PackedFunc( - [sptr_to_self, this](TVMArgs args, TVMRetValue* rv) { *rv = this->NumInputs(); }); - } else if (name == "get_input_pipeline_map") { - return PackedFunc([sptr_to_self, this](TVMArgs args, TVMRetValue* rv) { - if (String::CanConvertFrom(args[0])) { - *rv = this->GetInputPipeplineMap(args[0].operator String()); - } else { - LOG(FATAL) << "Function only support the input name value in the form of string"; - } - }); - } else if (name == "get_params_group_pipeline_map") { - return PackedFunc([sptr_to_self, this](TVMArgs args, TVMRetValue* rv) { - if (String::CanConvertFrom(args[0])) { - *rv = this->GetParamsGroupPipelineMap(args[0].operator String()); - } else { - LOG(FATAL) << "Function only support the input name value in the form of string"; - } - }); - } else if (name == "set_param") { - return PackedFunc([sptr_to_self, this](TVMArgs args, TVMRetValue* rv) { - if (String::CanConvertFrom(args[0]) && String::CanConvertFrom(args[1])) { - this->SetParam(args[0].operator String(), args[1].operator String(), args[2]); - } else { - LOG(FATAL) << "Function only support the parameter name and the key in the form of string"; - } - }); - } else if (name == "set_input") { - return PackedFunc([sptr_to_self, this](TVMArgs args, TVMRetValue* rv) { - if (String::CanConvertFrom(args[0])) { - this->SetInput(args[0].operator String(), args[1]); - } else { - LOG(FATAL) << "Function only support the input name value in the form of string"; - } - }); - } else if (name == "get_input") { - return PackedFunc([sptr_to_self, this](TVMArgs args, TVMRetValue* rv) { - if (String::CanConvertFrom(args[0])) { - *rv = this->GetInput(args[0].operator String()); - } else { - LOG(FATAL) << "Function only support the input name value in the form of string"; - } - }); - } else if (name == "get_output") { - return PackedFunc( - [sptr_to_self, this](TVMArgs args, TVMRetValue* rv) { *rv = this->GetOutput(); }); - } else if (name == "run") { - return PackedFunc([sptr_to_self, this](TVMArgs args, TVMRetValue* rv) { this->Run(); }); - } else if (name == "get_execute_count") { - return PackedFunc( - [sptr_to_self, this](TVMArgs args, TVMRetValue* rv) { *rv = this->GetExecutionCount(); }); - } else { - LOG(FATAL) << "Unknown packed function: " << name; - } -} -/*! - * brief Returns number of global inputs. - */ -int PipelineExecutor::NumInputs(void) { return input_connection_config_.GetInputNum(); } -/*! - * \brief set input to the runtime module. - * \param input_name The input name. - * \param data_in The input data. - */ -void PipelineExecutor::SetInput(std::string input_name, DLTensor* data_in) { - global_runtime_->SetPipelineInput(input_name, data_in); -} -/*! - * \brief get input from the runtime module. - * \param input_name The input name. - * \return Return the input data for a specific input name. - */ -NDArray PipelineExecutor::GetInput(std::string input_name) { - std::pair indexs = this->GetInputIndex(input_name); - if (indexs.first < 0 || indexs.first >= static_cast(runtimes_.size())) { - LOG(FATAL) << "input name " << input_name << " not found."; - } - return runtimes_[indexs.first]->GetInput(indexs.second); -} -/*! - * \brief Getting a module index via a input parameters group name. - * \param name The parameters group name. - * \return int The module index. - */ -int PipelineExecutor::GetParamModuleIndex(const std::string& name) { - return param_connection_config_[name]; -} -/*! - * \brief Using the global input name to get the index, and also get the input interface name - of corresponding subgraph from the input connection configuration. - * \param The global input name. - * \return Returning the index and the input interface name of corresponding subgraph. - */ -Array PipelineExecutor::GetInputPipeplineMap(std::string input_name) { - std::pair map = input_connection_config_[input_name]; - return {std::to_string(map.first), map.second}; -} - -/*! - * \brief Return the module index for the parameters group name. - * \param name The parameters group name. - * \return int The module index. - */ -int PipelineExecutor::GetParamsGroupPipelineMap(const std::string& name) { - return param_connection_config_[name]; -} - -/*!\brief Run the pipeline executor.*/ -void PipelineExecutor::Run() { pipeline_scheduler_.PipelineRun(runtimes_); } -/*! - * \brief return A list of global output data. - */ -Array PipelineExecutor::GetOutput(void) { return pipeline_scheduler_.PipelineGetOutput(); } -/*! - * \brief Use the mod_config information to create a graph runtime list. - * \param mod_config The config information that generates by the export library function call. - */ -std::vector PipelineExecutor::CreateGraphModules(const ModuleConfig& mod_config) { - const PackedFunc* graph_executor_create = Registry::Get("tvm.graph_executor.create"); - std::vector ret; - ret.resize(mod_config.size()); - for (auto config : mod_config) { - // Load library. - auto lib = Module::LoadFromFile(config.second.lib_name.c_str()); - - // Read json. - std::ifstream ifJson(config.second.json_name.c_str()); - if (ifJson.fail()) { - LOG(FATAL) << "json file not found: " << config.second.json_name; - } - const std::string json((std::istreambuf_iterator(ifJson)), - std::istreambuf_iterator()); - - // Create a graph executor. - std::istringstream istr(config.second.dev); - std::string str; - int device_type = 1, device_id = 0; - while (getline(istr, str, ';')) { - std::istringstream istr_dev(str); - std::string str_temp; - if (getline(istr_dev, str_temp)) { - device_type = stoi(str_temp); - } - if (getline(istr_dev, str_temp)) { - device_id = stoi(str_temp); - } - } - Module graph_module = (*graph_executor_create)(json, lib, device_type, device_id); - - // Load parameters. - TVMByteArray params_arr; - const char* params_file_name = config.second.params_name.c_str(); - std::ifstream if_param(params_file_name); - if (if_param.fail()) { - LOG(FATAL) << "params file not found: " << params_file_name; - } - const std::string params((std::istreambuf_iterator(if_param)), - std::istreambuf_iterator()); - params_arr.data = params.c_str(); - params_arr.size = params.length(); - auto load_params = graph_module.GetFunction("load_params"); - load_params(params_arr); - - // Put a graph executor module into the vector. - ret[config.first] = graph_module; - } - return ret; -} -/*! - * \brief Set a parameter into a graph module. - * \param param_group_name The parameters group name. - * \param param_key_name The parameter key name. - * \param data_in The parameter data. - */ -void PipelineExecutor::SetParam(std::string param_group_name, std::string param_key_name, - DLTensor* data_in) { - // Get the module index via the parameters group name. - int module_index = this->GetParamModuleIndex(param_group_name); - ICHECK(module_index >= 0 && module_index < static_cast(runtimes_.size())) - << "Parameter group name " << param_group_name << " does not exist."; - auto runtime = runtimes_[module_index]; - // Get the parameter index via the param key name - int index = runtime->GetInputIndex(param_key_name); - ICHECK(index >= 0) << "Parameter name " << param_key_name << " does not exist in module " - << module_index; - runtime->SetInput(index, data_in); -} -/*! - * \brief Return the input index and module index for a given input name. - * \param name The input name. - * \return std::pair A pair of module index and the input index. - */ -std::pair PipelineExecutor::GetInputIndex(const std::string& name) { - std::pair index = input_connection_config_[name]; - auto gruntime = runtimes_[index.first]; - return std::make_pair(index.first, gruntime->GetInputIndex(index.second)); -} -/*! - * \brief Getting the count of running pipeline. - */ -int PipelineExecutor::GetExecutionCount() { return runtimes_.back()->GetExecutionCount(); } -/*! - * \brief Initialize the pipeline executor with a list of modules to be pipelined - * and config in JSON format. - * \param modules The module list used for building the pipeline. - * \param pipeline_json The configuration of modules dependencies. - */ -void PipelineExecutor::Init(const std::vector& modules, const std::string& pipeline_json) { - ICHECK(!modules.empty()) << "The graph executor module list is empty."; - // Use JSONReader to load pipeline configuration. - std::istringstream is(pipeline_json); - dmlc::JSONReader reader(&is); - this->LoadConfig(&reader); - ICHECK(!pipeline_config_.Empty()) << "The pipeline config information is empty."; - num_outputs_ = pipeline_config_.GetGlobalOutputNum(); - // Initialize the pipeline function class used for pipeline thread pool management - // and schedule etc. This function returns a list of runtime. - global_runtime_ = - pipeline_scheduler_.PipelineInit(modules, pipeline_config_, input_connection_config_); - runtimes_ = global_runtime_->GetRuntimeList(); - return; -} - -Module PipelineExecutorCreate(const Array& m, const std::string& pipeline_json) { - ICHECK(!m.empty()) << "The module list is empty."; - auto exec = make_object(); - std::vector graph_modules; - for (auto mod : m) { - graph_modules.push_back(mod); - } - exec->Init(graph_modules, pipeline_json); - return Module(exec); -} - -Module PipelineExecutorLoad(const std::string& load_json, const std::string& pipeline_json) { - auto exec = make_object(); - std::istringstream is(load_json); - dmlc::JSONReader reader(&is); - ModuleConfig& mod_config = exec->LoadModuleConfig(&reader); - ICHECK(!mod_config.empty()) << "The module config is empty."; - std::vector modules = exec->CreateGraphModules(mod_config); - exec->Init(modules, pipeline_json); - return Module(exec); -} - -TVM_REGISTER_GLOBAL("tvm.pipeline_executor.create").set_body([](TVMArgs args, TVMRetValue* rv) { - *rv = PipelineExecutorCreate(args[0], args[1]); -}); - -TVM_REGISTER_GLOBAL("tvm.pipeline_executor.load").set_body([](TVMArgs args, TVMRetValue* rv) { - *rv = PipelineExecutorLoad(args[0], args[1]); -}); -} // namespace runtime -} // namespace tvm diff --git a/src/runtime/pipeline/pipeline_executor.h b/src/runtime/pipeline/pipeline_executor.h deleted file mode 100644 index d9058871e7b9..000000000000 --- a/src/runtime/pipeline/pipeline_executor.h +++ /dev/null @@ -1,210 +0,0 @@ -/* - * 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. - */ - -/*! - * \brief pipeline executor - * \file pipeline_executor.h - */ -#ifndef TVM_RUNTIME_PIPELINE_PIPELINE_EXECUTOR_H_ -#define TVM_RUNTIME_PIPELINE_PIPELINE_EXECUTOR_H_ - -#include -#include - -#include -#include -#include -#include -#include -#include -#include - -#include "pipeline_scheduler.h" -namespace tvm { -namespace runtime { -/*! - * \brief pipeline executor. - * This executor class use the module list and dependency configuration of modules as - * the parameters and executes these modules on heterogeneous targets in a pipeline - * parallel manner to improve throughput. - * - * This executor can be accessed by various language via TVM runtime PackedFunc API. - */ -class TVM_DLL PipelineExecutor : public ModuleNode { - public: - /*! - * \Return the type key of the executor. - */ - const char* type_key() const final { return "PipelineExecutor"; } - /*! - * \brief Initialize the pipeline executor with module array and JSON text. - * \param modules The module list used for building pipeline. - * \param pipeline_json The configuration of modules dependencies. - */ - void Init(const std::vector& modules, const std::string& pipeline_json); - /*! - * \brief Use the information of mod_config to create a list of graph executor. - * \param mod_config The configuration information generated by the library export function call. - */ - std::vector CreateGraphModules(const ModuleConfig& mod_config); - /*! - * \brief Give frontends an access to packed functions. - * \param name The name of the function. - * \param sptr_to_self The pointer to the module node. - * \return The corresponding packed function. - */ - virtual PackedFunc GetFunction(const String& name, const ObjectPtr& sptr_to_self); - /*! - * \brief Using the global input name to get the index, and also get the input interface name - of corresponding subgraph from the input connection configuration. - * \param The global input name. - * \return Returning the index and the input interface name of corresponding subgraph. - */ - Array GetInputPipeplineMap(std::string input_name); - /*! - * \brief This function return a module index for the global parameters group name. - * \param name The parameters group name. - * \return Returning a runtime module index. - */ - int GetParamsGroupPipelineMap(const std::string& name); - /*! - * \brief Use the input name to set the input data of pipeline executor. - * \param input_name The input name. - * \param data_in The input data. - */ - void SetInput(std::string input_name, DLTensor* data_in); - /*! - * \brief Use the input name to get the input data. - * \param input name The input name. - * \return Return input data. - */ - NDArray GetInput(std::string input_name); - /*! - * \brief Getting the count of running pipeline. - */ - int GetExecutionCount(); - /*! - * \brief Use the parameters group name to get the specific backend runtime then use - * the param_key_name to set param data for the said backend runtime. - * \param param_group_name The parameters group name. - * \param param_key_name The parameter key name. - * \param data_in The parameter value. - */ - void SetParam(std::string param_group_name, std::string param_key_name, DLTensor* data_in); - /*! - * \brief Get the number of outputs. - * - * \return The number of outputs. - */ - int NumOutputs() const { return num_outputs_; } - /*!\brief Run the pipeline executor.*/ - void Run(); - int NumInputs(); - /*! - * \brief Get a list output data. - * \return A list of output data. - */ - Array GetOutput(); - /*! - * \brief A pipeline params with a specific name correspond with the params of a specific - * backend module, this function return the module index for the params name. - * \param name The parameters group name. - * \return Return backend runtime module index. - */ - int GetParamModuleIndex(const std::string& name); - /*! - * \brief A pipeline input with a specific name correspond with a input of a specific - * backend module, this function return a module index and a input index in "pair" - * form for a input name. - * return Return a module index and a input index. - */ - std::pair GetInputIndex(const std::string& name); - /*!\brief Load the module files information.*/ - ModuleConfig& LoadModuleConfig(dmlc::JSONReader* reader) { - reader->BeginArray(); - while (reader->NextArrayItem()) { - std::string key; - reader->BeginObject(); - int mod_idx = -1; - std::string lib_name; - std::string json_name; - std::string params_name; - std::string dev; - while (reader->NextObjectItem(&key)) { - if (key == "mod_idx") { - reader->Read(&mod_idx); - } else if (key == "lib_name") { - reader->Read(&lib_name); - } else if (key == "json_name") { - reader->Read(&json_name); - } else if (key == "params_name") { - reader->Read(¶ms_name); - } else if (key == "dev") { - reader->Read(&dev); - } else { - LOG(FATAL) << "do not support key " << key; - } - } - ICHECK(mod_idx >= 0) << "Invalid mod_idx value " << mod_idx; - // Load the lib, json, and params information. - ICHECK(!lib_name.empty()) << "lib_name is empty."; - ICHECK(!json_name.empty()) << "json_name is empty."; - ICHECK(!params_name.empty()) << "params_name is empty."; - mod_config_[mod_idx] = GraphModuleLoadInfo(lib_name, json_name, params_name, dev); - } - return mod_config_; - } - - private: - /*!\brief The class used to execute and schedule the pipeline logic.*/ - PipelineScheduler pipeline_scheduler_; - /*!\brief The dependency information of each graph runtime module of the pipeline.*/ - ConfigPipelineExecution pipeline_config_; - /*!\brief The map of global input and subgraph input.*/ - InputConnectionConfig input_connection_config_; - /*!\brief The map includes global parameters groups and runtime modules.*/ - ParamConnectionConfig param_connection_config_; - /*!\brief The module information used to create the graph runtimes.*/ - ModuleConfig mod_config_; - /*!\brief How many outputs are in this pipeline executor.*/ - size_t num_outputs_ = 0; - /*!The list of backend runtime module.*/ - std::vector> runtimes_; - std::shared_ptr global_runtime_; - /*!\brief Json loader.*/ - void LoadConfig(dmlc::JSONReader* reader) { - reader->BeginObject(); - std::string key; - while (reader->NextObjectItem(&key)) { - if (key == "module_connection") { - reader->Read(&pipeline_config_); - } else if (key == "input_connection") { - reader->Read(&input_connection_config_); - } else if (key == "param_connection") { - reader->Read(¶m_connection_config_); - } else { - LOG(FATAL) << "do not support key " << key; - } - } - return; - } -}; -} // namespace runtime -} // namespace tvm -#endif // TVM_RUNTIME_PIPELINE_PIPELINE_EXECUTOR_H_ diff --git a/src/runtime/pipeline/pipeline_scheduler.cc b/src/runtime/pipeline/pipeline_scheduler.cc deleted file mode 100644 index bc5e060d849f..000000000000 --- a/src/runtime/pipeline/pipeline_scheduler.cc +++ /dev/null @@ -1,77 +0,0 @@ -/* - * 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 "pipeline_scheduler.h" - -#include -#include -#include -namespace tvm { -namespace runtime { -/*! - * \brief Initialize the pipeline. - * \param modules The list of graph executor modules. - * \param pipeline_conf The dependency information of each graph executor module. - */ -std::shared_ptr PipelineScheduler::PipelineInit( - const std::vector& modules, const ConfigPipelineExecution& pipeline_config, - const InputConnectionConfig& input_connection_config) { - std::vector> runtimes; - graph_modules_ = modules; - // Creating a list of runtimes. - for (size_t i = 0; i < graph_modules_.size(); i++) { - auto run_item = std::make_shared(graph_modules_[i], i); - runtimes.push_back(run_item); - } - // Creating the global runtime to represent the pipeline executor. - global_runtime_ = std::make_shared(GLOBAL_MODULE_INDEX); - // Initializing the data structures used by pipeline logic. - global_runtime_->InitializePipeline(input_connection_config, runtimes); - // Creating a list of NDArray in order to storage the outputs data. - auto global_output_map = pipeline_config.GetGlobalConfigOutputBindings(); - for (size_t i = 0; i < global_output_map.size(); i++) { - if (global_output_map.find(i) == global_output_map.end()) { - LOG(FATAL) << "Not find global output index " << i; - } - ModuleOutputPair& output_pair = global_output_map[i]; - NDArray output = runtimes[output_pair.first]->CreateFromOutput(output_pair.second); - output_arrays_.push_back(output); - } - // Initializing and then running the worker thread. - for (auto runtime : runtimes) { - runtime->InitializePipeline(pipeline_config, &runtimes, global_runtime_); - } - return global_runtime_; -} -/*! - * \brief Running pipeline logic. - * \param runtimes A list of backend runtime modules. - * \param pipeline_config The dependency configuration of each runtime module. - */ -void PipelineScheduler::PipelineRun(const std::vector>& runtimes) { - runtimes.front()->RunPipeline(); -} -/*! - * \brief Get a list of output. - */ -Array PipelineScheduler::PipelineGetOutput() { - bool ret = global_runtime_->GetOutput(&output_arrays_); - return ret ? output_arrays_ : Array{}; -} -} // namespace runtime -} // namespace tvm diff --git a/src/runtime/pipeline/pipeline_scheduler.h b/src/runtime/pipeline/pipeline_scheduler.h deleted file mode 100644 index 1141af26f57b..000000000000 --- a/src/runtime/pipeline/pipeline_scheduler.h +++ /dev/null @@ -1,67 +0,0 @@ -/* - * 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. - */ -#ifndef TVM_RUNTIME_PIPELINE_PIPELINE_SCHEDULER_H_ -#define TVM_RUNTIME_PIPELINE_PIPELINE_SCHEDULER_H_ -#include -#include -#include - -#include -#include -#include -#include - -#include "pipeline_struct.h" -namespace tvm { -namespace runtime { -/*! - * \brief The class that executes the pipeline logic,it is used to initialize the thread pool, - execute and schedule pipeline tasks, allocate and manage memory, etc. - */ -class PipelineScheduler { - public: - /*! - * \brief Initialize the pipeline. - * \param modules The list of graph executor module. - * \param pipeline_config The dependency information of each graph executor module. - */ - std::shared_ptr PipelineInit(const std::vector& modules, - const ConfigPipelineExecution& pipeline_config, - const InputConnectionConfig& input_connection_config); - /*! - * \brief Running the pipeline logic. - * \param runtimes A list of backend runtime modules. - */ - void PipelineRun(const std::vector>& runtimes); - /*! - * \brief Get a list of outputs. - */ - Array PipelineGetOutput(); - - private: - /*!\brief The list of graph executors.*/ - std::vector graph_modules_; - /*!\brief A list of NDArray used to storage outputs.*/ - Array output_arrays_; - /*!\brief The global runtime to represent the pipeline executor.*/ - std::shared_ptr global_runtime_; -}; -} // namespace runtime -} // namespace tvm -#endif // TVM_RUNTIME_PIPELINE_PIPELINE_SCHEDULER_H_ diff --git a/src/runtime/pipeline/pipeline_struct.h b/src/runtime/pipeline/pipeline_struct.h deleted file mode 100644 index 9f14d9163c7e..000000000000 --- a/src/runtime/pipeline/pipeline_struct.h +++ /dev/null @@ -1,1212 +0,0 @@ -/* - * 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. - */ -#ifndef TVM_RUNTIME_PIPELINE_PIPELINE_STRUCT_H_ -#define TVM_RUNTIME_PIPELINE_PIPELINE_STRUCT_H_ -#include -#include -#include -#include -#include -#include - -#include -#include -#include -#include -#include -#include -#include -#include -#include -#include - -#include "spsc_queue.h" -namespace tvm { -namespace runtime { -#define GLOBAL_MODULE_INDEX -1 -/*! - *\brief The function is used to build the binding configuration for a runtime. The first - * 'int' is the output index of the current runtime, the second 'int' is the index of child - * runtime, and the 'string' is the input name of child runtime. - */ -using BindingConfigParseFunc = std::function; -/*! - *\brief The 'pair' includes the module output index and the global output index. - * The first 'int' is the module output index, and the second 'int' is the global output index. - */ -using GlobalOutputPair = std::pair; -/*! - *\brief The pair includes the module index and the module output index. - * The first 'int' is the module index, and the second 'int' is the module output index. - */ -using ModuleOutputPair = std::pair; -/*! - *\brief The pair includes the runtime module index and the module input index. - * The first 'int' is the module index, and the second 'int' is the module input index. - */ -using ModuleInputPair = std::pair; -/*!\brief The runtime module interface type.*/ -enum InterfaceType { - INPUT = 0, - OUTPUT, -}; -/*!\brief The state of the pipeline.*/ -enum PipelineState { - STOPPED = 0, - RUNNING, - STOPPING, -}; -/*! - *\brief The structure includes the module index and the module output index. - */ -struct ModuleInterfaceID { - ModuleInterfaceID() { SetID(0, 0, INPUT); } - ModuleInterfaceID(int runtime_index, int runtime_interface_index, InterfaceType type = INPUT) { - SetID(runtime_index, runtime_interface_index, type); - } - /*! - * \brief Set the value of ID. - * \param runtime_index The index of runtime. - * \param runtime_interface_index The index of interface. - * \param type The type of the interface. - */ - void SetID(int runtime_index, int runtime_interface_index, InterfaceType type) { - runtime_idx = runtime_index; - runtime_interface_idx = runtime_interface_index; - interface_type = type; - } - int runtime_idx; - union { - /*!\brief The output interface index.*/ - int runtime_output_idx; - /*!\brief The input interface index.*/ - int runtime_input_idx; - /*!\brief The interface index.*/ - int runtime_interface_idx; - }; - /*!\brief The interface type*/ - InterfaceType interface_type; - ModuleInterfaceID& operator=(const struct ModuleInterfaceID& id) { - SetID(id.runtime_idx, id.runtime_interface_idx, id.interface_type); - return *this; - } - bool operator==(const struct ModuleInterfaceID& id) const { - return id.interface_type == interface_type && - id.runtime_interface_idx == runtime_interface_idx && id.runtime_idx == runtime_idx; - } -}; -/*!brief The hash function used to generate the hash value for the "ModuleInterfaceID" variable.*/ -struct ModuleIDHash { - bool operator()(const ModuleInterfaceID& id) const { - int offset = sizeof(std::size_t) / 3; - return id.interface_type | id.runtime_interface_idx << offset | id.runtime_idx << offset * 2; - } -}; -/*!\brief The data notification structure.*/ -class DataNotify { - private: - /*!\brief The 'contitional variable' is used to wait for notification.*/ - std::condition_variable notify_cv_; - /*!\brief The mutex is used to protect the 'conditional variable'.*/ - std::mutex mutex_; - /*!\brief Whether a data is ready or not.*/ - bool data_ready_ = false; - /*!\brief Whether the thread should exit or not.*/ - std::atomic exit_state_{false}; - /*!\brief The 'ModuleInterfaceID' of an interface which sent this notification.*/ - ModuleInterfaceID notification_source_; - - public: - /*! - * \brief Constructing the DataNotify class. - * \param source_interface_id The id of a runtime interface which is sending out the data - * notification. - */ - explicit DataNotify(ModuleInterfaceID source_interface_id) { - notification_source_ = source_interface_id; - } - /*! - * \brief Getting the notification target. - * \return The ID of the interface which is sending out the notification. - */ - ModuleInterfaceID GetNotifySource(void) { return notification_source_; } - /*! - *\brief Waiting for the notification. - *\return Returning the value 'false' when the notification is in a 'exit' state, else - * return true. - */ - bool Wait(void) { - std::unique_lock lock(mutex_); - notify_cv_.wait(lock, [&] { return this->data_ready_; }); - data_ready_ = false; - return !GetExitState(); - } - /*!brief Sending the notification in which the related data is ready.*/ - void Notify(void) { - { - std::lock_guard lock(mutex_); - data_ready_ = true; - } - notify_cv_.notify_one(); - } - /*!brief Sending the notification when the notification state changes into 'exit'.*/ - void ExitNotify(void) { - exit_state_.store(true, std::memory_order_release); - Notify(); - } - /*! - *\brief Getting the 'exit state'. - *\return Returning the value of 'exit_state_' - */ - bool GetExitState(void) { return exit_state_.load(std::memory_order_acquire); } -}; -/*!\brief The container used to store the forwarding data of the pipeline.*/ -class QueueData { - public: - explicit QueueData(DLTensor* data) { - if (data_ == data) { - LOG(FATAL) << "The value of 'data'(" << data << ") is the same as 'data_'(" << data_ << ")"; - } - data_ = data; - SetAsDataOwner(false); - } - QueueData() { SetAsDataOwner(true); } - /*!\brief Doing a deep copy for the 'QueueData' structure.*/ - QueueData& operator=(const QueueData& data) { - CreateCopyFrom(data.GetDLData()); - return *this; - } - QueueData& operator=(const NDArray& from) { - CreateCopyFrom(const_cast(from.operator->())); - return *this; - } - QueueData& operator=(const DLTensor* from) { - CreateCopyFrom(from); - return *this; - } - /*!\brief Create a deep copy of the 'DLTensor' data.*/ - DLTensor* CreateCopyFrom(const DLTensor* from) { - if (!from) { - LOG(FATAL) << "the 'from' pointer is a null pointer!"; - } - size_t fromLen = tvm::runtime::GetDataSize(*from); - size_t toLen = data_ ? tvm::runtime::GetDataSize(*data_) : 0; - if (fromLen != toLen) { - // If this container ownes the variable 'data_', then recreating the 'data_' variable. - if (IsDataOwner()) { - if (data_) { - TVMArrayFree(data_); - data_ = nullptr; - } - TVMArrayAlloc(from->shape, from->ndim, from->dtype.code, from->dtype.bits, - from->dtype.lanes, from->device.device_type, from->device.device_id, &data_); - } else { - LOG(FATAL) << "The 'from' data is not matched with the 'data_'."; - } - } - TVMArrayCopyFromTo(const_cast(from), data_, nullptr); - return data_; - } - /*!\brief Return a pointer to the 'DLTensor' data.*/ - DLTensor* GetDLData() const { return data_; } - ~QueueData() { - if (IsDataOwner() && data_) { - TVMArrayFree(data_); - data_ = nullptr; - } - } - - private: - /*!\brief Pointer to the forwarding data.*/ - DLTensor* data_ = nullptr; - /*!\brief Whether this container is the owner of the 'data_'.*/ - bool is_data_owner_ = false; - /*!\brief Set the current container as the owner of the 'data_'.*/ - void SetAsDataOwner(bool is_owner) { is_data_owner_ = is_owner; } - /*!Check whether the current container is the owner of the 'data_'.*/ - bool IsDataOwner() const { return is_data_owner_; } -}; -/*! - * \brief All binding information of an output interface. - */ -class ConfigBindings { - public: - /*!\brief Whether this binding is bound to the PipelineExecutor output interface.*/ - bool IsGlobalOutput() const { return global_output_index_ > -1; } - /*!\brief Getting the global output index in the current binding.*/ - int GetGlobalOutputIndex() const { return global_output_index_; } - /*!\brief Returning the binding configuration.*/ - std::unordered_map& Get() { return bindings_; } - /*! - * \brief Enumerating the binding configuration. - * \param parse_function The function is used to parse the binding configuration. - * \param output_idx The index of output interface is used for parsing. - */ - void VisitOutput(BindingConfigParseFunc parse_function, int output_idx) { - for (auto output : bindings_) { - parse_function(output_idx, output.first, output.second); - } - if (IsGlobalOutput()) { - parse_function(output_idx, GLOBAL_MODULE_INDEX, std::to_string(global_output_index_)); - } - } - /*! - * \brief Create a module interface map from JSONReader. - * \param reader JSON reader. - */ - void Load(dmlc::JSONReader* reader) { - reader->BeginArray(); - while (reader->NextArrayItem()) { - std::string key; - reader->BeginObject(); - std::string input_name; - int mod_idx = std::numeric_limits::min(); - // Whether the output binding is global. - bool global_binding = false; - while (reader->NextObjectItem(&key)) { - if (key == "mod_idx") { - reader->Read(&mod_idx); - } else if (key == "input_name") { - reader->Read(&input_name); - } else if (key == "global_output_index") { - // There should be only one global binding. - ICHECK(global_output_index_ < 0); - reader->Read(&global_output_index_); - // When the key value is 'global_output_index', it means that this output is bound to - // a global interface. - global_binding = true; - } else { - LOG(FATAL) << "do not support key " << key; - } - } - // When this output is bound to a global interface, check if the global interface index - // start from 0. - if (global_binding) { - ICHECK(global_output_index_ >= 0); - } else { - // When this output is bound to a graph executor module interface, check if the module - // index start from 0. - ICHECK(mod_idx >= 0); - bindings_[mod_idx] = input_name; - } - } - } - - private: - /*!\brief Output interface binding information, 'int' is the index of the module that - * uses this output data as the input interface data, 'string' is the input interface name - * of the module. - */ - std::unordered_map bindings_; - /*! The index value of the global interface to which the current output are bound.*/ - int global_output_index_ = std::numeric_limits::min(); -}; -/*! - * \brief The binding information of all outputs of a module. - */ -class ConfigRuntime { - public: - ConfigRuntime& operator=(const ConfigRuntime& output) { - output_binding_map_ = output.GetOutBindings(); - cpu_affinity_ = output.GetCPUAffinity(); - return *this; - } - - ConfigBindings& operator[](const int key) { - ICHECK(output_binding_map_.find(key) != output_binding_map_.end()); - return output_binding_map_[key]; - } - /*! - * \brief Store the CPU affinity settings. - * \param cpu_affinity The CPU affinity settings in the text form. - */ - void StoreCPUAffinity(std::string cpu_affinity) { cpu_affinity_ = cpu_affinity; } - /*! - * \brief Getting the setting of the cpu affinity. - * \param Returning the cpu affinity in text form. - */ - std::string GetCPUAffinity() const { return cpu_affinity_; } - /*! - * \brief Enumerating the output configuration. - * \param parse_function The callback function is used to parse the binding configeration. - */ - void VisitOutputConfig(BindingConfigParseFunc parse_function) { - for (auto output : output_binding_map_) { - output.second.VisitOutput(parse_function, output.first); - } - } - /*!brief Return the variable "output_binding_map_".*/ - std::unordered_map GetOutBindings() const { return output_binding_map_; } - /*! - *\brief This function is used to verify whether ConfigRuntime is successfully loaded. - *\return Return true to indicate that this class has not been successfully loaded. - */ - bool Empty() { return output_binding_map_.empty(); } - /*! - * \brief The pipeline outputs is the final outputs of pipeline, this function is used to - * get how many pipeline outputs are in this Outputmap - * \return Number of pipeline outputs. - */ - size_t GetGlobalOutputNum(void) const { - size_t num_output = 0; - for (auto bindings : output_binding_map_) { - num_output += bindings.second.IsGlobalOutput() ? 1 : 0; - } - return num_output; - } - /*! - *\brief Getting the map which includes the global outputs and the current module outputs. - *\return A list of "GlobalOutputPair". - */ - std::vector GetGlobalConfigOutputBindings(void) const { - std::vector ret; - for (auto bindings : output_binding_map_) { - if (bindings.second.IsGlobalOutput()) { - ret.push_back(GlobalOutputPair(bindings.first, bindings.second.GetGlobalOutputIndex())); - } - } - return ret; - } - /*! - * \brief Create an output binding map from JSONReader. - * \param reader Json reader. - */ - void Load(dmlc::JSONReader* reader) { - reader->BeginArray(); - while (reader->NextArrayItem()) { - std::string key; - reader->BeginObject(); - int output_idx = -1; - ConfigBindings binding; - while (reader->NextObjectItem(&key)) { - if (key == "output_idx") { - reader->Read(&output_idx); - } else if (key == "dependencies") { - reader->Read(&binding); - } else { - LOG(FATAL) << "do not support key " << key; - } - } - ICHECK(output_idx >= 0); - output_binding_map_[output_idx] = binding; - } - } - - private: - /*!\brief The map of output binding, 'int' is the output interface index.*/ - std::unordered_map output_binding_map_; - /*!\brief The cpu affinity setting for the tvm thread pool.*/ - std::string cpu_affinity_; -}; - -/*! - * \brief The binding or dependency information of each module output interface. - */ -class ConfigPipelineExecution { - public: - ConfigRuntime& operator[](int key) { - ICHECK(config_.find(key) != config_.end()); - return config_[key]; - } - /*Get the cpu affinity settings.*/ - std::string GetCPUAffinity(int runtime_idx) { - auto config = config_.find(runtime_idx); - if (config == config_.end()) { - LOG(FATAL) << "Do not finding the runtime " << runtime_idx; - } - auto config_runtime = config->second; - return config_runtime.GetCPUAffinity(); - } - /*! - * \brief Enumerating the binding configuration for a specified runtime. - * \param parse_function The callback function is used to parse the binding configuration. - * \param runtime_index The index of a runtime is used to parse the binding configuration. - */ - void VisitRuntimeOutputConfig(BindingConfigParseFunc parse_function, int runtime_index) { - auto config = config_.find(runtime_index); - if (config == config_.end()) { - LOG(FATAL) << "Do not finding the runtime " << runtime_index; - } - config->second.VisitOutputConfig(parse_function); - } - /* - *!\brief This function is used to verify whether config is loaded successfully. - * \return Return "true" to indicate that this class has not been successfully loaded. - */ - bool Empty() { return config_.empty(); } - /*! - *\brief Check if the module index existing in the "config". - */ - bool FindModuleInConfig(int mod_idx) { return config_.find(mod_idx) != config_.end(); } - /*! - * \brief Getting the number of global outputs. - * \return The number of outputs in the entire pipeline. - */ - size_t GetGlobalOutputNum() const { - size_t num_output = 0; - for (auto mod_output : config_) { - num_output += mod_output.second.GetGlobalOutputNum(); - } - return num_output; - } - /* - *!\brief Get the map of global outputs and module outputs. - */ - std::unordered_map GetGlobalConfigOutputBindings(void) const { - return global_output_map_; - } - /* - *!\brief Parsing the configuration. - */ - void ParseConfiguration(const std::unordered_map& config) { - if (config.empty()) { - LOG(FATAL) << "The Configuration loading not finish yet."; - } - for (auto mod_output : config) { - // Using the global output index as the key to create a map including global index and - // module output index. - const std::vector& global_output = - mod_output.second.GetGlobalConfigOutputBindings(); - - for (auto output : global_output) { - global_output_map_[output.second] = ModuleOutputPair(mod_output.first, output.first); - } - } - return; - } - /*! - * \brief Create a pipeline config from JSONReader. - * \param reader Json reader. - */ - void Load(dmlc::JSONReader* reader) { - reader->BeginArray(); - while (reader->NextArrayItem()) { - std::string key; - reader->BeginObject(); - int mod_idx = -1; - ConfigRuntime output; - std::string dev; - std::string cpu_affinity; - while (reader->NextObjectItem(&key)) { - if (key == "mod_idx") { - reader->Read(&mod_idx); - } else if (key == "dev") { - reader->Read(&dev); - } else if (key == "output") { - reader->Read(&output); - } else if (key == "cpu_affinity") { - reader->Read(&cpu_affinity); - } else { - LOG(FATAL) << "do not support key " << key; - } - } - ICHECK(mod_idx >= 0) << "Invalid mod_idx value " << mod_idx; - // Check if the output is successfully read. - ICHECK(!output.Empty()) << "Invalid output binding result."; - // Store the cpu affinity into the 'ConfigRuntime' structure. - output.StoreCPUAffinity(cpu_affinity); - // Build the mapping of mod_idx and "ConfigRuntime". - config_[mod_idx] = output; - } - // Doing the configuration parsing after the loading finished. - ParseConfiguration(config_); - } - - private: - /* - *!\brief The key is the module index, this variable records all module pipeline configuration - * information. - */ - std::unordered_map config_; - /* - *\brief The key is the global output index, and the map is including global outputs index and - * the module outputs pair. - */ - std::unordered_map global_output_map_; -}; - -struct InputConnectionConfig { - /*!\brief The key("string") is the name of global module input interfaces. The value("pair") - * includes the index of graph module and the name of a graph module input interface. - */ - std::unordered_map> input_connection; - /*!\brief The map includes the global input name and global input index.*/ - std::unordered_map input_name_index_map; - /*! - * \brief The map not only includes the runtime index ,but also the pair of global interface - * and runtime interface. - */ - std::unordered_map>> input_runtime_map; - std::pair operator[](const std::string key) { - if (input_connection.find(key) == input_connection.end()) { - LOG(FATAL) << "Not find the key " << key; - } - return input_connection[key]; - } - /*!\brief Returns the number of global inputs through the input_runtime_map list size.*/ - int GetInputNum() { return input_runtime_map.size(); } - - /*! - * \brief Getting the global input index through the input name. - * \param input_name The global input name. - */ - int GetInputIndex(std::string input_name) { - auto input_index_iter = input_name_index_map.find(input_name); - if (input_index_iter == input_name_index_map.end()) { - LOG(FATAL) << "Do not finding the input name! " << input_name; - } - return input_index_iter->second; - } - /*!\brief Enumerating the input binding configuration for a specified runtime.*/ - void VisitConfig(BindingConfigParseFunc parse_function, int runtime_index) { - auto config = input_runtime_map.find(runtime_index); - // Only do the processing when there are input configuration in the runtime. - if (config != input_runtime_map.end()) { - for (auto x : config->second) { - int input_index = GetInputIndex(x.first); - parse_function(input_index, runtime_index, x.second); - } - } - } - /*! - * \brief Create an input connection config from JSONReader. - * \param reader Json reader. - */ - void Load(dmlc::JSONReader* reader) { - int global_interface_index = 0; - reader->BeginArray(); - while (reader->NextArrayItem()) { - reader->BeginObject(); - std::string key; - std::string global_interface_name; - std::string module_interface_name; - int mod_idx = -1; - while (reader->NextObjectItem(&key)) { - if (key == "global_interface_name") { - reader->Read(&global_interface_name); - input_name_index_map[global_interface_name] = global_interface_index++; - } else if (key == "mod_idx") { - reader->Read(&mod_idx); - } else if (key == "module_interface_name") { - reader->Read(&module_interface_name); - } else { - LOG(FATAL) << "do not support key " << key; - } - } - ICHECK(mod_idx >= 0) << "Invalid mod_idx value " << mod_idx; - ICHECK(!global_interface_name.empty()) << "Invalid global interface name value"; - ICHECK(!module_interface_name.empty()) << "Invalid module interface name value"; - input_connection[global_interface_name] = make_pair(mod_idx, module_interface_name); - // Creating a map which not only includes the runtime index, but also the pair of gloal - // interface, and runtime interface. - input_runtime_map[mod_idx].push_back( - std::make_pair(global_interface_name, module_interface_name)); - } - } -}; - -/*! - * \brief A map includes global module parameters groups and graph modudles. - */ -struct ParamConnectionConfig { - /*!\brief Mapping from the name of a global module parameters group to the index of a runtime - * module. - */ - std::unordered_map param_connection; - bool Empty() { return param_connection.empty(); } - int operator[](const std::string key) { - if (param_connection.find(key) == param_connection.end()) { - LOG(FATAL) << "do not support key " << key; - } - return param_connection[key]; - } - /*! - * \brief Load from JSONReader. - * \param reader Json reader. - */ - void Load(dmlc::JSONReader* reader) { - reader->BeginArray(); - while (reader->NextArrayItem()) { - reader->BeginObject(); - std::string key; - std::string global_param_name; - int mod_idx = -1; - while (reader->NextObjectItem(&key)) { - if (key == "global_param_name") { - reader->Read(&global_param_name); - } else if (key == "mod_idx") { - reader->Read(&mod_idx); - } else { - LOG(FATAL) << "do not support key " << key; - } - } - ICHECK(mod_idx >= 0) << "Invalid module index value " << mod_idx; - ICHECK(!global_param_name.empty()) << "Invalid global parameter group name value"; - param_connection[global_param_name] = mod_idx; - } - } -}; -/*! - * \brief The single consumer single producer queue which is used to forward data between two - * interfaces of backend cores. - */ -using ForwardQueue = SPSCLockFreeQueue; -using ForwardQueueMap = - std::unordered_map, ModuleIDHash>; -/*!\brief The basic class for runtime.*/ -class BasicRuntime { - using ModuleInputPairList = std::vector, int>>; - - public: - explicit BasicRuntime(int runtime_idx) : runtime_idx_(runtime_idx) {} - /*!\brief Return the index of the current module.*/ - int GetModuleIndex() { return runtime_idx_; } - /*!\brief Setting the data into this runtime via the input index.*/ - virtual void SetInput(const int index, DLTensor* data_in) {} - /*! - * \brief Sending a notification when data is ready. - * \param input_index The index of an input interface which have data ready. - */ - virtual void ParentNotify(int input_index) {} - /*! - *\brief Creating a parent notification. - *\param input_index The input index of the 'current runtime'. - *\param parent_idx The index of 'parent runtime' which will send the notification. - *\param parent_output_idx The output index of the 'parent runtime' which will send - * the notification. - */ - void CreateParentsNotify(int input_index, int parent_idx, int parent_output_idx) { - if (parents_notify_.find(input_index) != parents_notify_.end()) { - LOG(FATAL) << "The notification associated with the input interface " << input_index - << " in runtime " << runtime_idx_ << " already been created!"; - return; - } - parents_notify_[input_index] = - std::make_shared(ModuleInterfaceID(parent_idx, parent_output_idx, OUTPUT)); - } - - protected: - /*!\brief The index of runtime indicates the runtime position in the pipeline.*/ - int runtime_idx_; - /*!\brief A list of runtime which depends on the current runtime.*/ - std::unordered_map children_; - /*!\brief The map includes the runtime input index and the notification data structure.*/ - std::unordered_map> parents_notify_; - /*! - * \brief There is a list of SPSC input queues in which the input interface would poll the - * data comed from other backend cores. - */ - std::unordered_map> input_queue_; - - /*! - * \brief A list of SPSC forward queues in which the parent interface will push the data to - * other backend cores. - */ - std::unordered_map forward_queue_; - /*!\brief The state of the pipeline.*/ - std::atomic pipeline_state_{STOPPED}; - /*! - * \brief Generate the ID of an input queue. - * \param runtime_index The index of backend runtime. - * \param interface_index The index of the interface. - * \param type The type of the interface. - */ - ModuleInterfaceID GenerateQueueID(int runtime_index, int interface_index, InterfaceType type) { - return ModuleInterfaceID(runtime_index, interface_index, type); - } - /*! - * \brief Forwarding the data into the child runtimes. - * \param forward_queue_map The map includes the id and the queue. - * \param child_runtime The child runtime. - * \param child_input_index The child runtime index. - * \param data The data is used for forwarding. - */ - bool ForwardData(const ForwardQueueMap* forward_queue_map, - std::shared_ptr child_runtime, int child_input_index, - const DLTensor* data) { - auto child_runtime_index = child_runtime->GetModuleIndex(); - auto queue_id = GenerateQueueID(child_runtime_index, child_input_index, INPUT); - if (forward_queue_map->find(queue_id) == forward_queue_map->end()) { - LOG(FATAL) << "Not find the associated queue of the runtime(" << child_runtime_index - << ").input(" << child_input_index << ") which is connected with runtime(" - << runtime_idx_; - } - auto forward_queue = forward_queue_map->at(queue_id); - // If the queue is full, keep try until the push get success or the pipeline run into - // a STOP state. - while (!forward_queue->Push(data)) { - if (PipelineIsStop()) { - LOG(INFO) << "The forwarding process is stopped after the pipeline status is changed" - << " into stop."; - return false; - } - } - child_runtime->ParentNotify(child_input_index); - return true; - } - /*! - * \brief Creating a forwarding queue for the pair of an output interface and an input interface. - * \param forward_inf_idx The index of an interface which will send the forwarding data. - * \param child_runtime The backend runtime which owns the input interface. - * \param input_index The index of an input interface. This interface will receive the - * forwarding data. - */ - void CreateForwardingQueue(int forward_inf_idx, std::shared_ptr child_runtime, - int input_index) { - auto queue_id = GenerateQueueID(child_runtime->GetModuleIndex(), input_index, INPUT); - // The forwarding queue map of a specified output interface. - auto& queue_map = forward_queue_[forward_inf_idx]; - if (queue_map.find(queue_id) != queue_map.end()) { - LOG(FATAL) << "The queue " << queue_id.runtime_idx << "." << queue_id.runtime_interface_idx - << " is already created!"; - return; - } - auto queue = std::make_shared(queue_id); - queue_map[queue_id] = queue; - // Use the created queue as the consumer queue for the input interface of this forwarding - // pair. - child_runtime->AppendInputQueue(input_index, queue); - } - /*! - * \brief Setting the consumer queue for the input interface. - * \param input_index The index of the input interface. - * \param queue The consumer queue. - */ - void AppendInputQueue(int input_index, std::shared_ptr queue) { - input_queue_[input_index] = queue; - } - /*!\brief Checking if the pipeline is stopped or stopping.*/ - const bool PipelineIsStop() const { - auto state = pipeline_state_.load(std::memory_order_acquire); - return state == STOPPING || state == STOPPED; - } -}; -/* - *!\brief Backend Runtime. - */ -class BackendRuntime : public BasicRuntime { - private: - /*!The cpu affinity settings for this runtime.*/ - std::string cpu_affinity_ = ""; - /*!\brief The Runtime module of a backend graph executor.*/ - Module module_; - /*\brief The thread is associated with the current runtime*/ - std::thread thread_; - /*!\brief The execution count of the 'RunPipeline' function. */ - uint32_t pipeline_execution_count_ = 0; - /*! - *\brief In order to transfer data from one backend runtime to another, we need a local - * tensor variable as a medium. "input_tensor_local_copy_" is a map including - * input data and local tensor vairable. - */ - std::unordered_map input_tensor_local_copy_; - /*!\brief The packed functions.*/ - tvm::runtime::PackedFunc set_input_; - tvm::runtime::PackedFunc get_input_; - tvm::runtime::PackedFunc get_output_; - tvm::runtime::PackedFunc get_num_output_; - tvm::runtime::PackedFunc get_num_inputs_; - tvm::runtime::PackedFunc get_input_index_; - tvm::runtime::PackedFunc run_; - /*!\brief The worker thread is used to execute the runtimes in pipeline.*/ - void StartWorkThread() { - SetPipelineState(RUNNING); - if (runtime_idx_ == 0) { - this->SetCPUAffinity(); - } else { - // Only launching the worker thread for the runtimes after the first runtime. - thread_ = std::thread([&]() { - this->SetCPUAffinity(); - while (!this->WaitAndLoadPipelineData()) { - if (!this->RunPipeline()) { - break; - } - } - VLOG(1) << "Runtime " << this->runtime_idx_ << " exit."; - }); - } - return; - } - /*!\brief Setting the state of the pipeline.*/ - void SetPipelineState(PipelineState state) { - pipeline_state_.store(state, std::memory_order_release); - } - /*!\brief Stopping the threads in pipeline.*/ - void StopPipeline() { - SetPipelineState(STOPPING); - for (auto notify : parents_notify_) { - notify.second->ExitNotify(); - } - if (thread_.joinable()) { - thread_.join(); - } - SetPipelineState(STOPPED); - } - /*! - * \brief Waiting for the internal forwarding data. - * \return Returning 'true' when getting a 'exit' notification otherwise returning 'false'. - */ - bool WaitAndLoadPipelineData() { - std::unordered_map> notifys = parents_notify_; - bool exit_notify = false; - while (!notifys.empty() && !exit_notify) { - auto notify = notifys.begin(); - // Breaking the loop when the notification is in the exit state. - if ((exit_notify = notify->second->GetExitState())) break; - // Getting the source which sends this notification. - auto target_input_interface_index = notify->first; - // Loading the binding data. - while (!this->LoadBindingData(target_input_interface_index)) { - // Waiting for the notification. - if (!notify->second->Wait()) { - exit_notify = true; - break; - } - } - notifys.erase(notify); - } - return exit_notify; - } - /*! - * \brief Loading the binding data. - * \param input_index The index of the interface which will receive the forwarding data. - * \return Returning 'true' when data is loaded successfully, otherwise returning 'false'. - */ - bool LoadBindingData(int input_index) { - if (input_queue_.find(input_index) == input_queue_.end()) { - LOG(FATAL) << "Not finding the associated input queue of the input " << input_index << " !"; - } - auto queue = input_queue_[input_index]; - QueueData data; - // TODO(huajsj): Doing the 'SetInput' inside the poll function to avoid one time data copy. - if (!queue->Poll(&data)) { - return false; - } - SetInput(input_index, data.GetDLData()); - return true; - } - /*! - * \brief Forwarding the output data into the child runtimes. - * \return bool Return false when the "PipelineIsStop" function returns true or this function - * reaches some errors. Otherwise, return true. - */ - bool ForwardingOutputDataToChildren(void) { - for (auto child : children_) { - auto output_idx = child.first; - if (forward_queue_.find(output_idx) == forward_queue_.end()) { - LOG(FATAL) << "Not find the forwarding queue map for output(" << output_idx << ")!"; - } - NDArray output = GetOutput(output_idx); - auto forward_queue_map = forward_queue_[output_idx]; - // Notifying the 'children runtime' that the forwarding data are ready. - for (auto module_pair : child.second) { - auto child_runtime = module_pair.first; - auto child_input_index = module_pair.second; - auto output_data = const_cast(output.operator->()); - if (!ForwardData(&forward_queue_map, child_runtime, child_input_index, output_data)) { - return false; - } - } - } - return true; - } - /*! - * \brief Copying from a given tensor and using 'CPU' as the device. - */ - inline DLTensor* CopyDLTensorToCPU(const DLTensor* from) { - DLTensor* ret = NULL; - TVMArrayAlloc(from->shape, from->ndim, from->dtype.code, from->dtype.bits, from->dtype.lanes, - kDLCPU, 0, &ret); - return ret; - } - /*!\brief Creating a new NDArray with same shape and data type as the given DLTensor.*/ - NDArray CreateNDArrayFromDLTensor(const DLTensor* from) { - std::vector shape; - for (int i = 0; i < from->ndim; i++) { - shape.push_back(from->shape[i]); - } - auto ndarray = NDArray::Empty(shape, from->dtype, from->device); - ndarray.CreateView(shape, from->dtype); - return ndarray; - } - /* - *\brief Copying data from one DLTensor to another. - */ - void CopyFromTo(DLTensor* from, DLTensor* to) { - // When the 'from' device and the 'to' device are not the same, we use a temporary CPU - // DLTensor as the bridge. - if (from->device.device_type != to->device.device_type && from->device.device_type != kDLCPU && - to->device.device_type != kDLCPU) { - DLTensor* dltensor_local = nullptr; - if (input_tensor_local_copy_.find(to) == input_tensor_local_copy_.end()) { - dltensor_local = CopyDLTensorToCPU(from); - input_tensor_local_copy_[to] = dltensor_local; - } else { - dltensor_local = input_tensor_local_copy_[to]; - } - TVMArrayCopyFromTo(from, dltensor_local, nullptr); - from = dltensor_local; - } - - TVMArrayCopyFromTo(from, to, nullptr); - } - /*!\brief Setting the cpu affinity for the tvm threads pool in the current BackendRuntime.*/ - void SetCPUAffinity(void) { - if (cpu_affinity_.empty()) { - return; - } - auto affinity_mode = tvm::runtime::threading::ThreadGroup::kSpecifyThreadShareAllCore; - std::istringstream istr(cpu_affinity_); - std::string affinity; - std::vector cpus; - while (getline(istr, affinity, ',')) { - cpus.push_back(std::stoi(affinity)); - } - tvm::runtime::threading::Configure(affinity_mode, 0, cpus); - } - - public: - BackendRuntime(Module mod, int mod_idx) : BasicRuntime(mod_idx), module_(mod) { - get_input_index_ = module_.GetFunction("get_input_index"); - get_num_output_ = module_.GetFunction("get_num_outputs"); - get_num_inputs_ = module_.GetFunction("get_num_inputs"); - set_input_ = module_.GetFunction("set_input"); - get_input_ = module_.GetFunction("get_input"); - get_output_ = module_.GetFunction("get_output"); - run_ = module_.GetFunction("run"); - } - ~BackendRuntime() { - for (auto data : input_tensor_local_copy_) { - TVMArrayFree(data.second); - } - StopPipeline(); - } - /*! - * \brief Getting the times of using pipeline function. - * \return The times of using pipeline function. - */ - int GetExecutionCount() const { return pipeline_execution_count_; } - /*! - * \brief Initializing data structures for the pipeline execution. - * \param config The pipeline configueration. - * \param runtimes A list of BackendRuntime. - */ - void InitializePipeline(ConfigPipelineExecution config, - std::vector>* runtimes, - std::shared_ptr global_runtime) { - // Getting the current BackendRuntime's cpu affinity setting. - cpu_affinity_ = config.GetCPUAffinity(runtime_idx_); - // Getting the 'binding configuration' for each child runtime. - config.VisitRuntimeOutputConfig( - [&](int output_idx, int child_idx, std::string child_input_name) { - std::shared_ptr child_runtime = nullptr; - int input_index; - if (GLOBAL_MODULE_INDEX == child_idx) { - int global_output_index = std::stoi(child_input_name); - input_index = global_output_index; - child_runtime = global_runtime; - } else { - int runtime_idx_max = runtimes->size(); - if (child_idx < 0 || child_idx >= runtime_idx_max) { - LOG(FATAL) << "The runtime index " << child_idx << " is out of the range."; - } - auto runtime = runtimes->at(child_idx); - ICHECK(runtime->GetModuleIndex() == child_idx); - input_index = runtime->GetInputIndex(child_input_name); - if (input_index < 0) { - LOG(FATAL) << "Can not find the input " << input_index << "in runtime " << child_idx; - } - child_runtime = runtime; - } - ICHECK(child_runtime != nullptr); - children_[output_idx].push_back(std::make_pair(child_runtime, input_index)); - child_runtime->CreateParentsNotify(input_index, runtime_idx_, output_idx); - VLOG(1) << " parent_idx.output:" << runtime_idx_ << "." << output_idx - << " child.input:" << child_idx << "." << input_index; - // Creating the pipeline forwarding queue. - this->CreateForwardingQueue(output_idx, child_runtime, input_index); - }, - runtime_idx_); - - StartWorkThread(); - } - /*! - * \brief Notifying an input is ready. - * \param input_index The index of 'input interface' which is ready for data. - */ - void ParentNotify(int input_index) { - auto notify = parents_notify_.find(input_index); - if (notify == parents_notify_.end()) { - LOG(FATAL) << "Can not find the input for index " << input_index << " in runtime" - << runtime_idx_; - return; - } - notify->second->Notify(); - } - /*!\brief Creating a NDArray containing same shape and data type with a module output. */ - NDArray CreateFromOutput(int idx) { - NDArray data = get_output_(idx); - return CreateNDArrayFromDLTensor(const_cast(data.operator->())); - } - /*!\brief Return the number of output*/ - int NumOutputs() const { return get_num_output_(); } - /*!\brief Return the number of input*/ - int NumInputs() const { return get_num_inputs_(); } - /*!\brief Setting the data to this runtime via input index.*/ - void SetInput(const int index, DLTensor* data_in) { - NDArray input = get_input_(index); - DLTensor* dltensor_input = const_cast(input.operator->()); - CopyFromTo(data_in, dltensor_input); - } - /*!\brief Setting the data to the current runtime moduel via the input name. */ - void SetInput(const std::string name, DLTensor* data_in) { - int index = this->GetInputIndex(name); - SetInput(index, data_in); - } - /*!\brief Getting the input data via the input index.*/ - NDArray GetInput(int index) const { return get_input_(index); } - /*!\bief Getting the input data via the input name.*/ - int GetInputIndex(const std::string& name) { return get_input_index_(name); } - /*!\brief Using the output index to get the module output.*/ - NDArray GetOutput(int index) { return get_output_(index); } - /*!\brief Running the runtime.*/ - void Run() { run_(); } - /*! - * \brief Running the runtime in the pipeline mode. - * \return Returning false if the forwarding function failed. Otherwise, returning true.; - */ - bool RunPipeline() { - Run(); - bool ret = ForwardingOutputDataToChildren(); - pipeline_execution_count_++; - return ret; - } -}; -/*! - * \brief This global runtime represents the pipeline executor and exposes the input and output - * interface. - */ -class GlobalRuntime : public BasicRuntime { - public: - explicit GlobalRuntime(int runtime_idx) : BasicRuntime(runtime_idx) {} - /**/ - std::vector> GetRuntimeList() { return runtimes_; } - /*!\brief Push the data into the queue for the current runtime.*/ - void SetPipelineInput(const std::string input_name, DLTensor* data_in) { - auto input_index = input_config_.GetInputIndex(input_name); - auto child_iter = children_.find(input_index); - if (child_iter == children_.end()) { - return; - } - auto forward_queue_map = forward_queue_[input_index]; - // Notifying the 'children runtime' that the forwarding data are ready. - for (auto module_pair : child_iter->second) { - auto child_runtime = module_pair.first; - auto child_input_index = module_pair.second; - // No need to go through the forward queue when the runtime is the first one. - if (child_runtime->GetModuleIndex() == 0) { - child_runtime->SetInput(child_input_index, data_in); - } else { - if (!ForwardData(&forward_queue_map, child_runtime, child_input_index, data_in)) { - return; - } - } - } - return; - } - /*!\brief Whether the output data is ready.*/ - bool DataIsReady(bool wait_data) { - bool data_ready = true; - for (auto queue_pair : input_queue_) { - auto queue = queue_pair.second; - if (queue->Empty()) { - data_ready = false; - break; - } - } - if (!data_ready && wait_data) { - // TODO(huajsj): Waitting the data ready message. - } - return data_ready; - } - /*!\brief Get the output data.*/ - bool GetOutput(Array* outputs, bool wait_data = false) { - if (!DataIsReady(wait_data)) { - return false; - } - for (auto queue_pair : input_queue_) { - auto output_index = queue_pair.first; - auto queue = queue_pair.second; - QueueData data(const_cast(((*outputs)[output_index]).operator->())); - if (!queue->Poll(&data)) { - LOG(FATAL) << "There is no data in the data queue, it should not happen!"; - } - } - return true; - } - /*!\brief Initialized the data structures for pipeline.*/ - void InitializePipeline(InputConnectionConfig input_config, - const std::vector> runtimes) { - input_config_ = input_config; - runtimes_ = runtimes; - for (auto child_runtime : runtimes) { - int runtime_idx = child_runtime->GetModuleIndex(); - input_config.VisitConfig( - [&](int input_index, int child_idx, std::string child_input_name) { - auto child_input_index = child_runtime->GetInputIndex(child_input_name); - if (child_input_index < 0) { - LOG(FATAL) << "Can not find the input " << child_input_name << "in runtime " - << child_idx; - } - children_[input_index].push_back(std::make_pair(child_runtime, child_input_index)); - // Only create notify and queue for the runtime after the first runtime. - if (runtime_idx != 0) { - child_runtime->CreateParentsNotify(input_index, GLOBAL_MODULE_INDEX, - child_input_index); - // Creating the pipeline forwarding queue. - this->CreateForwardingQueue(input_index, child_runtime, child_input_index); - } - }, - runtime_idx); - } - } - - private: - std::vector> runtimes_; - InputConnectionConfig input_config_; -}; -/*! - * \brief The information used to initialize the graph executor module, the information - * come from the export library function call. - */ -struct GraphModuleLoadInfo { - GraphModuleLoadInfo(const std::string& lib, const std::string& json, const std::string& params, - const std::string& device) - : lib_name(lib), json_name(json), params_name(params), dev(device) {} - GraphModuleLoadInfo() { ; } - std::string lib_name; - std::string json_name; - std::string params_name; - std::string dev; -}; -/*! The Module information of each module.The 'int' is module index. */ -using ModuleConfig = std::unordered_map; -}; // namespace runtime -}; // namespace tvm -#endif // TVM_RUNTIME_PIPELINE_PIPELINE_STRUCT_H_ diff --git a/src/runtime/pipeline/spsc_queue.h b/src/runtime/pipeline/spsc_queue.h deleted file mode 100644 index 17313909f204..000000000000 --- a/src/runtime/pipeline/spsc_queue.h +++ /dev/null @@ -1,83 +0,0 @@ -/* - * 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. - */ -#ifndef TVM_RUNTIME_PIPELINE_SPSC_QUEUE_H_ -#define TVM_RUNTIME_PIPELINE_SPSC_QUEUE_H_ -#include -#include -/*!\brief A single producer and single consumer lock free queue. - */ -template -class SPSCLockFreeQueue { - public: - explicit SPSCLockFreeQueue(IDType id) : id_(id) {} - /*A read barrier enforcing the CPU to performe the reads before this barrier.*/ - inline void read_barrier() { std::atomic_thread_fence(std::memory_order_acquire); } - /*A write barrier enforcing the CPU to performe the writes before this barrier.*/ - inline void write_barrier() { std::atomic_thread_fence(std::memory_order_release); } - /*!\brief Checking whether the queue is full.*/ - bool Full() { - read_barrier(); - return ((tail_ + 1) % len_) == head_; - } - /*!brief Checking whether the queue is empty.*/ - bool Empty() { - read_barrier(); - return head_ == tail_; - } - /*! - * \brief Pushing the data into the queue. Only a single producer will call this function. - * \param data The data which is pushed into the queue. - * \return Return false when the queue is full. Otherwise, return true. - */ - template - bool Push(const data_type& data) { - if (Full()) return false; - queue_[tail_] = data; - write_barrier(); - tail_ = (tail_ + 1) % len_; - return true; - } - /*! - * \brief Poll the data from the front of the queue. Only the single consumer will call this - * function. - * \param data A pointer to the structure which stores the polled data.. - * \return Returning false when the queue is empty. Otherwise, return true. - */ - template - bool Poll(data_type* data) { - if (Empty()) return false; - *data = queue_[head_]; - write_barrier(); - head_ = (head_ + 1) % len_; - return true; - } - - private: - /*!\brief The pointer points to the first slot with valid data in the queue.*/ - size_t head_ = 0; - /*!\brief The end of the queue at which elements are added.*/ - size_t tail_ = 0; - /*!\brief The length of the queue.*/ - size_t len_ = QueueLength; - /*!\brief The queue used to store the data.*/ - SlotType queue_[QueueLength]; - /*!\brief The ID of the queue.*/ - IDType id_; -}; -#endif // TVM_RUNTIME_PIPELINE_SPSC_QUEUE_H_ diff --git a/src/runtime/thread_storage_scope.h b/src/runtime/thread_storage_scope.h index d1af2cb701a0..02481d1f7ce7 100644 --- a/src/runtime/thread_storage_scope.h +++ b/src/runtime/thread_storage_scope.h @@ -24,7 +24,6 @@ #ifndef TVM_RUNTIME_THREAD_STORAGE_SCOPE_H_ #define TVM_RUNTIME_THREAD_STORAGE_SCOPE_H_ -#include #include #include diff --git a/src/runtime/vm/bytecode.cc b/src/runtime/vm/bytecode.cc deleted file mode 100644 index dc52e8c8f01d..000000000000 --- a/src/runtime/vm/bytecode.cc +++ /dev/null @@ -1,694 +0,0 @@ -/* - * 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 src/runtime/vm/bytecode.cc - * \brief The bytecode for Relay virtual machine. - */ - -#include -#include -#include - -#include - -namespace tvm { -namespace runtime { -namespace vm { - -Instruction::Instruction() {} - -template -static T* Duplicate(T* src, Index size) { - auto dst = new T[size]; - std::copy(src, src + size, dst); - return dst; -} - -Instruction::Instruction(const Instruction& instr) { - this->op = instr.op; - this->dst = instr.dst; - - switch (instr.op) { - case Opcode::Move: - this->from = instr.from; - return; - case Opcode::Fatal: - return; - case Opcode::Ret: - this->result = instr.result; - return; - case Opcode::AllocTensor: - this->alloc_tensor.storage = instr.alloc_tensor.storage; - this->alloc_tensor.offset = instr.alloc_tensor.offset; - this->alloc_tensor.ndim = instr.alloc_tensor.ndim; - this->alloc_tensor.shape = - Duplicate(instr.alloc_tensor.shape, instr.alloc_tensor.ndim); - this->alloc_tensor.dtype = instr.alloc_tensor.dtype; - return; - case Opcode::AllocTensorReg: - this->alloc_tensor_reg.storage = instr.alloc_tensor_reg.storage; - this->alloc_tensor_reg.offset = instr.alloc_tensor_reg.offset; - this->alloc_tensor_reg.shape_register = instr.alloc_tensor_reg.shape_register; - this->alloc_tensor_reg.dtype = instr.alloc_tensor_reg.dtype; - return; - case Opcode::AllocADT: - this->constructor_tag = instr.constructor_tag; - this->num_fields = instr.num_fields; - this->datatype_fields = Duplicate(instr.datatype_fields, instr.num_fields); - return; - case Opcode::AllocClosure: - this->clo_index = instr.clo_index; - this->num_freevar = instr.num_freevar; - this->free_vars = Duplicate(instr.free_vars, instr.num_freevar); - return; - case Opcode::InvokePacked: - this->packed_index = instr.packed_index; - this->arity = instr.arity; - this->output_size = instr.output_size; - this->packed_args = Duplicate(instr.packed_args, instr.arity); - return; - case Opcode::InvokeClosure: - this->closure = instr.closure; - this->num_closure_args = instr.num_closure_args; - this->closure_args = Duplicate(instr.closure_args, instr.num_closure_args); - return; - case Opcode::Invoke: - this->func_index = instr.func_index; - this->num_args = instr.num_args; - this->invoke_args_registers = Duplicate(instr.invoke_args_registers, instr.num_args); - return; - case Opcode::If: - this->if_op = instr.if_op; - return; - case Opcode::LoadConst: - this->const_index = instr.const_index; - this->device_index = instr.device_index; - return; - case Opcode::LoadConsti: - this->load_consti = instr.load_consti; - return; - case Opcode::GetField: - this->object = instr.object; - this->field_index = instr.field_index; - return; - case Opcode::GetTag: - this->get_tag = instr.get_tag; - return; - case Opcode::Goto: - this->pc_offset = instr.pc_offset; - return; - case Opcode::AllocStorage: - this->alloc_storage.allocation_size = instr.alloc_storage.allocation_size; - this->alloc_storage.alignment = instr.alloc_storage.alignment; - this->alloc_storage.dtype_hint = instr.alloc_storage.dtype_hint; - this->alloc_storage.device_index = instr.alloc_storage.device_index; - this->alloc_storage.ndim = instr.alloc_storage.ndim; - if (this->alloc_storage.ndim > 0) { - this->alloc_storage.shape = - Duplicate(instr.alloc_storage.shape, instr.alloc_storage.ndim); - } - return; - case Opcode::ShapeOf: - this->shape_of.tensor = instr.shape_of.tensor; - return; - case Opcode::ReshapeTensor: - this->reshape_tensor = instr.reshape_tensor; - return; - case Opcode::DeviceCopy: - this->device_copy = instr.device_copy; - return; - case Opcode::KillRegister: - return; - default: - std::ostringstream out; - out << "Invalid instruction " << static_cast(instr.op); - throw std::runtime_error(out.str()); - } -} - -template -static inline void FreeIf(T* t) { - if (t != nullptr) { - delete t; - } -} - -Instruction& Instruction::operator=(const Instruction& instr) { - this->op = instr.op; - this->dst = instr.dst; - - switch (instr.op) { - case Opcode::Move: - this->from = instr.from; - return *this; - case Opcode::Fatal: - return *this; - case Opcode::LoadConsti: - this->load_consti = instr.load_consti; - return *this; - case Opcode::Ret: - this->result = instr.result; - return *this; - case Opcode::AllocTensor: - this->alloc_tensor.storage = this->alloc_tensor.storage; - this->alloc_tensor.offset = instr.alloc_tensor.offset; - this->alloc_tensor.ndim = instr.alloc_tensor.ndim; - this->alloc_tensor.shape = - Duplicate(instr.alloc_tensor.shape, instr.alloc_tensor.ndim); - this->alloc_tensor.dtype = instr.alloc_tensor.dtype; - return *this; - case Opcode::AllocTensorReg: - this->alloc_tensor_reg.storage = instr.alloc_tensor_reg.storage; - this->alloc_tensor_reg.offset = instr.alloc_tensor_reg.offset; - this->alloc_tensor_reg.shape_register = instr.alloc_tensor_reg.shape_register; - this->alloc_tensor_reg.dtype = instr.alloc_tensor_reg.dtype; - return *this; - case Opcode::AllocADT: - this->constructor_tag = instr.constructor_tag; - this->num_fields = instr.num_fields; - FreeIf(this->datatype_fields); - this->datatype_fields = Duplicate(instr.datatype_fields, instr.num_fields); - return *this; - case Opcode::AllocClosure: - this->clo_index = instr.clo_index; - this->num_freevar = instr.num_freevar; - FreeIf(this->free_vars); - this->free_vars = Duplicate(instr.free_vars, instr.num_freevar); - return *this; - case Opcode::InvokePacked: - this->packed_index = instr.packed_index; - this->arity = instr.arity; - this->output_size = instr.output_size; - FreeIf(this->packed_args); - this->packed_args = Duplicate(instr.packed_args, instr.arity); - return *this; - case Opcode::InvokeClosure: - this->closure = instr.closure; - this->num_closure_args = instr.num_closure_args; - FreeIf(this->closure_args); - this->closure_args = Duplicate(instr.closure_args, instr.num_closure_args); - return *this; - case Opcode::Invoke: - this->func_index = instr.func_index; - this->num_args = instr.num_args; - FreeIf(this->invoke_args_registers); - this->invoke_args_registers = Duplicate(instr.invoke_args_registers, instr.num_args); - return *this; - case Opcode::If: - this->if_op = instr.if_op; - return *this; - case Opcode::LoadConst: - this->const_index = instr.const_index; - this->device_index = instr.device_index; - return *this; - case Opcode::GetField: - this->object = instr.object; - this->field_index = instr.field_index; - return *this; - case Opcode::GetTag: - this->get_tag = instr.get_tag; - return *this; - case Opcode::Goto: - this->pc_offset = instr.pc_offset; - return *this; - case Opcode::AllocStorage: - this->alloc_storage.allocation_size = instr.alloc_storage.allocation_size; - this->alloc_storage.alignment = instr.alloc_storage.alignment; - this->alloc_storage.dtype_hint = instr.alloc_storage.dtype_hint; - this->alloc_storage.device_index = instr.alloc_storage.device_index; - this->alloc_storage.ndim = instr.alloc_storage.ndim; - if (this->alloc_storage.ndim > 0) { - this->alloc_storage.shape = - Duplicate(instr.alloc_storage.shape, instr.alloc_storage.ndim); - } - return *this; - case Opcode::ShapeOf: - this->shape_of.tensor = instr.shape_of.tensor; - return *this; - case Opcode::ReshapeTensor: - this->reshape_tensor = instr.reshape_tensor; - return *this; - case Opcode::DeviceCopy: - this->device_copy = instr.device_copy; - return *this; - case Opcode::KillRegister: - return *this; - default: - std::ostringstream out; - out << "Invalid instruction " << static_cast(instr.op); - throw std::runtime_error(out.str()); - } -} - -Instruction::~Instruction() { - switch (this->op) { - case Opcode::Move: - case Opcode::Ret: - case Opcode::AllocTensorReg: - case Opcode::If: - case Opcode::LoadConst: - case Opcode::GetField: - case Opcode::GetTag: - case Opcode::Goto: - case Opcode::LoadConsti: - case Opcode::ShapeOf: - case Opcode::ReshapeTensor: - case Opcode::DeviceCopy: - case Opcode::Fatal: - case Opcode::KillRegister: - return; - case Opcode::AllocStorage: - if (this->alloc_storage.ndim > 0) { - delete[] this->alloc_storage.shape; - } - return; - case Opcode::AllocTensor: - delete[] this->alloc_tensor.shape; - return; - case Opcode::AllocADT: - delete[] this->datatype_fields; - return; - case Opcode::AllocClosure: - delete[] this->free_vars; - return; - case Opcode::InvokePacked: - delete[] this->packed_args; - return; - case Opcode::InvokeClosure: - delete[] this->closure_args; - return; - case Opcode::Invoke: - delete[] this->invoke_args_registers; - return; - default: - std::ostringstream out; - LOG(FATAL) << "Invalid instruction " << static_cast(this->op); - } -} - -Instruction Instruction::Ret(RegName result) { - Instruction instr; - instr.op = Opcode::Ret; - instr.result = result; - return instr; -} - -Instruction Instruction::Fatal() { - Instruction instr; - instr.op = Opcode::Fatal; - return instr; -} - -Instruction Instruction::InvokePacked(Index packed_index, Index arity, Index output_size, - const std::vector& args) { - Instruction instr; - instr.op = Opcode::InvokePacked; - instr.packed_index = packed_index; - instr.arity = arity; - instr.output_size = output_size; - instr.packed_args = new RegName[arity]; - for (Index i = 0; i < arity; ++i) { - instr.packed_args[i] = args[i]; - } - return instr; -} - -Instruction Instruction::AllocTensor(RegName storage, RegName offset, - const std::vector& shape, DLDataType dtype, - RegName dst) { - Instruction instr; - instr.op = Opcode::AllocTensor; - instr.dst = dst; - instr.alloc_tensor.storage = storage; - instr.alloc_tensor.offset = offset; - instr.alloc_tensor.ndim = shape.size(); - instr.alloc_tensor.shape = new int64_t[shape.size()]; - for (size_t i = 0; i < shape.size(); ++i) { - instr.alloc_tensor.shape[i] = shape[i]; - } - instr.alloc_tensor.dtype = dtype; - return instr; -} - -Instruction Instruction::AllocTensorReg(RegName storage, RegName offset, RegName shape_register, - DLDataType dtype, RegName dst) { - Instruction instr; - instr.op = Opcode::AllocTensorReg; - instr.dst = dst; - instr.alloc_tensor_reg.storage = storage; - instr.alloc_tensor_reg.offset = offset; - instr.alloc_tensor_reg.shape_register = shape_register; - instr.alloc_tensor_reg.dtype = dtype; - return instr; -} - -Instruction Instruction::AllocStorage(RegName size, Index alignment, DLDataType dtype_hint, - Index device_index, const std::vector& shape, - RegName dst) { - Instruction instr; - instr.op = Opcode::AllocStorage; - instr.dst = dst; - instr.alloc_storage.allocation_size = size; - instr.alloc_storage.alignment = alignment; - instr.alloc_storage.dtype_hint = dtype_hint; - instr.alloc_storage.device_index = device_index; - instr.alloc_storage.ndim = static_cast(shape.size()); - if (instr.alloc_storage.ndim > 0) { - instr.alloc_storage.shape = new int64_t[shape.size()]; - for (size_t i = 0; i < shape.size(); ++i) { - instr.alloc_storage.shape[i] = shape[i]; - } - } - return instr; -} - -Instruction Instruction::ShapeOf(RegName tensor, RegName dst) { - Instruction instr; - instr.op = Opcode::ShapeOf; - instr.dst = dst; - instr.shape_of.tensor = tensor; - return instr; -} - -Instruction Instruction::ReshapeTensor(RegName tensor, RegName newshape, RegName dst) { - Instruction instr; - instr.op = Opcode::ReshapeTensor; - instr.dst = dst; - instr.reshape_tensor.tensor = tensor; - instr.reshape_tensor.newshape = newshape; - return instr; -} - -Instruction Instruction::DeviceCopy(RegName src, Index src_device_index, Index dst_device_index, - RegName dst) { - Instruction instr; - instr.op = Opcode::DeviceCopy; - instr.dst = dst; - instr.device_copy.src = src; - instr.device_copy.src_device_index = src_device_index; - instr.device_copy.dst_device_index = dst_device_index; - return instr; -} - -Instruction Instruction::KillRegister(RegName dst) { - Instruction instr; - instr.op = Opcode::KillRegister; - instr.dst = dst; - return instr; -} - -Instruction Instruction::AllocADT(Index tag, Index num_fields, - const std::vector& datatype_fields, RegName dst) { - Instruction instr; - instr.op = Opcode::AllocADT; - instr.dst = dst; - instr.constructor_tag = tag; - instr.num_fields = num_fields; - instr.datatype_fields = new RegName[num_fields]; - for (Index i = 0; i < num_fields; ++i) { - instr.datatype_fields[i] = datatype_fields[i]; - } - return instr; -} - -Instruction Instruction::AllocClosure(Index func_index, Index free_vars, - const std::vector& free_var_register, RegName dst) { - Instruction instr; - instr.op = Opcode::AllocClosure; - instr.dst = dst; - instr.clo_index = func_index; - instr.num_freevar = free_vars; - instr.free_vars = new RegName[instr.num_freevar]; - for (Index i = 0; i < instr.num_freevar; ++i) { - instr.free_vars[i] = free_var_register[i]; - } - return instr; -} - -Instruction Instruction::GetField(RegName object, Index field_index, RegName dst) { - Instruction instr; - instr.op = Opcode::GetField; - instr.dst = dst; - instr.object = object; - instr.field_index = field_index; - return instr; -} - -Instruction Instruction::GetTag(RegName object, RegName dst) { - Instruction instr; - instr.op = Opcode::GetTag; - instr.dst = dst; - instr.get_tag.object = object; - return instr; -} - -Instruction Instruction::If(RegName test, RegName target, Index true_branch, Index false_branch) { - Instruction instr; - instr.op = Opcode::If; - instr.if_op.test = test; - instr.if_op.target = target; - instr.if_op.true_offset = true_branch; - instr.if_op.false_offset = false_branch; - return instr; -} - -Instruction Instruction::Goto(Index pc_offset) { - Instruction instr; - instr.op = Opcode::Goto; - instr.pc_offset = pc_offset; - return instr; -} - -Instruction Instruction::Invoke(Index func_index, const std::vector& args_registers, - RegName dst) { - Instruction instr; - instr.op = Opcode::Invoke; - instr.dst = dst; - instr.func_index = func_index; - instr.num_args = args_registers.size(); - instr.invoke_args_registers = new RegName[instr.num_args]; - for (Index i = 0; i < instr.num_args; ++i) { - instr.invoke_args_registers[i] = args_registers[i]; - } - return instr; -} - -Instruction Instruction::InvokeClosure(RegName closure, const std::vector& args, - RegName dst) { - Instruction instr; - instr.op = Opcode::InvokeClosure; - instr.dst = dst; - instr.closure = closure; - instr.num_closure_args = args.size(); - instr.closure_args = new RegName[args.size()]; - for (size_t i = 0; i < args.size(); ++i) { - instr.closure_args[i] = args[i]; - } - return instr; -} - -Instruction Instruction::LoadConst(Index const_index, Index device_index, RegName dst) { - Instruction instr; - instr.op = Opcode::LoadConst; - instr.dst = dst; - instr.const_index = const_index; - instr.device_index = device_index; - return instr; -} - -Instruction Instruction::LoadConsti(Index val, RegName dst) { - Instruction instr; - instr.op = Opcode::LoadConsti; - instr.dst = dst; - instr.load_consti.val = val; - return instr; -} - -Instruction Instruction::Move(RegName src, RegName dst) { - Instruction instr; - instr.op = Opcode::Move; - instr.dst = dst; - instr.from = src; - return instr; -} - -void DLDatatypePrint(std::ostream& os, const DLDataType& dtype) { - switch (dtype.code) { - case kDLInt: - os << "int"; - break; - case kDLUInt: - os << "uint"; - break; - case kDLFloat: - os << "float"; - break; - case kDLBfloat: - os << "bfloat"; - break; - } - - os << int(dtype.bits); - if (dtype.lanes != 1) { - os << "x" << dtype.lanes; - } -} - -template -std::string StrJoin(T* items, int offset, int cnt, std::string delim = ", ") { - if (cnt == 0) { - return ""; - } - std::ostringstream oss; - oss << items[offset]; - for (int i = 1; i < cnt; ++i) { - oss << delim << items[offset + i]; - } - return oss.str(); -} - -void InstructionPrint(std::ostream& os, const Instruction& instr) { - switch (instr.op) { - case Opcode::Move: { - os << "move $" << instr.dst << " $" << instr.from; - break; - } - case Opcode::Ret: { - os << "ret $" << instr.result; - break; - } - case Opcode::Fatal: { - os << "fatal"; - break; - } - case Opcode::InvokePacked: { - os << "invoke_packed PackedFunc[" << instr.packed_index << "] (in: $" - << StrJoin(instr.packed_args, 0, instr.arity - instr.output_size, ", $") - << ", out: $" - << StrJoin(instr.packed_args, instr.arity - instr.output_size, instr.output_size, - ", $") - << ")"; - break; - } - case Opcode::AllocTensor: { - os << "alloc_tensor $" << instr.dst << " $" << instr.alloc_tensor.storage << " $" - << instr.alloc_tensor.offset << " [" - << StrJoin(instr.alloc_tensor.shape, 0, instr.alloc_tensor.ndim) << "] "; - DLDatatypePrint(os, instr.alloc_tensor.dtype); - break; - } - case Opcode::AllocTensorReg: { - os << "alloc_tensor_reg $" << instr.dst << " $" << instr.alloc_tensor_reg.storage << " $" - << instr.alloc_tensor_reg.offset << " $" << instr.alloc_tensor_reg.shape_register << " "; - DLDatatypePrint(os, instr.alloc_tensor_reg.dtype); - break; - } - case Opcode::AllocADT: { - os << "alloc_data $" << instr.dst << " tag(" << instr.constructor_tag << ") [$" - << StrJoin(instr.datatype_fields, 0, instr.num_fields, ",$") << "]"; - break; - } - case Opcode::AllocClosure: { - os << "alloc_closure $" << instr.dst << " VMFunc[" << instr.clo_index << "]($" - << StrJoin(instr.free_vars, 0, instr.num_freevar, ",$") << ")"; - break; - } - case Opcode::If: { - os << "if " - << "$" << instr.if_op.test << " $" << instr.if_op.target << " " << instr.if_op.true_offset - << " " << instr.if_op.false_offset; - break; - } - case Opcode::Invoke: { - os << "invoke $" << instr.dst << " VMFunc[" << instr.func_index << "]($" - << StrJoin(instr.invoke_args_registers, 0, instr.num_args, ",$") << ")"; - break; - } - case Opcode::InvokeClosure: { - os << "invoke_closure $" << instr.dst << " $" << instr.closure << "($" - << StrJoin(instr.closure_args, 0, instr.num_closure_args, ",$") << ")"; - break; - } - case Opcode::LoadConst: { - os << "load_const $" << instr.dst << " Const[" << instr.const_index << "] " - << instr.device_index; - break; - } - case Opcode::LoadConsti: { - os << "load_consti $" << instr.dst << " " << instr.load_consti.val; - break; - } - case Opcode::GetField: { - os << "get_field $" << instr.dst << " $" << instr.object << "[" << instr.field_index << "]"; - break; - } - case Opcode::GetTag: { - os << "get_tag $" << instr.dst << " $" << instr.get_tag.object; - break; - } - case Opcode::Goto: { - os << "goto " << instr.pc_offset; - break; - } - case Opcode::AllocStorage: { - os << "alloc_storage $" << instr.dst << " "; - if (instr.alloc_storage.ndim > 0) { - os << "[" << StrJoin(instr.alloc_storage.shape, 0, instr.alloc_storage.ndim) - << "] "; - } else { - os << "$" << instr.alloc_storage.allocation_size << " " << instr.alloc_storage.alignment - << " "; - } - os << DLDataType2String(instr.alloc_storage.dtype_hint) << " " - << instr.alloc_storage.device_index; - break; - } - case Opcode::ShapeOf: { - os << "shape_of $" << instr.dst << " $" << instr.shape_of.tensor; - break; - } - case Opcode::ReshapeTensor: { - os << "reshape_tensor $" << instr.dst << " $" << instr.reshape_tensor.tensor << " $" - << instr.reshape_tensor.newshape; - break; - } - case Opcode::DeviceCopy: { - os << "device_copy $" << instr.dst << " $" << instr.device_copy.src << " " - << instr.device_copy.dst_device_index << " " << instr.device_copy.src_device_index; - break; - } - case Opcode::KillRegister: { - os << "kill_register $" << instr.dst; - break; - } - default: - LOG(FATAL) << "should never hit this case" << static_cast(instr.op); - break; - } -} - -std::ostream& operator<<(std::ostream& os, const Instruction& instr) { - InstructionPrint(os, instr); - return os; -} - -} // namespace vm -} // namespace runtime -} // namespace tvm diff --git a/src/runtime/vm/executable.cc b/src/runtime/vm/executable.cc deleted file mode 100644 index 161c6dbfbd76..000000000000 --- a/src/runtime/vm/executable.cc +++ /dev/null @@ -1,1077 +0,0 @@ -/* - * 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 tvm/runtime/vm/executable.cc - * \brief The implementation of a virtual machine executable APIs. - */ - -#include -#include -#include -#include -#include -#include - -#include -#include -#include -#include -#include -#include -#include -#include - -#include "../file_utils.h" -#include "../library_module.h" -#include "serialize_utils.h" - -namespace tvm { -namespace runtime { -namespace vm { - -#define STREAM_CHECK(val, section) \ - ICHECK(val) << "Invalid VM file format in the " << section << " section." \ - << "\n"; - -// Helper to serialize a vm instruction. -VMInstructionSerializer SerializeInstruction(const Instruction& instr); -// Helper to deserialize a serialized vm instruction. -Instruction DeserializeInstruction(const VMInstructionSerializer& instr); - -const VMFunction& Executable::GetVMFunctionWithName(const std::string& func_name) const { - auto it = global_map.find(func_name); - ICHECK(it != global_map.end()) << "Cannot find function " << func_name << " in executable"; - return functions[it->second]; -} - -Module Executable::VMLoadExecutable() { - auto vm = make_object(); - vm->LoadExecutable(GetObjectPtr(this)); - return Module(vm); -} - -int Executable::GetFunctionArity(std::string func_name) const { - const auto& func = GetVMFunctionWithName(func_name); - return func.params.size(); -} - -std::string Executable::GetFunctionParameterName(std::string func_name, int index) const { - const auto& func = GetVMFunctionWithName(func_name); - ICHECK_LT(index, func.params.size()) << "Invalid parameter index"; - return func.params[index]; -} - -std::string Executable::GetBytecode() const { - std::ostringstream oss; - - for (size_t i = 0; i < functions.size(); ++i) { - const auto& func = functions[i]; - // Print the header of the function format. - oss << "VM Function[" << i << "]: " << func.name << "("; - bool first = true; - for (const auto& param : func.params) { - if (!first) { - oss << ", "; - } - oss << param; - first = false; - } - oss << ")" << std::endl; - oss << "# reg file size = " << func.register_file_size << std::endl; - oss << "# instruction count = " << func.instructions.size() << std::endl; - - // Print the instructions of a `VMFunction`. - // The part after ";" is the instruction in text format. - oss << "opcode, fields # inst(text):" << std::endl; - for (size_t idx = 0; idx < func.instructions.size(); ++idx) { - const auto& instr = func.instructions[idx]; - const auto& serialized_instr = SerializeInstruction(instr); - std::ostringstream line; - line << std::setw(2) << idx << ": " << serialized_instr.opcode << " "; - for (auto it : serialized_instr.fields) { - line << it << " "; - } - oss << std::setw(40) << std::setfill(' ') << std::left << line.str(); - oss << " # " << instr; - if (oss.str().back() != '\n') oss << std::endl; - } - oss << std::endl; - } - - return oss.str(); -} - -std::string Executable::GetConstants() const { - std::ostringstream oss; - for (size_t i = 0; i < constants.size(); ++i) { - const auto& constant = constants[i]; - auto ndarray = Downcast(constant); - oss << "VM Const[" << i - << "]: " << RuntimeObject2String(ndarray, virtual_devices[host_device_index].first) - << " on device index " << const_device_indexes[i] << std::endl; - } - return oss.str(); -} - -std::string Executable::GetVirtualDevices() const { - std::ostringstream oss; - for (size_t i = 0; i < virtual_devices.size(); ++i) { - const auto& [device, scope] = virtual_devices[i]; - oss << "VM VirtualDevice[" << i << "]: device type " << device.device_type << ", id " - << device.device_id << " and mem_scope " << scope << std::endl; - } - return oss.str(); -} - -std::string Executable::GetPrimitives() const { - std::ostringstream os; - std::vector> entries; - entries.reserve(primitive_map.size()); - for (const auto& kv : primitive_map) { - entries.emplace_back(kv.second, kv.first); - } - std::sort(entries.begin(), entries.end(), - [](const std::pair& left, const std::pair& right) { - return left.first < right.first; - }); - for (const auto& entry : entries) { - os << "VM PackedFunc[" << entry.first << "]: " << entry.second << std::endl; - } - return os.str(); -} - -std::string Executable::Stats() const { - std::ostringstream oss; - oss << "Relay VM executable statistics:" << std::endl; - - // Get the number of constants and the shape of each of them. - oss << " Constant shapes (# " << constants.size() << "): ["; - for (const auto& it : constants) { - const auto constant = Downcast(it); - const auto& shape = constant.Shape(); - - // Scalar - if (shape.empty()) { - oss << "scalar, "; - continue; - } - - oss << "["; - for (auto s : shape) { - oss << s << ", "; - } - oss.seekp(-2, oss.cur); - oss << "], " << std::endl; - } - if (!constants.empty()) oss.seekp(-2, oss.cur); - oss << "]" << std::endl; - - // Get the number of globals and the name of each of them. - oss << " Globals (#" << global_map.size() << "): ["; - for (const auto& it : global_map) { - oss << "(\"" << it.first << "\", " << it.second << ")" - << ", "; - } - if (!global_map.empty()) oss.seekp(-2, oss.cur); - oss << "]" << std::endl; - - // Get the number of primitive ops and the name of each of them. - oss << " Primitive ops (#" << primitive_map.size() << "): ["; - std::vector prim_ops; - for (const auto& it : primitive_map) { - auto packed_index = static_cast(it.second); - if (prim_ops.size() <= packed_index) { - prim_ops.resize(packed_index + 1); - } - prim_ops[packed_index] = it.first; - } - for (const auto& it : prim_ops) { - oss << it << ", "; - } - if (!prim_ops.empty()) oss.seekp(-2, oss.cur); - oss << "]" << std::endl; - - return oss.str(); -} - -void SaveHeader(dmlc::Stream* strm) { - uint64_t header = kTVMVMBytecodeMagic; - strm->Write(header); - std::string version = TVM_VERSION; - strm->Write(version); -} - -TVMByteArray Executable::Save() { - // Initialize the stream object. - code_.clear(); - dmlc::MemoryStringStream strm(&code_); - - // Save header - SaveHeader(&strm); - - // Save virtual devices section. - SaveVirtualDevicesSection(&strm); - - // Global section. - SaveGlobalSection(&strm); - - // Constant section. - SaveConstantSection(&strm); - - // Primitive names. - SavePrimitiveOpNames(&strm); - - // Code section. - SaveCodeSection(&strm); - - TVMByteArray arr; - arr.data = code_.c_str(); - arr.size = code_.length(); - return arr; -} - -void Executable::SaveVirtualDevicesSection(dmlc::Stream* strm) { - strm->Write(virtual_devices); - strm->Write(host_device_index); -} - -Map Executable::GetLateBoundConstants(int64_t byte_limit) { - ICHECK(late_bound_constant_names.empty()); - late_bound_constant_names.reserve(constants.size()); - Map map; - size_t total_late_bound_bytes = 0; - for (size_t const_index = 0; const_index < constants.size(); ++const_index) { - const auto ndarray = Downcast(constants[const_index]); - ICHECK(ndarray.defined()) << "Undefined constant at index " << const_index; - int64_t num_bytes = runtime::GetDataSize(*ndarray.operator->()); - if (num_bytes < byte_limit) { - // Leave as immediate. - late_bound_constant_names.emplace_back(nullptr); - continue; - } - total_late_bound_bytes += num_bytes; - std::ostringstream os; - os << "const_" << const_index; - String name = os.str(); - map.Set(name, Downcast(std::move(constants[const_index]))); - late_bound_constant_names.emplace_back(std::move(name)); - } - VLOG(1) << "moved " << map.size() << " constants of " << total_late_bound_bytes - << " bytes (out of " << constants.size() << " overall) to be late-bound"; - return map; -} - -void Executable::MoveLateBoundConstantsToStream(dmlc::Stream* stream, int64_t byte_limit) { - Map map = GetLateBoundConstants(byte_limit); - runtime::SaveParams(stream, map); -} - -void Executable::MoveLateBoundConstantsToFile(const std::string& path, int64_t byte_limit) { - tvm::runtime::SimpleBinaryFileStream stream(path, "wb"); - MoveLateBoundConstantsToStream(&stream, byte_limit); -} - -void Executable::LoadLateBoundConstantsFromStream(dmlc::Stream* stream) { - if (late_bound_constant_names.empty()) { - VLOG(1) << "Found no late-bound constants to load"; - return; - } - ICHECK_EQ(late_bound_constant_names.size(), constants.size()); - Map map = runtime::LoadParams(stream); - VLOG(1) << "loaded " << map.size() << " late-bound constants"; - LoadLateBoundConstantsFromMap(map); -} - -void Executable::LoadLateBoundConstantsFromMap(Map map) { - for (size_t const_index = 0; const_index < constants.size(); ++const_index) { - if (!late_bound_constant_names[const_index].defined()) { - ICHECK(constants[const_index].defined()) - << "Undefined immediate constant at index " << const_index; - continue; - } - const String& name = late_bound_constant_names[const_index]; - ICHECK(!constants[const_index].defined()) << "Unexpected constant at index " << const_index; - auto itr = map.find(name); - ICHECK(itr != map.end()) << "No binding for late-bound constant at index " << const_index - << " with name '" << name << "'"; - constants[const_index] = (*itr).second; - map.erase(name); - } - late_bound_constant_names.clear(); - ICHECK(map.empty()) << "Have " << map.size() << " unused late-bound constants"; -} - -void Executable::LoadLateBoundConstantsFromFile(const std::string& path) { - tvm::runtime::SimpleBinaryFileStream stream(path, "rb"); - LoadLateBoundConstantsFromStream(&stream); -} - -void Executable::SaveGlobalSection(dmlc::Stream* strm) { - std::vector> globals(this->global_map.begin(), - this->global_map.end()); - auto comp = [](const std::pair& a, const std::pair& b) { - return a.second < b.second; - }; - std::sort(globals.begin(), globals.end(), comp); - - std::vector glbs; - for (const auto& it : globals) { - glbs.push_back(it.first); - } - strm->Write(glbs); -} - -namespace { -// Tags to distinguish immediate vs late-bound constants in constants table bytestream. -constexpr uint32_t kImmediateConstTag = 0; -constexpr uint32_t kLateBoundConstTag = 1; -} // namespace - -void Executable::SaveConstantSection(dmlc::Stream* stream) { - // Save the overall number of constants. - stream->Write(static_cast(constants.size())); - - for (size_t const_index = 0; const_index < constants.size(); ++const_index) { - if (late_bound_constant_names.empty() || !late_bound_constant_names[const_index].defined()) { - // Tag immediate constants by 0. - stream->Write(kImmediateConstTag); - // Write as DLTensor. - const auto ndarray = Downcast(constants[const_index]); - ICHECK(ndarray.defined()); - runtime::SaveDLTensor(stream, ndarray.operator->()); - VLOG(1) << "save " << const_index << " as immediate"; - } else { - // Tag late-bound constants by 1. - const String& name = late_bound_constant_names[const_index]; - ICHECK(!constants[const_index].defined()); - stream->Write(kLateBoundConstTag); - // Write a string. - stream->Write(std::string(name)); - VLOG(1) << "save " << const_index << " as late-bound"; - } - } - - VLOG(1) << "saved " << constants.size() << " constants"; - - // Save the const to device index mapping. - stream->Write(const_device_indexes); -} - -void Executable::LoadConstantSection(dmlc::Stream* stream) { - uint64_t sz; - // Load the overall number of constants. - STREAM_CHECK(stream->Read(&sz, sizeof(sz)), "constants table size"); - size_t size = static_cast(sz); - - VLOG(1) << "loading " << size << " constants"; - - constants.resize(size); - late_bound_constant_names.resize(size); - bool any_late_bound = false; - - // Load each of the constants. - for (size_t const_index = 0; const_index < size; const_index++) { - uint32_t tag; - STREAM_CHECK(stream->Read(&tag, sizeof(tag)), "constant tag"); - if (tag == kImmediateConstTag) { - // Immediate constants tagged by 0. - VLOG(1) << "load " << const_index << " as immediate"; - runtime::NDArray ndarray; - STREAM_CHECK(ndarray.Load(stream), "constant tensor"); - constants[const_index] = std::move(ndarray); - late_bound_constant_names[const_index] = String(ObjectPtr(nullptr)); - } else if (tag == kLateBoundConstTag) { - // Late-bound constants tagged by 1. - VLOG(1) << "load " << const_index << " as late-bound"; - std::string name; - STREAM_CHECK(stream->Read(&name), "late-bound constant name"); - constants[const_index] = NDArray(nullptr); - late_bound_constant_names[const_index] = std::move(name); - any_late_bound = true; - } else { - STREAM_CHECK(false, "constant tag"); - } - } - - if (!any_late_bound) { - late_bound_constant_names.clear(); - } - - // Load the const to device index mapping. - std::vector indexes; - indexes.reserve(size); - STREAM_CHECK(stream->Read(&indexes), "constant devices"); - ICHECK_EQ(size, indexes.size()); - const_device_indexes = std::move(indexes); -} - -void Executable::SavePrimitiveOpNames(dmlc::Stream* strm) { - std::vector primitive_names; - for (const auto& it : this->primitive_map) { - auto packed_index = static_cast(it.second); - if (primitive_names.size() <= packed_index) { - primitive_names.resize(packed_index + 1); - } - primitive_names[packed_index] = it.first; - } - strm->Write(primitive_names); - std::map> primitive_attrs; - for (const auto& it : this->op_attrs) { - auto packed_index = static_cast(it.first); - std::map attrs; - for (const auto& elem : it.second) { - // TODO(tkonolige): cannot serialize ObjectRefs with dmlc's serializer, so we just serialize - // strings for now - if (elem.second.as()) { - attrs[elem.first] = Downcast(elem.second); - } - } - primitive_attrs[packed_index] = attrs; - } - strm->Write(primitive_attrs); -} - -// Serialize a virtual machine instruction. It creates a list that contains the -// hash, opcode, and all fields of an instruction. -// -// For example, the function signature used to create an `AllocTensor` -// instruction is: -// Instruction AllocTensor(std::vector shape, DLDataType dtype, RegName dst) -// -// The serialized form will be: -// `hash 5 dtype.code dtype.bits dtype.lanes ndim dst_register val1 val2 ... valn` -// -// where hash is the hash of serialized instruction that is computed internally -// by the `VMInstructionExecutable`. It is used for sanity check before decoding. -// 5 shows opcode of `AllocTensor`, `(dtype.code dtype.bits dtype.lanes)` -// represents a `DLDataType`, `ndim` is the number of dimensions, `dst_register` -// is the destination register, and the rest of it together indicates the shape -// of the tensor to be allocated. -VMInstructionSerializer SerializeInstruction(const Instruction& instr) { - std::vector fields; - // Save the opcode. - VLOG(2) << "Serializing: " << instr << std::endl; - switch (instr.op) { - case Opcode::Move: { - // Number of fields = 2 - fields.assign({instr.from, instr.dst}); - break; - } - case Opcode::Ret: { - // Number of fields = 1 - fields.push_back(instr.result); - break; - } - case Opcode::Fatal: { - // Number of fields = 0 - break; - } - case Opcode::InvokePacked: { - // Number of fields = 3 + instr.arity - // Note that arity includes both input arguments and outputs. We will - // put all the `arity` number of fields in the end for serialization. - fields.assign({instr.packed_index, instr.arity, instr.output_size}); - // Save the args. - fields.insert(fields.end(), instr.packed_args, instr.packed_args + instr.arity); - break; - } - case Opcode::AllocTensor: { - // Number of fields = 7 + instr.alloc_tensor.ndim - fields.push_back(instr.alloc_tensor.storage); - fields.push_back(instr.alloc_tensor.offset); - // Save `DLDataType` and the dst register. - const auto& dtype = instr.alloc_tensor.dtype; - fields.push_back(dtype.code); - fields.push_back(dtype.bits); - fields.push_back(dtype.lanes); - - // The number of dimensions is not needed for constructing an - // `AllocTensor` instruction as it equals to the length of the `shape` - // vector. However, we save it to conveniently deserialize the instruction - // because we will know how many fields are needed by the `shape` argument. - fields.push_back(instr.alloc_tensor.ndim); - fields.push_back(instr.dst); - - // Save the shape of the tensor. - // Note that this field is rotated to the end of the list. - fields.insert(fields.end(), instr.alloc_tensor.shape, - instr.alloc_tensor.shape + instr.alloc_tensor.ndim); - break; - } - case Opcode::AllocTensorReg: { - // Number of fields = 7 - fields.push_back(instr.alloc_tensor_reg.storage); - fields.push_back(instr.alloc_tensor_reg.offset); - fields.push_back(instr.alloc_tensor_reg.shape_register); - // Save `DLDataType` and the dst register. - const auto& dtype = instr.alloc_tensor_reg.dtype; - fields.push_back(dtype.code); - fields.push_back(dtype.bits); - fields.push_back(dtype.lanes); - fields.push_back(instr.dst); - break; - } - case Opcode::AllocStorage: { - fields.push_back(instr.alloc_storage.allocation_size); - fields.push_back(instr.alloc_storage.alignment); - // Save `DLDataType` and the dst register. - const auto& dtype = instr.alloc_storage.dtype_hint; - fields.push_back(dtype.code); - fields.push_back(dtype.bits); - fields.push_back(dtype.lanes); - fields.push_back(instr.alloc_storage.device_index); - fields.push_back(instr.alloc_storage.ndim); - fields.push_back(instr.dst); - - // Save the shape of the tensor. - // Note that this field is rotated to the end of the list. - fields.insert(fields.end(), instr.alloc_storage.shape, - instr.alloc_storage.shape + instr.alloc_storage.ndim); - break; - } - case Opcode::AllocADT: { - // Number of fields = 3 + instr.num_fields - fields.assign({instr.constructor_tag, instr.num_fields, instr.dst}); - - // Save the fields. - fields.insert(fields.end(), instr.datatype_fields, instr.datatype_fields + instr.num_fields); - break; - } - case Opcode::AllocClosure: { - // Number of fields = 3 + instr.num_freevar - fields.assign({instr.clo_index, instr.num_freevar, instr.dst}); - - // Save the free vars. - fields.insert(fields.end(), instr.free_vars, instr.free_vars + instr.num_freevar); - break; - } - case Opcode::If: { - // Number of fields = 4 - fields.assign({instr.if_op.test, instr.if_op.target, instr.if_op.true_offset, - instr.if_op.false_offset}); - break; - } - case Opcode::Invoke: { - // Number of fields = 3 + instr.num_args - fields.assign({instr.func_index, instr.num_args, instr.dst}); - - // Save the args. - fields.insert(fields.end(), instr.invoke_args_registers, - instr.invoke_args_registers + instr.num_args); - break; - } - case Opcode::InvokeClosure: { - // Number of fields = 3 + instr.num_closure_args - fields.assign({instr.closure, instr.num_closure_args, instr.dst}); - - // Save the args. - fields.insert(fields.end(), instr.closure_args, instr.closure_args + instr.num_closure_args); - break; - } - case Opcode::LoadConst: { - // Number of fields = 3 - fields.assign({instr.const_index, instr.device_index, instr.dst}); - break; - } - case Opcode::LoadConsti: { - // Number of fields = 2 - fields.assign({instr.load_consti.val, instr.dst}); - break; - } - case Opcode::GetField: { - // Number of fields = 3 - fields.assign({instr.object, instr.field_index, instr.dst}); - break; - } - case Opcode::GetTag: { - // Number of fields = 2 - fields.assign({instr.get_tag.object, instr.dst}); - break; - } - case Opcode::Goto: { - // Number of fields = 1 - fields.push_back(instr.pc_offset); - break; - } - case Opcode::ShapeOf: { - // Number of fields = 2 - fields.assign({instr.shape_of.tensor, instr.dst}); - break; - } - case Opcode::ReshapeTensor: { - // Number of fields = 3 - fields.assign({instr.reshape_tensor.tensor, instr.reshape_tensor.newshape, instr.dst}); - break; - } - case Opcode::DeviceCopy: { - // Number of fields = 4 - fields.assign({instr.device_copy.src, instr.device_copy.src_device_index, - instr.device_copy.dst_device_index, instr.dst}); - break; - } - case Opcode::KillRegister: { - fields.assign({instr.dst}); - break; - } - default: - LOG(FATAL) << "Invalid opcode" << static_cast(instr.op); - break; - } - - return VMInstructionSerializer(static_cast(instr.op), fields); -} - -void Executable::SaveCodeSection(dmlc::Stream* strm) { - // Save the number of functions. - strm->Write(static_cast(this->functions.size())); - for (const auto& func : this->functions) { - // Save the function info. - VMFunctionSerializer func_format(func.name, func.register_file_size, func.instructions.size(), - func.params, func.param_device_indexes); - func_format.Save(strm); - - // Serialize each instruction. - for (const auto& instr : func.instructions) { - const auto& serialized_instr = SerializeInstruction(instr); - serialized_instr.Save(strm); - } - } -} - -void LoadHeader(dmlc::Stream* strm) { - // Check header. - uint64_t header; - STREAM_CHECK(strm->Read(&header), "header"); - STREAM_CHECK(header == kTVMVMBytecodeMagic, "header"); - - // Check version. - std::string version; - STREAM_CHECK(strm->Read(&version), "version"); - STREAM_CHECK(version == TVM_VERSION, "version"); -} - -runtime::Module Executable::GetLib() const { - ICHECK_LE(this->imports_.size(), 1) - << "The kernel library must be imported as the only module in an Executable"; - - if (this->imports().size() == 0) { - return Module(nullptr); - } else { - return this->imports_[0]; - } -} - -void Executable::SetLib(const runtime::Module& lib) { - ICHECK(lib.defined()) << "the provided library can not be null"; - - ICHECK_EQ(this->imports_.size(), 0) - << "A VMExecutable should never have more than one import inside an the executable, \n" - << "the first import should *always* be the library containing" - << "the platform specific kernel code"; - - this->Import(lib); -} - -runtime::Module Executable::Load(const std::string& code, const runtime::Module lib) { - auto exec = make_object(); - - // Support null-initialization of lib, to enable initialization during - // deserialization before we have deserialized the imports. - if (lib.defined()) { - exec->SetLib(lib); - } - - exec->code_ = code; - dmlc::MemoryStringStream strm(&exec->code_); - - // Load header. - LoadHeader(&strm); - - // Virtual devices section - exec->LoadVirtualDevicesSection(&strm); - - // Global section. - exec->LoadGlobalSection(&strm); - - // Constant section. - exec->LoadConstantSection(&strm); - - // Primitive names that will be invoked by `InvokePacked` instructions. - exec->LoadPrimitiveOpNames(&strm); - - // Code section. - exec->LoadCodeSection(&strm); - - return runtime::Module(exec); -} - -void Executable::LoadVirtualDevicesSection(dmlc::Stream* strm) { - STREAM_CHECK(strm->Read(&virtual_devices), "virtual_device"); - STREAM_CHECK(strm->Read(&host_device_index), "virtual_device"); - ICHECK(host_device_index >= 0 && host_device_index < static_cast(virtual_devices.size())); -} - -void Executable::LoadGlobalSection(dmlc::Stream* strm) { - std::vector globals; - STREAM_CHECK(strm->Read(&globals), "global"); - for (size_t i = 0; i < globals.size(); i++) { - this->global_map.insert({globals[i], i}); - } -} - -void Executable::LoadPrimitiveOpNames(dmlc::Stream* strm) { - std::vector primitive_names; - STREAM_CHECK(strm->Read(&primitive_names), "primitive name"); - for (size_t i = 0; i < primitive_names.size(); i++) { - this->primitive_map.insert({primitive_names[i], i}); - } - - std::map> primitive_attrs; - STREAM_CHECK(strm->Read(&primitive_attrs), "primitive attrs"); - for (const auto& fn : primitive_attrs) { - std::vector> attrs; - for (const auto& elem : fn.second) { - attrs.push_back({elem.first, String(elem.second)}); - } - this->op_attrs[fn.first] = Map(attrs.begin(), attrs.end()); - } -} - -// Extract the `cnt` number of fields started at `start` from the list -// `instr_fields`. -inline std::vector ExtractFields(const std::vector& instr_fields, Index start, - Index cnt) { - ICHECK_LE(static_cast(start + cnt), instr_fields.size()); - std::vector ret; - for (auto i = start; i < start + cnt; i++) { - ret.push_back(instr_fields[i]); - } - return ret; -} - -Instruction DeserializeInstruction(const VMInstructionSerializer& instr) { - Opcode opcode = static_cast(instr.opcode); - switch (opcode) { - case Opcode::Move: { - // Number of fields = 2 - DCHECK_EQ(instr.fields.size(), 2U); - return Instruction::Move(instr.fields[0], instr.fields[1]); - } - case Opcode::Ret: { - // Number of fields = 1 - DCHECK_EQ(instr.fields.size(), 1U); - return Instruction::Ret(instr.fields[0]); - } - case Opcode::Fatal: { - // Number of fields = 0 - DCHECK(instr.fields.empty()); - return Instruction::Fatal(); - } - case Opcode::InvokePacked: { - // Number of fields = 3 + instr.arity - DCHECK_GE(instr.fields.size(), 3U); - DCHECK_EQ(instr.fields.size(), 3U + static_cast(instr.fields[1])); - - Index packed_index = instr.fields[0]; - Index arity = instr.fields[1]; - Index output_size = instr.fields[2]; - std::vector args = ExtractFields(instr.fields, 3, arity); - return Instruction::InvokePacked(packed_index, arity, output_size, args); - } - case Opcode::AllocTensor: { - // Number of fields = 7 + instr.alloc_tensor.ndim - DCHECK_GE(instr.fields.size(), 7U); - DCHECK_EQ(instr.fields.size(), 7U + static_cast(instr.fields[5])); - - RegName storage_reg = instr.fields[0]; - RegName offset = instr.fields[1]; - - DLDataType dtype; - dtype.code = instr.fields[2]; - dtype.bits = instr.fields[3]; - dtype.lanes = instr.fields[4]; - - Index ndim = instr.fields[5]; - RegName dst = instr.fields[6]; - - std::vector shape = ExtractFields(instr.fields, 7, ndim); - - return Instruction::AllocTensor(storage_reg, offset, shape, dtype, dst); - } - case Opcode::AllocTensorReg: { - // Number of fields = 7 - DCHECK_EQ(instr.fields.size(), 7U); - - RegName storage_reg = instr.fields[0]; - RegName offset = instr.fields[1]; - Index shape_register = instr.fields[2]; - - DLDataType dtype; - dtype.code = instr.fields[3]; - dtype.bits = instr.fields[4]; - dtype.lanes = instr.fields[5]; - - RegName dst = instr.fields[6]; - - return Instruction::AllocTensorReg(storage_reg, offset, shape_register, dtype, dst); - } - case Opcode::AllocADT: { - // Number of fields = 3 + instr.num_fields - DCHECK_GE(instr.fields.size(), 3U); - DCHECK_EQ(instr.fields.size(), 3U + static_cast(instr.fields[1])); - - Index constructor_tag = instr.fields[0]; - Index num_fields = instr.fields[1]; - RegName dst = instr.fields[2]; - std::vector fields = ExtractFields(instr.fields, 3, num_fields); - - return Instruction::AllocADT(constructor_tag, num_fields, fields, dst); - } - case Opcode::AllocClosure: { - // Number of fields = 3 + instr.num_freevar - DCHECK_GE(instr.fields.size(), 3U); - DCHECK_EQ(instr.fields.size(), 3U + static_cast(instr.fields[1])); - - Index clo_index = instr.fields[0]; - Index num_freevar = instr.fields[1]; - RegName dst = instr.fields[2]; - std::vector free_vars = ExtractFields(instr.fields, 3, num_freevar); - - return Instruction::AllocClosure(clo_index, num_freevar, free_vars, dst); - } - case Opcode::AllocStorage: { - // Number of fields = 9 - DCHECK_GE(instr.fields.size(), 9U); - Index allocation_size = instr.fields[0]; - Index alignment = instr.fields[1]; - - DLDataType dtype; - dtype.code = instr.fields[2]; - dtype.bits = instr.fields[3]; - dtype.lanes = instr.fields[4]; - - Index device_type = instr.fields[5]; - Index ndim = instr.fields[6]; - RegName dst = instr.fields[7]; - std::vector shape = ExtractFields(instr.fields, 8, ndim); - - return Instruction::AllocStorage(allocation_size, alignment, dtype, device_type, shape, dst); - } - case Opcode::If: { - // Number of fields = 4 - DCHECK_EQ(instr.fields.size(), 4U); - Index test = instr.fields[0]; - Index target = instr.fields[1]; - Index true_offset = instr.fields[2]; - Index false_offset = instr.fields[3]; - - return Instruction::If(test, target, true_offset, false_offset); - } - case Opcode::Invoke: { - // Number of fields = 3 + instr.num_args - DCHECK_GE(instr.fields.size(), 3U); - DCHECK_EQ(instr.fields.size(), 3U + static_cast(instr.fields[1])); - - Index func_index = instr.fields[0]; - Index num_args = instr.fields[1]; - RegName dst = instr.fields[2]; - std::vector args = ExtractFields(instr.fields, 3, num_args); - - return Instruction::Invoke(func_index, args, dst); - } - case Opcode::InvokeClosure: { - // Number of fields = 3 + instr.num_closure_args - DCHECK_GE(instr.fields.size(), 3U); - DCHECK_EQ(instr.fields.size(), 3U + static_cast(instr.fields[1])); - - Index closure = instr.fields[0]; - Index num_closure_args = instr.fields[1]; - RegName dst = instr.fields[2]; - std::vector args = ExtractFields(instr.fields, 3, num_closure_args); - - return Instruction::InvokeClosure(closure, args, dst); - } - case Opcode::LoadConst: { - // Number of fields = 3 - DCHECK_EQ(instr.fields.size(), 3U); - return Instruction::LoadConst(instr.fields[0], instr.fields[1], instr.fields[2]); - } - case Opcode::LoadConsti: { - // Number of fields = 2 - DCHECK_EQ(instr.fields.size(), 2U); - return Instruction::LoadConsti(instr.fields[0], instr.fields[1]); - } - case Opcode::GetField: { - // Number of fields = 3 - DCHECK_EQ(instr.fields.size(), 3U); - return Instruction::GetField(instr.fields[0], instr.fields[1], instr.fields[2]); - } - case Opcode::GetTag: { - // Number of fields = 2 - DCHECK_EQ(instr.fields.size(), 2U); - return Instruction::GetTag(instr.fields[0], instr.fields[1]); - } - case Opcode::Goto: { - // Number of fields = 1 - DCHECK_EQ(instr.fields.size(), 1U); - return Instruction::Goto(instr.fields[0]); - } - case Opcode::ShapeOf: { - // Number of fields = 2 - DCHECK_EQ(instr.fields.size(), 2U); - return Instruction::ShapeOf(instr.fields[0], instr.fields[1]); - } - case Opcode::ReshapeTensor: { - // Number of fields = 3 - DCHECK_EQ(instr.fields.size(), 3U); - return Instruction::ReshapeTensor(instr.fields[0], instr.fields[1], instr.fields[2]); - } - case Opcode::DeviceCopy: { - // Number of fields = 4 - DCHECK_EQ(instr.fields.size(), 4U); - return Instruction::DeviceCopy(instr.fields[0], instr.fields[1], instr.fields[2], - instr.fields[3]); - } - case Opcode::KillRegister: { - DCHECK_EQ(instr.fields.size(), 1U); - return Instruction::KillRegister(instr.fields[0]); - } - default: - LOG(FATAL) << "Invalid opcode" << instr.opcode; - } -} - -void Executable::LoadCodeSection(dmlc::Stream* strm) { - // Load the number of functions. - uint64_t sz; - STREAM_CHECK(strm->Read(&sz, sizeof(sz)), "code"); - - size_t num_funcs = static_cast(sz); - this->functions.resize(num_funcs); - for (size_t i = 0; i < num_funcs; i++) { - // Load the function info. - VMFunctionSerializer loaded_func; - STREAM_CHECK(loaded_func.Load(strm), "code/function"); - - // Load the instructions. - std::vector instructions; - for (size_t j = 0; j < loaded_func.num_instructions; j++) { - VMInstructionSerializer instr; - std::vector instr_fields; - STREAM_CHECK(instr.Load(strm), "code/instruction"); - instructions.push_back(DeserializeInstruction(instr)); - } - - // Create the VM function. - VMFunction vm_func = - VMFunction(loaded_func.name, loaded_func.params, instructions, - loaded_func.register_file_size, loaded_func.param_device_indexes); - auto it = this->global_map.find(loaded_func.name); - ICHECK(it != this->global_map.end()); - ICHECK_LE(it->second, this->global_map.size()); - this->functions[it->second] = vm_func; - } -} - -void Executable::SaveToBinary(dmlc::Stream* stream) { - auto code_bytes = this->Save(); - std::string code(code_bytes.data, code_bytes.size); - stream->Write(code); - - ICHECK(this->imports()[0].defined()) << "the library must be imported before serialization"; -} - -Module ExecutableLoadBinary(void* strm) { - dmlc::Stream* stream = static_cast(strm); - std::string code; - stream->Read(&code); - auto exec = Executable::Load(code, Module()); - return exec; -} - -void Executable::SaveToFile(const String& path, const String& format) { - tvm::runtime::SimpleBinaryFileStream stream(path, "wb"); - SaveToBinary(&stream); -} - -TVM_REGISTER_GLOBAL("runtime.module.loadbinary_VMExecutable").set_body_typed(ExecutableLoadBinary); - -// Load module from module. -Module ExecutableLoadFile(const std::string& file_name, const String& format) { - tvm::runtime::SimpleBinaryFileStream stream(file_name, "rb"); - auto exec = ExecutableLoadBinary(reinterpret_cast(&stream)); - return exec; -} - -TVM_REGISTER_GLOBAL("runtime.module.loadfile_VMExecutable").set_body_typed(ExecutableLoadFile); - -TVM_REGISTER_GLOBAL("runtime.GetNumOfGlobals").set_body([](TVMArgs args, TVMRetValue* rv) { - runtime::Module mod = args[0]; - const auto* exec = dynamic_cast(mod.operator->()); - ICHECK(exec); - *rv = static_cast(exec->global_map.size()); -}); - -TVM_REGISTER_GLOBAL("runtime.GetGlobalFields").set_body([](TVMArgs args, TVMRetValue* rv) { - runtime::Module mod = args[0]; - const auto* exec = dynamic_cast(mod.operator->()); - ICHECK(exec); - int idx = args[1]; - std::vector> globals(exec->global_map.begin(), - exec->global_map.end()); - auto comp = [](const std::pair& a, const std::pair& b) { - return a.second < b.second; - }; - std::sort(globals.begin(), globals.end(), comp); - ICHECK_LT(idx, globals.size()); - *rv = globals[idx].first; -}); - -TVM_REGISTER_GLOBAL("runtime.GetNumOfPrimitives").set_body([](TVMArgs args, TVMRetValue* rv) { - runtime::Module mod = args[0]; - const auto* exec = dynamic_cast(mod.operator->()); - ICHECK(exec); - *rv = static_cast(exec->primitive_map.size()); -}); - -TVM_REGISTER_GLOBAL("runtime.GetPrimitiveFields").set_body([](TVMArgs args, TVMRetValue* rv) { - runtime::Module mod = args[0]; - const auto* exec = dynamic_cast(mod.operator->()); - ICHECK(exec); - int idx = args[1]; - ICHECK_GE(idx, 0); - ICHECK_LT(idx, exec->primitive_map.size()); - - for (const auto& it : exec->primitive_map) { - if (idx == static_cast(it.second)) { - *rv = it.first; - break; - } - } -}); - -TVM_REGISTER_GLOBAL("runtime.Load_Executable") - .set_body_typed([](std::string code, runtime::Module lib) { - return Executable::Load(code, lib); - }); - -} // namespace vm -} // namespace runtime -} // namespace tvm diff --git a/src/runtime/vm/profiler/vm.cc b/src/runtime/vm/profiler/vm.cc deleted file mode 100644 index 7df6b928a32e..000000000000 --- a/src/runtime/vm/profiler/vm.cc +++ /dev/null @@ -1,230 +0,0 @@ -/* - * 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 src/runtime/vm/profiler/vm.cc - * \brief The Relay debug virtual machine. - */ - -#include "vm.h" - -#include -#include -#include - -#include -#include -#include -#include -#include -#include -#include -#include -#include - -namespace tvm { -namespace runtime { -namespace vm { - -PackedFunc VirtualMachineDebug::GetFunction(const String& name, - const ObjectPtr& sptr_to_self) { - if (name == "profile") { - return TypedPackedFunc)>( - [sptr_to_self, this](String arg_name, Array collectors) { - std::vector devices; - for (auto dev : devices_) { - if (dev.device_type > 0) { - devices.push_back(dev); - } - } - - // We cannot send Arrays over rpc, so in order to support profiling - // on remotes, we accept a nullptr for collectors. - if (collectors.defined()) { - std::vector cs(collectors.begin(), collectors.end()); - prof_ = profiling::Profiler(devices, cs, {{String("Executor"), String("VM")}}); - } else { - prof_ = profiling::Profiler(devices, {}, {{String("Executor"), String("VM")}}); - } - - auto invoke = VirtualMachine::GetFunction("invoke", sptr_to_self); - // warmup - for (int i = 0; i < 3; i++) { - invoke(arg_name); - } - - prof_.operator*().Start(); - invoke(arg_name); - prof_.operator*().Stop(); - auto report = prof_.operator*().Report(); - prof_ = std::nullopt; // releases hardware counters - return report; - }); - } else if (name == "profile_rpc") { - // We cannot return a Report over RPC because TVM RPC mechanism only - // supports a subset of Object classes. Instead we serialize it on the - // remote (here) and deserialize it on the other end. - return TypedPackedFunc([sptr_to_self, this](std::string arg_name) { - PackedFunc profile = GetFunction("profile", sptr_to_self); - profiling::Report report = profile(arg_name, Array()); - return report->AsJSON(); - }); - } else { - return VirtualMachine::GetFunction(name, sptr_to_self); - } -} - -void VirtualMachineDebug::LoadExecutable(const ObjectPtr& exec) { - VirtualMachine::LoadExecutable(exec); - for (auto kv : exec_->primitive_map) { - packed_index_map_[kv.second] = kv.first; - } -} - -void VirtualMachineDebug::OpStartHook(Instruction instr) { - if (prof_ && prof_.operator*().IsRunning()) { - if (instr.op == Opcode::LoadConst) { - Device dev = GetDevice(exec_->const_device_indexes[instr.const_index]); - prof_.operator*().StartCall("VM::LoadConst", dev, {}); - } else if (instr.op == Opcode::DeviceCopy) { - Device dst_dev = GetDevice(instr.device_copy.dst_device_index); - prof_.operator*().StartCall("VM::DeviceCopy", dst_dev, {}); - } else if (instr.op == Opcode::ReshapeTensor) { - prof_.operator*().StartCall("VM::ReshapeTensor", devices_[exec_->host_device_index], {}); - } else if (instr.op == Opcode::AllocTensor) { - auto shape = std::vector(instr.alloc_tensor.ndim); - - for (uint32_t i = 0; i < instr.alloc_tensor.ndim; ++i) { - shape[i] = instr.alloc_tensor.shape[i]; - } - auto storage_obj = ReadRegister(instr.alloc_tensor.storage); - auto storage = Downcast(storage_obj); - prof_.operator*().StartCall( - "VM::AllocTensor", storage->buffer.device, - {{"Argument Shapes", profiling::ShapeString(shape, instr.alloc_tensor.dtype)}}); - } else if (instr.op == Opcode::AllocTensorReg) { - auto storage_obj = ReadRegister(instr.alloc_tensor_reg.storage); - auto storage = Downcast(storage_obj); - Device cpu_dev = GetDevice(exec_->host_device_index); - auto shape_obj = ReadRegister(instr.alloc_tensor_reg.shape_register); - NDArray shape_tensor = Downcast(shape_obj).CopyTo(cpu_dev); - prof_.operator*().StartCall( - "VM::AllocTensorReg", storage->buffer.device, - {{"Argument Shapes", - profiling::ShapeString(shape_tensor, instr.alloc_tensor_reg.dtype)}}); - } else if (instr.op == Opcode::AllocStorage) { - std::ostringstream shape; - if (instr.alloc_storage.ndim > 0) { - std::string shape_str = "["; - for (uint32_t i = 0; i < instr.alloc_storage.ndim; ++i) { - if (i > 0) { - shape_str += ", "; - } - shape_str += std::to_string(instr.alloc_storage.shape[i]); - } - shape_str += "]"; - shape << DLDataType2String(instr.alloc_storage.dtype_hint) << shape_str; - } else { - auto size = LoadScalarInt(instr.alloc_storage.allocation_size); - shape << DLDataType2String(instr.alloc_storage.dtype_hint) << "[" << size << "]"; - } - Device dev = GetDevice(instr.alloc_storage.device_index); - prof_.operator*().StartCall("VM::AllocStorage", dev, - {{"VM::Argument Shapes", String(shape.str())}}); - } else { - prof_.operator*().StartCall("VM::UnknownOp", GetDevice(exec_->host_device_index), {}); - } - } -} - -void VirtualMachineDebug::OpStopHook() { - if (prof_ && prof_.operator*().IsRunning()) { - prof_.operator*().StopCall(); - } -} - -void VirtualMachineDebug::InvokePacked(Index packed_index, const PackedFunc& func, Index arg_count, - Index output_size, const std::vector& args) { - ICHECK(exec_); - ICHECK(!devices_.empty()) << "Device has not been initialized yet."; - if (prof_ && prof_.operator*().IsRunning()) { - // The device of any input of the operator is used for synchronization. - ICHECK_GT(arg_count, 0U); - ObjectRef arg = args[0]; - while (arg->IsInstance()) { - ADT adt = Downcast(arg); - arg = adt[0]; - } - ICHECK(arg->IsInstance()); - auto nd_array = Downcast(arg); - auto dev = nd_array->device; - - // get argument sizes - std::vector shapes; - for (Index i = 0; i < arg_count; i++) { - if (const auto* obj = args[i].as()) { - for (size_t fi = 0; fi < obj->size; ++fi) { - auto o = (*obj)[fi]; - shapes.push_back(Downcast(o)); - } - } else { - shapes.push_back(Downcast(args[i])); - } - } - - std::unordered_map metrics; - - ICHECK(exec_->op_attrs.find(packed_index) != exec_->op_attrs.end()) - << packed_index_map_[packed_index] << " not found in op attrs"; - - auto& op_attrs = exec_->op_attrs.at(packed_index); - for (auto p : op_attrs) { - if (std::string(p.first).find("layout") != std::string::npos) { - metrics[p.first] = p.second; - } - } - auto it = op_attrs.find("hash"); - if (it != op_attrs.end()) { - metrics["Hash"] = Downcast((*it).second); - } - metrics["Argument Shapes"] = profiling::ShapeString(shapes); - - prof_.operator*().StartCall(packed_index_map_[packed_index], dev, metrics); - } - VirtualMachine::InvokePacked(packed_index, func, arg_count, output_size, args); - if (prof_ && prof_.operator*().IsRunning()) { - prof_.operator*().StopCall(); - } -} - -runtime::Module CreateVirtualMachineDebug(Executable* exec) { - auto vm = make_object(); - vm->LoadExecutable(GetObjectPtr(exec)); - return runtime::Module(vm); -} - -TVM_REGISTER_GLOBAL("runtime._VirtualMachineDebug").set_body([](TVMArgs args, TVMRetValue* rv) { - runtime::Module mod = args[0]; - auto* exec = dynamic_cast(mod.operator->()); - *rv = CreateVirtualMachineDebug(exec); -}); - -} // namespace vm -} // namespace runtime -} // namespace tvm diff --git a/src/runtime/vm/profiler/vm.h b/src/runtime/vm/profiler/vm.h deleted file mode 100644 index a91869454e3b..000000000000 --- a/src/runtime/vm/profiler/vm.h +++ /dev/null @@ -1,65 +0,0 @@ -/* - * 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 src/runtime/vm/profiler/vm.h - * \brief The Relay debug virtual machine. - */ - -#ifndef TVM_RUNTIME_VM_PROFILER_VM_H_ -#define TVM_RUNTIME_VM_PROFILER_VM_H_ - -#include -#include - -#include -#include -#include -#include -#include - -namespace tvm { -namespace runtime { -namespace vm { - -class VirtualMachineDebug : public VirtualMachine { - public: - VirtualMachineDebug() : VirtualMachine(), prof_({}) {} - - PackedFunc GetFunction(const String& name, const ObjectPtr& sptr_to_self) final; - - void LoadExecutable(const ObjectPtr& exec) final; - - ~VirtualMachineDebug() {} - - private: - void InvokePacked(Index packed_index, const PackedFunc& func, Index arg_count, Index output_size, - const std::vector& args) final; - void OpStartHook(Instruction instr) final; - void OpStopHook() final; - - std::unordered_map packed_index_map_; - std::optional prof_; -}; - -} // namespace vm -} // namespace runtime -} // namespace tvm - -#endif // TVM_RUNTIME_VM_PROFILER_VM_H_ diff --git a/src/runtime/vm/serialize_utils.h b/src/runtime/vm/serialize_utils.h deleted file mode 100644 index 04a79c9b0210..000000000000 --- a/src/runtime/vm/serialize_utils.h +++ /dev/null @@ -1,168 +0,0 @@ -/* - * 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 src/runtime/vm/serialize_utils.h - * \brief Definitions of helpers for serializing and deserializing a Relay VM. - */ -#ifndef TVM_RUNTIME_VM_SERIALIZE_UTILS_H_ -#define TVM_RUNTIME_VM_SERIALIZE_UTILS_H_ - -#include -#include - -#include -#include -#include - -#include "../../support/utils.h" - -namespace tvm { -namespace runtime { -namespace vm { - -/*! \brief The magic number for the serialized VM bytecode file */ -constexpr uint64_t kTVMVMBytecodeMagic = 0xD225DE2F4214151D; - -template -static inline uint64_t VectorHash(uint64_t key, const std::vector& values) { - for (const auto& it : values) { - key = support::HashCombine(key, it); - } - return key; -} - -// A struct to hold the funciton info in the code section. -struct VMFunctionSerializer { - /*! \brief The name of the VMFunction. */ - std::string name; - /*! \brief The number of registers used by the VMFunction. */ - Index register_file_size; - /*! \brief The number of instructions in the VMFunction. */ - size_t num_instructions; - /*! \brief The parameters of the VMFunction. */ - std::vector params; - /*! \brief The index for the devices holding each parameter of the VMFunction. */ - std::vector param_device_indexes; - - VMFunctionSerializer() = default; - - VMFunctionSerializer(const std::string& name, Index register_file_size, size_t num_instructions, - const std::vector& params, - const std::vector& param_device_indexes) - : name(name), - register_file_size(register_file_size), - num_instructions(num_instructions), - params(params), - param_device_indexes(param_device_indexes) {} - - /*! - * \brief Load the serialized function header. - * \param strm The stream used to load data. - * \return True if successful. Otherwise, false. - */ - bool Load(dmlc::Stream* strm) { - std::vector func_info; - if (!strm->Read(&func_info)) return false; - ICHECK_EQ(func_info.size(), 3U) << "Failed to decode the vm function." - << "\n"; - name = func_info[0]; - register_file_size = std::stoll(func_info[1]); - // Get the number of instructions. - num_instructions = static_cast(std::stoll(func_info[2])); - if (!strm->Read(¶ms)) return false; - if (!strm->Read(¶m_device_indexes)) return false; - return true; - } - - /*! - * \brief Save the VM function header into the serialized form. - * \param strm The stream used to save data. - */ - void Save(dmlc::Stream* strm) const { - std::vector func_info; - func_info.push_back(name); - func_info.push_back(std::to_string(register_file_size)); - func_info.push_back(std::to_string(num_instructions)); - strm->Write(func_info); - strm->Write(params); - strm->Write(param_device_indexes); - } -}; - -struct VMInstructionSerializer { - /*! \brief The opcode of the instruction. */ - Index opcode; - /*! \brief The fields of the instruction. */ - std::vector fields; - - VMInstructionSerializer() = default; - - VMInstructionSerializer(Index opcode, const std::vector& fields) - : opcode(opcode), fields(fields) {} - - /*! - * \brief Compute the hash of the serialized instruction. - * \return The hash that combines the opcode and all fields of the VM - * instruction. - */ - Index Hash() const { - uint64_t key = static_cast(opcode); - key = VectorHash(key, fields); - return key; - } - - /*! - * \brief Load the serialized instruction. - * \param strm The stream used to load data. - * \return True if successful. Otherwise, false. - */ - bool Load(dmlc::Stream* strm) { - std::vector instr; - if (!strm->Read(&instr)) return false; - ICHECK_GE(instr.size(), 2U); - Index loaded_hash = instr[0]; - opcode = instr[1]; - - for (size_t i = 2; i < instr.size(); i++) { - fields.push_back(instr[i]); - } - - Index hash = Hash(); - ICHECK_EQ(loaded_hash, hash) << "Found mismatch in hash for opcode: " << opcode << "\n"; - return true; - } - - /*! - * \brief Save the instruction into the serialized form. - * \param strm The stream used to save data. - */ - void Save(dmlc::Stream* strm) const { - Index hash = Hash(); - std::vector serialized({hash, opcode}); - serialized.insert(serialized.end(), fields.begin(), fields.end()); - strm->Write(serialized); - } -}; - -} // namespace vm -} // namespace runtime -} // namespace tvm - -#endif // TVM_RUNTIME_VM_SERIALIZE_UTILS_H_ diff --git a/src/runtime/vm/vm.cc b/src/runtime/vm/vm.cc deleted file mode 100644 index dfde076bfc30..000000000000 --- a/src/runtime/vm/vm.cc +++ /dev/null @@ -1,1034 +0,0 @@ -/* - * 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 src/runtime/vm/vm.cc - * \brief The Relay virtual machine runtime. - */ - -#include -#include -#include -#include -#include -#include -#include -#include - -#include -#include -#include -#include -#include - -#include "../file_utils.h" - -using namespace tvm::runtime; - -namespace tvm { -namespace runtime { -namespace vm { - -TVM_REGISTER_OBJECT_TYPE(VMClosureObj); - -VMClosure::VMClosure(size_t func_index, std::vector free_vars) { - auto ptr = make_object(); - ptr->func_index = func_index; - ptr->free_vars = std::move(free_vars); - data_ = std::move(ptr); -} - -void VMFunctionPrint(std::ostream& os, const VMFunction& vm_func) { - os << vm_func.name << ": " << std::endl; - for (size_t i = 0; i < vm_func.instructions.size(); ++i) { - os << i << ": " << vm_func.instructions[i] << ";" << std::endl; - } -} - -std::ostream& operator<<(std::ostream& os, const VMFunction& vm_func) { - VMFunctionPrint(os, vm_func); - return os; -} - -inline ObjectRef CopyTo(ObjectRef src, const DLDevice& dev, Optional mem_scope = NullOpt) { - if (src->IsInstance()) { - auto nd_array = Downcast(src); - // TODO(mbs): Should respect device id also. - // TODO(vvchernov): it still does not work for different device id - // due to simple implementation of Get() and AllocDataSpace() methods - // see tvm/src/runtime/c_runtime_api.cc: L139 - // tvm/src/runtime/cpu_device_api.cc: L47 - if (nd_array->device.device_type != dev.device_type || - nd_array->device.device_id != dev.device_id) { - VLOG(2) << "copying from " << nd_array->device.device_type << "[" - << nd_array->device.device_id << "] to " << dev.device_type << "[" << dev.device_id - << "]"; - return nd_array.CopyTo(dev, mem_scope); - } - return src; - } else { - ICHECK(src->IsInstance()) - << "VM data must be NDArray or a list of NDArray, but received: " << src->_type_key; - std::vector ret; - ADT adt = Downcast(src); - for (size_t i = 0; i < adt.size(); i++) { - ret.push_back(CopyTo(adt[i], dev, mem_scope)); - } - return ADT(adt->tag, ret.begin(), ret.end()); - } -} - -ShapeTuple ToShape(NDArray shape_tensor) { - std::vector shape; - auto rank = shape_tensor.Shape().size(); - auto dtype = shape_tensor.DataType(); - - // For 0-rank shapes we need to allocate a single scalar. - if (rank == 0) { - return shape; - } - - // Otherwise we should be rank-1, and we will extract the number of dimensions - // for the output vector. - ICHECK_EQ(rank, 1U) << "shape tensor should be a k-length vector, found " << rank; - int64_t ndim = shape_tensor.Shape().at(0); - shape.resize(ndim); - - const DLTensor* dl_tensor = shape_tensor.operator->(); - if (dtype.is_int() && dtype.bits() == 32 && dtype.lanes() == 1) { - int32_t* dims = reinterpret_cast(dl_tensor->data); - shape.assign(dims, dims + ndim); - } else if (dtype.is_int() && dtype.bits() == 64 && dtype.lanes() == 1) { - int64_t* dims = reinterpret_cast(dl_tensor->data); - shape.assign(dims, dims + ndim); - } else { - LOG(FATAL) << "invalid shape tensor datatype: " << dtype; - } - - return ShapeTuple(shape); -} - -void VirtualMachine::OpStartHook(Instruction instr) {} -void VirtualMachine::OpStopHook() {} - -PackedFunc VirtualMachine::GetFunction(const String& name, const ObjectPtr& sptr_to_self) { - if (name == "invoke") { - return PackedFunc([sptr_to_self, this](TVMArgs args, TVMRetValue* rv) { - ICHECK(exec_) << "The executable is not created yet."; - - std::string func_name = args[0]; - auto git = exec_->global_map.find(func_name); - ICHECK(git != exec_->global_map.end()) - << "Cannot find function " << func_name << " in the executable"; - auto func = exec_->functions[git->second]; - if (func.params.empty()) { - *rv = Invoke(func, {}); - } else { - auto it = inputs_.find(func_name); - ICHECK(it != inputs_.end()) << "Input has not been set for function " << func_name; - const std::vector& input_args = it->second; - if (set_outputs_enabled_.count(func_name) && set_outputs_enabled_[func_name]) { - ICHECK(outputs_.count(func_name)) - << "Outputs have not been set for function " << func_name; - *rv = Invoke(func, input_args, outputs_[func_name]); - outputs_[func_name].clear(); - set_outputs_enabled_[func_name] = false; - } else { - *rv = Invoke(func, input_args); - } - } - }); - } else if (name == "invoke_stateful") { - // TODO(tkonolige, jroesch, tqchen): invoke_stateful and get_output are - // stop-gap measure to allow using vm over a remote connection. - return PackedFunc([sptr_to_self, this](TVMArgs args, TVMRetValue* rv) { - PackedFunc invoke = GetFunction("invoke", sptr_to_self); - TVMRetValue rv_; - invoke.CallPacked(args, &rv_); - }); - } else if (name == "invoke_return_to_device") { - return PackedFunc([sptr_to_self, this](TVMArgs args, TVMRetValue* rv) { - Device host{static_cast(args[1].operator int()), args[2].operator int()}; - - SetInput(args[0].operator std::string(), args, 3); - PackedFunc invoke = GetFunction("invoke", sptr_to_self); - TVMRetValue rv_; - invoke.CallPacked(args, &rv_); // Invoke only uses the first arg, so the rest of the args - // should not cause an issue - if (rv_.type_code() == kTVMObjectHandle) { - ADT adt = Downcast(rv_.operator ObjectRef()); - std::vector transfered; - for (size_t i = 0; i < adt.size(); i++) { - transfered.push_back(CopyTo(adt[i], host)); - } - *rv = ADT(adt.tag(), transfered); - } else { - *rv = CopyTo(rv_, host); - } - }); - } else if (name == "get_output") { - return TypedPackedFunc([this](int64_t index) { - if (this->return_register_.as()) { - return Downcast(Downcast(this->return_register_)[index]); - } else { - CHECK_EQ(index, 0) << "VM output contains only one item, but you are trying to get the " - << index << "th."; - return Downcast(this->return_register_); - } - }); - } else if (name == "get_num_outputs") { - return TypedPackedFunc([this]() -> int64_t { - // single output is an NDArray not an ADT - if (this->return_register_.as()) { - return Downcast(this->return_register_).size(); - } else { - return 1; - } - }); - } else if (name == "get_input_index") { - return TypedPackedFunc( - [this](std::string input_name, std::string func_name) { - return GetInputIndexFromVMFunction(func_name, input_name); - }); - } else if (name == "init") { - return PackedFunc([sptr_to_self, this](TVMArgs args, TVMRetValue* rv) { - ICHECK_EQ(args.size() % 3, 0); - std::vector devices; - std::vector alloc_types; - for (int i = 0; i < args.size() / 3; ++i) { - Device dev; - int device_type = args[i * 3]; - dev.device_type = DLDeviceType(device_type); - dev.device_id = args[i * 3 + 1]; - int type = args[i * 3 + 2]; - devices.push_back(dev); - alloc_types.push_back(AllocatorType(type)); - } - this->Init(devices, alloc_types); - }); - } else if (name == "set_input") { - return PackedFunc( - [sptr_to_self, this](TVMArgs args, TVMRetValue* rv) { SetInput(args[0], args, 1); }); - } else if (name == "set_one_input") { - return PackedFunc([sptr_to_self, this](TVMArgs args, TVMRetValue* rv) { - ICHECK_EQ(args.size(), 3) << "The expected number of arguments is 3 " - << "(func_name, index or name, tensor)"; - SetOneInput(args[0], args[1], args[2]); - }); - } else if (name == "set_outputs") { - return PackedFunc( - [sptr_to_self, this](TVMArgs args, TVMRetValue* rv) { SetOutputs(args[0], args); }); - } else if (name == "load_late_bound_consts") { - return PackedFunc([this](TVMArgs args, TVMRetValue* rv) { - CHECK_EQ(args.size(), 1); - std::string path = args[0]; - exec_->LoadLateBoundConstantsFromFile(path); - }); - } else { - LOG(FATAL) << "Unknown packed function: " << name; - } -} - -void VirtualMachine::SetInput(std::string func_name, TVMArgs args, int offset) { - const auto& vm_func = CheckAndGetVMFunction(func_name); - size_t params_num = vm_func.params.size(); - ICHECK_EQ(args.size() - offset, params_num) - << "The number of provided parameters doesn't match the number of arguments"; - std::vector func_args(params_num); - for (int i = offset; i < args.size(); ++i) { - int index = i - offset; - Device dev = GetDevice(vm_func.param_device_indexes[index]); - SetInputTensorWithIndex(func_args, args[i], index, dev); - } - inputs_.erase(func_name); - inputs_.emplace(func_name, func_args); -} - -void VirtualMachine::SetOneInput(std::string func_name, const TVMArgValue& tag, - const TVMArgValue& tensor) { - const auto& vm_func = CheckAndGetVMFunction(func_name); - size_t params_num = vm_func.params.size(); - - int inp_index = 0; - if (tag.type_code() == kTVMArgInt) { - inp_index = tag; - } else if (tag.type_code() == kTVMStr) { - inp_index = static_cast(GetInputIndexFromName(vm_func.params, tag)); - } else { - LOG(FATAL) << "The type of input tensor tag (" << tag.type_code() - << ") doesn't match integer or string"; - } - ICHECK_LT(inp_index, params_num); - - CreateInputsOrCheckSize(func_name, params_num); - Device dev = GetDevice(vm_func.param_device_indexes[inp_index]); - SetInputTensorWithIndex(inputs_[func_name], tensor, inp_index, dev); -} - -void VirtualMachine::SetOutputs(std::string func_name, TVMArgs args) { - set_outputs_enabled_[func_name] = true; - size_t outputs_size = args.size(); - // First args is func_name - ICHECK_GT(outputs_size, 1) << "There is no output arguments set"; - - std::vector func_args(outputs_size - 1); - for (size_t i = 1; i < outputs_size; ++i) { - // TODO(vvchernov): device? - func_args[i - 1] = TensorFromTVMArgValueToObjectRef(args[i]); - } - outputs_.erase(func_name); - outputs_.emplace(func_name, func_args); -} - -void VirtualMachine::PrintInfoAndSetInputArgs(const VMFunction& func, - const std::vector& args) { - VLOG(2) << "Executing Function: " << std::endl << func; - for (int i = 0; i < static_cast(devices_.size()); ++i) { - VLOG(2) << "Device " << i << " has device type " << devices_[i].device_type << " and device id " - << devices_[i].device_id - << (i == exec_->host_device_index ? " (using as host device)" : ""); - } - - InvokeGlobal(func, args); -} - -void VirtualMachine::SetOutputTensorsToRegister(const std::string& func_name, - const std::vector& outputs) { - size_t size = outputs.size(); - - if (output_tensor_reg_indices_[func_name].empty()) { - output_tensor_reg_indices_[func_name] = GetOutputTensorRegIndices(); - } - auto& reg_indices = output_tensor_reg_indices_[func_name]; - ICHECK_EQ(reg_indices.size(), size) - << "Number of outside output tensors should be equal to model outputs number"; - size_t i = 0; - for (auto it = reg_indices.begin(); it != reg_indices.end(); ++it, ++i) { - WriteRegister(*it, outputs[i]); - } -} - -ObjectRef VirtualMachine::TensorFromTVMArgValueToObjectRef(const TVMArgValue& output_tensor) const { - if (output_tensor.type_code() == kTVMDLTensorHandle) { - DLTensor* dl_tensor = output_tensor; - return NDArray::FromExternalDLTensor(*dl_tensor); - } else if (output_tensor.type_code() == kTVMNDArrayHandle) { - return output_tensor.AsObjectRef(); - } else { - LOG(FATAL) << "It supports tensor of DLTensor or NDArray type only! Given type is " - << output_tensor.type_code(); - } - return ObjectRef(); -} - -int64_t VirtualMachine::GetInputIndexFromVMFunction(const std::string& func_name, - const std::string& input_name) const { - const auto& vm_func = CheckAndGetVMFunction(func_name); - return GetInputIndexFromName(vm_func.params, input_name); -} - -int64_t VirtualMachine::GetInputIndexFromName(const std::vector& params, - const std::string& input_name) const { - // TODO(vvchernov): excess integer type? - for (uint64_t i = 0; i < params.size(); i++) { - if (input_name == params[i]) { - return static_cast(i); - } - } - return static_cast(-1); -} - -const VMFunction& VirtualMachine::CheckAndGetVMFunction(const std::string& func_name) const { - ICHECK(exec_) << "The executable is not created yet."; - return exec_->GetVMFunctionWithName(func_name); -} - -void VirtualMachine::CreateInputsOrCheckSize(const std::string& func_name, size_t size) { - if (inputs_.count(func_name)) { - ICHECK_EQ(inputs_[func_name].size(), size) - << "The size of function" << func_name - << " doesn't match the number of provided parameters"; - } else { - std::vector func_args(size); - inputs_.emplace(func_name, func_args); - } -} - -void VirtualMachine::SetInputTensorWithIndex(std::vector& tensors, - const TVMArgValue& inp_tensor, int index, Device dev) { - if (inp_tensor.type_code() == kTVMDLTensorHandle) { - if (NDArray::AbilityOfZeroCopyForDLTensor(inp_tensor, dev)) { - tensors[index] = NDArray::FromExternalDLTensor(*inp_tensor); - } else { - tensors[index] = NDArray::NewFromDLTensor(inp_tensor, dev); - } - } else { - tensors[index] = CopyTo(inp_tensor, dev); - } -} - -inline Device VirtualMachine::GetDevice(Index device_index) const { - ICHECK_GE(devices_.size(), device_index) << "invalid device index: " << device_index; - return devices_[device_index]; -} - -inline Allocator* VirtualMachine::GetAllocator(Index device_index) const { - ICHECK_GE(allocators_.size(), device_index) << "invalid device index: " << device_index; - return allocators_[device_index]; -} - -void VirtualMachine::PushFrame(Index arg_count, Index ret_pc, const VMFunction& vm_func) { - auto frame = VMFrame(ret_pc, func_index_, arg_count, code_, vm_func.register_file_size); - frames_.push_back(frame); -} - -Index VirtualMachine::PopFrame() { - ICHECK_GT(frames_.size(), 0); - const VMFrame& fr = frames_.back(); - func_index_ = fr.func_index; - code_ = fr.code; - pc_ = fr.pc; - auto call_stack_size = frames_.size(); - frames_.pop_back(); - return call_stack_size; -} - -void VirtualMachine::InvokeGlobal(const VMFunction& func, const std::vector& args) { - VLOG(2) << "Invoking global " << func.name << " with " << args.size() << " args"; - - PushFrame(func.params.size(), this->pc_ + 1, func); - for (size_t i = 0; i < args.size(); ++i) { - WriteRegister(i, args[i]); - VLOG(2) << "arg " << i << " = " - << RuntimeObject2String(args[i], GetDevice(exec_->host_device_index)); - } - - code_ = func.instructions.data(); - pc_ = 0; -} - -ObjectRef VirtualMachine::Invoke(const VMFunction& func, const std::vector& args) { - PrintInfoAndSetInputArgs(func, args); - RunLoop(); - return return_register_; -} - -ObjectRef VirtualMachine::Invoke(const std::string& name, const std::vector& args) { - ICHECK(exec_) << "The executable has not been created yet."; - auto it = exec_->global_map.find(name); - ICHECK(it != exec_->global_map.end()) << "Cannot find function " << name << " in the executable"; - Index func_index = it->second; - VLOG(2) << "Invoke Global " << name << " at index " << func_index; - return Invoke(exec_->functions[func_index], args); -} - -ObjectRef VirtualMachine::Invoke(const VMFunction& func, const std::vector& input_args, - const std::vector& output_args) { - PrintInfoAndSetInputArgs(func, input_args); - SetOutputTensorsToRegister(func.name, output_args); - RunLoop(output_tensor_reg_indices_[func.name]); - return return_register_; -} - -void VirtualMachine::InvokePacked(Index packed_index, const PackedFunc& func, Index arg_count, - Index output_size, const std::vector& args) { - size_t arity = 0; - for (Index i = 0; i < arg_count; i++) { - if (const auto* obj = args[i].as()) { - arity += obj->size; - } else { - ++arity; - } - } - - std::vector values(arity); - std::vector codes(arity); - runtime::TVMArgsSetter setter(values.data(), codes.data()); - int idx = 0; - bool is_empty_output = false; - for (Index i = 0; i < arg_count; i++) { - if (const auto* dt_cell = args[i].as()) { - for (size_t fi = 0; fi < dt_cell->size; ++fi) { - auto obj = (*dt_cell)[fi]; - auto nd_array = Downcast(obj); - setter(idx++, nd_array); - } - } else { - auto nd_array = Downcast(args[i]); - // We can safely skip CallPacked if there is only one - // output and it is empty. - if (i == arg_count - 1 && output_size == 1) { - for (const auto& dim : nd_array.Shape()) { - if (!dim) { - is_empty_output = true; - break; - } - } - } - setter(idx++, nd_array); - } - } - - if (!is_empty_output) { - TVMRetValue rv; - func.CallPacked(TVMArgs(values.data(), codes.data(), arity), &rv); - } -} - -void VirtualMachine::LoadExecutable(const ObjectPtr& exec) { - ICHECK(exec) << "The executable is not created yet."; - ICHECK(exec->late_bound_constant_names.empty()) - << "Need to load late-bound-constants before creating VM"; - exec_ = exec; - - runtime::Module lib = exec_->GetLib(); - - ICHECK(exec_->primitive_map.empty() || lib.operator->()) - << "If the executable has declared primitive functions, the " - << "generated kernel library must non-be null."; - - for (const auto& it : exec_->primitive_map) { - const auto& packed_name = it.first; - auto packed_index = static_cast(it.second); - if (packed_funcs_.size() <= packed_index) { - packed_funcs_.resize(packed_index + 1); - } - tvm::runtime::PackedFunc pf = lib.GetFunction(packed_name, /*query_imports=*/true); - ICHECK(pf != nullptr) << "Cannot find function in module: " << packed_name; - packed_funcs_[packed_index] = pf; - } - for (size_t i = 0; i < packed_funcs_.size(); ++i) { - ICHECK(packed_funcs_[i] != nullptr) << "Packed function " << i << " is not initialized"; - } -} - -void VirtualMachine::Init(const std::vector& physical_devices, - const std::vector& alloc_types) { - ICHECK_EQ(physical_devices.size(), alloc_types.size()); - - // Find a physical device to represent each virtual device the VM code requires. - // (Recall the VM instructions refer to devices by "device index" into this vector of - // virtual devices.) - const size_t num_virtual_devices = exec_->virtual_devices.size(); - devices_.reserve(num_virtual_devices); - allocators_.reserve(num_virtual_devices); - - for (size_t device_index = 0; device_index < num_virtual_devices; ++device_index) { - // We'll retain the legacy behaviour and just match by device type. - // TODO(mbs): Generalize. - DLDeviceType virtual_device_type = exec_->virtual_devices[device_index].first.device_type; - auto itr = std::find_if(physical_devices.begin(), physical_devices.end(), - [virtual_device_type](const Device& physical_device) { - return physical_device.device_type == virtual_device_type; - }); - CHECK(itr != physical_devices.end()) - << "Unable to find a physical device (from among the " << physical_devices.size() - << " given) to match the virtual device with device type " << virtual_device_type; - const size_t i = std::distance(physical_devices.begin(), itr); - devices_.push_back(*itr); - allocators_.push_back(MemoryManager::GetOrCreateAllocator(*itr, alloc_types[i])); - } -} - -inline void VirtualMachine::WriteRegister(Index r, const ObjectRef& val) { - frames_.back().register_file[r] = val; -} - -ObjectRef VirtualMachine::ReadRegister(Index r) const { return frames_.back().register_file[r]; } - -int64_t VirtualMachine::LoadScalarInt(Index r) const { - int64_t result = 0; - const auto& obj = ReadRegister(r); - NDArray array = Downcast(CopyTo(obj, GetDevice(exec_->host_device_index))); - - switch (array->dtype.bits) { - case 1: { - result = reinterpret_cast(array->data)[0]; - break; - } - case 8: { - result = reinterpret_cast(array->data)[0]; - break; - } - case 16: { - result = reinterpret_cast(array->data)[0]; - break; - } - case 32: { - result = reinterpret_cast(array->data)[0]; - break; - } - case 64: { - result = reinterpret_cast(array->data)[0]; - break; - } - default: - LOG(FATAL) << "Unknown scalar int type: " << DLDataType2String(array->dtype); - } - return result; -} - -Index VirtualMachine::GetResultRegisterIndex() const { - Index op_index = 0; - while (code_[op_index].op != Opcode::Ret) { - ++op_index; - } - - return code_[op_index].result; -} - -void VirtualMachine::CalculatePreResultOpIndex(Index res_index) { - if (preresult_op_index_ == -1) { - preresult_op_index_ = 0; - while (code_[preresult_op_index_].dst != res_index) { - ++preresult_op_index_; - } - } -} - -std::vector VirtualMachine::GetOutputTensorRegIndices() { - std::vector reg_indices; - Index res_index = GetResultRegisterIndex(); - CalculatePreResultOpIndex(res_index); - auto& preres_instr = code_[preresult_op_index_]; - auto op_code = preres_instr.op; - if (op_code == Opcode::AllocTensor) { - reg_indices.emplace_back(res_index); - } else if (op_code == Opcode::AllocADT) { - for (Index i = 0; i < preres_instr.num_fields; ++i) { - reg_indices.push_back(preres_instr.datatype_fields[i]); - } - } else if (op_code == Opcode::ReshapeTensor) { - reg_indices.push_back(preres_instr.reshape_tensor.tensor); - } else { - LOG(FATAL) << "Operation " << size_t(op_code) << " is not supported for set_outputs method"; - } - return reg_indices; -} - -void VirtualMachine::RunLoop(const std::vector& output_tensor_reg_indices) { - ICHECK(this->exec_); - ICHECK(this->code_); - pc_ = 0; - Index frame_start = frames_.size(); - while (true) { - main_loop: - auto const& instr = code_[this->pc_]; - VLOG(2) << "Executing(" << pc_ << "): " << instr; - - switch (instr.op) { - case Opcode::Move: { - ObjectRef from_obj; - from_obj = ReadRegister(instr.from); - WriteRegister(instr.dst, from_obj); - pc_++; - goto main_loop; - } - case Opcode::Fatal: { - throw std::runtime_error("VM encountered fatal error"); - } - case Opcode::LoadConst: { - bool is_not_cached = const_pool_.size() <= static_cast(instr.const_index) || - !const_pool_[instr.const_index].defined(); - if (is_not_cached) { - OpStartHook(instr); - } - auto constant_obj = exec_->constants[instr.const_index]; - // We cache the allocated object in the constant pool. To measure, the - // first iteration will set the pool up. The other iterations will - // directly reuse the allocated objects. - if (const_pool_.size() <= static_cast(instr.const_index)) { - const_pool_.resize(instr.const_index + 1); - } - - if (!const_pool_[instr.const_index].defined()) { - auto& [dev, mem_scope] = - exec_->virtual_devices[exec_->const_device_indexes[instr.const_index]]; - const_pool_[instr.const_index] = CopyTo(constant_obj, dev, String(mem_scope)); - } - WriteRegister(instr.dst, const_pool_[instr.const_index]); - if (is_not_cached) { - OpStopHook(); - } - pc_++; - goto main_loop; - } - case Opcode::LoadConsti: { - auto tensor = NDArray::Empty({1}, {kDLInt, 64, 1}, GetDevice(exec_->host_device_index)); - reinterpret_cast(tensor->data)[0] = instr.load_consti.val; - WriteRegister(instr.dst, tensor); - pc_++; - goto main_loop; - } - case Opcode::Invoke: { - std::vector args; - for (Index i = 0; i < instr.num_args; ++i) { - args.push_back(ReadRegister(instr.invoke_args_registers[i])); - } - InvokeGlobal(exec_->functions[instr.func_index], args); - frames_.back().caller_return_register = instr.dst; - goto main_loop; - } - case Opcode::InvokePacked: { - ICHECK_LE(instr.packed_index, packed_funcs_.size()); - const auto& func = packed_funcs_[instr.packed_index]; - const auto& arity = instr.arity; - std::vector args; - for (Index i = 0; i < arity; ++i) { - auto arg = ReadRegister(instr.packed_args[i]); - args.push_back(arg); -#if TVM_LOG_DEBUG - if (i < arity) { - const bool is_input = i < arity - instr.output_size; - VLOG(2) << (is_input ? "input" : "placeholder") << " arg " << i << " = " - << RuntimeObject2String(arg, GetDevice(exec_->host_device_index), - /*show_contents=*/is_input); - } -#endif - } - - // We no longer need to write the registers back, we write directly - // through the registers mutably. - InvokePacked(instr.packed_index, func, arity, instr.output_size, args); - -#if TVM_LOG_DEBUG - for (Index i = arity - instr.output_size; i < arity; ++i) { - auto arg = ReadRegister(instr.packed_args[i]); - VLOG(2) << "output arg " << i << " = " - << RuntimeObject2String(arg, GetDevice(exec_->host_device_index)); - } -#endif - - pc_++; - goto main_loop; - } - case Opcode::InvokeClosure: { - auto object = ReadRegister(instr.closure); - const auto* closure = object.as(); - ICHECK(closure); - std::vector args; - for (auto free_var : closure->free_vars) { - args.push_back(free_var); - } - for (Index i = 0; i < instr.num_closure_args; ++i) { - args.push_back(ReadRegister(instr.closure_args[i])); - } - InvokeGlobal(exec_->functions[closure->func_index], args); - frames_.back().caller_return_register = instr.dst; - goto main_loop; - } - case Opcode::GetField: { - auto object = ReadRegister(instr.object); - const auto& tuple = Downcast(object); - auto field = tuple[instr.field_index]; - WriteRegister(instr.dst, field); - pc_++; - goto main_loop; - } - case Opcode::GetTag: { - auto object = ReadRegister(instr.get_tag.object); - const auto& adt = Downcast(object); - auto tag = adt.tag(); - auto tag_tensor = NDArray::Empty({1}, {kDLInt, 32, 1}, GetDevice(exec_->host_device_index)); - reinterpret_cast(tag_tensor->data)[0] = tag; - WriteRegister(instr.dst, tag_tensor); - pc_++; - goto main_loop; - } - case Opcode::Goto: { - pc_ += instr.pc_offset; - goto main_loop; - } - case Opcode::If: { - int32_t test_val = LoadScalarInt(instr.if_op.test); - int32_t target_val = LoadScalarInt(instr.if_op.target); - - if (test_val == target_val) { - ICHECK_NE(instr.if_op.true_offset, 0); - pc_ += instr.if_op.true_offset; - } else { - ICHECK_NE(instr.if_op.false_offset, 0); - pc_ += instr.if_op.false_offset; - } - - goto main_loop; - } - case Opcode::AllocTensor: { - OpStartHook(instr); - if (!output_tensor_reg_indices.empty() && FindIndex(output_tensor_reg_indices, instr.dst)) { - WriteAllocatedTensorFromOutside(instr); - } else { - WriteAllocatedTensor(instr); - } - OpStopHook(); - pc_++; - goto main_loop; - } - case Opcode::AllocTensorReg: { - OpStartHook(instr); - Device cpu_dev = GetDevice(exec_->host_device_index); - auto shape_obj = ReadRegister(instr.alloc_tensor_reg.shape_register); - NDArray shape_tensor = Downcast(CopyTo(shape_obj, cpu_dev)); - auto shape = ToShape(shape_tensor); - auto storage_obj = ReadRegister(instr.alloc_tensor_reg.storage); - auto storage = Downcast(storage_obj); - auto offset = LoadScalarInt(instr.alloc_tensor.offset); - auto obj = storage->AllocNDArray(offset, shape, instr.alloc_tensor_reg.dtype); - VLOG(2) << "allocated " - << RuntimeObject2String(obj, GetDevice(exec_->host_device_index), - /*show_contents=*/false); - - WriteRegister(instr.dst, obj); - OpStopHook(); - pc_++; - goto main_loop; - } - case Opcode::AllocADT: { - std::vector fields; - for (Index i = 0; i < instr.num_fields; ++i) { - fields.push_back(ReadRegister(instr.datatype_fields[i])); - } - ObjectRef obj = ADT(instr.constructor_tag, fields); - WriteRegister(instr.dst, obj); - pc_++; - goto main_loop; - } - case Opcode::AllocClosure: { - std::vector free_vars; - for (Index i = 0; i < instr.num_freevar; i++) { - free_vars.push_back(ReadRegister(instr.free_vars[i])); - } - WriteRegister(instr.dst, VMClosure(instr.func_index, free_vars)); - pc_++; - goto main_loop; - } - case Opcode::AllocStorage: { - OpStartHook(instr); - - auto storage_obj = SimpleObjAllocator().make_object(); - Allocator* allocator = GetAllocator(instr.alloc_storage.device_index); - Device device = devices_[instr.alloc_storage.device_index]; - ICHECK(allocator) << "Did you forget to init the VirtualMachine with devices?"; - - if (instr.alloc_storage.ndim > 0) { - std::string shape = "["; - for (uint32_t i = 0; i < instr.alloc_storage.ndim; ++i) { - if (i > 0) { - shape += ", "; - } - shape += std::to_string(instr.alloc_storage.shape[i]); - } - shape += "]"; - std::string mem_scope = exec_->virtual_devices[instr.alloc_storage.device_index].second; - VLOG(2) << "allocating with ndims=" << instr.alloc_storage.ndim << ", shape=" << shape - << ", dtype_hint=" << DLDataType2String(instr.alloc_storage.dtype_hint) - << ", device_index=" << instr.alloc_storage.device_index - << ", memory_scope=" << mem_scope; - - std::vector shape_; - shape_.resize(instr.alloc_storage.ndim); - shape_.assign(instr.alloc_storage.shape, - instr.alloc_storage.shape + instr.alloc_storage.ndim); - storage_obj->buffer = allocator->Alloc(device, ShapeTuple(shape_), - instr.alloc_storage.dtype_hint, mem_scope); - storage_obj->allocator = allocator; - } else { - auto size = LoadScalarInt(instr.alloc_storage.allocation_size); - auto alignment = instr.alloc_storage.alignment; - VLOG(2) << "allocating with allocation_size=" << size << ", alignment=" << alignment - << ", dtype_hint=" << DLDataType2String(instr.alloc_storage.dtype_hint) - << ", device_index=" << instr.alloc_storage.device_index; - storage_obj->buffer = - allocator->Alloc(device, size, alignment, instr.alloc_storage.dtype_hint); - storage_obj->allocator = allocator; - } - Storage storage(storage_obj); - WriteRegister(instr.dst, storage); - OpStopHook(); - pc_++; - goto main_loop; - } - case Opcode::ShapeOf: { - auto input = ReadRegister(instr.shape_of.tensor); - NDArray input_array = Downcast(input); - int ndim = input_array->ndim; - auto out_tensor = - NDArray::Empty({ndim}, {kDLInt, 64, 1}, GetDevice(exec_->host_device_index)); - for (int i = 0; i < ndim; ++i) { - reinterpret_cast(out_tensor->data)[i] = input_array->shape[i]; - } - VLOG(2) << "shape = " - << RuntimeObject2String(out_tensor, GetDevice(exec_->host_device_index)); - WriteRegister(instr.dst, out_tensor); - pc_++; - goto main_loop; - } - case Opcode::Ret: { - // If we have hit the point from which we started - // running, we should return to the caller breaking - // the dispatch loop. - return_register_ = ReadRegister(instr.result); - auto caller_return_register = frames_.back().caller_return_register; - - if (PopFrame() == frame_start) { - return; - // Otherwise we are just returning from a local call. - } else { - WriteRegister(caller_return_register, return_register_); - goto main_loop; - } - } - case Opcode::ReshapeTensor: { - OpStartHook(instr); - Device cpu_dev = GetDevice(exec_->host_device_index); - auto tensor_obj = ReadRegister(instr.reshape_tensor.tensor); - NDArray tensor_arr = Downcast(tensor_obj); - // Read the shape from shape tensor - auto shape_obj = ReadRegister(instr.reshape_tensor.newshape); - NDArray shape_tensor = Downcast(CopyTo(shape_obj, cpu_dev)); - const DLTensor* dl_tensor = shape_tensor.operator->(); - ICHECK_EQ(dl_tensor->dtype.code, 0u); - ICHECK_EQ(dl_tensor->dtype.bits, 64u); - int64_t* dims = reinterpret_cast(dl_tensor->data); - int64_t ndim = shape_tensor->shape[0]; - std::vector shape(dims, dims + ndim); - // Reshape the input tensor - auto out_tensor = tensor_arr.CreateView(shape, tensor_arr->dtype); - VLOG(2) << "reshaped " - << RuntimeObject2String(tensor_obj, GetDevice(exec_->host_device_index)) << " to " - << RuntimeObject2String(out_tensor, GetDevice(exec_->host_device_index)); - WriteRegister(instr.dst, out_tensor); - OpStopHook(); - pc_++; - goto main_loop; - } - case Opcode::DeviceCopy: { - OpStartHook(instr); - auto tensor_src = ReadRegister(instr.device_copy.src); - NDArray src_data = Downcast(tensor_src); - Device actual_src_dev = src_data->device; - Device inst_src_dev = GetDevice(instr.device_copy.src_device_index); - ICHECK_EQ(actual_src_dev.device_type, inst_src_dev.device_type); - ICHECK_EQ(actual_src_dev.device_id, inst_src_dev.device_id); - Device dst_dev = GetDevice(instr.device_copy.dst_device_index); - auto mem_scope = exec_->virtual_devices[instr.device_copy.dst_device_index].second; - - NDArray dst_data = src_data.CopyTo(dst_dev, String(mem_scope)); - WriteRegister(instr.dst, dst_data); - OpStopHook(); - pc_++; - goto main_loop; - } - case Opcode::KillRegister: { - OpStartHook(instr); - WriteRegister(instr.dst, ObjectRef()); - OpStopHook(); - pc_++; - goto main_loop; - } - default: - LOG(FATAL) << "Unknown instruction opcode: " << int(instr.op); - } - } -} - -void VirtualMachine::WriteAllocatedTensor(const Instruction& instr) { - auto shape = std::vector(instr.alloc_tensor.ndim); - - for (uint32_t i = 0; i < instr.alloc_tensor.ndim; ++i) { - shape[i] = instr.alloc_tensor.shape[i]; - } - - auto storage_obj = ReadRegister(instr.alloc_tensor.storage); - auto offset = LoadScalarInt(instr.alloc_tensor.offset); - auto storage = Downcast(storage_obj); - auto obj = storage->AllocNDArray(offset, shape, instr.alloc_tensor.dtype); - VLOG(2) << "allocated " - << RuntimeObject2String(obj, GetDevice(exec_->host_device_index), - /*show_contents=*/false); - - WriteRegister(instr.dst, obj); -} - -void VirtualMachine::WriteAllocatedTensorFromOutside(const Instruction& instr) { - // External tensor(s) has been already written to the register (instr.dst) - auto ex_arr = Downcast(ReadRegister(instr.dst)); - auto ex_shape = ex_arr.Shape(); - auto ex_size = ex_shape.size(); - auto ex_dtype = ex_arr->dtype; - - auto in_size = instr.alloc_tensor.ndim; - auto in_dtype = instr.alloc_tensor.dtype; - ICHECK_EQ(TypeEqual(in_dtype, ex_dtype), true) - << "Data types mismatching for internal and external output tensors"; - - bool size_check = false; - if (ex_size != in_size) { - size_check = true; - } else { - for (size_t i = 0; i < in_size; ++i) { - if (ex_shape[i] != instr.alloc_tensor.shape[i]) { - size_check = true; - break; - } - } - } - - if (size_check) { - // Match element number - size_t in_el_num = 1, ex_el_num = 1; - for (size_t i = 0; i < ex_size; ++i) { - ex_el_num *= ex_shape[i]; - } - for (size_t i = 0; i < in_size; ++i) { - in_el_num *= instr.alloc_tensor.shape[i]; - } - ICHECK_EQ(in_el_num, ex_el_num) - << "Element number mismatching of internal and external output tensors"; - if (code_[preresult_op_index_].op == Opcode::ReshapeTensor) { - int64_t* dims = instr.alloc_tensor.shape; - std::vector ref_shape(dims, dims + int64_t(in_size)); - auto reshaped_tensor = ex_arr.CreateView(ref_shape, ex_dtype); - WriteRegister(instr.dst, reshaped_tensor); - } else { - LOG(FATAL) << "Internal and external output tensor shapes are mismatched"; - } - } -} - -bool VirtualMachine::FindIndex(const std::vector& indices, Index val) const { - auto it = std::find(indices.begin(), indices.end(), val); - return it != indices.end(); -} - -runtime::Module CreateVirtualMachine(Executable* exec) { - auto vm = make_object(); - vm->LoadExecutable(GetObjectPtr(exec)); - return runtime::Module(vm); -} - -TVM_REGISTER_GLOBAL("runtime._VirtualMachine").set_body([](TVMArgs args, TVMRetValue* rv) { - runtime::Module mod = args[0]; - auto* exec = dynamic_cast(mod.operator->()); - *rv = CreateVirtualMachine(exec); -}); - -} // namespace vm -} // namespace runtime -} // namespace tvm diff --git a/src/script/ir_builder/ir/frame.cc b/src/script/ir_builder/ir/frame.cc index 60a35ee010ec..2b02a80e3eaf 100644 --- a/src/script/ir_builder/ir/frame.cc +++ b/src/script/ir_builder/ir/frame.cc @@ -39,7 +39,7 @@ void IRModuleFrameNode::ExitWithScope() { IRBuilder builder = IRBuilder::Current(); ICHECK(!builder->result.defined()) << "ValueError: Builder.result has already been set"; auto dict_attrs = DictAttrs(attrs); - builder->result = tvm::IRModule(func_map, {}, {}, {}, dict_attrs, global_infos); + builder->result = tvm::IRModule(func_map, {}, dict_attrs, global_infos); } TVM_REGISTER_NODE_TYPE(IRModuleFrameNode); diff --git a/src/script/ir_builder/ir/ir.cc b/src/script/ir_builder/ir/ir.cc index 0fb4b256351b..8abf1a650e40 100644 --- a/src/script/ir_builder/ir/ir.cc +++ b/src/script/ir_builder/ir/ir.cc @@ -56,7 +56,7 @@ GlobalVar DeclFunction(const String& func_name, const BaseFunc& func_signature) auto gvar_type = [&]() -> Type { if (auto prim_func = func_signature.as()) { Array arg_types = prim_func->params.Map([](const auto& var) { return GetType(var); }); - return FuncType(arg_types, prim_func->ret_type, {}, {}); + return FuncType(arg_types, prim_func->ret_type); } return {}; diff --git a/src/script/printer/ir/ir.cc b/src/script/printer/ir/ir.cc index 5295cf2e41a1..ac3bc5584c66 100644 --- a/src/script/printer/ir/ir.cc +++ b/src/script/printer/ir/ir.cc @@ -142,26 +142,6 @@ TVM_STATIC_IR_FUNCTOR(IRDocsifier, vtable) return IR(d, "Op")->Call({LiteralDoc::Str(op->name, p->Attr("name"))}); }); -TVM_STATIC_IR_FUNCTOR(IRDocsifier, vtable) - .set_dispatch("", [](TypeVar var, ObjectPath p, IRDocsifier d) -> Doc { - return IR(d, "TypeVar") - ->Call({LiteralDoc::Str(var->name_hint, p->Attr("name_hint")), // - LiteralDoc::Str(TypeKind2String(var->kind), p->Attr("kind"))}); - }); - -TVM_STATIC_IR_FUNCTOR(IRDocsifier, vtable) - .set_dispatch( // - "", [](GlobalTypeVar var, ObjectPath p, IRDocsifier d) -> Doc { - return IR(d, "GlobalTypeVar") - ->Call({LiteralDoc::Str(var->name_hint, p->Attr("name_hint")), - LiteralDoc::Str(TypeKind2String(var->kind), p->Attr("kind"))}); - }); - -TVM_STATIC_IR_FUNCTOR(IRDocsifier, vtable) - .set_dispatch("", [](RelayRefType ref, ObjectPath p, IRDocsifier d) -> Doc { - return IR(d, "RelayRef")->Call({d->AsDoc(ref->value, p->Attr("value"))}); - }); - TVM_STATIC_IR_FUNCTOR(IRDocsifier, vtable) .set_dispatch("", [](TensorType type, ObjectPath p, IRDocsifier d) -> Doc { return IR(d, "TensorType") @@ -173,17 +153,11 @@ TVM_STATIC_IR_FUNCTOR(IRDocsifier, vtable) .set_dispatch("", [](FuncType func_type, ObjectPath p, IRDocsifier d) -> Doc { return IR(d, "FuncType") ->Call({ - d->AsDoc(func_type->type_params, p->Attr("type_params")), d->AsDoc(func_type->arg_types, p->Attr("arg_types")), d->AsDoc(func_type->ret_type, p->Attr("ret_type")), }); }); -TVM_STATIC_IR_FUNCTOR(IRDocsifier, vtable) - .set_dispatch("", [](IncompleteType ty, ObjectPath p, IRDocsifier d) -> Doc { - return IR(d, "IncompleteType")->Call({}); - }); - TVM_STATIC_IR_FUNCTOR(IRDocsifier, vtable) .set_dispatch("ir", [](Range range, ObjectPath p, IRDocsifier d) -> Doc { return IR(d, "Range") @@ -202,13 +176,9 @@ std::string ReprPrintIRModule(const ObjectRef& mod, const PrinterConfig& cfg) { return ReprPrintIR(mod, cfg); } -TVM_SCRIPT_REPR(TypeVarNode, ReprPrintIR); -TVM_SCRIPT_REPR(GlobalTypeVarNode, ReprPrintIR); TVM_SCRIPT_REPR(GlobalVarNode, ReprPrintIR); TVM_SCRIPT_REPR(DictAttrsNode, ReprPrintIR); -TVM_SCRIPT_REPR(RelayRefTypeNode, ReprPrintIR); TVM_SCRIPT_REPR(FuncTypeNode, ReprPrintIR); -TVM_SCRIPT_REPR(IncompleteTypeNode, ReprPrintIR); TVM_SCRIPT_REPR(RangeNode, ReprPrintIR); TVM_SCRIPT_REPR(IRModuleNode, ReprPrintIRModule); diff --git a/src/script/printer/legacy_repr.cc b/src/script/printer/legacy_repr.cc index 01fb514c497e..084a86a6f5c7 100644 --- a/src/script/printer/legacy_repr.cc +++ b/src/script/printer/legacy_repr.cc @@ -16,6 +16,7 @@ * specific language governing permissions and limitations * under the License. */ +#include #include #include #include @@ -174,12 +175,6 @@ TVM_STATIC_IR_FUNCTOR(ReprLegacyPrinter, vtable) (*p) << "TupleTypeNode(" << node->fields << ")"; }); -TVM_STATIC_IR_FUNCTOR(ReprLegacyPrinter, vtable) - .set_dispatch([](const ObjectRef& ref, ReprLegacyPrinter* p) { - auto* node = static_cast(ref.get()); - (*p) << "IncompleteTypeNode(" << node->kind << ", " << node << ")"; - }); - TVM_STATIC_IR_FUNCTOR(ReprLegacyPrinter, vtable) .set_dispatch([](const ObjectRef& node, ReprLegacyPrinter* p) { auto* op = static_cast(node.get()); @@ -198,29 +193,10 @@ TVM_STATIC_IR_FUNCTOR(ReprLegacyPrinter, vtable) (*p) << "IRModule(" << node->functions << ")"; }); -TVM_STATIC_IR_FUNCTOR(ReprLegacyPrinter, vtable) - .set_dispatch([](const ObjectRef& ref, ReprLegacyPrinter* p) { - auto* node = static_cast(ref.get()); - (*p) << "TypeVar(" << node->name_hint << ", " << node->kind << ")"; - }); - -TVM_STATIC_IR_FUNCTOR(ReprLegacyPrinter, vtable) - .set_dispatch([](const ObjectRef& ref, ReprLegacyPrinter* p) { - auto* node = static_cast(ref.get()); - (*p) << "GlobalTypeVar(" << node->name_hint << ", " << node->kind << ")"; - }); - TVM_STATIC_IR_FUNCTOR(ReprLegacyPrinter, vtable) .set_dispatch([](const ObjectRef& ref, ReprLegacyPrinter* p) { auto* node = static_cast(ref.get()); - (*p) << "FuncType(" << node->type_params << ", " << node->arg_types << ", " << node->ret_type - << ", " << node->type_constraints << ")"; - }); - -TVM_STATIC_IR_FUNCTOR(ReprLegacyPrinter, vtable) - .set_dispatch([](const ObjectRef& ref, ReprLegacyPrinter* p) { - auto* node = static_cast(ref.get()); - (*p) << "RelayRefTypeNode(" << node->value << ")"; + (*p) << "FuncType(" << node->arg_types << ", " << node->ret_type << ")"; }); } // namespace tvm diff --git a/src/script/printer/tir/usmp.cc b/src/script/printer/tir/usmp.cc deleted file mode 100644 index 5695c2a8076b..000000000000 --- a/src/script/printer/tir/usmp.cc +++ /dev/null @@ -1,58 +0,0 @@ -/* - * 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 "./utils.h" - -namespace tvm { -namespace script { -namespace printer { - -TVM_STATIC_IR_FUNCTOR(IRDocsifier, vtable) - .set_dispatch( - "", [](tir::usmp::AllocatedPoolInfo node, ObjectPath p, IRDocsifier d) -> Doc { - return IR(d, "AllocatedPoolInfo") - ->Call({}, {"pool_info", "allocated_size"}, - {d->AsDoc(node->pool_info, p->Attr("pool_info")), - d->AsDoc(node->allocated_size, p->Attr("allocated_size"))}); - }); - -TVM_STATIC_IR_FUNCTOR(IRDocsifier, vtable) - .set_dispatch("", - [](ConstantPoolInfo node, ObjectPath p, IRDocsifier d) -> Doc { - return IR(d, "ConstantPoolInfo") - ->Call( - {d->AsDoc(node->constant_info_array, - p->Attr("constant_info_array"))}); - }); - -TVM_STATIC_IR_FUNCTOR(IRDocsifier, vtable) - .set_dispatch("", [](ConstantInfo node, ObjectPath p, IRDocsifier d) -> Doc { - return IR(d, "ConstantInfo") - ->Call({d->AsDoc(node->name_hint, p->Attr("name_hint"))}, - {"byte_offset", "data"}, - {d->AsDoc(node->byte_offset, p->Attr("byte_offset")), - d->AddMetadata(node->data)}); - }); - -} // namespace printer -} // namespace script -} // namespace tvm diff --git a/src/support/ffi_testing.cc b/src/support/ffi_testing.cc index 52ffedda8030..ffc12f2b4c49 100644 --- a/src/support/ffi_testing.cc +++ b/src/support/ffi_testing.cc @@ -150,7 +150,7 @@ constexpr const char* FrontendTestModuleNode::kAddFunctionName; PackedFunc FrontendTestModuleNode::GetFunction(const String& name, const ObjectPtr& sptr_to_self) { if (name == kAddFunctionName) { - return TypedPackedFunc( + return runtime::TypedPackedFunc( [this, sptr_to_self](std::string func_name, PackedFunc pf) { CHECK_NE(func_name, kAddFunctionName) << "func_name: cannot be special function " << kAddFunctionName; diff --git a/src/support/libinfo.cc b/src/support/libinfo.cc index 1b986adcd551..d55181aab746 100644 --- a/src/support/libinfo.cc +++ b/src/support/libinfo.cc @@ -119,14 +119,6 @@ #define TVM_INFO_USE_STACKVM_RUNTIME "NOT-FOUND" #endif -#ifndef TVM_INFO_USE_GRAPH_EXECUTOR -#define TVM_INFO_USE_GRAPH_EXECUTOR "NOT-FOUND" -#endif - -#ifndef TVM_INFO_USE_PROFILER -#define TVM_INFO_USE_PROFILER "NOT-FOUND" -#endif - #ifndef TVM_INFO_USE_OPENMP #define TVM_INFO_USE_OPENMP "NOT-FOUND" #endif @@ -159,10 +151,6 @@ #define TVM_INFO_HIDE_PRIVATE_SYMBOLS "NOT-FOUND" #endif -#ifndef TVM_INFO_USE_TF_TVMDSOOP -#define TVM_INFO_USE_TF_TVMDSOOP "NOT-FOUND" -#endif - #ifndef TVM_INFO_USE_FALLBACK_STL_MAP #define TVM_INFO_USE_FALLBACK_STL_MAP "NOT-FOUND" #endif @@ -310,7 +298,6 @@ TVM_DLL Map GetLibInfo() { {"SUMMARIZE", TVM_INFO_SUMMARIZE}, {"TVM_CXX_COMPILER_PATH", TVM_CXX_COMPILER_PATH}, {"USE_ALTERNATIVE_LINKER", TVM_INFO_USE_ALTERNATIVE_LINKER}, - {"USE_AOT_EXECUTOR", TVM_INFO_USE_AOT_EXECUTOR}, {"USE_ARM_COMPUTE_LIB_GRAPH_EXECUTOR", TVM_INFO_USE_ARM_COMPUTE_LIB_GRAPH_EXECUTOR}, {"USE_ARM_COMPUTE_LIB", TVM_INFO_USE_ARM_COMPUTE_LIB}, {"USE_BLAS", TVM_INFO_USE_BLAS}, @@ -331,8 +318,6 @@ TVM_DLL Map GetLibInfo() { {"USE_AMX", TVM_INFO_USE_AMX}, {"USE_DNNL", TVM_INFO_USE_DNNL}, {"USE_FALLBACK_STL_MAP", TVM_INFO_USE_FALLBACK_STL_MAP}, - {"USE_GRAPH_EXECUTOR_CUDA_GRAPH", TVM_INFO_USE_GRAPH_EXECUTOR_CUDA_GRAPH}, - {"USE_GRAPH_EXECUTOR", TVM_INFO_USE_GRAPH_EXECUTOR}, {"USE_GTEST", TVM_INFO_USE_GTEST}, {"USE_HEXAGON", TVM_INFO_USE_HEXAGON}, {"USE_HEXAGON_RPC", TVM_INFO_USE_HEXAGON_RPC}, @@ -357,8 +342,6 @@ TVM_DLL Map GetLibInfo() { {"USE_OPENCL_GTEST", TVM_INFO_USE_OPENCL_GTEST}, {"USE_OPENMP", TVM_INFO_USE_OPENMP}, {"USE_PAPI", TVM_INFO_USE_PAPI}, - {"USE_PROFILER", TVM_INFO_USE_PROFILER}, - {"USE_PT_TVMDSOOP", TVM_INFO_USE_PT_TVMDSOOP}, {"USE_RANDOM", TVM_INFO_USE_RANDOM}, {"USE_RELAY_DEBUG", TVM_INFO_USE_RELAY_DEBUG}, {"TVM_DEBUG_WITH_ABI_CHANGE", TVM_INFO_TVM_DEBUG_WITH_ABI_CHANGE}, @@ -377,7 +360,6 @@ TVM_DLL Map GetLibInfo() { {"USE_TENSORFLOW_PATH", TVM_INFO_USE_TENSORFLOW_PATH}, {"USE_TENSORRT_CODEGEN", TVM_INFO_USE_TENSORRT_CODEGEN}, {"USE_TENSORRT_RUNTIME", TVM_INFO_USE_TENSORRT_RUNTIME}, - {"USE_TF_TVMDSOOP", TVM_INFO_USE_TF_TVMDSOOP}, {"USE_TFLITE", TVM_INFO_USE_TFLITE}, {"USE_THREADS", TVM_INFO_USE_THREADS}, {"USE_THRUST", TVM_INFO_USE_THRUST}, diff --git a/src/support/scalars.cc b/src/support/scalars.cc index 4ba505922b21..3c5239f84af9 100644 --- a/src/support/scalars.cc +++ b/src/support/scalars.cc @@ -24,7 +24,6 @@ #include "./scalars.h" -#include "tvm/relay/expr.h" #include "tvm/runtime/builtin_fp16.h" namespace tvm { @@ -39,15 +38,6 @@ static const DataType kFloat32 = DataType::Float(32); static const DataType kFloat64 = DataType::Float(64); static const DataType kBool = DataType::Bool(); -bool IsSimpleScalarDtype(DataType dtype) { - return dtype == kInt16 || dtype == kInt32 || dtype == kInt64 || dtype == kFloat16 || - dtype == kFloat32 || dtype == kFloat64 || dtype == kBool; -} - -bool IsSimpleScalar(const relay::ConstantNode* constant_node) { - return constant_node->is_scalar() && IsSimpleScalarDtype(DataType(constant_node->data->dtype)); -} - runtime::NDArray IntImmToNDArray(const IntImm& int_imm) { DLDevice dev = {DLDeviceType::kDLCPU, 0}; auto data = runtime::NDArray::Empty({}, int_imm->dtype, dev); diff --git a/src/support/scalars.h b/src/support/scalars.h index 2b34914565ed..05763f8044bf 100644 --- a/src/support/scalars.h +++ b/src/support/scalars.h @@ -26,21 +26,13 @@ #define TVM_SUPPORT_SCALARS_H_ #include -#include #include "tvm/ir/expr.h" -#include "tvm/relay/expr.h" #include "tvm/runtime/ndarray.h" namespace tvm { namespace support { -/*! \brief Returns true if a tensor of empty shape and given dtype is considered a Relay scalar. */ -bool IsSimpleScalarDtype(DataType dtype); - -/*! \brief Returns true if \p constant_node is a float/int/bool scalar. */ -bool IsSimpleScalar(const relay::ConstantNode* constant_node); - /*! \brief Returns NDArray 'scalar' for given TIR immediate. */ runtime::NDArray IntImmToNDArray(const IntImm& int_imm); runtime::NDArray FloatImmToNDArray(const FloatImm& float_imm); diff --git a/src/target/codegen.cc b/src/target/codegen.cc index afdbf841ad7d..49c250b8dc64 100644 --- a/src/target/codegen.cc +++ b/src/target/codegen.cc @@ -61,10 +61,6 @@ runtime::Module Build(IRModule mod, Target target) { .value()) { mod = tir::transform::SkipAssert()(mod); } - auto target_attr_map = tvm::TargetKind::GetAttrMap("TIRToRuntime"); - if (target_attr_map.count(target->kind)) { - return target_attr_map[target->kind](mod, target); - } // the build function. std::string build_f_name = "target.build." + target->kind->name; diff --git a/src/target/compilation_config.cc b/src/target/compilation_config.cc deleted file mode 100644 index a7f708f12a15..000000000000 --- a/src/target/compilation_config.cc +++ /dev/null @@ -1,314 +0,0 @@ -/* - * 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 tvm/target/compilation_config.cc - * \brief Implementation of \p CompilationConfig for collecting \p Targets. - */ - -#include -#include - -namespace tvm { - -TVM_REGISTER_NODE_TYPE(CompilationConfigNode); - -void CompilationConfigNode::VisitAttrs(AttrVisitor* v) { - v->Visit("host_target", &host_target); - v->Visit("primitive_targets", &primitive_targets); - v->Visit("default_primitive_virtual_device", &default_primitive_virtual_device); - v->Visit("host_virtual_device", &host_virtual_device); - v->Visit("optional_homogenous_target", &optional_homogeneous_target); - // NOTE: The virtual_device_cache_ is not accessible via FFI. -} - -Target CompilationConfigNode::FindPrimitiveTargetForDeviceOrFail(DLDeviceType device_type) const { - ICHECK_GT(device_type, 0) << "Invalid device type"; - auto itr = std::find_if( - primitive_targets.begin(), primitive_targets.end(), - [device_type](const Target& target) { return target->GetTargetDeviceType() == device_type; }); - if (itr == primitive_targets.end()) { - std::stringstream msg; - msg << "No target is specified for device type " << device_type - << ". The available device types and targets are:" << std::endl; - for (const auto& target : primitive_targets) { - msg << " " << target->GetTargetDeviceType() << "-> " << target->ToDebugString() << std::endl; - } - LOG(FATAL) << msg.str(); - } - return *itr; -} - -Optional CompilationConfigNode::FindPrimitiveTargetForKind( - const std::string& kind_name) const { - Optional opt_kind = TargetKind::Get(kind_name); - if (!opt_kind.defined()) { - VLOG(1) << "No such target kind for '" << kind_name << "'"; - return {}; - } - auto itr = - std::find_if(primitive_targets.begin(), primitive_targets.end(), - [kind_name](const Target& target) { return target->kind->name == kind_name; }); - if (itr == primitive_targets.end()) { - VLOG(1) << "No target available matching kind '" << kind_name << "'"; - return {}; - } - return *itr; -} - -Target CompilationConfigNode::CanonicalTarget(const Target& target) const { - // Fast path -- object identity. - if (target == host_target) { - return target; - } - for (const auto& primitive_target : primitive_targets) { - if (target == primitive_target) { - return target; - } - } - // Slow path -- structural equality. We have so few targets it does not seem worth building an - // index. - if (StructuralEqual()(target, host_target)) { - return host_target; - } - for (const auto& primitive_target : primitive_targets) { - if (StructuralEqual()(target, primitive_target)) { - return primitive_target; - } - } - // No match. - return target; -} - -VirtualDevice CompilationConfigNode::CanonicalVirtualDevice( - const VirtualDevice& virtual_device) const { - // Targets need special handling. - Target target = virtual_device->target; - if (target.defined()) { - // It is possible the given target object was constructed by the user, but was then - // rewritten on the way into the CompilationConfig. So 'canonicalize' it by replacing - // the given target with one structurally equal to one already known in the config if - // possible. - Target canon_target = CanonicalTarget(target); - if (canon_target != target) { - VLOG(1) << "Canonicalized target " << canon_target->ToDebugString(); - } - target = canon_target; - } else if (virtual_device->device_type() != kInvalidDeviceType) { - // Since no target was given, choose one with a matching device type. - // This is the one place where we allow device types to imply targets. - target = FindPrimitiveTargetForDeviceOrFail(virtual_device->device_type()); - VLOG(1) << "Defaulted to target " << target->ToDebugString(); - } - // else: the target will remain unknown. - - // Redirect to an existing structurally equal virtual device. - return virtual_device_cache_.Unique(VirtualDevice(virtual_device->device_type(), - virtual_device->virtual_device_id, target, - virtual_device->memory_scope)); -} - -void CompilationConfigNode::Init(const transform::PassContext& pass_ctx, - const Array& raw_targets) { - VLOG_CONTEXT << "CompilationConfig"; - CHECK_GT(raw_targets.size(), 0U) << "Require at least one target"; - - // - // Decide on the host target. - // - - // Any targets which could act as a host? - auto hosting_itr = std::find_if(raw_targets.begin(), raw_targets.end(), [](const Target& target) { - // TODO(tvm-team): The kDLHexagon device can act as a host. We can remove kDLHexagon - // here once we refactored kDLHexagon to kDLCPU. - return target->GetTargetDeviceType() == kDLCPU || target->GetTargetDeviceType() == kDLHexagon; - }); - - // Any targets with their host field set? - auto has_host_itr = std::find_if(raw_targets.begin(), raw_targets.end(), - [](const Target& target) { return target->host.defined(); }); - - if (has_host_itr != raw_targets.end()) { - // RULE A: If any raw target has a host, use the first such host for all the primitive - // targets. - host_target = Target((*has_host_itr)->GetHost().value(), /*host=*/Target()); - VLOG(1) << "The target " << (*has_host_itr)->ToDebugString() << " supplies a host target " - << host_target->ToDebugString() << " of device type " - << host_target->GetTargetDeviceType(); - } else if (hosting_itr != raw_targets.end()) { - // RULE B: If any raw target is for a device which could be a host then use the first such as - // the host. - host_target = Target(*hosting_itr, /*host=*/Target()); - VLOG(1) << "Using target " << host_target->ToDebugString() << " of CPU-like device type " - << host_target->GetTargetDeviceType() << " as the host target"; - } else { - // RULE C: Otherwise, create a default CPU host target. - host_target = MakeDefaultCPUTarget(); - VLOG(1) << "Created a default target " << host_target->ToDebugString() << " of device type " - << host_target->GetTargetDeviceType() << " for the host target"; - } - ICHECK(host_target.defined()); - ICHECK(!host_target->host.defined()); - - if (host_target->GetTargetDeviceType() != kDLCPU) { - // I think we're on thin ice here until we've audited the code base for assumed CPU hosts. - VLOG(1) << "The host target is not a CPU. This is probably not going to work."; - } - - // - // Establish the host VirtualDevice. - // - host_virtual_device = virtual_device_cache_.Unique( - VirtualDevice(static_cast(host_target->GetTargetDeviceType()), - /*virtual_device_id=*/0, host_target)); - ICHECK(host_virtual_device.defined()); - ICHECK(host_virtual_device->target.defined()); - - // - // Now that we've settled on a host, we can set it as the host on all the raw targets. - // - primitive_targets.clear(); - primitive_targets.reserve(raw_targets.size()); - for (const auto& raw_target : raw_targets) { - if (raw_target->host.defined() && !StructuralEqual()(raw_target->host, host_target)) { - VLOG(1) << "The target " << raw_target->ToDebugString() - << " already has a host which disagrees with the desired host target. It " - << "will be overridden."; - } - primitive_targets.push_back(Target(raw_target, host_target)); - } - ICHECK_GT(primitive_targets.size(), 0U); - - // - // Check the primitive_targets are ordered correctly re Target::IsExternalCodegenFor, - // and make sure no two targets share a kind name. - // - - // TODO(mbs): We could just sort the list, but given all the implicit defaulting for backwards - // compat it seems we should avoid making this any more magical than necessary. But revisit - // if usability suffers. - std::unordered_set primitive_target_device_types; - std::unordered_set kind_names; - for (const auto& target : primitive_targets) { - primitive_target_device_types.emplace(static_cast(target->GetTargetDeviceType())); - CHECK(kind_names.emplace(target->kind->name).second) << "Multiple targets have been given" - "for the same device kind '" - << target->kind->name << "'"; - } - for (DLDeviceType device_type : primitive_target_device_types) { - Target first_primitive_target; - for (const auto& current_primitive_target : primitive_targets) { - if (current_primitive_target->GetTargetDeviceType() != device_type) { - continue; - } - if (!first_primitive_target.defined()) { - first_primitive_target = current_primitive_target; - // Note it is valid to have only one external codegen target. - } else { - CHECK(current_primitive_target.IsExternalCodegenFor(first_primitive_target)) - << "When given multiple targets for the device type " << device_type - << " the first must be for non external codegen, and all subsequent must be for " - "external codegen. However have been given first " - << first_primitive_target->ToDebugString() << " and subsequent " - << current_primitive_target->ToDebugString(); - } - } - } - - // - // Decide on the default device type for primitives. - // - DLDeviceType default_primitive_device_type; - Optional opt_fallback_dev = pass_ctx->GetConfig("relay.fallback_device_type"); - if (opt_fallback_dev) { - // RULE D: Respect the PassContext setting if given. - const int64_t v = opt_fallback_dev.value()->value; - CHECK_GT(v, 0) - << "The 'relay.fallback_device_type' pass attribute is set to an invalid device type " << v; - default_primitive_device_type = static_cast(v); - VLOG(1) << "Using the 'relay.fallback_device_type' pass attribute " - << default_primitive_device_type - << " as the default device type for all primitive operations"; - } else if (primitive_target_device_types.size() == 1) { - // RULE E: Since only one device in use there's no choice to make. - default_primitive_device_type = *primitive_target_device_types.begin(); - VLOG(1) << "All primitive targets have the device type " << default_primitive_device_type - << " so that is also the default device type for all primitive operations."; - } else { - // RULE F: Fallback to CPU. - default_primitive_device_type = kDLCPU; - VLOG(1) << "Using " << default_primitive_device_type - << " as the default device type for all primitive operations"; - } - - // - // Establish the default primitive VirtualDevice, choosing a known Target to match the device - // type. We do not create a default target, it must already exist as a primitive target. - // - default_primitive_virtual_device = CanonicalVirtualDevice( - VirtualDevice::ForDeviceType(default_primitive_device_type, /*virtual_device_id=*/0)); - - ICHECK(default_primitive_virtual_device.defined()); - ICHECK(default_primitive_virtual_device->target.defined()); - - // Legacy: Some passes only support homogenous compilation and expect the target to be - // given by the global target context. Make this easy to detect. - optional_homogeneous_target = - primitive_targets.size() == 1 ? *primitive_targets.begin() : Target(); -} - -/* static */ Target CompilationConfigNode::MakeDefaultCPUTarget() { - if (runtime::Registry::Get("codegen.LLVMModuleCreate")) { - // LLVM is available. - // TODO(mbs): More robust extension mechanism? - return Target("llvm"); - } else { - // LLVM is not available. - // TODO(mbs): Already deprecated? - return Target("stackvm"); - } -} - -TVM_STATIC_IR_FUNCTOR(ReprPrinter, vtable) - .set_dispatch([](const ObjectRef& ref, ReprPrinter* p) { - auto* node = ref.as(); - p->stream << "Primitive targets:"; - for (const auto& target : node->primitive_targets) { - p->stream << std::endl - << " " << target->GetTargetDeviceType() << " |-> " << target->ToDebugString(); - } - p->stream << std::endl - << "Default primitive virtual device: " << node->default_primitive_virtual_device; - p->stream << std::endl << "Host virtual device: " << node->host_virtual_device; - }); - -CompilationConfig::CompilationConfig(const transform::PassContext& pass_ctx, - const Array& raw_targets) { - auto node = make_object(); - node->Init(pass_ctx, raw_targets); - data_ = std::move(node); -} - -TVM_REGISTER_GLOBAL("target.MakeCompilationConfig") - .set_body_typed([](const transform::PassContext& pass_ctx, - const Array& raw_targets) -> CompilationConfig { - return CompilationConfig(pass_ctx, raw_targets); - }); - -} // namespace tvm diff --git a/src/target/func_registry_generator.cc b/src/target/func_registry_generator.cc deleted file mode 100644 index d679bf379b62..000000000000 --- a/src/target/func_registry_generator.cc +++ /dev/null @@ -1,49 +0,0 @@ -/* - * 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. - */ - -/*! - * Defines functions that generate FuncRegistry structs for C runtime. - * \file func_registry_generator.cc - */ - -#include "func_registry_generator.h" - -#include - -namespace tvm { -namespace target { - -std::string GenerateFuncRegistryNames(const Array& function_names) { - std::stringstream ss; - - unsigned char function_nums[sizeof(uint16_t)]; - *reinterpret_cast(function_nums) = function_names.size(); - for (auto f : function_nums) { - ss << f; - } - - for (auto f : function_names) { - ss << f << '\0'; - } - - return ss.str(); -} - -} // namespace target -} // namespace tvm diff --git a/src/target/func_registry_generator.h b/src/target/func_registry_generator.h deleted file mode 100644 index 8d2af305a0e4..000000000000 --- a/src/target/func_registry_generator.h +++ /dev/null @@ -1,44 +0,0 @@ -/* - * 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. - */ - -/*! - * Defines functions that generate FuncRegistry structs for C runtime. - * \file func_registry_generator.h - */ -#ifndef TVM_TARGET_FUNC_REGISTRY_GENERATOR_H_ -#define TVM_TARGET_FUNC_REGISTRY_GENERATOR_H_ - -#include -#include - -#include -#include - -using tvm::runtime::Array; -using tvm::runtime::String; - -namespace tvm { -namespace target { - -std::string GenerateFuncRegistryNames(const Array& function_names); - -} // namespace target -} // namespace tvm - -#endif // TVM_TARGET_FUNC_REGISTRY_GENERATOR_H_ diff --git a/src/target/llvm/codegen_cpu.cc b/src/target/llvm/codegen_cpu.cc index 8fdedeb78687..5e30695517a1 100644 --- a/src/target/llvm/codegen_cpu.cc +++ b/src/target/llvm/codegen_cpu.cc @@ -59,8 +59,6 @@ #include #include -#include "../func_registry_generator.h" -#include "../metadata_utils.h" #include "llvm_instance.h" namespace tvm { @@ -974,377 +972,6 @@ llvm::Value* CodeGenCPU::RuntimeTVMParallelBarrier() { return GetContextPtr(gv_tvm_parallel_barrier_); } -/*! \brief Defines LLVM Types for each Metadata member type. */ -struct MetadataLlvmTypes { - llvm::Type* t_float64; - llvm::Type* t_uint8; - llvm::Type* t_int64; - llvm::Type* t_bool; - llvm::Type* t_cstring; - llvm::Type* t_void_p; - llvm::StructType* t_data_type; - - /*! \brief Maps a MetadataBase subclass' type_key to its corresponding LLVM StructType. */ - ::std::unordered_map structs_by_type_key; -}; - -class MetadataTypeDefiner : public AttrVisitor { - public: - MetadataTypeDefiner(llvm::LLVMContext* ctx, struct MetadataLlvmTypes* llvm_types) - : ctx_{ctx}, llvm_types_{llvm_types} {} - - void Visit(const char* key, double* value) final { - elements_.emplace_back(llvm_types_->t_float64); - } - void Visit(const char* key, int64_t* value) final { - elements_.emplace_back(llvm_types_->t_int64); - } - void Visit(const char* key, uint64_t* value) final { - elements_.emplace_back(llvm_types_->t_int64); - } - void Visit(const char* key, int* value) final { elements_.emplace_back(llvm_types_->t_int64); } - void Visit(const char* key, bool* value) final { elements_.emplace_back(llvm_types_->t_bool); } - void Visit(const char* key, std::string* value) final { - elements_.emplace_back(llvm_types_->t_cstring); - } - void Visit(const char* key, void** value) final { elements_.emplace_back(llvm_types_->t_void_p); } - void Visit(const char* key, DataType* value) final { - elements_.emplace_back(llvm_types_->t_data_type); - } - void Visit(const char* key, runtime::NDArray* value) final { - elements_.emplace_back(llvm_types_->t_int64); - elements_.emplace_back(llvm_types_->t_void_p); - } - - private: - void VisitMetadataBase(runtime::metadata::MetadataBase metadata) { - elements_.emplace_back(llvm::PointerType::getUnqual( - llvm::StructType::create(*ctx_, metadata->get_c_struct_name()))); - if (visited_.find(metadata->get_c_struct_name()) != visited_.end()) { - return; - } - - if (to_visit_.find(metadata->get_c_struct_name()) != to_visit_.end()) { - return; - } - to_visit_[metadata->get_c_struct_name()] = metadata; - } - - public: - using MetadataKind = runtime::metadata::MetadataKind; - - void VisitArray(const runtime::metadata::MetadataArrayNode* arr) { - switch (arr->kind) { - case MetadataKind::kUint64: // LLVM encodes signed and unsigned with same types. - case MetadataKind::kInt64: - elements_.emplace_back(llvm::PointerType::getUnqual(llvm_types_->t_int64)); - break; - case MetadataKind::kBool: - elements_.emplace_back(llvm::PointerType::getUnqual(llvm_types_->t_bool)); - break; - case MetadataKind::kString: - elements_.emplace_back(llvm::PointerType::getUnqual(llvm_types_->t_cstring)); - break; - case MetadataKind::kHandle: - CHECK(false) << "Do not support handle"; - break; - case MetadataKind::kMetadata: - if (llvm_types_->structs_by_type_key.count(arr->type_key)) { - elements_.emplace_back( - llvm::PointerType::getUnqual(llvm_types_->structs_by_type_key[arr->type_key])); - } - break; - default: - CHECK(false) << "Unsupported metadata kind " << arr->kind; - break; - } - } - - void Visit(const char* key, ObjectRef* value) final { - const runtime::metadata::MetadataArrayNode* arr = - value->as(); - if (arr != nullptr) { - VisitArray(arr); - } else { - elements_.emplace_back( - llvm::PointerType::getUnqual(llvm_types_->structs_by_type_key[(*value)->GetTypeKey()])); - } - } - - void DefineType(runtime::metadata::MetadataBase metadata) { - ICHECK(elements_.empty()); - ReflectionVTable::Global()->VisitAttrs(metadata.operator->(), this); - llvm_types_->structs_by_type_key[metadata->GetTypeKey()] = - llvm::StructType::create(*ctx_, elements_, metadata->get_c_struct_name()); - elements_.clear(); - } - - llvm::LLVMContext* ctx_; - struct MetadataLlvmTypes* llvm_types_; - ::std::unordered_set<::std::string> visited_; - ::std::unordered_map<::std::string, runtime::metadata::MetadataBase> to_visit_; - ::std::vector elements_; -}; - -class MetadataSerializerLLVM : public AttrVisitor { - using MetadataKind = runtime::metadata::MetadataKind; - - public: - MetadataSerializerLLVM(CodeGenLLVM* codegen, struct MetadataLlvmTypes* llvm_types) - : codegen_{codegen}, llvm_types_{llvm_types} {} - - void Visit(const char* key, double* value) final { - elements_.back().emplace_back(llvm::ConstantFP::get(llvm_types_->t_float64, *value)); - } - void Visit(const char* key, int64_t* value) final { - elements_.back().emplace_back(llvm::ConstantInt::get( - llvm_types_->t_int64, static_cast(*value), true /* isSigned */)); - } - void Visit(const char* key, uint64_t* value) final { - elements_.back().emplace_back( - llvm::ConstantInt::get(llvm_types_->t_int64, *value, false /* isSigned */)); - } - void Visit(const char* key, int* value) final { - elements_.back().emplace_back( - llvm::ConstantInt::get(llvm_types_->t_int64, *value, true /* isSigned */)); - } - void Visit(const char* key, bool* value) final { - elements_.back().emplace_back(llvm::ConstantInt::get( - llvm_types_->t_uint8, static_cast(*value), false /* isSigned */)); - } - void Visit(const char* key, std::string* value) final { - elements_.back().emplace_back(codegen_->GetConstString(*value)); - } - void Visit(const char* key, void** value) final { - CHECK(false) << "Do not support serializing void*"; - } - void Visit(const char* key, DataType* value) final { - elements_.back().emplace_back(llvm::ConstantStruct::get( - llvm_types_->t_data_type, - {llvm::ConstantInt::get(llvm_types_->t_uint8, value->code(), false /* isSigned */), - llvm::ConstantInt::get(llvm_types_->t_uint8, value->bits(), false /* isSigned */), - llvm::ConstantInt::get(llvm_types_->t_uint8, value->lanes(), false /* isSigned */)})); - } - - // Serializing NDArray as tuple of len, data - void Visit(const char* key, runtime::NDArray* value) final { - std::string bytes; - dmlc::MemoryStringStream stream(&bytes); - value->Save(&stream); - elements_.back().emplace_back( - llvm::ConstantInt::get(llvm_types_->t_int64, bytes.length(), true /* isSigned */)); - elements_.back().emplace_back(codegen_->GetConstString(bytes)); - } - - void VisitMetadata(runtime::metadata::MetadataBase metadata) { - elements_.emplace_back(std::vector()); - ReflectionVTable::Global()->VisitAttrs(metadata.operator->(), this); - auto struct_elements = elements_.back(); - elements_.pop_back(); - auto struct_ty = llvm_types_->structs_by_type_key[metadata->GetTypeKey()]; - ICHECK(struct_ty != nullptr) << "Did not find LLVM StructType* for type_key=" - << metadata->GetTypeKey(); - CHECK_EQ(struct_elements.size(), struct_ty->getNumElements()); - auto out = llvm::ConstantStruct::get(struct_ty, struct_elements); - if (elements_.size() > 0) { - elements_.back().push_back(out); - } else { - last_production_ = out; - } - } - - void VisitArray(const runtime::metadata::MetadataArrayNode* arr) { - llvm::Type* element_type; - switch (arr->kind) { - case MetadataKind::kInt64: - element_type = llvm_types_->t_int64; - break; - case MetadataKind::kUint64: - element_type = llvm_types_->t_int64; - break; - case MetadataKind::kBool: - element_type = llvm_types_->t_uint8; - break; - case MetadataKind::kString: - element_type = llvm_types_->t_cstring; - break; - case MetadataKind::kMetadata: { - element_type = llvm_types_->structs_by_type_key[arr->type_key]; - ICHECK(element_type != nullptr) - << "Did not find LLVM StructType* for type_key=" << arr->type_key; - break; - } - default: - LOG(FATAL) << "unknown metadata kind " << arr->kind; - break; - } - - elements_.emplace_back(std::vector()); - for (auto o : arr->array) { - if (o->IsInstance()) { - double value = Downcast(o)->value; - Visit(nullptr, &value); - } - if (o->IsInstance()) { - auto value = Downcast(o)->value; - Visit(nullptr, &value); - } else if (o->IsInstance()) { - ::std::string value = Downcast(o); - Visit(nullptr, &value); - } else { - // nested array not possible. - VisitMetadata(Downcast(o)); - } - } - auto array = elements_.back(); - elements_.pop_back(); - CHECK(element_type != nullptr); - auto arr_ty = llvm::ArrayType::get(element_type, array.size()); - auto llvm_arr = llvm::ConstantArray::get(arr_ty, array); - - if (elements_.size() > 0) { - elements_.back().emplace_back( - codegen_->GetGlobalConstant(llvm_arr, "", llvm::GlobalValue::PrivateLinkage)); - } else { - last_production_ = llvm_arr; - } - } - - void Visit(const char* key, ObjectRef* value) final { - const runtime::metadata::MetadataArrayNode* arr = - value->as(); - if (arr != nullptr) { - VisitArray(arr); - return; - } - - runtime::metadata::MetadataBase metadata = Downcast(*value); - VisitMetadata(metadata); - } - - llvm::Constant* Serialize(runtime::metadata::MetadataBase metadata) { - Visit(nullptr, &metadata); - ICHECK(last_production_); - return codegen_->GetGlobalConstant(last_production_); - } - - CodeGenLLVM* codegen_; - MetadataLlvmTypes* llvm_types_; - llvm::LLVMContext* ctx_; - llvm::Module* module_; - std::vector> elements_; - llvm::Constant* last_production_; -}; - -void CodeGenCPU::DefineMetadata(runtime::metadata::Metadata metadata) { - llvm::LLVMContext* ctx = llvm_target_->GetContext(); - MetadataLlvmTypes llvm_types{ - t_float64_ /* t_float64 */, - llvm::Type::getInt8Ty(*ctx) /* t_uint8 */, - t_int64_ /* t_int64 */, - llvm::Type::getInt8Ty(*ctx) /* t_bool */, - llvmGetPointerTo(t_char_, 0) /* t_cstring */, - t_void_p_ /* t_void_p */, - llvm::StructType::create(*ctx, {t_int8_, t_int8_, t_int8_}, "DLDataType") /* t_data_type */, - }; - - // create sample ConstantInfoMetadata instance for MetadataTypeDefiner - std::string bytes; - runtime::NDArray ci = runtime::NDArray::Empty({0}, DataType::UInt(8), Device{kDLCPU}); - dmlc::MemoryStringStream stream(&bytes); - ci.Save(&stream); - TVMConstantInfo di = - TVMConstantInfo{"default-none", 0, static_cast(bytes.size()), bytes.c_str()}; - - std::vector queue; - queue.push_back(runtime::metadata::ConstantInfoMetadata(&di)); - - metadata::DiscoverComplexTypesVisitor discover_complex{&queue}; - discover_complex.Discover(metadata); - - MetadataTypeDefiner definer{ctx, &llvm_types}; - for (auto md : queue) { - if (md.defined()) { - definer.DefineType(md); - } - } - - MetadataSerializerLLVM serializer{this, &llvm_types}; - auto metadata_constant_gv = serializer.Serialize(metadata); - - function_ = - llvm::Function::Create(ftype_tvm_backend_packed_c_func_, llvm::Function::ExternalLinkage, - runtime::symbol::tvm_get_c_metadata, module_.get()); - SetTargetAttributes(function_); - function_->setCallingConv(llvm::CallingConv::C); - function_->setDLLStorageClass(llvm::GlobalValue::DLLStorageClassTypes::DLLExportStorageClass); - - llvm::BasicBlock* entry_point_entry = llvm::BasicBlock::Create(*ctx, "entry", function_); - builder_->SetInsertPoint(entry_point_entry); - - auto ret_values_p = builder_->CreateBitCast(GetArg(function_, 3), llvmGetPointerTo(t_void_p_, 0)); - builder_->CreateStore(builder_->CreateBitCast(metadata_constant_gv, t_void_p_), ret_values_p); - - auto ret_tcode = builder_->CreateBitCast(GetArg(function_, 4), llvmGetPointerTo(t_int_, 0)); - builder_->CreateStore(llvm::ConstantInt::get(t_int_, kTVMOpaqueHandle), ret_tcode); - - builder_->CreateRet(ConstInt32(0)); -} - -void CodeGenCPU::DefineFunctionRegistry(Array func_names) { - ICHECK(system_lib_prefix_.defined()) - << "Loading of --system-lib modules is yet to be defined for C runtime"; - Array symbols; - std::vector funcs; - for (auto sym : func_names) { - symbols.push_back(sym); - auto* sym_func = - llvm::Function::Create(ftype_tvm_backend_packed_c_func_, llvm::GlobalValue::ExternalLinkage, - sym.operator std::string(), module_.get()); - - funcs.emplace_back(sym_func); - } - llvm::ArrayType* t_tvm_crt_func_ptrs = - llvm::ArrayType::get(llvmGetPointerTo(ftype_tvm_backend_packed_c_func_, 0), funcs.size()); -#if TVM_LLVM_VERSION >= 200 - llvm::DataLayout layout(module_.get()->getDataLayout()); -#else - llvm::DataLayout layout(module_.get()); -#endif - - llvm::GlobalVariable* func_registry_ptrs = new llvm::GlobalVariable( - *module_, t_tvm_crt_func_ptrs, true, llvm::GlobalValue::InternalLinkage, - llvm::ConstantArray::get(t_tvm_crt_func_ptrs, funcs), "_tvm_func_registry_ptrs"); - - uint64_t align = layout.getTypeAllocSize(llvmGetPointerTo(ftype_tvm_backend_packed_c_func_, 0)); -#if TVM_LLVM_VERSION >= 100 - func_registry_ptrs->setAlignment(llvm::Align(align)); -#else - func_registry_ptrs->setAlignment(align); -#endif - llvm::GlobalVariable* func_registry = new llvm::GlobalVariable( - *module_, t_tvm_crt_func_registry_, true, llvm::GlobalVariable::InternalLinkage, - llvm::ConstantStruct::get( - t_tvm_crt_func_registry_, - {GetConstString(::tvm::target::GenerateFuncRegistryNames(symbols)), - llvm::ConstantExpr::getBitCast(func_registry_ptrs, - llvmGetPointerTo(ftype_tvm_backend_packed_c_func_, 0))}), - "_tvm_crt_func_registry"); - llvm::GlobalVariable* module = new llvm::GlobalVariable( - *module_, t_tvm_crt_module_, true, llvm::GlobalValue::InternalLinkage, - llvm::ConstantStruct::get(t_tvm_crt_module_, {func_registry}), "_tvm_crt_module"); - - // Now build TVMSystemLibEntryPoint. - llvm::FunctionType* ftype = llvm::FunctionType::get(t_void_p_, {}, false); - function_ = llvm::Function::Create(ftype, llvm::Function::ExternalLinkage, - "TVMSystemLibEntryPoint", module_.get()); - SetTargetAttributes(function_); - llvm::BasicBlock* entry_point_entry = - llvm::BasicBlock::Create(*llvm_target_->GetContext(), "entry", function_); - builder_->SetInsertPoint(entry_point_entry); - builder_->CreateRet(builder_->CreateBitCast(module, t_void_p_)); -} - void CodeGenCPU::AddStartupFunction() { if (!target_c_runtime_) { llvm::FunctionType* ftype = llvm::FunctionType::get(t_void_, {}, false); diff --git a/src/target/llvm/codegen_cpu.h b/src/target/llvm/codegen_cpu.h index 91fe1bc18631..182dae81ce15 100644 --- a/src/target/llvm/codegen_cpu.h +++ b/src/target/llvm/codegen_cpu.h @@ -77,18 +77,6 @@ class CodeGenCPU : public CodeGenLLVM { llvm::Value* CreateCallExtern(Type ret_type, String global_symbol, const Array& args, bool skip_first_arg) override; - /*! - * \brief A CPU-specific function to create the FuncRegistry. - * \param func_names List of functions to be included, in order. - */ - void DefineFunctionRegistry(Array func_names); - - /*! - * \brief Serialize the metadata object as data, and implement get_c_metadata function. - * \param metadata The metadata which should be serialized. - */ - void DefineMetadata(runtime::metadata::Metadata metadata); - protected: void AddStartupFunction() final; // meta data diff --git a/src/target/llvm/codegen_llvm.cc b/src/target/llvm/codegen_llvm.cc index e27bf33e69d6..1f9cc09b1e2f 100644 --- a/src/target/llvm/codegen_llvm.cc +++ b/src/target/llvm/codegen_llvm.cc @@ -104,7 +104,6 @@ #include "../../arith/pattern_match.h" #include "../build_common.h" -#include "../func_registry_generator.h" #include "codegen_params.h" #include "llvm_instance.h" diff --git a/src/target/llvm/llvm_module.cc b/src/target/llvm/llvm_module.cc index 0f026f39684c..3a30674d4c0f 100644 --- a/src/target/llvm/llvm_module.cc +++ b/src/target/llvm/llvm_module.cc @@ -23,8 +23,6 @@ */ #ifdef TVM_LLVM_VERSION -#include "llvm_module.h" - #include #include #include @@ -56,11 +54,9 @@ #include #include #include -#include #include #include #include -#include #include #include #include @@ -80,7 +76,6 @@ #include "../../runtime/file_utils.h" #include "../../runtime/library_module.h" -#include "../func_registry_generator.h" #include "codegen_blob.h" #include "codegen_cpu.h" #include "codegen_llvm.h" @@ -328,19 +323,11 @@ void LLVMModuleNode::Init(const IRModule& mod, const Target& target) { std::unique_ptr cg = CodeGenLLVM::Create(llvm_target.get()); std::string entry_func; - relay::Runtime runtime = - mod->GetAttr(tvm::attr::kRuntime).value_or(relay::Runtime::Create("cpp")); Optional system_lib_prefix = mod->GetAttr(tvm::attr::kSystemLibPrefix); - if (!system_lib_prefix && runtime->GetAttr("system-lib").value_or(Bool(false))) { - system_lib_prefix = ""; - } - - bool target_c_runtime = runtime->name == "crt"; for (auto kv : mod->functions) { if (!kv.second->IsInstance()) { - // (@jroesch): we relax constraints here, Relay functions will just be ignored. DLOG(INFO) << "Can only lower IR Module with PrimFuncs, but got " << kv.second->GetTypeKey(); continue; } @@ -361,8 +348,7 @@ void LLVMModuleNode::Init(const IRModule& mod, const Target& target) { // ICHECK(funcs.size() > 0); // TODO(tqchen): remove the entry function behavior as it does not // makes sense when we start to use multiple modules. - cg->Init("TVMMod", llvm_target.get(), system_lib_prefix, system_lib_prefix.defined(), - target_c_runtime); + cg->Init("TVMMod", llvm_target.get(), system_lib_prefix, system_lib_prefix.defined(), false); cg->SetFastMathFlags(llvm_target->GetFastMathFlags()); cg->AddFunctionsOrdered(mod->functions.begin(), mod->functions.end()); @@ -795,89 +781,6 @@ TVM_REGISTER_GLOBAL("codegen.codegen_blob") return runtime::Module(n); }); -runtime::Module CreateLLVMCppMetadataModule(runtime::metadata::Metadata metadata, Target target, - tvm::relay::Runtime runtime) { - auto llvm_instance = std::make_unique(); - With llvm_target(*llvm_instance, target); - - Optional system_lib_prefix = NullOpt; - if (runtime->GetAttr("system-lib").value_or(Bool(false))) { - system_lib_prefix = ""; - } - - auto cg = std::make_unique(); - - cg->Init("TVMMetadataMod", llvm_target.get(), system_lib_prefix, system_lib_prefix.defined(), - /*target_c_runtime=*/false); - - cg->DefineMetadata(metadata); - auto mod = cg->Finish(); - llvm_target->SetTargetMetadata(mod.get()); - mod->addModuleFlag(llvm::Module::Override, "Debug Info Version", llvm::DEBUG_METADATA_VERSION); - - mod->addModuleFlag( - llvm::Module::Override, "Dwarf Version", - llvm_target->GetOrCreateTargetMachine()->getTargetTriple().isOSDarwin() ? 2 : 4); - - auto n = make_object(); - n->Init(std::move(mod), std::move(llvm_instance)); - n->SetJITEngine(llvm_target->GetJITEngine()); - - auto meta_mod = MetadataModuleCreate(metadata); - meta_mod->Import(runtime::Module(n)); - return meta_mod; -} - -runtime::Module CreateLLVMCrtMetadataModule(const Array& modules, Target target, - tvm::relay::Runtime runtime) { - Array func_names; - for (runtime::Module mod : modules) { - auto pf_funcs = mod.GetFunction("get_func_names"); - if (pf_funcs != nullptr) { - Array func_names_ = pf_funcs(); - for (const auto& fname : func_names_) { - func_names.push_back(fname); - } - } - } - - auto llvm_instance = std::make_unique(); - With llvm_target(*llvm_instance, target); - - Optional system_lib_prefix = NullOpt; - if (runtime->GetAttr("system-lib").value_or(Bool(false))) { - system_lib_prefix = ""; - } - - bool target_c_runtime = runtime->name == "crt"; - ICHECK(system_lib_prefix.defined() && target_c_runtime) - << "For LLVM C-runtime metadata module, must include --system-lib and --runtime=c; " - << "got target: " << target->str(); - auto cg = std::make_unique(); - cg->Init("TVMMetadataMod", llvm_target.operator->(), system_lib_prefix, - system_lib_prefix.defined(), target_c_runtime); - - cg->DefineFunctionRegistry(func_names); - auto mod = cg->Finish(); - llvm_target->SetTargetMetadata(mod.get()); - mod->addModuleFlag(llvm::Module::Override, "Debug Info Version", llvm::DEBUG_METADATA_VERSION); - - mod->addModuleFlag( - llvm::Module::Override, "Dwarf Version", - llvm_target->GetOrCreateTargetMachine()->getTargetTriple().isOSDarwin() ? 2 : 4); - - auto n = make_object(); - n->Init(std::move(mod), std::move(llvm_instance)); - n->SetJITEngine(llvm_target->GetJITEngine()); - for (auto m : modules) { - n->Import(m); - } - return runtime::Module(n); -} - -TVM_REGISTER_GLOBAL("runtime.CreateLLVMCrtMetadataModule") - .set_body_typed(CreateLLVMCrtMetadataModule); - } // namespace codegen } // namespace tvm diff --git a/src/target/llvm/llvm_module.h b/src/target/llvm/llvm_module.h index 66492f8152e5..716a0a6c7007 100644 --- a/src/target/llvm/llvm_module.h +++ b/src/target/llvm/llvm_module.h @@ -27,7 +27,6 @@ #ifdef TVM_LLVM_VERSION -#include #include #include #include @@ -36,11 +35,7 @@ namespace tvm { namespace codegen { -runtime::Module CreateLLVMCppMetadataModule(runtime::metadata::Metadata metadata, Target target, - tvm::relay::Runtime runtime); - -runtime::Module CreateLLVMCrtMetadataModule(const Array& modules, Target target, - tvm::relay::Runtime runtime); +runtime::Module CreateLLVMCppMetadataModule(runtime::metadata::Metadata metadata, Target target); } // namespace codegen } // namespace tvm diff --git a/src/target/metadata.cc b/src/target/metadata.cc deleted file mode 100644 index 35df3ada0000..000000000000 --- a/src/target/metadata.cc +++ /dev/null @@ -1,53 +0,0 @@ -/* - * 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 metadata.cc - * \brief Implementations of the compiler extensions for Metadata. - */ - -#include "metadata.h" - -#include - -namespace tvm { -namespace target { -namespace metadata { - -TVM_REGISTER_REFLECTION_VTABLE(VisitableMetadataNode, - ::tvm::detail::ReflectionTrait) - .set_creator([](const std::string&) -> ObjectPtr { - return ::tvm::runtime::make_object(); - }); - -TVM_REGISTER_REFLECTION_VTABLE(VisitableTensorInfoNode, - ::tvm::detail::ReflectionTrait) - .set_creator([](const std::string&) -> ObjectPtr { - return ::tvm::runtime::make_object(); - }); - -TVM_REGISTER_REFLECTION_VTABLE(VisitableConstantInfoMetadataNode, - ::tvm::detail::ReflectionTrait) - .set_creator([](const std::string&) -> ObjectPtr { - return ::tvm::runtime::make_object(); - }); - -} // namespace metadata -} // namespace target -} // namespace tvm diff --git a/src/target/metadata.h b/src/target/metadata.h deleted file mode 100644 index b761f7ff2bbb..000000000000 --- a/src/target/metadata.h +++ /dev/null @@ -1,260 +0,0 @@ -/* - * 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 tvm/target/metadata.h - * \brief Extends Metadata for use in the compiler. - */ -#ifndef TVM_TARGET_METADATA_H_ -#define TVM_TARGET_METADATA_H_ -#include -#include - -#include -#include -#include - -namespace tvm { -namespace target { -namespace metadata { - -/*! - * \brief Subclass of MetadataNode that implements the VisitAttrs reflection method. - * - * This implementation (and other such Visitable subclasses) is compiled into libtvm.so, but not - * libtvm_runtime.so, because reflection is not supported in libtvm_runtime.so over code size - * concerns. It is used during compilation by the generic metadata code-generators. - */ -class VisitableMetadataNode : public ::tvm::runtime::metadata::MetadataNode { - public: - explicit VisitableMetadataNode(const struct ::TVMMetadata* data) : MetadataNode{data} {} - VisitableMetadataNode() : MetadataNode{nullptr} {} - - void VisitAttrs(AttrVisitor* v) { - int64_t version_cpp{version()}; - v->Visit("version", &version_cpp); - auto inputs_array = Array(); - auto inputs_accessor = inputs(); - inputs_array.reserve(num_inputs()); - for (int64_t i = 0; i < num_inputs(); ++i) { - inputs_array.push_back(::tvm::runtime::metadata::TensorInfo{inputs_accessor[i]}); - } - ::tvm::runtime::metadata::MetadataArray inputs_metadata_array{ - inputs_array, ::tvm::runtime::metadata::MetadataKind::kMetadata, - ::tvm::runtime::metadata::TensorInfoNode::_type_key}; - v->Visit("inputs", &inputs_metadata_array); - int64_t num_inputs_cpp = num_inputs(); - v->Visit("num_inputs", &num_inputs_cpp); - auto outputs_array = Array(); - auto outputs_accessor = outputs(); - outputs_array.reserve(num_outputs()); - for (int64_t i = 0; i < num_outputs(); ++i) { - outputs_array.push_back(::tvm::runtime::metadata::TensorInfo{outputs_accessor[i]}); - } - ::tvm::runtime::metadata::MetadataArray outputs_metadata_array{ - outputs_array, ::tvm::runtime::metadata::MetadataKind::kMetadata, - ::tvm::runtime::metadata::TensorInfoNode::_type_key}; - v->Visit("outputs", &outputs_metadata_array); - int64_t num_outputs_cpp = num_outputs(); - v->Visit("num_outputs", &num_outputs_cpp); - auto pools_array = Array(); - auto pools_accessor = workspace_pools(); - pools_array.reserve(num_workspace_pools()); - for (int64_t i = 0; i < num_workspace_pools(); ++i) { - pools_array.push_back(::tvm::runtime::metadata::TensorInfo{pools_accessor[i]}); - } - ::tvm::runtime::metadata::MetadataArray workspace_pools_metadata_array{ - pools_array, ::tvm::runtime::metadata::MetadataKind::kMetadata, - ::tvm::runtime::metadata::TensorInfoNode::_type_key}; - v->Visit("workspace_pools", &workspace_pools_metadata_array); - int64_t num_workspace_pools_cpp = num_workspace_pools(); - v->Visit("num_workspace_pools", &num_workspace_pools_cpp); - - auto consts_array = Array(); - auto consts_accessor = constant_pools(); - consts_array.reserve(num_constant_pools()); - for (int64_t i = 0; i < num_constant_pools(); ++i) { - consts_array.push_back(::tvm::runtime::metadata::ConstantInfoMetadata{consts_accessor[i]}); - } - - int64_t num_const_pools_cpp = num_constant_pools(); - ::tvm::runtime::metadata::MetadataArray constant_pools_metadata_array{ - consts_array, ::tvm::runtime::metadata::MetadataKind::kMetadata, - ::tvm::runtime::metadata::ConstantInfoMetadataNode::_type_key}; - v->Visit("constant_pools", &constant_pools_metadata_array); - v->Visit("num_constant_pools", &num_const_pools_cpp); - ::std::string mod_name_cpp{data()->mod_name}; - v->Visit("mod_name", &mod_name_cpp); - } -}; - -/*! - * \brief Subclass of MetadataNode which also owns the backing C structures. - * - * This class (and other InMemory subclasses) are used during compilation to instantiate Metadata - * instances whose storage lives outside of .rodata. This class exists because the Module returned - * from tvm.relay.build must also be ready to run inference. - */ -class InMemoryMetadataNode : public ::tvm::target::metadata::VisitableMetadataNode { - public: - InMemoryMetadataNode() - : InMemoryMetadataNode(0 /* version */, {} /* inputs */, {} /* outputs */, - {} /* workspace_pools */, {} /* constant_pools */, "" /* mod_name */) { - } - InMemoryMetadataNode(int64_t version, - const ::std::vector<::tvm::runtime::metadata::TensorInfo>& inputs, - const ::std::vector<::tvm::runtime::metadata::TensorInfo>& outputs, - const ::std::vector<::tvm::runtime::metadata::TensorInfo>& workspace_pools, - const ::std::vector<::tvm::ConstantInfo>& constant_pools, - const ::tvm::runtime::String mod_name) - : VisitableMetadataNode{&storage_}, - inputs_{new struct TVMTensorInfo[inputs.size()]}, - inputs_objs_{inputs}, - outputs_{new struct TVMTensorInfo[outputs.size()]}, - outputs_objs_{outputs}, - workspace_pools_{new struct TVMTensorInfo[workspace_pools.size()]}, - workspace_pools_objs_{workspace_pools}, - constant_pools_{new struct TVMConstantInfo[constant_pools.size()]}, - constant_pools_objs_{constant_pools}, - mod_name_{mod_name}, - storage_{version, nullptr, 0ull, nullptr, 0ull, - nullptr, 0ull, nullptr, 0ull, mod_name_.c_str()} { - storage_.inputs = inputs_.get(); - storage_.num_inputs = inputs.size(); - for (unsigned int i = 0; i < inputs.size(); ++i) { - inputs_.get()[i] = *inputs[i]->data(); - } - storage_.outputs = outputs_.get(); - storage_.num_outputs = outputs.size(); - for (unsigned int i = 0; i < outputs.size(); ++i) { - outputs_.get()[i] = *outputs[i]->data(); - } - storage_.workspace_pools = workspace_pools_.get(); - storage_.num_workspace_pools = workspace_pools.size(); - for (unsigned int i = 0; i < workspace_pools.size(); ++i) { - workspace_pools_.get()[i] = *workspace_pools[i]->data(); - } - storage_.constant_pools = constant_pools_.get(); - storage_.num_constant_pools = constant_pools.size(); - for (size_t i = 0; i < constant_pools.size(); ++i) { - constant_pools_.get()[i].name_hint = constant_pools[i]->name_hint.c_str(); - constant_pools_.get()[i].byte_offset = constant_pools[i]->byte_offset.IntValue(); - - std::string bytes; - dmlc::MemoryStringStream stream(&bytes); - auto data = constant_pools[i]->data; - data.Save(&stream); - // Allocated mem freed in destructor - constant_pools_.get()[i].data_len = bytes.size(); - char* a = reinterpret_cast(malloc(bytes.size())); - constant_pools_.get()[i].data_bytes = a; - memcpy(a, bytes.c_str(), bytes.size()); - } - } - - ~InMemoryMetadataNode() { - // frees allocated mem for const_objs_ - for (int i = 0; i < storage_.num_constant_pools; ++i) { - free(const_cast(constant_pools_.get()[i].data_bytes)); - } - } - - private: - ::std::unique_ptr inputs_; - std::vector<::tvm::runtime::metadata::TensorInfo> inputs_objs_; - ::std::unique_ptr outputs_; - std::vector<::tvm::runtime::metadata::TensorInfo> outputs_objs_; - ::std::unique_ptr workspace_pools_; - std::vector<::tvm::runtime::metadata::TensorInfo> workspace_pools_objs_; - ::std::unique_ptr constant_pools_; - std::vector<::tvm::ConstantInfo> constant_pools_objs_; - ::std::string mod_name_; - struct ::TVMMetadata storage_; -}; - -class VisitableTensorInfoNode : public ::tvm::runtime::metadata::TensorInfoNode { - public: - explicit VisitableTensorInfoNode(const struct ::TVMTensorInfo* data) : TensorInfoNode{data} {} - VisitableTensorInfoNode() : TensorInfoNode{nullptr} {} - - void VisitAttrs(AttrVisitor* v) { - ::std::string name_cpp{data()->name}; - v->Visit("name", &name_cpp); - auto shape_array = Array(); - auto shape_accessor = shape(); - shape_array.reserve(num_shape()); - for (int64_t i = 0; i < num_shape(); ++i) { - shape_array.push_back(::tvm::Integer{static_cast(shape_accessor[i])}); - } - ::tvm::runtime::metadata::MetadataArray shape_metadata_array{ - shape_array, ::tvm::runtime::metadata::MetadataKind::kInt64, nullptr}; - v->Visit("shape", &shape_metadata_array); - int64_t num_shape_cpp = num_shape(); - v->Visit("num_shape", &num_shape_cpp); - ::tvm::runtime::DataType dtype_cpp{dtype()}; - v->Visit("dtype", &dtype_cpp); - } -}; - -class InMemoryTensorInfoNode : public ::tvm::target::metadata::VisitableTensorInfoNode { - public: - InMemoryTensorInfoNode() : InMemoryTensorInfoNode("", {}, ::tvm::runtime::DataType(0, 0, 0)) {} - InMemoryTensorInfoNode(const ::tvm::runtime::String& name, const ::std::vector& shape, - ::tvm::runtime::DataType dtype) - : VisitableTensorInfoNode{&storage_}, - name_{name}, - shape_{new int64_t[shape.size()]()}, - storage_{name_.c_str(), nullptr, 0, dtype} { - storage_.shape = shape_.get(); - storage_.num_shape = shape.size(); - for (unsigned int i = 0; i < shape.size(); ++i) { - shape_.get()[i] = shape[i]; - } - } - - private: - ::std::string name_; - ::std::unique_ptr shape_; - struct ::TVMTensorInfo storage_; -}; - -class VisitableConstantInfoMetadataNode - : public ::tvm::runtime::metadata::ConstantInfoMetadataNode { - public: - explicit VisitableConstantInfoMetadataNode(const struct ::TVMConstantInfo* data) - : ConstantInfoMetadataNode{data} {} - VisitableConstantInfoMetadataNode() : ConstantInfoMetadataNode{nullptr} {} - - void VisitAttrs(AttrVisitor* v) { - ::std::string name_cpp{name_hint()}; - v->Visit("name_hint", &name_cpp); - - uint64_t byte_offset_cpp{byte_offset()}; - v->Visit("byte_offset", &byte_offset_cpp); - - ::tvm::runtime::NDArray data_cpp = data(); - v->Visit("data", &data_cpp); - } -}; - -} // namespace metadata -} // namespace target -} // namespace tvm - -#endif // TVM_TARGET_METADATA_H_ diff --git a/src/target/metadata_module.cc b/src/target/metadata_module.cc deleted file mode 100644 index c8c099171c96..000000000000 --- a/src/target/metadata_module.cc +++ /dev/null @@ -1,252 +0,0 @@ -/* - * 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 metadata_module.cc - * \brief Defines functions that build MetadataModules for C++ and C runtimes. - */ -#include "metadata_module.h" - -#include - -#include -#include - -#include "../runtime/const_loader_module.h" -#include "../runtime/meta_data.h" -#include "llvm/llvm_module.h" -#include "source/source_module.h" - -namespace tvm { -namespace codegen { - -static runtime::metadata::Metadata ConvertMetaData( - relay::backend::ExecutorCodegenMetadata metadata); - -static runtime::Module CreateCrtMetadataModule( - runtime::Module target_module, Target target, relay::Runtime runtime, relay::Executor executor, - relay::backend::ExecutorCodegenMetadata metadata, - Array non_crt_exportable_modules, - Array crt_exportable_modules, - const std::unordered_map& const_var_ndarray) { - if (!non_crt_exportable_modules.empty()) { - std::string non_exportable_modules; - for (unsigned int i = 0; i < non_crt_exportable_modules.size(); i++) { - if (i > 0) { - non_exportable_modules += ", "; - } - auto mod = non_crt_exportable_modules[i]; - auto pf_sym = mod.GetFunction("get_symbol"); - if (pf_sym != nullptr) { - non_exportable_modules += pf_sym().operator std::string(); - } else { - non_exportable_modules += - std::string{"(module type_key="} + mod->type_key() + std::string{")"}; - } - } - CHECK(false) << "These " << non_crt_exportable_modules.size() - << " modules are not exportable to C-runtime: " << non_exportable_modules; - } - - if (target->kind->name == "c") { - runtime::metadata::Metadata aot_metadata; - if (executor->GetAttr("interface-api", tvm::String("packed")) == "packed") { - aot_metadata = ConvertMetaData(metadata); - } - - crt_exportable_modules.push_back(target_module); - target_module = CreateCSourceCrtMetadataModule(crt_exportable_modules, target, runtime, - metadata, aot_metadata); - } else if (target->kind->name == "llvm") { -#ifdef TVM_LLVM_VERSION - crt_exportable_modules.push_back(target_module); - target_module = CreateLLVMCrtMetadataModule(crt_exportable_modules, target, runtime); -#else // TVM_LLVM_VERSION - LOG(FATAL) << "TVM was not built with LLVM enabled."; -#endif // TVM_LLVM_VERSION - } - - return target_module; -} - -// TODO(areusch,masahi): Unify metadata representation and remove the need for this function -static runtime::metadata::Metadata ConvertMetaData( - relay::backend::ExecutorCodegenMetadata metadata) { - ICHECK(metadata.defined()); - ICHECK_NOTNULL(metadata->pool_inputs); - - std::vector inputs; - for (size_t i = 0; i < metadata->inputs.size(); ++i) { - auto v = metadata->inputs[i]; - auto ttype = metadata->input_tensor_types[i]; - inputs.push_back( - runtime::metadata::TensorInfo(make_object( - v->name_hint, relay::backend::ShapeToJSON(ttype->shape), ttype->dtype))); - } - - std::vector outputs; - auto output_ttypes = metadata->output_tensor_types; - for (size_t i = 0; i < output_ttypes.size(); ++i) { - auto ttype = output_ttypes[i]; - std::stringstream name; - name << "output" << i; - outputs.push_back( - runtime::metadata::TensorInfo(make_object( - name.str(), relay::backend::ShapeToJSON(ttype->shape), ttype->dtype))); - } - - std::vector pools; - for (size_t i = 0; i < metadata->pools.size(); ++i) { - auto var = metadata->pools[i]; - auto api = metadata->pool_inputs.value()[var]; - if (api->pool_info.as()) { - pools.push_back( - runtime::metadata::TensorInfo(make_object( - var->name_hint, std::vector{api->allocated_size.IntValue()}, - tvm::runtime::DataType{kDLUInt, 8, 1}))); - } - } - - std::vector consts; - for (const auto& kv : metadata->pool_inputs.value()) { - const auto& api = kv.second; - if (const auto* pi = api->pool_info.as()) { - if (pi->is_internal) { - for (const auto ci : pi->constant_info_array) { - consts.emplace_back(ci->name_hint, ci->byte_offset, ci->data); - } - } - } - } - auto n = make_object( - runtime::metadata::kMetadataVersion, inputs, outputs, pools, consts, metadata->mod_name); - - return runtime::metadata::Metadata(std::move(n)); -} - -static runtime::Module CreateCppMetadataModule( - runtime::Module target_module, Target target, relay::Runtime runtime, - relay::backend::ExecutorCodegenMetadata metadata, - const std::unordered_map>& const_vars_by_symbol, - Array non_crt_exportable_modules, - Array crt_exportable_modules, - const std::unordered_map& const_var_ndarray) { - if (!non_crt_exportable_modules.empty()) { - runtime::Module const_loader_mod = - runtime::ConstLoaderModuleCreate(const_var_ndarray, const_vars_by_symbol); - const_loader_mod.Import(target_module); - for (const auto& it : non_crt_exportable_modules) { - const_loader_mod.Import(it); - } - target_module = const_loader_mod; - } - - if (metadata.defined()) { - runtime::metadata::Metadata runtime_metadata = ConvertMetaData(metadata); - - if (metadata->executor == runtime::kTvmExecutorAot && runtime->name == relay::kTvmRuntimeCpp) { - if (target->kind->name == "c") { - auto metadata_module = CreateCSourceCppMetadataModule(runtime_metadata); - metadata_module->Import(target_module); - target_module = metadata_module; -#ifdef TVM_LLVM_VERSION // defining TVM_LLVM_VERSION indicates TVM was compiled with USE_LLVM ON. - } else if (target->kind->name == "llvm") { - auto metadata_module = CreateLLVMCppMetadataModule(runtime_metadata, target, runtime); - metadata_module->Import(target_module); - target_module = metadata_module; -#endif // TVM_LLVM_VERSION - } else { - CHECK(false) << "Don't know how to create MetadataModule for target type " << target->str(); - } - } - } - - return target_module; -} - -/*! - * \brief Create a metadata module wrapper. The helper is used by different - * codegens, such as graph executor codegen and the vm compiler. - * - * \param params The metadata for initialization of all modules. - * \param target_module the internal module that is compiled by tvm. - * \param ext_modules The external modules that needs to be imported inside the metadata - * module(s). - * \param target The target that all the modules are compiled for - * \return The created metadata module that manages initialization of metadata. - */ -runtime::Module CreateMetadataModule( - const std::unordered_map& const_var_ndarray, - tvm::runtime::Module target_module, const Array& ext_modules, Target target, - tvm::relay::Runtime runtime, tvm::relay::Executor executor, - relay::backend::ExecutorCodegenMetadata metadata) { - // Here we split modules into two groups: - // 1. Those modules which can be exported to C-runtime. These are DSO-exportable - // (i.e. llvm or c) modules which return nothing from get_const_vars(). - // 2. Other modules. - Array crt_exportable_modules; - Array non_crt_exportable_modules; - - bool is_targeting_crt = runtime->name == "crt"; - - // Wrap all submodules in the initialization wrapper. - std::unordered_map> const_vars_by_symbol; - for (tvm::runtime::Module mod : ext_modules) { - auto pf_sym = mod.GetFunction("get_symbol"); - auto pf_var = mod.GetFunction("get_const_vars"); - std::vector symbol_const_vars; - if (pf_sym != nullptr && pf_var != nullptr) { - String symbol = pf_sym(); - Array variables = pf_var(); - for (size_t i = 0; i < variables.size(); i++) { - symbol_const_vars.push_back(variables[i].operator std::string()); - } - ICHECK_EQ(const_vars_by_symbol.count(symbol), 0U) << "Found duplicated symbol: " << symbol; - const_vars_by_symbol[symbol] = symbol_const_vars; - } - // We only need loading of serialized constant data - // if there are constants present and required by the - // runtime module to be initialized by the binary - // metadata module. If not rest of the modules are - // wrapped in c-source metadata module. - - // TODO(@manupa-arm) : we should be able to use csource_metadata - // if the variables are empty when all the runtime modules implement get_func_names - if (symbol_const_vars.empty() && is_targeting_crt && mod->IsDSOExportable() && - (target->kind->name == "c" || target->kind->name == "llvm")) { - crt_exportable_modules.push_back(mod); - } else { - non_crt_exportable_modules.push_back(mod); - } - } - - if (is_targeting_crt) { - return CreateCrtMetadataModule(target_module, target, runtime, executor, metadata, - non_crt_exportable_modules, crt_exportable_modules, - const_var_ndarray); - } else { - return CreateCppMetadataModule(target_module, target, runtime, metadata, const_vars_by_symbol, - non_crt_exportable_modules, crt_exportable_modules, - const_var_ndarray); - } -} - -} // namespace codegen - -} // namespace tvm diff --git a/src/target/metadata_module.h b/src/target/metadata_module.h deleted file mode 100644 index daeaf212c992..000000000000 --- a/src/target/metadata_module.h +++ /dev/null @@ -1,63 +0,0 @@ -/* - * 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 metadata_module.h - * \brief Declares functions that build MetadataModules for C++ and C runtimes. - */ - -#ifndef TVM_TARGET_METADATA_MODULE_H_ -#define TVM_TARGET_METADATA_MODULE_H_ - -#include -#include -#include -#include -#include - -#include -#include - -#include "../relay/backend/utils.h" - -namespace tvm { -namespace codegen { - -/*! - * \brief Create a metadata module wrapper. The helper is used by different - * codegens, such as graph executor codegen and the vm compiler. - * - * \param params The metadata for initialization of all modules. - * \param target_module the internal module that is compiled by tvm. - * \param ext_modules The external modules that needs to be imported inside the metadata - * module(s). - * \param target The target that all the modules are compiled for - * \param runtime The runtime to codegen for - * \param metadata Module metadata - * \return The created metadata module that manages initialization of metadata. - */ -runtime::Module CreateMetadataModule( - const std::unordered_map& params, runtime::Module target_module, - const Array& ext_modules, Target target, tvm::relay::Runtime runtime, - tvm::relay::Executor executor, relay::backend::ExecutorCodegenMetadata metadata); - -} // namespace codegen -} // namespace tvm - -#endif // TVM_TARGET_METADATA_MODULE_H_ diff --git a/src/target/metadata_utils.cc b/src/target/metadata_utils.cc deleted file mode 100644 index db17d1862846..000000000000 --- a/src/target/metadata_utils.cc +++ /dev/null @@ -1,155 +0,0 @@ -/* - * 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 tvm/target/metadata_utils.cc - * \brief Defines utility functions and classes for emitting metadata. - */ -#include "metadata_utils.h" - -namespace tvm { -namespace codegen { -namespace metadata { - -std::string AddressFromParts(const std::vector& parts) { - std::stringstream ss; - for (unsigned int i = 0; i < parts.size(); ++i) { - if (i > 0) { - ss << "_"; - } - ss << parts[i]; - } - return ss.str(); -} - -DiscoverArraysVisitor::DiscoverArraysVisitor(std::vector* queue) : queue_{queue} {} - -void DiscoverArraysVisitor::Visit(const char* key, double* value) {} -void DiscoverArraysVisitor::Visit(const char* key, int64_t* value) {} -void DiscoverArraysVisitor::Visit(const char* key, uint64_t* value) {} -void DiscoverArraysVisitor::Visit(const char* key, int* value) {} -void DiscoverArraysVisitor::Visit(const char* key, bool* value) {} -void DiscoverArraysVisitor::Visit(const char* key, std::string* value) {} -void DiscoverArraysVisitor::Visit(const char* key, DataType* value) {} -void DiscoverArraysVisitor::Visit(const char* key, runtime::NDArray* value) {} -void DiscoverArraysVisitor::Visit(const char* key, void** value) {} - -void DiscoverArraysVisitor::Visit(const char* key, ObjectRef* value) { - address_parts_.push_back(key); - if (value->as() != nullptr) { - auto metadata = Downcast(*value); - const runtime::metadata::MetadataArrayNode* arr = - value->as(); - if (arr != nullptr) { - for (unsigned int i = 0; i < arr->array.size(); i++) { - ObjectRef o = arr->array[i]; - if (o.as() != nullptr) { - std::stringstream ss; - ss << i; - address_parts_.push_back(ss.str()); - runtime::metadata::MetadataBase metadata = Downcast(o); - ReflectionVTable::Global()->VisitAttrs(metadata.operator->(), this); - address_parts_.pop_back(); - } - } - - queue_->push_back(std::make_tuple(AddressFromParts(address_parts_), - Downcast(metadata))); - } else { - ReflectionVTable::Global()->VisitAttrs(metadata.operator->(), this); - } - } - address_parts_.pop_back(); -} - -void DiscoverComplexTypesVisitor::Visit(const char* key, double* value) {} -void DiscoverComplexTypesVisitor::Visit(const char* key, int64_t* value) {} -void DiscoverComplexTypesVisitor::Visit(const char* key, uint64_t* value) {} -void DiscoverComplexTypesVisitor::Visit(const char* key, int* value) {} -void DiscoverComplexTypesVisitor::Visit(const char* key, bool* value) {} -void DiscoverComplexTypesVisitor::Visit(const char* key, std::string* value) {} -void DiscoverComplexTypesVisitor::Visit(const char* key, DataType* value) {} -void DiscoverComplexTypesVisitor::Visit(const char* key, runtime::NDArray* value) {} -void DiscoverComplexTypesVisitor::Visit(const char* key, void** value) {} - -bool DiscoverComplexTypesVisitor::DiscoverType(std::string type_key) { - VLOG(2) << "DiscoverType " << type_key; - auto position_it = type_key_to_position_.find(type_key); - if (position_it != type_key_to_position_.end()) { - return false; - } - - queue_->emplace_back(tvm::runtime::metadata::MetadataBase()); - type_key_to_position_[type_key] = queue_->size() - 1; - return true; -} - -void DiscoverComplexTypesVisitor::DiscoverInstance(runtime::metadata::MetadataBase md) { - auto position_it = type_key_to_position_.find(md->GetTypeKey()); - ICHECK(position_it != type_key_to_position_.end()) - << "DiscoverInstance requires that DiscoverType has already been called: type_key=" - << md->GetTypeKey(); - - int queue_position = (*position_it).second; - if (!(*queue_)[queue_position].defined() && md.defined()) { - VLOG(2) << "DiscoverInstance " << md->GetTypeKey() << ":" << md; - (*queue_)[queue_position] = md; - } -} - -void DiscoverComplexTypesVisitor::Visit(const char* key, ObjectRef* value) { - ICHECK_NOTNULL(value->as()); - - auto metadata = Downcast(*value); - const runtime::metadata::MetadataArrayNode* arr = - value->as(); - - if (arr == nullptr) { - VLOG(2) << "No array, object-traversing " << metadata->GetTypeKey(); - ReflectionVTable::Global()->VisitAttrs(metadata.operator->(), this); - DiscoverType(metadata->GetTypeKey()); - DiscoverInstance(metadata); - return; - } - - if (arr->kind != tvm::runtime::metadata::MetadataKind::kMetadata) { - return; - } - - bool needs_instance = DiscoverType(arr->type_key); - for (unsigned int i = 0; i < arr->array.size(); i++) { - tvm::runtime::metadata::MetadataBase o = - Downcast(arr->array[i]); - if (needs_instance) { - DiscoverInstance(o); - needs_instance = false; - } - ReflectionVTable::Global()->VisitAttrs(o.operator->(), this); - } -} - -void DiscoverComplexTypesVisitor::Discover(runtime::metadata::MetadataBase metadata) { - ReflectionVTable::Global()->VisitAttrs(metadata.operator->(), this); - DiscoverType(metadata->GetTypeKey()); - DiscoverInstance(metadata); -} - -} // namespace metadata -} // namespace codegen -} // namespace tvm diff --git a/src/target/metadata_utils.h b/src/target/metadata_utils.h deleted file mode 100644 index f21de2986e33..000000000000 --- a/src/target/metadata_utils.h +++ /dev/null @@ -1,146 +0,0 @@ -/* - * 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 tvm/target/metadata_utils.h - * \brief Declares utilty functions and classes for emitting metadata. - */ -#ifndef TVM_TARGET_METADATA_UTILS_H_ -#define TVM_TARGET_METADATA_UTILS_H_ - -#include -#include -#include - -#include -#include -#include -#include - -#include "metadata.h" - -namespace tvm { -namespace codegen { -namespace metadata { - -/*! - * \brief Construct a unique string "address" for a struct member from a vector of pieces. - * - * In codegen, it is frequently necessary to assemble a C-style identifier for an - * otherwise-anonymous member of Metadata. For instance, suppose Metadata declares an array: - * struct TVMMetadata { - * int64_t* shape; - * }; - * - * In order to properly initialize this struct, the array must be declared separately with a global - * name. This function produces such a name, here termed "address." - * - * \param parts A vector of pieces, typically the struct member names which identify the path to - * this member. - * \return The joined pieces. - */ -std::string AddressFromParts(const std::vector& parts); - -/*! - * \brief A prefix in metadata symbol names. - * This prefix is typically given to AddressFromParts as the 0th item in parts. - */ -static constexpr const char* kMetadataGlobalSymbol = "kTvmgenMetadata"; - -/*! - * \brief Post-order traverse metadata to discover arrays which need to be forward-defined. - */ -class DiscoverArraysVisitor : public AttrVisitor { - public: - /*! \brief Models a single array discovered in this visitor. - * Conatains two fields: - * 0. An address which uniquely identifies the array in this Metadata instance. - * 1. The discovered MetadataArray. - */ - using DiscoveredArray = std::tuple; - explicit DiscoverArraysVisitor(std::vector* queue); - - void Visit(const char* key, double* value) final; - void Visit(const char* key, int64_t* value) final; - void Visit(const char* key, uint64_t* value) final; - void Visit(const char* key, int* value) final; - void Visit(const char* key, bool* value) final; - void Visit(const char* key, std::string* value) final; - void Visit(const char* key, DataType* value) final; - void Visit(const char* key, runtime::NDArray* value) final; - void Visit(const char* key, void** value) final; - - void Visit(const char* key, ObjectRef* value) final; - - private: - /*! \brief The queue to be filled with discovered arrays. */ - std::vector* queue_; - - /*! \brief Tracks the preceding address pieces. */ - std::vector address_parts_; -}; - -/*! - * \brief Post-order traverse Metadata to discover all complex types which need to be - * forward-defined. This visitor finds one defined() MetadataBase instance for each unique subclass - * present inside Metadata in the order in which the subclass was first discovered. - */ -class DiscoverComplexTypesVisitor : public AttrVisitor { - public: - /*! \brief Construct a new instance. - * \param queue An ordered map which holds the - */ - explicit DiscoverComplexTypesVisitor(std::vector* queue) - : queue_{queue} { - int i = 0; - for (auto q : *queue) { - type_key_to_position_[q->GetTypeKey()] = i++; - } - } - - void Visit(const char* key, double* value) final; - void Visit(const char* key, int64_t* value) final; - void Visit(const char* key, uint64_t* value) final; - void Visit(const char* key, int* value) final; - void Visit(const char* key, bool* value) final; - void Visit(const char* key, std::string* value) final; - void Visit(const char* key, DataType* value) final; - void Visit(const char* key, runtime::NDArray* value) final; - void Visit(const char* key, void** value) final; - - void Visit(const char* key, ObjectRef* value) final; - - void Discover(runtime::metadata::MetadataBase metadata); - - private: - bool DiscoverType(std::string type_key); - - void DiscoverInstance(runtime::metadata::MetadataBase md); - - std::vector* queue_; - - /*! \brief map type_index to index in queue_. */ - std::unordered_map type_key_to_position_; -}; - -} // namespace metadata -} // namespace codegen -} // namespace tvm - -#endif // TVM_TARGET_METADATA_UTILS_H_ diff --git a/src/target/source/codegen_c_host.cc b/src/target/source/codegen_c_host.cc index 2e059eeee520..9b062ef488bb 100644 --- a/src/target/source/codegen_c_host.cc +++ b/src/target/source/codegen_c_host.cc @@ -22,8 +22,6 @@ */ #include "codegen_c_host.h" -#include -#include #include #include @@ -35,7 +33,6 @@ #include "../../support/str_escape.h" #include "../build_common.h" -#include "../func_registry_generator.h" #include "codegen_params.h" namespace tvm { @@ -460,22 +457,6 @@ runtime::Module BuildCHost(IRModule mod, Target target) { cg.AddFunction(gvar, prim_func, emit_fwd_func_decl); } - // NOTE: it's possible that kRuntime attr is not attached when the mod was built with tvm.build(). - // See issue #10373. - auto opt_runtime = mod->GetAttr(tvm::attr::kRuntime); - relay::Runtime runtime; - if (opt_runtime.get() != nullptr) { - runtime = opt_runtime.value(); - } else { - runtime = relay::Runtime::Create("cpp", {}); - } - - bool has_aot_executor_fn = std::any_of( - funcs.begin(), funcs.end(), [&](const auto& kv) { return is_aot_executor_fn(kv.second); }); - if (has_aot_executor_fn && runtime->name == relay::kTvmRuntimeCpp) { - cg.InitGlobalContext(); - } - if (target->GetAttr("system-lib").value_or(Bool(false))) { ICHECK_EQ(target->GetAttr("runtime").value_or(""), "c") << "c target only supports generating C runtime SystemLibs"; diff --git a/src/target/source/codegen_source_base.h b/src/target/source/codegen_source_base.h index e2312ddb778e..a416e3fcae31 100644 --- a/src/target/source/codegen_source_base.h +++ b/src/target/source/codegen_source_base.h @@ -26,7 +26,6 @@ #define TVM_TARGET_SOURCE_CODEGEN_SOURCE_BASE_H_ #include -#include #include #include #include @@ -162,12 +161,11 @@ runtime::Module CSourceModuleCreate(const String& code, const String& fmt, * \param target_module The main TIR-lowered internal runtime module * \param modules All the external modules that needs to be imported inside the metadata module(s). * \param target The target that all the modules are compiled for - * \param metadata Metadata which should be exported to the runtime. * \return The wrapped module. */ runtime::Module CreateMetadataModule( const std::unordered_map& params, runtime::Module target_module, - const Array& ext_modules, Target target, runtime::metadata::Metadata metadata); + const Array& ext_modules, Target target); /*! * \brief Create a source module for viewing and limited saving for device. @@ -181,16 +179,6 @@ runtime::Module DeviceSourceModuleCreate( std::string data, std::string fmt, std::unordered_map fmap, std::string type_key, std::function fget_source = nullptr); -/*! - * \brief Wrap the submodules that are to be wrapped in a c-source metadata module for C runtime. - * \param modules The modules to be wrapped. - * \param target the target the modules are compiled for. - * \param metadata the metadata needed for code generation. - * \return The wrapped module. - */ -runtime::Module CreateCSourceCrtMetadataModule(const Array& modules, Target target, - runtime::metadata::Metadata metadata); - } // namespace codegen } // namespace tvm #endif // TVM_TARGET_SOURCE_CODEGEN_SOURCE_BASE_H_ diff --git a/src/target/source/codegen_vhls.cc b/src/target/source/codegen_vhls.cc index aa7a32320c5e..e4ea1db347cc 100644 --- a/src/target/source/codegen_vhls.cc +++ b/src/target/source/codegen_vhls.cc @@ -23,7 +23,6 @@ #include "codegen_vhls.h" #include -#include #include "../../runtime/opencl/sdaccel/sdaccel_module.h" #include "../build_common.h" diff --git a/src/target/source/codegen_vhls.h b/src/target/source/codegen_vhls.h index 32ddce1b3a30..d8ba2b687496 100644 --- a/src/target/source/codegen_vhls.h +++ b/src/target/source/codegen_vhls.h @@ -28,8 +28,6 @@ #include #include -#include - #include "codegen_c.h" namespace tvm { diff --git a/src/target/source/interface_c.cc b/src/target/source/interface_c.cc deleted file mode 100644 index 8529b8b1301c..000000000000 --- a/src/target/source/interface_c.cc +++ /dev/null @@ -1,320 +0,0 @@ -/* - * 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 interface_c.cc - * \brief Generates a C interface header for a given modules inputs and outputs - */ - -#include -#include -#include -#include -#include -#include -#include - -#include -#include - -#include "../../relay/backend/name_transforms.h" -#include "codegen_params.h" - -namespace tvm { -namespace codegen { - -using runtime::PackedFunc; -using namespace tvm::relay::backend; - -class InterfaceCNode : public runtime::ModuleNode { - public: - InterfaceCNode(std::string module_name, Array inputs, Array outputs, - Array pools, - Map io_pool_allocations, Array devices, - int workspace_size, Map input_sizes, - Map output_sizes) - : module_name_(module_name), - inputs_(inputs), - outputs_(outputs), - devices_(devices), - pools_(FilterExternalPools(pools)), - io_pool_allocations_(io_pool_allocations), - workspace_size_(workspace_size), - input_sizes_(input_sizes), - output_sizes_(output_sizes) {} - const char* type_key() const final { return "h"; } - - String GetSource(const String& format) final { - std::stringstream code; - - EmitUpperHeaderGuard(code); - - // Emit macros for input sizes - for (auto const& it : input_sizes_) { - std::string input_name = SanitizeName(it.first); - std::string input_macro_name = input_name + "_size"; - int input_size = it.second->value; - EmitIntegerValueMacro(code, "Input tensor " + input_name + " size (in bytes)", - input_macro_name, input_size); - } - - // Emit macros for output sizes - for (auto const& it : output_sizes_) { - std::string output_name = SanitizeName(it.first); - std::string output_macro_name = output_name + "_size"; - int output_size = it.second->value; - EmitIntegerValueMacro(code, "Output tensor " + output_name + " size (in bytes)", - output_macro_name, output_size); - } - - EmitBrief(code, "Input tensor pointers"); - EmitStruct(code, "inputs", inputs_); - EmitBrief(code, "Output tensor pointers"); - EmitStruct(code, "outputs", outputs_); - - if (!devices_.empty()) { - EmitBrief(code, "Device context pointers"); - EmitStruct(code, "devices", devices_); - } - if (!pools_.empty()) { - EmitBrief(code, "Workspace pool pointers"); - Array pool_names; - for (const tir::usmp::AllocatedPoolInfo pool : pools_) { - pool_names.push_back(pool->pool_info->pool_name); - } - EmitStruct(code, "workspace_pools", pool_names); - } - - if (!io_pool_allocations_.empty()) { - std::string inputs_struct = ToCVariableStyle(PrefixGeneratedName({module_name_, "inputs"})); - EmitMapIOToPoolsFunction(code, inputs_struct, "map_inputs", inputs_); - std::string outputs_struct = ToCVariableStyle(PrefixGeneratedName({module_name_, "outputs"})); - EmitMapIOToPoolsFunction(code, outputs_struct, "map_outputs", outputs_); - } - - EmitRunFunction(code); - // Emit workspace - EmitIntegerValueMacro(code, "Workspace size", "WORKSPACE_SIZE", workspace_size_); - // Emit memory pool sizes - for (const tir::usmp::AllocatedPoolInfo pool : pools_) { - String pool_name = pool->pool_info->pool_name; - Integer pool_size = pool->allocated_size; - if (const auto* pool_info = pool->pool_info.as()) { - EmitConstantPool(code, SanitizeName(pool_name) + " initialization data", pool_info); - } else { - EmitIntegerValueMacro(code, SanitizeName(pool_name) + " size", - SanitizeName(pool_name) + _macro_workspace_pool_size_postfix, - pool_size->value); - } - } - EmitLowerHeaderGuard(code); - - return code.str(); - } - - PackedFunc GetFunction(const String& name, const ObjectPtr& sptr_to_self) final { - return PackedFunc(); - } - - private: - constexpr static const char* _macro_workspace_pool_size_postfix = "_WORKSPACE_POOL_SIZE"; - constexpr static const char* _macro_constant_pool_size_postfix = "_CONSTANT_POOL_SIZE"; - constexpr static const char* _macro_constant_pool_data_postfix = "_CONSTANT_POOL_DATA"; - - void EmitUpperHeaderGuard(std::stringstream& code_stream) { - std::string header_guard_name = ToCConstantStyle(PrefixGeneratedName({module_name_, "H"})); - code_stream << "#ifndef " << header_guard_name << "_\n" - << "#define " << header_guard_name << "_\n" - << "#include \n\n" - << "#ifdef __cplusplus\n" - << "extern \"C\" {\n" - << "#endif\n\n"; - } - - void EmitLowerHeaderGuard(std::stringstream& code_stream) { - std::string header_guard_name = ToCConstantStyle(PrefixGeneratedName({module_name_, "H"})); - code_stream << "\n#ifdef __cplusplus\n" - << "}\n" - << "#endif\n\n" - << "#endif // " << header_guard_name << "_\n"; - } - - void EmitBrief(std::stringstream& code_stream, const std::string& description) { - code_stream << "/*!\n" - << " * \\brief " << description << " for TVM module \"" << module_name_ << "\" \n" - << " */\n"; - } - - void EmitStruct(std::stringstream& code_stream, const std::string& suffix, - Array properties) { - std::string struct_name = ToCVariableStyle(PrefixGeneratedName({module_name_, suffix})); - code_stream << "struct " << struct_name << " {\n"; - - std::vector sanitized_properties; - for (const String& property : properties) { - std::string sanitized_property = SanitizeName(property); - ICHECK(std::find(sanitized_properties.begin(), sanitized_properties.end(), - sanitized_property) == sanitized_properties.end()) - << "Sanitized input tensor name clash" << sanitized_property; - code_stream << " void* " << sanitized_property << ";\n"; - sanitized_properties.push_back(sanitized_property); - } - code_stream << "};\n\n"; - } - - void EmitIntegerValueMacro(std::stringstream& code_stream, const std::string& brief_description, - const std::string& macro_name, int macro_value) { - EmitBrief(code_stream, brief_description); - std::string macro_name_prefixed = - ToCConstantStyle(PrefixGeneratedName({module_name_, macro_name})); - code_stream << "#define " << macro_name_prefixed << " " << macro_value << "\n"; - } - - void EmitConstantPool(std::stringstream& code_, const std::string& brief_description, - const ConstantPoolInfoNode* pool_info) { - EmitBrief(code_, brief_description); - std::string name_prefixed = - ToCConstantStyle(PrefixGeneratedName({module_name_, SanitizeName(pool_info->pool_name)})); - - if (pool_info->constant_info_array.size() > 0) { - std::vector const_info_vec(pool_info->constant_info_array.begin(), - pool_info->constant_info_array.end()); - std::sort(const_info_vec.begin(), const_info_vec.end(), - [](const ConstantInfo& a, const ConstantInfo& b) { - return a->byte_offset->value < b->byte_offset->value; - }); - int64_t accumulated_pool_len = - const_info_vec.back()->byte_offset.IntValue() + - runtime::GetDataSize(*const_info_vec.back()->data.operator->()); - const auto& accumulated_pool = runtime::NDArray::Empty( - {accumulated_pool_len}, DataType::UInt(8), const_info_vec.back()->data->device); - for (const auto& const_info : const_info_vec) { - const auto& data = const_info->data; - const auto& offs = const_info->byte_offset; - data.CopyToBytes(static_cast(accumulated_pool->data) + offs.IntValue(), - runtime::GetDataSize(*data.operator->())); - } - - code_ << "#define " << name_prefixed << _macro_constant_pool_size_postfix << " " - << accumulated_pool_len << "\n"; - code_ << "#define " << name_prefixed << _macro_constant_pool_data_postfix << " \\\n"; - codegen::NDArrayDataToC(accumulated_pool, 4, code_, "\\\n"); - code_ << '\n'; - - } else { - LOG(FATAL) << "No constant data in constant pool found " << GetRef(pool_info); - } - } - - void EmitRunFunction(std::stringstream& code_stream) { - std::string run_function = ToCVariableStyle(PrefixGeneratedName({module_name_, "run"})); - std::string inputs_struct = ToCVariableStyle(PrefixGeneratedName({module_name_, "inputs"})); - std::string outputs_struct = ToCVariableStyle(PrefixGeneratedName({module_name_, "outputs"})); - std::string devices_struct = ToCVariableStyle(PrefixGeneratedName({module_name_, "devices"})); - std::string pools_struct = - ToCVariableStyle(PrefixGeneratedName({module_name_, "workspace_pools"})); - - code_stream << "/*!\n" - << " * \\brief entrypoint function for TVM module \"" << module_name_ << "\"\n"; - if (io_pool_allocations_.empty()) { - code_stream << " * \\param inputs Input tensors for the module \n"; - code_stream << " * \\param outputs Output tensors for the module \n"; - } - - if (!pools_.empty()) { - code_stream << " * \\param workspace_pools Workspace memory pool pointers for the module \n"; - } - if (!devices_.empty()) { - code_stream << " * \\param devices Device context pointers for the module \n"; - } - - code_stream << " */\n" - << "int32_t " << run_function << "(\n"; - - std::stringstream call_args_ss; - if (io_pool_allocations_.empty()) { - call_args_ss << " struct " << inputs_struct << "* inputs,\n"; - call_args_ss << " struct " << outputs_struct << "* outputs,\n"; - } - if (!pools_.empty()) { - call_args_ss << " struct " << pools_struct << "* workspace_pools,\n"; - } - if (!devices_.empty()) { - call_args_ss << " struct " << devices_struct << "* devices,\n"; - } - std::string call_args_str = call_args_ss.str(); - call_args_str.pop_back(); - call_args_str.pop_back(); - code_stream << call_args_str << "\n);\n"; - } - - void EmitMapIOToPoolsFunction(std::stringstream& code_stream, const std::string& struct_type, - const std::string& function_name, - const Array& tensor_names) { - code_stream << "/*!\n" - << " * \\brief Maps I/O inside the workspace pools for TVM module \"" - << module_name_ << "\"\n" - << " * \\param workspace_pools Workspace memory pool struct for the module \n" - << " * \\return I/O tensor struct for the module \n"; - std::string map_function = ToCVariableStyle(PrefixGeneratedName({module_name_, function_name})); - code_stream << " */\n" - << "struct " << struct_type << " " << map_function << "(\n"; - std::string pools_struct = - ToCVariableStyle(PrefixGeneratedName({module_name_, "workspace_pools"})); - code_stream << " struct " << pools_struct << "* workspace_pools\n"; - code_stream << ");\n\n"; - } - - Array FilterExternalPools( - const Array& pools) { - Array external_pools; - for (tir::usmp::AllocatedPoolInfo pool : pools) { - if (!pool->pool_info->is_internal) { - external_pools.push_back(pool); - } - } - return external_pools; - } - - std::string module_name_; - Array inputs_; - Array outputs_; - Array devices_; - Array pools_; - Map io_pool_allocations_; - int workspace_size_; - Map input_sizes_; - Map output_sizes_; -}; - -runtime::Module InterfaceCCreate(std::string module_name, Array inputs, - Array outputs, Array pools, - Map io_pool_allocations, - Array devices, int workspace_size, - Map input_sizes, - Map output_sizes) { - auto n = make_object(module_name, inputs, outputs, pools, io_pool_allocations, - devices, workspace_size, input_sizes, output_sizes); - return runtime::Module(n); -} - -TVM_REGISTER_GLOBAL("runtime.InterfaceCCreate").set_body_typed(InterfaceCCreate); - -} // namespace codegen -} // namespace tvm diff --git a/src/target/source/source_module.cc b/src/target/source/source_module.cc index 1877d3da8e63..fd16c3c85a02 100644 --- a/src/target/source/source_module.cc +++ b/src/target/source/source_module.cc @@ -21,10 +21,8 @@ * \file source_module.cc * \brief Source code module, only for viewing */ -#include "source_module.h" #include -#include #include #include #include @@ -33,22 +31,13 @@ #include #include -#include #include #include #include -#include #include -#include "../../relay/backend/name_transforms.h" #include "../../runtime/file_utils.h" -#include "../../support/str_escape.h" -#include "../func_registry_generator.h" -#include "../metadata.h" -#include "../metadata_utils.h" -#include "codegen_params.h" #include "codegen_source_base.h" -#include "tvm/relay/executor.h" namespace tvm { namespace codegen { @@ -206,827 +195,6 @@ class ConcreteCodegenSourceBase : public CodeGenSourceBase { } }; -class CSourceCrtMetadataModuleNode : public runtime::ModuleNode { - public: - CSourceCrtMetadataModuleNode(const Array& func_names, const std::string& fmt, - Target target, relay::Runtime runtime, - relay::backend::ExecutorCodegenMetadata metadata) - : fmt_(fmt), - func_names_(func_names), - target_(target), - runtime_(runtime), - metadata_(metadata) { - CreateSource(); - } - const char* type_key() const final { return "c"; } - - String GetSource(const String& format) final { return code_.str(); } - - String GetFormat() override { return fmt_; } - PackedFunc GetFunction(const String& name, const ObjectPtr& sptr_to_self) final { - return PackedFunc(); - } - - void SaveToFile(const String& file_name, const String& format) final { - std::string fmt = GetFileFormat(file_name, format); - std::string meta_file = GetMetaFilePath(file_name); - if (fmt == "c" || fmt == "cc" || fmt == "cpp") { - auto code_str = code_.str(); - ICHECK_NE(code_str.length(), 0); - SaveBinaryToFile(file_name, code_str); - } else { - ICHECK_EQ(fmt, fmt_) << "Can only save to format=" << fmt_; - } - } - - int GetPropertyMask() const override { return runtime::ModulePropertyMask::kDSOExportable; } - - bool ImplementsFunction(const String& name, bool query_imports) final { - return std::find(func_names_.begin(), func_names_.end(), name) != func_names_.end(); - } - - protected: - std::stringstream code_; - std::string fmt_; - Array func_names_; - Target target_; - relay::Runtime runtime_; - relay::backend::ExecutorCodegenMetadata metadata_; - ConcreteCodegenSourceBase codegen_c_base_; - - void CreateFuncRegistry() { - code_ << "#include \n"; - for (const auto& fname : func_names_) { - code_ << "#ifdef __cplusplus\n"; - code_ << "extern \"C\"\n"; - code_ << "#endif\n"; - code_ << "TVM_DLL int32_t " << fname.data(); - code_ << "(TVMValue* args, int* type_code, int num_args, TVMValue* out_value, " - "int* out_type_code, void* resource_handle);\n"; - } - code_ << "static TVMBackendPackedCFunc _tvm_func_array[] = {\n"; - for (auto f : func_names_) { - code_ << " (TVMBackendPackedCFunc)" << f << ",\n"; - } - code_ << "};\n"; - auto registry = target::GenerateFuncRegistryNames(func_names_); - code_ << "static const TVMFuncRegistry _tvm_func_registry = {\n" - << " \"" << ::tvm::support::StrEscape(registry.data(), registry.size(), true) << "\"," - << " _tvm_func_array,\n" - << "};\n"; - } - - void GenerateCrtSystemLib() { - code_ << "static const TVMModule _tvm_system_lib = {\n" - << " &_tvm_func_registry,\n" - << "};\n" - << "const TVMModule* TVMSystemLibEntryPoint(void) {\n" - << " return &_tvm_system_lib;\n" - << "}\n"; - } - - String GenerateDLTensorStructWrapper(String reference_arg) { - code_ << "DLTensor " << reference_arg << "_dltensor = {\n"; - code_ << ".data = &" << reference_arg << "\n"; - code_ << "};\n"; - code_ << "TVMValue " << reference_arg << "_tvm_value = {\n"; - code_ << ".v_handle = &" << reference_arg << "_dltensor\n"; - code_ << "};\n"; - return reference_arg + "_tvm_value"; - } - - void GenerateInternalBuffers() { - if (metadata_->pool_inputs.defined()) { - for (const auto& kv : metadata_->pool_inputs.value()) { - tir::usmp::AllocatedPoolInfo allocated_pool_info = kv.second; - if (allocated_pool_info->pool_info->is_internal) { - if (const auto* pool_info = allocated_pool_info->pool_info.as()) { - GenerateConstantBuffer(pool_info, allocated_pool_info->allocated_size->value); - } else { - GenerateWorkspaceBuffer(allocated_pool_info->pool_info.as(), - allocated_pool_info->allocated_size->value); - } - } - } - } - } - - void GenerateIOWorkspaceMapFunction(const std::string& struct_type, - const std::string& function_name, - const Array& tensor_names) { - std::string map_function = runtime::get_name_mangled(metadata_->mod_name, function_name); - code_ << "struct " << struct_type << " " << map_function << "(\n"; - std::string pools_struct = runtime::get_name_mangled(metadata_->mod_name, "workspace_pools"); - code_ << " struct " << pools_struct << "* workspace_pools\n"; - code_ << "\n){\n"; - code_ << "struct " << struct_type << " ret = {\n"; - for (const String& name : tensor_names) { - tir::usmp::PoolAllocation pool_allocation = metadata_->io_pool_allocations[name]; - code_ << "\t." << name << " = " - << "&((uint8_t*)workspace_pools->" << pool_allocation->pool_info->pool_name << ")[" - << pool_allocation->byte_offset << "],\n"; - } - code_ << "};\n"; - code_ << "return ret;\n"; - code_ << "}\n\n"; - } - - void GenerateConstantBuffer(const ConstantPoolInfoNode* pool_info, size_t allocated_size) { - size_t ord = 0; - if (pool_info->constant_info_array.size() > 0) { - // Pool is RO, form an initialized struct - code_ << "__attribute__((section(\".rodata.tvm\"), "; - code_ << "))\n"; - code_ << "static const struct " << pool_info->pool_name << " {\n"; - // emit struct field names - std::vector const_info_vec(pool_info->constant_info_array.begin(), - pool_info->constant_info_array.end()); - std::sort(const_info_vec.begin(), const_info_vec.end(), - [](const ConstantInfo& a, const ConstantInfo& b) { - return a->byte_offset->value < b->byte_offset->value; - }); - for (const auto& const_info : const_info_vec) { - const auto& data = const_info->data; - const auto& offs = const_info->byte_offset; - int64_t num_elements = std::accumulate(data.Shape().begin(), data.Shape().end(), 1, - std::multiplies()); - code_ << " "; - codegen_c_base_.PrintType(data.DataType(), code_); - code_ << " " << const_info->name_hint << "[" << num_elements << "] __attribute__((" - << (ord++ ? "packed, " : "") << "aligned(" << metadata_->constant_alignment << ")));"; - code_ << " // " << num_elements * data.DataType().bytes() - << " bytes, aligned offset: " << offs << "\n"; - } - code_ << "} " << pool_info->pool_name << " = {\n"; - - // emit struct field initialization data - for (const auto& const_info : const_info_vec) { - code_ << " ." << const_info->name_hint << " = {\n"; - codegen::NDArrayDataToC(const_info->data, 4, code_); - code_ << " },\n"; - } - code_ << "};"; - code_ << "// of total size " << allocated_size << " bytes\n"; - } else { - LOG(FATAL) << "No constant data in constant pool found " << GetRef(pool_info); - } - } - - void GenerateWorkspaceBuffer(const WorkspacePoolInfoNode* pool_info, size_t allocated_size) { - code_ << "__attribute__((section(\".bss.noinit.tvm\"), "; - code_ << "aligned(" << metadata_->workspace_alignment << ")))\n"; - code_ << "static uint8_t " << pool_info->pool_name << "["; - code_ << allocated_size << "];\n"; - } - - bool IsInternalWorkspaceBuffer(const tir::Var& pool_var) { - if (metadata_->pool_inputs.defined()) { - Map allocated_pool_infos = - metadata_->pool_inputs.value(); - if (allocated_pool_infos.find(pool_var) != allocated_pool_infos.end()) { - tir::usmp::AllocatedPoolInfo allocate_pool_info = allocated_pool_infos[pool_var]; - if (allocate_pool_info->pool_info->is_internal) { - return true; - } - } - } - return false; - } - - void GenerateEntrypointForUnpackedAPI(const std::string& entrypoint_name, - const std::string& run_func) { - code_ << "TVM_DLL int32_t " << run_func << "("; - - { - std::stringstream call_args_ss; - if (metadata_->io_pool_allocations.empty()) { - for (const tir::Var& input_var : metadata_->inputs) { - if (input_var->type_annotation.defined()) { - codegen_c_base_.PrintType(input_var->type_annotation, call_args_ss); - } else { - codegen_c_base_.PrintType(input_var.dtype(), call_args_ss); - } - call_args_ss << " " << input_var->name_hint << ","; - } - for (unsigned int i = 0; i < metadata_->outputs.size(); ++i) { - call_args_ss << "void* output" << i << ","; - } - } - for (const tir::Var& pool_var : metadata_->pools) { - if (pool_var->type_annotation.defined()) { - codegen_c_base_.PrintType(pool_var->type_annotation, call_args_ss); - } else { - codegen_c_base_.PrintType(pool_var.dtype(), call_args_ss); - } - call_args_ss << " " << pool_var->name_hint << ","; - } - std::string call_args_str = call_args_ss.str(); - call_args_str.pop_back(); - code_ << call_args_str; - } - - code_ << ");\n"; - code_ << "int32_t " << entrypoint_name; - code_ << "(void* args, void* type_code, int num_args, void* out_value, void* " - "out_type_code, void* resource_handle) {\n"; - code_ << "return " << run_func << "("; - - { - std::stringstream call_args_ss; - if (metadata_->io_pool_allocations.empty()) { - for (unsigned int i = 0; i < metadata_->inputs.size(); ++i) { - call_args_ss << "((DLTensor*)(((TVMValue*)args)[" << i << "].v_handle))[0].data,"; - } - for (unsigned int i = 0; i < metadata_->outputs.size(); ++i) { - int j = metadata_->inputs.size() + i; - call_args_ss << "((DLTensor*)(((TVMValue*)args)[" << j << "].v_handle))[0].data,"; - } - } - for (const tir::Var& pool_var : metadata_->pools) { - if (IsInternalWorkspaceBuffer(pool_var)) { - call_args_ss << "&" << metadata_->pool_inputs.value()[pool_var]->pool_info->pool_name - << ","; - } - } - std::string call_args_str = call_args_ss.str(); - call_args_str.pop_back(); - code_ << call_args_str; - code_ << ");\n"; - code_ << "}\n"; - } - } - - std::unordered_map GenerateRunFuncToEntryPointArgMap() { - std::unordered_map run_func_to_entry_point_args; - int entrypoint_arg_count = 0; - int run_func_arg_count = 0; - - if (metadata_->io_pool_allocations.empty()) { - for (unsigned int i = 0; i < metadata_->inputs.size(); i++) { - run_func_to_entry_point_args[run_func_arg_count] = Integer(entrypoint_arg_count); - entrypoint_arg_count++; - run_func_arg_count++; - } - for (unsigned int i = 0; i < metadata_->outputs.size(); i++) { - run_func_to_entry_point_args[run_func_arg_count] = Integer(entrypoint_arg_count); - entrypoint_arg_count++; - run_func_arg_count++; - } - } - for (const tir::Var& pool_var : metadata_->pools) { - if (IsInternalWorkspaceBuffer(pool_var)) { - tir::usmp::AllocatedPoolInfo allocated_pool_info = metadata_->pool_inputs.value()[pool_var]; - run_func_to_entry_point_args[run_func_arg_count] = - allocated_pool_info->pool_info->pool_name; - run_func_arg_count++; - } - } - return run_func_to_entry_point_args; - } - - void GenerateEntrypointForPackedAPI(const std::string& entrypoint_name, - const std::string& run_func) { - code_ << "TVM_DLL int32_t " << run_func; - code_ << "(TVMValue* args, int* type_code, int num_args, TVMValue* out_value, int* " - "out_type_code, void* resource_handle);\n\n"; - - code_ << "int32_t " << entrypoint_name; - code_ << "(TVMValue* args, int* type_code, int num_args, TVMValue* out_value, int* " - "out_type_code, void* resource_handle) {\n"; - - // We are creating a copy of the set of pointers - size_t number_of_io_tensors = metadata_->inputs.size() + metadata_->outputs.size() + - metadata_->pools.size() - metadata_->io_pool_allocations.size(); - code_ << "TVMValue tensors[" << number_of_io_tensors << "];\n"; - - std::unordered_map run_func_to_entry_point_args = - GenerateRunFuncToEntryPointArgMap(); - for (unsigned int i = 0; i < number_of_io_tensors; i++) { - if (run_func_to_entry_point_args.find(i) != run_func_to_entry_point_args.end()) { - if (run_func_to_entry_point_args[i]->IsInstance()) { - String pool_name = Downcast(run_func_to_entry_point_args[i]); - String pool_name_tvmv = GenerateDLTensorStructWrapper(pool_name); - code_ << "tensors[" << i << "] = " << pool_name_tvmv << ";\n"; - } else { - code_ << "tensors[" << i << "] = ((TVMValue*)args)[" << run_func_to_entry_point_args[i] - << "];\n"; - } - } - } - - code_ << "return " << run_func; - code_ << "((void*)tensors, type_code, num_args, out_value, out_type_code, resource_handle);\n"; - code_ << "}\n"; - } - - static int isNotAlnum(char c) { return !std::isalnum(c); } - - void GenerateCInterfaceEntrypoint(const std::string& entrypoint_name, const std::string& run_func, - const std::string& mod_name) { - code_ << "#include <" << mod_name << ".h>\n"; - if (!metadata_->io_pool_allocations.empty()) { - const std::string input_struct_type = - runtime::get_name_mangled(metadata_->mod_name, "inputs"); - Array input_tensor_names; - for (const tir::Var& input_var : metadata_->inputs) { - input_tensor_names.push_back(input_var->name_hint); - } - GenerateIOWorkspaceMapFunction(input_struct_type, "map_inputs", input_tensor_names); - const std::string output_struct_type = - runtime::get_name_mangled(metadata_->mod_name, "outputs"); - GenerateIOWorkspaceMapFunction(output_struct_type, "map_outputs", metadata_->outputs); - } - code_ << "TVM_DLL int32_t " << run_func << "("; - { - std::stringstream call_args_ss; - if (metadata_->io_pool_allocations.empty()) { - for (const tir::Var& input_var : metadata_->inputs) { - if (input_var->type_annotation.defined()) { - codegen_c_base_.PrintType(input_var->type_annotation, call_args_ss); - } else { - codegen_c_base_.PrintType(input_var.dtype(), call_args_ss); - } - call_args_ss << " " << tvm::runtime::SanitizeName(input_var->name_hint) << ","; - } - for (unsigned int i = 0; i < metadata_->outputs.size(); ++i) { - call_args_ss << "void* output" << i << ","; - } - } - for (const tir::Var& pool_var : metadata_->pools) { - if (pool_var->type_annotation.defined()) { - codegen_c_base_.PrintType(pool_var->type_annotation, call_args_ss); - } else { - codegen_c_base_.PrintType(pool_var.dtype(), call_args_ss); - } - call_args_ss << " " << pool_var->name_hint << ","; - } - for (const String& device : metadata_->devices) { - call_args_ss << "void* " << device << ","; - } - std::string call_args_str = call_args_ss.str(); - call_args_str.pop_back(); - code_ << call_args_str; - } - - code_ << ");\n"; - code_ << "int32_t " << entrypoint_name << "("; - { - std::stringstream call_args_ss; - if (metadata_->io_pool_allocations.empty()) { - call_args_ss << "struct " << runtime::get_name_mangled(mod_name, "inputs") << "* inputs,"; - call_args_ss << "struct " << runtime::get_name_mangled(mod_name, "outputs") << "* outputs,"; - } - if (!metadata_->pools.empty()) { - bool is_external_pools_present = false; - for (tir::Var pool_var : metadata_->pools) { - if (!IsInternalWorkspaceBuffer(pool_var)) { - is_external_pools_present = true; - break; - } - } - if (is_external_pools_present) { - call_args_ss << "struct " << runtime::get_name_mangled(mod_name, "workspace_pools") - << "* workspace_pools,"; - } - } - if (!metadata_->devices.empty()) { - call_args_ss << "struct " << runtime::get_name_mangled(mod_name, "devices") << "* devices,"; - } - std::string call_args_str = call_args_ss.str(); - call_args_str.pop_back(); - code_ << call_args_str; - } - - code_ << ") {" - << "return " << run_func << "("; - - { - std::stringstream call_args_ss; - if (metadata_->io_pool_allocations.empty()) { - for (const auto& input : metadata_->inputs) { - call_args_ss << "inputs->" << tvm::runtime::SanitizeName(input->name_hint) << ","; - } - for (const auto& output : metadata_->outputs) { - call_args_ss << "outputs->" << tvm::runtime::SanitizeName(output); - call_args_ss << ","; - } - } - - for (const tir::Var& pool_var : metadata_->pools) { - call_args_ss << "((uint8_t*)"; - String pool_name = metadata_->pool_inputs.value()[pool_var]->pool_info->pool_name; - if (IsInternalWorkspaceBuffer(pool_var)) { - call_args_ss << "&" << pool_name; - } else { - call_args_ss << "workspace_pools->" << tvm::runtime::SanitizeName(pool_name); - } - call_args_ss << "),"; - } - for (const String& device : metadata_->devices) { - call_args_ss << "devices->" << device << ","; - } - std::string call_args_str = call_args_ss.str(); - call_args_str.pop_back(); - code_ << call_args_str; - } - code_ << ");\n"; - code_ << "}\n"; - } - - void GenerateAOTDescriptor() { - const std::string run_func_suffix = ::tvm::runtime::symbol::tvm_module_main; - const std::string tvm_entrypoint_suffix = ::tvm::runtime::symbol::tvm_entrypoint_suffix; - const std::string run_func_mangled = - runtime::get_name_mangled(metadata_->mod_name, run_func_suffix); - const std::string entrypoint_mangled = - runtime::get_name_mangled(metadata_->mod_name, tvm_entrypoint_suffix); - const std::string network_mangled = runtime::get_name_mangled(metadata_->mod_name, "network"); - - code_ << "#include \"tvm/runtime/c_runtime_api.h\"\n"; - code_ << "#ifdef __cplusplus\n"; - code_ << "extern \"C\" {\n"; - code_ << "#endif\n"; - - GenerateInternalBuffers(); - - if (metadata_->unpacked_api) { - if (metadata_->interface_api == "c") { - GenerateCInterfaceEntrypoint(entrypoint_mangled, run_func_mangled, metadata_->mod_name); - } else { - GenerateEntrypointForUnpackedAPI(entrypoint_mangled, run_func_mangled); - } - } else { - ICHECK_EQ(metadata_->interface_api, "packed") - << "Packed interface required for packed operators"; - GenerateEntrypointForPackedAPI(entrypoint_mangled, run_func_mangled); - } - - code_ << "#ifdef __cplusplus\n"; - code_ << "}\n"; - code_ << "#endif\n"; - } - - void CreateSource() { - if (runtime_->GetAttr("system-lib").value_or(Bool(false)) && !func_names_.empty()) { - CreateFuncRegistry(); - GenerateCrtSystemLib(); - } - if (metadata_.defined() && metadata_->executor == runtime::kTvmExecutorAot) { - GenerateAOTDescriptor(); - } - code_ << ";"; - } -}; - -class MetadataSerializer : public AttrVisitor { - public: - static constexpr const char* kGlobalSymbol = "kTvmgenMetadata"; - using MetadataKind = ::tvm::runtime::metadata::MetadataKind; - - MetadataSerializer() : is_first_item_{true} {} - - void WriteComma() { - if (is_first_item_) { - is_first_item_ = false; - } else { - code_ << ", " << std::endl; - } - } - - void WriteKey(const char* key) { - if (key != nullptr) { - code_ << " /* " << key << "*/"; - } - } - - void Visit(const char* key, double* value) final { - WriteComma(); - code_.setf(std::ios::hex | std::ios::showbase | std::ios::fixed | std::ios::scientific, - std::ios::basefield | std::ios::showbase | std::ios::floatfield); - code_ << *value; - WriteKey(key); - } - - void Visit(const char* key, int64_t* value) final { - WriteComma(); - code_ << *value << "L"; - WriteKey(key); - } - - void Visit(const char* key, uint64_t* value) final { - WriteComma(); - code_ << *value << "UL"; - WriteKey(key); - } - void Visit(const char* key, int* value) final { - WriteComma(); - code_ << *value; - WriteKey(key); - } - void Visit(const char* key, bool* value) final { - WriteComma(); - code_ << *value; - WriteKey(key); - } - void Visit(const char* key, std::string* value) final { - WriteComma(); - code_ << "\"" << *value << "\""; - WriteKey(key); - } - void Visit(const char* key, void** value) final { - WriteComma(); - code_ << *value; - WriteKey(key); - } - void Visit(const char* key, DataType* value) final { - WriteComma(); - code_ << "{" << value->code() << ", " << value->bits() << ", " << value->lanes() << "}"; - WriteKey(key); - } - - // Serialiding NDArray as tuple of len, data - void Visit(const char* key, runtime::NDArray* value) final { - WriteComma(); - std::string bytes; - dmlc::MemoryStringStream stream(&bytes); - value->Save(&stream); - // Serializing length of the data of NDArray - code_ << stream.Tell(); - WriteComma(); - // Serializing NDArray as bytestream - code_ << "\""; - std::stringstream ss; - char buf[6] = {0}; - for (uint8_t c : bytes) { - snprintf(buf, sizeof(buf), "\\x%02x", c); - ss << buf; - } - std::string as_bytes(ss.str()); - code_ << as_bytes; - code_ << "\"\n"; - } - - void VisitArray(runtime::metadata::MetadataArray array) { - auto old_is_first_item = is_first_item_; - is_first_item_ = true; - for (unsigned int i = 0; i < array->array.size(); ++i) { - ObjectRef o = array->array[i]; - - switch (array->kind) { - case MetadataKind::kUint64: { - int64_t i = Downcast(o).IntValue(); - CHECK_GT(i, 0) - << "Metadata is of type uint64_t, but array type contains a negative number"; - uint64_t ui = static_cast(i); - Visit(nullptr, &ui); - continue; - } - case MetadataKind::kInt64: { - int64_t i = Downcast(o).IntValue(); - Visit(nullptr, &i); - continue; - } - case MetadataKind::kBool: { - bool b = Downcast(o); - Visit(nullptr, &b); - break; - } - case MetadataKind::kString: { - std::string s = Downcast(o); - Visit(nullptr, &s); - break; - } - case MetadataKind::kHandle: - CHECK(false) << "Don't know how to serialize handle"; - break; - - case MetadataKind::kMetadata: { - runtime::metadata::MetadataBase metadata = Downcast(o); - std::stringstream i_str; - i_str << i; - address_.push_back(i_str.str()); - Visit(nullptr, &metadata); - address_.pop_back(); - break; - } - default: - CHECK(false) << "Unknown MetadataKind for array: " << array->kind; - break; - } - is_first_item_ = false; - } - is_first_item_ = old_is_first_item; - } - - void Visit(const char* key, ObjectRef* value) final { - const runtime::metadata::MetadataArrayNode* arr = - value->as(); - if (arr != nullptr) { - WriteComma(); - if (key != nullptr) { - address_.push_back(key); - } - code_ << metadata::AddressFromParts(address_); - if (key != nullptr) { - address_.pop_back(); - } - return; - } - - runtime::metadata::MetadataBase metadata = Downcast(*value); - if (key != nullptr) { // NOTE: outermost call passes nullptr key - address_.push_back(key); - } - WriteComma(); - code_ << "{\n"; - is_first_item_ = true; - ReflectionVTable::Global()->VisitAttrs(metadata.operator->(), this); - code_ << "}\n"; - if (key != nullptr) { // NOTE: outermost call passes nullptr key - address_.pop_back(); - } - } - - private: - void EmitCType(const runtime::metadata::MetadataArrayNode* arr, std::ostream& os) { - switch (arr->kind) { - case MetadataKind::kUint64: - os << "uint64_t"; - break; - case MetadataKind::kInt64: - os << "int64_t"; - break; - case MetadataKind::kBool: - os << "bool"; - break; - case MetadataKind::kString: - os << "const char*"; - break; - case MetadataKind::kHandle: - os << "void*"; - break; - case MetadataKind::kMetadata: - os << "struct " << arr->get_element_c_struct_name(); - break; - default: - CHECK(false) << "Unknown kind in MetadataArray: " << arr->kind - << " (struct_name=" << arr->get_c_struct_name() << ")"; - break; - } - } - - public: - void CodegenMetadata(::tvm::runtime::metadata::Metadata metadata) { - decl_ << "#include " << std::endl - << "#include " << std::endl - << "#include " << std::endl; - std::vector queue; - metadata::DiscoverArraysVisitor array_discover{&queue}; - array_discover.Visit(metadata::kMetadataGlobalSymbol, &metadata); - - for (auto item : queue) { - auto struct_address = std::get<0>(item); - address_.push_back(struct_address); - - auto arr = std::get<1>(item); - - // Prepend const with everything except C-string, which needs appending. - code_ << "static "; - if (arr->kind != MetadataKind::kString) { - code_ << "const "; - } - EmitCType(arr.operator->(), code_); - if (arr->kind == MetadataKind::kString) { - code_ << " const"; - } - code_ << " " << struct_address << "[" << arr->array.size() << "] = {" << std::endl; - is_first_item_ = true; - - VisitArray(arr); - address_.pop_back(); - code_ << "};" << std::endl; - } - - // Finally, emit overall struct. - address_.push_back(metadata::kMetadataGlobalSymbol); - code_ << "static const struct TVMMetadata " << metadata::AddressFromParts(address_) << "[1] = {" - << std::endl; - Visit(nullptr, &metadata); - code_ << "};" << std::endl; - address_.pop_back(); - } - - std::string GetOutput() { return decl_.str() + code_.str(); } - - private: - std::vector address_; - std::stringstream decl_; - std::stringstream code_; - bool is_first_item_; - std::unordered_set generated_struct_decls_; - std::vector is_defining_struct_; -}; - -namespace { -runtime::Module CreateAotMetadataModule(runtime::metadata::Metadata aot_metadata, - bool is_c_runtime) { - MetadataSerializer serializer; - serializer.CodegenMetadata(aot_metadata); - std::stringstream lookup_func; - std::string get_c_metadata_func_name; - - // NOTE: mangling is not needed in the c++ runtime because the function - // name is looked-up via LibraryModule. - // TODO(alanmacd): unify these two approaches - - if (is_c_runtime == true) { - get_c_metadata_func_name = runtime::get_name_mangled( - aot_metadata->mod_name(), ::tvm::runtime::symbol::tvm_get_c_metadata); - } else { - get_c_metadata_func_name = ::tvm::runtime::symbol::tvm_get_c_metadata; - } - - lookup_func << "#ifdef __cplusplus\n" - << "extern \"C\"\n" - << "#endif\n"; - - lookup_func << "TVM_DLL int32_t " << get_c_metadata_func_name - << "(TVMValue* arg_values, int* arg_tcodes, int " - "num_args, TVMValue* ret_values, int* ret_tcodes, void* resource_handle) {" - << std::endl; - lookup_func << " ret_values[0].v_handle = (void*) &" << MetadataSerializer::kGlobalSymbol - << ";" << std::endl; - lookup_func << " ret_tcodes[0] = kTVMOpaqueHandle;" << std::endl; - lookup_func << " return 0;" << std::endl; - lookup_func << "};" << std::endl; - std::vector func_names{get_c_metadata_func_name}; - return CSourceModuleCreate(serializer.GetOutput() + lookup_func.str(), "c", func_names, - Array()); -} -} // namespace - -runtime::Module CreateCSourceCrtMetadataModule(const Array& modules, Target target, - relay::Runtime runtime, - relay::backend::ExecutorCodegenMetadata metadata, - runtime::metadata::Metadata aot_metadata) { - Array final_modules(modules); - Array func_names; - - if (metadata.defined()) { - if (metadata->executor == "aot") { - if (aot_metadata.defined()) { - final_modules.push_back(CreateAotMetadataModule(aot_metadata, true)); - } - - // add the run function (typically "tvmgen_default_run") to function registry - // when using AOT executor - std::string run_func = runtime::get_name_mangled(metadata->mod_name, "run"); - func_names.push_back(run_func); - } - } - - for (runtime::Module mod : final_modules) { - auto pf_funcs = mod.GetFunction("get_func_names"); - if (pf_funcs != nullptr) { - Array func_names_ = pf_funcs(); - for (const auto& fname : func_names_) { - func_names.push_back(fname); - } - } - } - - auto n = make_object(func_names, "c", target, runtime, metadata); - auto csrc_metadata_module = runtime::Module(n); - for (const auto& mod : final_modules) { - csrc_metadata_module.Import(mod); - } - - return std::move(csrc_metadata_module); -} - -runtime::Module CreateCSourceCppMetadataModule(runtime::metadata::Metadata metadata) { - MetadataSerializer serializer; - serializer.CodegenMetadata(metadata); - std::stringstream lookup_func; - lookup_func << "#ifdef __cplusplus\n" - << "extern \"C\"\n" - << "#endif\n"; - - lookup_func << "TVM_DLL int32_t " << ::tvm::runtime::symbol::tvm_get_c_metadata - << "(TVMValue* arg_values, int* arg_tcodes, int " - "num_args, TVMValue* ret_values, int* ret_tcodes, void* resource_handle) {" - << std::endl; - lookup_func << " ret_values[0].v_handle = (void*) &" << metadata::kMetadataGlobalSymbol << ";" - << std::endl; - lookup_func << " ret_tcodes[0] = kTVMOpaqueHandle;" << std::endl; - lookup_func << " return 0;" << std::endl; - lookup_func << "};" << std::endl; - - auto mod = MetadataModuleCreate(metadata); - mod->Import(CreateAotMetadataModule(metadata, false)); - return mod; -} - // supports limited save without cross compile class DeviceSourceModuleNode final : public runtime::ModuleNode { public: @@ -1090,14 +258,5 @@ TVM_REGISTER_GLOBAL("runtime.CSourceModuleCreate") return CSourceModuleCreate(code, fmt, func_names, const_vars); }); -TVM_REGISTER_GLOBAL("runtime.CreateCSourceCrtMetadataModule") - .set_body_typed([](const Array& modules, Target target, - relay::Runtime runtime) { - // Note that we don't need metadata when we compile a single operator - return CreateCSourceCrtMetadataModule(modules, target, runtime, - relay::backend::ExecutorCodegenMetadata(), - runtime::metadata::Metadata()); - }); - } // namespace codegen } // namespace tvm diff --git a/src/target/source/source_module.h b/src/target/source/source_module.h deleted file mode 100644 index e01445ce2ca5..000000000000 --- a/src/target/source/source_module.h +++ /dev/null @@ -1,63 +0,0 @@ -/* - * 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 source_module.h - * \brief Source code module - */ - -#ifndef TVM_TARGET_SOURCE_SOURCE_MODULE_H_ -#define TVM_TARGET_SOURCE_SOURCE_MODULE_H_ - -#include -#include -#include -#include - -#include "../../relay/backend/utils.h" -#include "../../runtime/meta_data.h" - -namespace tvm { -namespace codegen { - -/*! - - * \brief Wrap the submodules that are to be wrapped in a c-source metadata module for C runtime. - * \param modules The modules to be wrapped. - * \param target the target the modules are compiled for. - * \param runtime the runtime to code generate against - * \param metadata Compiler-generated metadata exported to runtime. - * \param aot_metadata If supplied, metadata for the AOTExecutor module. - * \return The wrapped module. - */ -runtime::Module CreateCSourceCrtMetadataModule(const Array& modules, Target target, - relay::Runtime runtime, - relay::backend::ExecutorCodegenMetadata metadata, - runtime::metadata::Metadata aot_metadata); - -/*! - * \brief Create C++-runtime targeted metadata module for "c" backend. - * \param metadata Compiler-generated metadata. - */ -runtime::Module CreateCSourceCppMetadataModule(runtime::metadata::Metadata metadata); - -} // namespace codegen -} // namespace tvm - -#endif // TVM_TARGET_SOURCE_SOURCE_MODULE_H_ diff --git a/src/target/target.cc b/src/target/target.cc index a8337b58ae9b..84c8bc3126cd 100644 --- a/src/target/target.cc +++ b/src/target/target.cc @@ -613,20 +613,6 @@ Target::Target(TargetKind kind, Optional host, String tag, Array is_external_codegen_map = - TargetKind::GetAttrMap(tvm::attr::kIsExternalCodegen); - TargetKindAttrMap relay_to_tir_map = - TargetKind::GetAttrMap(tvm::attr::kRelayToTIR); - return is_external_codegen_map.get(get()->kind, Bool(false)) || - relay_to_tir_map.count(get()->kind); -} - -bool Target::IsExternalCodegenFor(const Target& that) const { - return get()->GetTargetDeviceType() == that->GetTargetDeviceType() && IsExternalCodegen() && - !that.IsExternalCodegen(); -} - std::vector TargetNode::GetKeys() const { std::vector result; for (auto& expr : keys) { diff --git a/src/target/target_kind.cc b/src/target/target_kind.cc index e12c18e5ac73..3605831da1c7 100644 --- a/src/target/target_kind.cc +++ b/src/target/target_kind.cc @@ -288,7 +288,6 @@ TVM_REGISTER_TARGET_KIND("llvm", kDLCPU) .set_default_keys({"cpu"}) // Force the external codegen kind attribute to be registered, even if no external // codegen targets are enabled by the TVM build. - .set_attr(tvm::attr::kIsExternalCodegen, runtime::Bool(false)) .set_target_parser(tvm::target::parsers::cpu::ParseTarget); // Note regarding the "cl-opt" attribute: diff --git a/src/tir/analysis/calculate_allocated_memory.cc b/src/tir/analysis/calculate_allocated_memory.cc index 70e82a605369..04d3fc729c71 100644 --- a/src/tir/analysis/calculate_allocated_memory.cc +++ b/src/tir/analysis/calculate_allocated_memory.cc @@ -28,7 +28,6 @@ #include #include #include -#include #include #include diff --git a/src/tir/analysis/calculate_workspace.cc b/src/tir/analysis/calculate_workspace.cc deleted file mode 100644 index a667e2354b9b..000000000000 --- a/src/tir/analysis/calculate_workspace.cc +++ /dev/null @@ -1,97 +0,0 @@ -/* - * 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 tir/analysis/calculate_workspace.cc - * \brief Calculate any intermediary memory required by PrimFuncs. - */ -#include -#include -#include -#include -#include -#include - -namespace tvm { -namespace tir { - -template -class WorkspaceCalculator : public StmtExprVisitor { - public: - WorkspaceCalculator() = default; - size_t operator()(const PrimFunc& func); - size_t byte_alignment = tvm::runtime::kDefaultWorkspaceAlignment; - - private: - void VisitStmt_(const T* op) override; - size_t GetByteAlignedSize(Integer non_aligned_size); - size_t CalculateExtentsSize(const DataType& dtype, const Array& extents); - size_t current_size = 0; - size_t max_size = 0; -}; - -template -size_t WorkspaceCalculator::operator()(const PrimFunc& func) { - this->VisitStmt(func->body); - return this->max_size; -} - -template -size_t WorkspaceCalculator::GetByteAlignedSize(Integer non_aligned_size) { - return non_aligned_size.defined() - ? ((non_aligned_size.IntValue() + byte_alignment - 1) / byte_alignment) * - byte_alignment - : 0; -} - -template -void WorkspaceCalculator::VisitStmt_(const T* op) { - auto size = GetByteAlignedSize(usmp::CalculateExtentsSize(op)); - current_size += size; - if (current_size > max_size) { - max_size = current_size; - } - StmtExprVisitor::VisitStmt(op->body); - current_size -= size; -} - -size_t CalculateConstantBytes(const PrimFunc& func, const Integer& byte_alignment) { - WorkspaceCalculator wc; - wc.byte_alignment = byte_alignment->value; - return wc(func); -} - -size_t CalculateWorkspaceBytes(const PrimFunc& func, const Integer& byte_alignment) { - WorkspaceCalculator wc; - wc.byte_alignment = byte_alignment->value; - return wc(func); -} - -TVM_REGISTER_GLOBAL("tir.analysis.calculate_constant_bytes") - .set_body_typed([](PrimFunc func, Integer constant_byte_alignment) { - return static_cast(CalculateConstantBytes(func, constant_byte_alignment)); - }); - -TVM_REGISTER_GLOBAL("tir.analysis.calculate_workspace_bytes") - .set_body_typed([](PrimFunc func, Integer workspace_byte_alignment) { - return static_cast(CalculateWorkspaceBytes(func, workspace_byte_alignment)); - }); - -} // namespace tir -} // namespace tvm diff --git a/src/tir/analysis/device_constraint_utils.cc b/src/tir/analysis/device_constraint_utils.cc deleted file mode 100644 index 40df8b65c295..000000000000 --- a/src/tir/analysis/device_constraint_utils.cc +++ /dev/null @@ -1,498 +0,0 @@ -/* - * 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 tir/analysis/apply_device_constraints.cc - * \brief Applies device-related constraints to \p PrimFunc parameters. - * - * This is used by the \p PlanDevices pass to flow device-constraints *into* \p PrimFuncs. - * - * Currently only applies memory scope constraints into \p Buffer data pointer - * storage scopes. Aliased ('matched') buffers take on any scope introduced on - * the buffer they alias. However currently does not attempt to flow constraints into - * allocated buffers. - */ - -#include "./device_constraint_utils.h" - -#include -#include -#include -#include - -namespace tvm { -namespace tir { -namespace { - -/*! - * \brief Returns the \p PointerTypeNode for \p buffer, or nullptr if \p buffer does not describe a - * pointer. - */ -const PointerTypeNode* PointerInBuffer(const tir::Buffer& buffer) { - return buffer->data->type_annotation.defined() - ? buffer->data->type_annotation.as() - : nullptr; -} - -/*! - * \brief Returns the parameter variable and corresponding buffer at or after \p - * *current_primfunc_param_index in \p prim_func. Will skip over any non-pointer parameters. This - * can be used to find the parameter matching a tensor type in a flattened Relay function parameter - * or result. - */ -std::pair FindPointerParam(const tir::PrimFunc& prim_func, - size_t* current_primfunc_param_index) { - while (true) { - ICHECK_LT(*current_primfunc_param_index, prim_func->params.size()); - const tir::Var& param = prim_func->params[*current_primfunc_param_index]; - auto itr = prim_func->buffer_map.find(param); - if (itr == prim_func->buffer_map.end()) { - VLOG(2) << "no buffer map entry for '" << param->name_hint << "'"; - ++*current_primfunc_param_index; - continue; - } - const auto* pointer_type_node = PointerInBuffer((*itr).second); - if (pointer_type_node == nullptr) { - VLOG(2) << "not a pointer type for '" << param->name_hint << "'"; - ++*current_primfunc_param_index; - continue; - } - VLOG(2) << "using PrimFunc param '" << param->name_hint << "'"; - return *itr; - } -} - -/*! - * \brief Check fails if any parameter at or after \p *current_primfunc_param_index in \p prim_func - * is for a pointer type. This can be used to check all \p prim_func parameters have been accounted - * for when using \p FindPointerParam above. - */ -void CheckNoRemainingPointerParams(const tir::PrimFunc& prim_func, - size_t* current_primfunc_param_index) { - while (*current_primfunc_param_index < prim_func->params.size()) { - const tir::Var& param = prim_func->params[*current_primfunc_param_index]; - auto itr = prim_func->buffer_map.find(param); - if (itr == prim_func->buffer_map.end()) { - VLOG(1) << "no buffer map entry for '" << param->name_hint << "'"; - ++*current_primfunc_param_index; - continue; - } - const auto* pointer_type_node = PointerInBuffer((*itr).second); - ICHECK(pointer_type_node == nullptr); - ++*current_primfunc_param_index; - } -} - -/*! - * \brief Returns the (consistent) constraint to use for a Relay parameter of \p type, - * using \p prim_func parameters at or after \p *current_primfunc_param_index. Currently - * only memory scope is extracted. Fails if constraints are not consistent, ie \p type is a tuple - * type and the \p prim_func is attempting to map different fields of that tuple to different memory - * scopes. Returns the fully unconstrained \p VirtualDevice if no memory scopes constraints arise - * from the \p prim_func, ie all storage scope strings in pointer types are empty. - */ -VirtualDevice ConsistentParamConstraint(const tir::PrimFunc& prim_func, const Type& type, - size_t* current_primfunc_param_index) { - std::string memory_scope; // default empty => no constraint - for (size_t i = 0; i < relay::FlattenTupleType(type).size(); ++i) { - std::pair kv = FindPointerParam(prim_func, current_primfunc_param_index); - const tir::Buffer& buffer = kv.second; - const auto* pointer_type_node = buffer->data->type_annotation.as(); - const MemoryScope& buffer_memory_scope = pointer_type_node->storage_scope; - if (memory_scope.empty()) { - memory_scope = buffer_memory_scope; - } else if (buffer_memory_scope.empty()) { - // No constraint. - } else { - // Tuples must be homogenous on their VirtualDevice and thus memory scope. - ICHECK_EQ(buffer_memory_scope, memory_scope); - } - ++*current_primfunc_param_index; - } - return VirtualDevice::ForMemoryScope(memory_scope); -} - -/*! - * \brief Insert into param_constraints an entry for each parameter of \p prim_func starting from - * \p *current_primfunc_param_index for the flattened form of a Rleay parameters of \p type. Each - * entry maps to \p virtual_device. - */ -void InsertParamConstraints( - const tir::PrimFunc& prim_func, const Type& type, const VirtualDevice& virtual_device, - size_t* current_primfunc_param_index, - std::unordered_map* param_constraints) { - for (size_t i = 0; i < relay::FlattenTupleType(type).size(); ++i) { - std::pair kv = FindPointerParam(prim_func, current_primfunc_param_index); - param_constraints->emplace(kv.first.get(), virtual_device); - ++*current_primfunc_param_index; - } -} - -/*! - * \brief Apply the memory scope constraints to the \p Buffers and data \p Vars of a \p PrimFunc. - * - * All definitional occurrences of buffer Vars are rewritten to capture memory scopes in their - * PointerTypes: - * - Buffer::data (if the buffer itself is a definitional occurrence) - * - AllocateNode::buffer_var - * - FUTURE: LetStmtNode::var if aliasing a buffer data var. - * - * All referential occurrences of buffer Vars are replaced with their new definitions: - * - LoadNode::buffer_var - * - StoreNode::buffer_var - * - * Similarly all definitional occurrences of Buffers are rewritten to account for any new memory - * scopes: - * - PrimFuncNode::buffer_map keys. - * - BlockNode::match_buffers.buffer - * - FUTURE: BlockNode::alloc_buffers? - * - * And all referential occurrences of Buffers are replaced with their new definitions: - * - BufferLoadNode::buffer - * - BufferStoreNode::buffer - * - BufferRealizeNode::buffer - * - PrefetchNode::buffer - * - BufferRegionNode:buffer - * - BlockNode.match_buffers.source.buffer - * - BlockNode::{reads, writes}.buffer - * - * CAUTION: We assume strict sharing of Buffer objects and do not attempt to rewrite the bodies - * of referential buffers. - * - * CAUTION: EXPERIMENTAL: We don't yet account for all buffers and pointer types. - */ -class ApplyDeviceConstraintsMutator : public StmtExprMutator { - public: - ApplyDeviceConstraintsMutator() = default; - - /*! - * \brief Returns \p prim_func written to capture the memory scope constraints in \p - * param_constraints for each pointer \p prim_func parameter. Returns \p prim_func unchanged if no - * memory scopes needed to change. - */ - PrimFunc Rewrite(const PrimFunc& prim_func, const FuncType& relay_func_type, - const Array& arg_and_result_virtual_devices) { - size_t current_primfunc_param_index = 0; - std::unordered_map param_constraints; - - // For each Relay function parameter... - for (size_t i = 0; i < relay_func_type->arg_types.size(); ++i) { - const Type& param_type = relay_func_type->arg_types[i]; - const VirtualDevice& param_virtual_device = arg_and_result_virtual_devices[i]; - InsertParamConstraints(prim_func, param_type, param_virtual_device, - ¤t_primfunc_param_index, ¶m_constraints); - } - - // For the Relay function result... - const Type& ret_type = relay_func_type->ret_type; - const VirtualDevice& ret_virtual_device = arg_and_result_virtual_devices.back(); - InsertParamConstraints(prim_func, ret_type, ret_virtual_device, ¤t_primfunc_param_index, - ¶m_constraints); - - // Make sure we accounted for all prim_func parameters. - CheckNoRemainingPointerParams(prim_func, ¤t_primfunc_param_index); - - // Start with a copy of the current prim_func buffer map. - Map new_buffer_map(prim_func->buffer_map.begin(), prim_func->buffer_map.end()); - bool any_change = false; - - // For each constrained parameter... - for (const auto& kv : param_constraints) { - const tir::Var param = GetRef(kv.first); - const VirtualDevice& virtual_device = kv.second; - const tir::Buffer& buffer = prim_func->buffer_map[param]; - // Rewrite the buffer to account for constraint. - const Buffer new_buffer = RewriteBuffer(buffer, virtual_device); - if (!new_buffer.same_as(buffer)) { - any_change = true; - } - new_buffer_map.Set(param, new_buffer); - } - // Make sure we have accounted for all prim_func parameters. - CheckNoRemainingPointerParams(prim_func, ¤t_primfunc_param_index); - - // Apply data variable and buffer substitutions to the prim_func body. These will have been - // accumulated from processing the parameters above. - Stmt new_body = VisitStmt(prim_func->body); - if (!new_body.same_as(prim_func->body)) { - any_change = true; - } - - // We are done with the substitutions. - var_subst_.clear(); - buffer_subst_.clear(); - - if (any_change) { - return PrimFunc(prim_func->params, std::move(new_body), prim_func->ret_type, - std::move(new_buffer_map), prim_func->attrs, prim_func->span); - } else { - return prim_func; - } - } - - private: - PrimExpr VisitExpr_(const VarNode* var_node) final { return Subst(var_node); } - - PrimExpr VisitExpr_(const BufferLoadNode* buffer_load_node) final { - BufferLoad new_buffer_load = - Downcast(StmtExprMutator::VisitExpr_(buffer_load_node)); - Buffer new_buffer = Subst(new_buffer_load->buffer.get()); - if (!new_buffer.same_as(new_buffer_load->buffer)) { - return BufferLoad(new_buffer, new_buffer_load->indices, new_buffer_load->predicate, - new_buffer_load->span); - } - return std::move(new_buffer_load); - } - - Stmt VisitStmt_(const LetStmtNode* let_stmt_node) final { - // TODO(mbs): If the let-bound var is aliasing an existing buffer data var we need to - // rewrite it. - return StmtExprMutator::VisitStmt_(let_stmt_node); - } - - Stmt VisitStmt_(const AttrStmtNode* attr_stmt_node) final { - AttrStmt new_attr_stmt = Downcast(StmtExprMutator::VisitStmt_(attr_stmt_node)); - // remap node if a var - if (const auto* var_node = new_attr_stmt->node.as()) { - Var new_var = Subst(var_node); - if (!new_var.same_as(new_attr_stmt->node)) { - return AttrStmt(new_var, new_attr_stmt->attr_key, new_attr_stmt->value, - new_attr_stmt->body); - } - } - return std::move(new_attr_stmt); - } - - // ForNode default ok since loop_var never of PointerType - - // WhileNode default ok - - Stmt VisitStmt_(const AllocateNode* allocate_node) final { - // TODO(mbs): What memory scope should we assign to the new pointer? - return StmtExprMutator::VisitStmt_(allocate_node); - } - - Stmt VisitStmt_(const BufferStoreNode* buffer_store_node) final { - BufferStore new_buffer_store = - Downcast(StmtExprMutator::VisitStmt_(buffer_store_node)); - Buffer new_buffer = Subst(new_buffer_store->buffer.get()); - if (!new_buffer.same_as(new_buffer_store->buffer)) { - return BufferStore(new_buffer, new_buffer_store->value, new_buffer_store->indices, - new_buffer_store->predicate, new_buffer_store->span); - } - return std::move(new_buffer_store); - } - - Stmt VisitStmt_(const BufferRealizeNode* buffer_realize_node) final { - BufferRealize new_buffer_realize = - Downcast(StmtExprMutator::VisitStmt_(buffer_realize_node)); - Buffer new_buffer = Subst(new_buffer_realize->buffer.get()); - if (!new_buffer.same_as(new_buffer_realize->buffer)) { - return BufferRealize(new_buffer, new_buffer_realize->bounds, new_buffer_realize->condition, - new_buffer_realize->body, new_buffer_realize->span); - } - return std::move(new_buffer_realize); - } - - // IfThenElseNode default ok - // AssertStmtNode default ok - // ProducerStoreNode default ok (though does not visit producer) - // ProducerRealizeNode default ok (though does not visit producer) - - Stmt VisitStmt_(const PrefetchNode* prefetch_node) final { - Prefetch new_prefetch = Downcast(StmtExprMutator::VisitStmt_(prefetch_node)); - Buffer new_buffer = Subst(new_prefetch->buffer.get()); - if (!new_buffer.same_as(new_prefetch->buffer)) { - return Prefetch(new_buffer, prefetch_node->bounds, prefetch_node->span); - } - return std::move(new_prefetch); - } - - // SeqStmtNode default ok - // EvaluateNode default ok - - BufferRegion VisitItem(const BufferRegionNode* buffer_region_node) { - Buffer new_buffer = Subst(buffer_region_node->buffer.get()); - if (!new_buffer.same_as(buffer_region_node->buffer)) { - return BufferRegion(new_buffer, buffer_region_node->region); - } - return GetRef(buffer_region_node); - } - - MatchBufferRegion VisitItem(const MatchBufferRegionNode* match_buffer_region_node) { - // The source field has a referential occurrence of the buffer. Apply the buffer substitution - // to that. - BufferRegion new_source = VisitItem(match_buffer_region_node->source.get()); - // The buffer field however is a definitional occurrence, aliased on top of the source. - // Transfer any memory scope from the source to the destination. - Optional opt_virtual_device = GetBufferConstraint(new_source->buffer); - tir::Buffer new_buffer; - if (opt_virtual_device.defined()) { - new_buffer = RewriteBuffer(match_buffer_region_node->buffer, opt_virtual_device.value()); - } else { - new_buffer = match_buffer_region_node->buffer; - } - if (!new_buffer.same_as(match_buffer_region_node->buffer) || - !new_source.same_as(match_buffer_region_node->source)) { - return MatchBufferRegion(new_buffer, new_source); - } - return GetRef(match_buffer_region_node); - } - - template - Array VisitItems(const Array& items) { - return items.Map([this](T item) -> T { return VisitItem(item.get()); }); - } - - Stmt VisitStmt_(const BlockNode* block_node) final { - Block new_block = Downcast(StmtExprMutator::VisitStmt_(block_node)); - Array new_reads = VisitItems(new_block->reads); - Array new_writes = VisitItems(new_block->writes); - // TODO(mbs): What memory scope should we assign to the new buffers? - Array new_match_buffers = VisitItems(new_block->match_buffers); - if (!new_reads.same_as(new_block->reads) || new_writes.same_as(new_block->writes) || - new_match_buffers.same_as(new_block->match_buffers)) { - return Block(new_block->iter_vars, std::move(new_reads), std::move(new_writes), - new_block->name_hint, new_block->body, new_block->init, new_block->alloc_buffers, - std::move(new_match_buffers), new_block->annotations, new_block->span); - } - return std::move(new_block); - } - - // BlockRealizeNode default ok - - /*! Applies \p var_subst_ substitution to \p var_node. */ - Var Subst(const VarNode* var_node) const { - auto itr = var_subst_.find(var_node); - return itr == var_subst_.end() ? GetRef(var_node) : itr->second; - } - - /*! Applies \p buffer_subst_ substitution to \p buffer. */ - Buffer Subst(const BufferNode* buffer_node) const { - auto itr = buffer_subst_.find(buffer_node); - return itr == buffer_subst_.end() ? GetRef(buffer_node) : itr->second; - } - - /*! - * \brief Rewrites \p buffer so as to follow the constraints in \p virtual_device - * (currently just memory scope). - * - * Updates both the var_subst_ and buffer_subst_ to capture the rewrite, but - * also returns the new buffer. - */ - Buffer RewriteBuffer(const Buffer& buffer, const VirtualDevice& virtual_device) { - ICHECK(buffer->data->type_annotation.defined()); - const auto* pointer_type_node = buffer->data->type_annotation.as(); - ICHECK(pointer_type_node); - if (pointer_type_node->storage_scope == virtual_device->memory_scope) { - // No change. - return buffer; - } - PointerType new_pointer_type(pointer_type_node->element_type, virtual_device->memory_scope); - Var new_data(buffer->data->name_hint, new_pointer_type, buffer->data->span); - var_subst_.emplace(buffer->data.get(), new_data); - Buffer new_buffer = buffer; - new_buffer.CopyOnWrite()->data = new_data; - buffer_subst_.emplace(buffer.get(), new_buffer); - return new_buffer; - } - - /*! - * \brief Returns the VirtualDevice capturing any memory scope in \p buffer. Returns nullptr if - * buffer's data var does not have a type annotation of \p PointerType. Returns the fully - * unconstrained \p VirtualDevice if no memory scope is given. - */ - static Optional GetBufferConstraint(const tir::Buffer& buffer) { - const auto* pointer_type_node = PointerInBuffer(buffer); - return pointer_type_node == nullptr - ? Optional() - : VirtualDevice::ForMemoryScope(pointer_type_node->storage_scope); - } - - /*! - * \brief Maps each \p Buffer::data \p Var to its constrained equivalent. - */ - std::unordered_map var_subst_; - - /*! - * \brief Maps each \p Buffer to its constrained equivalent. - */ - std::unordered_map buffer_subst_; -}; - -} // namespace - -Array GetPrimFuncArgAndResultConstraints(const tir::PrimFunc& prim_func, - const FuncType& relay_func_type) { - // Build the implied domain (in terms of the function's Relay type) implied by any memory scope - // constrains in the function's buffers, for both arguments and results. - Array virtual_devices; - virtual_devices.reserve(relay_func_type->arg_types.size() + 1); - - // For each Relay function parameter... - size_t current_primfunc_param_index = 0; - for (const auto& param_type : relay_func_type->arg_types) { - VirtualDevice param_virtual_device = - ConsistentParamConstraint(prim_func, param_type, ¤t_primfunc_param_index); - virtual_devices.push_back(param_virtual_device); - } - - // For the Relay function result... - const Type& ret_type = relay_func_type->ret_type; - VirtualDevice ret_virtual_device = - ConsistentParamConstraint(prim_func, ret_type, ¤t_primfunc_param_index); - virtual_devices.push_back(ret_virtual_device); - - // Make sure all parameters of the prim_func have been accounted for. - CheckNoRemainingPointerParams(prim_func, ¤t_primfunc_param_index); - - return virtual_devices; -} - -TVM_REGISTER_GLOBAL("tir.analysis.GetPrimFuncArgAndResultMemoryConstraints") - .set_body_typed([](const PrimFunc& prim_func, const FuncType& relay_func_type) { - Array memory_scopes; - memory_scopes.reserve(relay_func_type->type_params.size() + 1); - for (const auto& virtual_device : - GetPrimFuncArgAndResultConstraints(prim_func, relay_func_type)) { - memory_scopes.push_back(virtual_device->memory_scope); - } - return memory_scopes; - }); - -PrimFunc ApplyPrimFuncArgAndResultConstraints( - const PrimFunc& prim_func, const FuncType& relay_func_type, - const Array& arg_and_result_virtual_devices) { - return ApplyDeviceConstraintsMutator().Rewrite(prim_func, relay_func_type, - arg_and_result_virtual_devices); -} - -TVM_REGISTER_GLOBAL("tir.analysis.ApplyPrimFuncArgAndResultMemoryConstraints") - .set_body_typed([](const PrimFunc& prim_func, const FuncType& relay_func_type, - const Array& arg_and_result_memory_scopes) { - Array virtual_devices; - virtual_devices.reserve(arg_and_result_memory_scopes.size()); - for (const auto& memory_scope : arg_and_result_memory_scopes) { - virtual_devices.push_back(VirtualDevice::ForMemoryScope(memory_scope)); - } - return ApplyPrimFuncArgAndResultConstraints(prim_func, relay_func_type, virtual_devices); - }); - -} // namespace tir -} // namespace tvm diff --git a/src/tir/analysis/device_constraint_utils.h b/src/tir/analysis/device_constraint_utils.h deleted file mode 100644 index 717bf5280c00..000000000000 --- a/src/tir/analysis/device_constraint_utils.h +++ /dev/null @@ -1,98 +0,0 @@ -/* - * 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 tir/analysis/device_constraint_utils.cc - * \brief Utilities for extracting and applying device-related constraints to \p PrimFunc - * parameters. - * - * These utilities are used by the \p PlanDevices pass to extract memory (aka 'storage') scope - * information from \p PrimFuncs and convert them back into \p VirtualDevice form w.r.t. the - * original Relay type of the \p PrimFunc (ie before flattening of tuple arguments/results and - * conversion to destination-passing style aka DPS). - * - * A utility is also supplied to go the other way: impose memory scopes on \p PrimFunc parameters. - * However that's still in EXPERIMENTAL form. - * - * We may extend these utilities to also gather/apply layout information should we add that to - * \p VirtualDevice. - */ - -#ifndef TVM_TIR_ANALYSIS_DEVICE_CONSTRAINT_UTILS_H_ -#define TVM_TIR_ANALYSIS_DEVICE_CONSTRAINT_UTILS_H_ - -#include -#include - -namespace tvm { -namespace tir { - -/*! - * A Relay Function with type: - * \code - * fn((Tensor[...], Tensor[...]), Tensor[...]) -> (Tensor[...], Tensor[...]) - * ^ ^ ^ ^ ^ - * a b c d e - * \endcode - * will be represented by a TIR PrimFunc in flattened and DPS form with at least 5 argument a..e. - * \code - * primfn(a: handle, b: handle, c: handle, d: handle, e: handle) { - * buffers = { ... } - * buffer_map = { ... } - * ... - * } - * \endcode - * - * Each such PrimFunc argument will me mapped to a \p Buffer who's underlying \p data \p Var - * has a \p PointerType. - * - * The PrimFunc may have additional non-pointer arguments, eg for: - * - scalar inputs and tensor dimensions - * - device contexts - * Those should be ignored here since they have no counterpart in the Relay Function. - * - * We'll need helpers to map on-the-fly between the Relay and TIR view of functions. - */ - -/*! - * \brief Returns the \p VirtualDevices capturing the memory (aka storage) scope constraints for all - * the arguments and result of \p prim_func. However the result will be w.r.t. the \p prim_func's - * representation as a Relay \p Function of \p relay_func_type_ before lowering and conversion to - * DPS. - */ -Array GetPrimFuncArgAndResultConstraints(const tir::PrimFunc& prim_func, - const FuncType& relay_func_type); - -/* - * \brief Returns \p prim_func written to capture the memory (aka storage) scope constraints - * for each of the \p prim_func's parameters given by \p arg_and_result_virtual_devices. However, - * \p arg_and_result_virtual_devices should be w.r.t. the \p prim_func's representation as a Relay - * \p Function of \p relay_func_type before lowering and conversion to DPS. - * - * CAUTION: This is experimental. The resulting \p PrimFunc may not have fully accounted for all - * new memory scopes. - */ -PrimFunc ApplyPrimFuncArgAndResultConstraints( - const PrimFunc& prim_func, const FuncType& relay_func_type, - const Array& arg_and_result_virtual_devices); - -} // namespace tir -} // namespace tvm - -#endif // TVM_TIR_ANALYSIS_DEVICE_CONSTRAINT_UTILS_H_ diff --git a/src/tir/ir/function.cc b/src/tir/ir/function.cc index 2c94b9d8646b..509a53f9ea98 100644 --- a/src/tir/ir/function.cc +++ b/src/tir/ir/function.cc @@ -105,7 +105,7 @@ FuncType PrimFuncNode::func_type_annotation() const { for (auto param : this->params) { param_types.push_back(GetType(param)); } - return FuncType(param_types, ret_type, {}, {}); + return FuncType(param_types, ret_type); } TVM_REGISTER_NODE_TYPE(PrimFuncNode); diff --git a/src/tir/transforms/bind_params.cc b/src/tir/transforms/bind_params.cc index 0b71b2e8fa34..66d0fb61661b 100644 --- a/src/tir/transforms/bind_params.cc +++ b/src/tir/transforms/bind_params.cc @@ -24,7 +24,6 @@ */ #include #include -#include #include #include #include @@ -34,11 +33,6 @@ #include #include -#include -#include -#include - -#include "../../runtime/thread_storage_scope.h" #include "ir_utils.h" namespace tvm { diff --git a/src/tir/transforms/default_gpu_schedule.cc b/src/tir/transforms/default_gpu_schedule.cc index 6d0542257309..ea521d696836 100644 --- a/src/tir/transforms/default_gpu_schedule.cc +++ b/src/tir/transforms/default_gpu_schedule.cc @@ -94,12 +94,10 @@ IRModule MarkScheduled(const IRModule& mod) { } } - return IRModule(result, // functions - mod->type_definitions, // type_definitions - mod->import_set_, // import_set - mod->source_map, // map - mod->attrs, // attrs - mod->global_infos); // global_infos + return IRModule(result, // functions + mod->source_map, // map + mod->attrs, // attrs + mod->global_infos); // global_infos } bool IsScheduledOnGPU(const BaseFunc& func) { diff --git a/src/tir/transforms/install_debug_spans.cc b/src/tir/transforms/install_debug_spans.cc deleted file mode 100644 index ea61378ccccc..000000000000 --- a/src/tir/transforms/install_debug_spans.cc +++ /dev/null @@ -1,162 +0,0 @@ -/* - * 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 install_debug_spans.cc - * \brief Prints TIR code in memory and replaces all spans in the module with - the location to which the ops would be printed - */ - -#include "./install_debug_spans.h" - -#include - -#include -#include - -#include "../../relay/printer/tir_text_printer_debug.h" -#include "ir_utils.h" - -namespace tvm { -namespace tir { - -Stmt DebugInfoInstaller::InstallInfo(const std::string& name, const Stmt& stmt) { - DebugInfoInstaller installer(stmt, name + ".tir"); - return installer.VisitStmt(stmt); -} - -DebugInfoInstaller::DebugInfoInstaller(const Stmt& stmt, const std::string& filename) { - // Determine the line that each stmt/expr will be printed on - tvm::relay::TIRTextPrinterDebug printer(false); - - // Fill in the stmts and exprs' line info - auto result = printer.Print(stmt).str(); - - // Create map of the stmt/expr -> its line number in the output to later - // create new spans for each stmt/expr - const auto& stmts = printer.GetStmtsByLine(); - VLOG(0) << "Debug printer found " << stmts.size() << " stmts after printing"; - for (const auto& line : stmts) { - stmt_lines_[std::get<0>(line)] = std::get<1>(line); - } - - const auto& exprs = printer.GetExprsByLine(); - VLOG(0) << "Debug printer found " << exprs.size() << " exprs after printing"; - for (const auto& line : exprs) { - expr_lines_[std::get<0>(line)] = std::get<1>(line); - } - - // Output the printed TIR to the specified file - VLOG(0) << "Outputting TIR to " << filename; - filename_ = std::move(filename); - std::ofstream out(filename_); - out << result; - out.close(); -} - -PrimExpr DebugInfoInstaller::VisitExpr(const PrimExpr& expr) { - PrimExpr result = expr; - result = StmtExprMutator::VisitExpr(result); - return result; -} - -Stmt DebugInfoInstaller::VisitStmt(const Stmt& stmt) { - Stmt result = stmt; - result = StmtExprMutator::VisitStmt(result); - return result; -} - -Span DebugInfoInstaller::MaybeSpan(const StmtNode* op) { - auto entry = stmt_lines_.find(op); - if (entry == stmt_lines_.end()) { - return Span(); - } else { - size_t column = 0; - size_t line = entry->second; - return Span(SourceName::Get(filename_), line, line, column, column); - } -} - -Span DebugInfoInstaller::MaybeSpan(const PrimExprNode* op) { - auto entry = expr_lines_.find(op); - if (entry == expr_lines_.end()) { - return Span(); - } else { - size_t column = 0; - size_t line = entry->second; - return Span(SourceName::Get(filename_), line, line, column, column); - } -} - -#define X(TypeName) \ - PrimExpr DebugInfoInstaller::VisitExpr_(const TypeName##Node* op) { \ - auto new_expr = StmtExprMutator::VisitExpr_(op); \ - auto new_type = Downcast(new_expr); \ - auto new_node = new_type.CopyOnWrite(); \ - new_node->span = MaybeSpan(op); \ - return new_type; \ - } -TVM_TIR_TRANSFORMS_INSTALL_DEBUG_SPANS_SUPPORTED_EXPRS -#undef X - -#define X(TypeName) \ - Stmt DebugInfoInstaller::VisitStmt_(const TypeName##Node* op) { \ - Stmt new_stmt = StmtExprMutator::VisitStmt_(op); \ - auto new_type = Downcast(new_stmt); \ - auto new_node = new_type.CopyOnWrite(); \ - new_node->span = MaybeSpan(op); \ - return new_type; \ - } -TVM_TIR_TRANSFORMS_INSTALL_DEBUG_SPANS_SUPPORTED_STMTS -#undef X - -namespace transform { - -Pass InstallDebugSpans() { - auto pass_func = [](IRModule mod, PassContext ctx) { - Map external_host_functions; - for (const auto& [gvar, base_func] : mod->functions) { - if (auto opt = base_func.as()) { - auto prim_func = opt.value(); - if (IsHostFunc(prim_func).value_or(false) && - prim_func->GetAttr(tvm::attr::kGlobalSymbol)) { - external_host_functions.Set(gvar, prim_func); - } - } - } - - ICHECK_EQ(external_host_functions.size(), 1) - << "Debug info can only be added to IRModules with a single host function"; - - for (auto [gvar, prim_func] : external_host_functions) { - auto name = prim_func->GetAttr(tvm::attr::kGlobalSymbol).value(); - prim_func.CopyOnWrite()->body = DebugInfoInstaller::InstallInfo(name, prim_func->body); - mod.CopyOnWrite()->Update(gvar, prim_func); - } - - return mod; - }; - return tvm::transform::CreateModulePass(pass_func, 0, "tir.InstallDebugSpans", {}); -} - -TVM_REGISTER_GLOBAL("tir.transform.InstallDebugSpans").set_body_typed(InstallDebugSpans); - -} // namespace transform -} // namespace tir -} // namespace tvm diff --git a/src/tir/transforms/install_debug_spans.h b/src/tir/transforms/install_debug_spans.h deleted file mode 100644 index 40f3e07940cf..000000000000 --- a/src/tir/transforms/install_debug_spans.h +++ /dev/null @@ -1,131 +0,0 @@ -/* - * 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 install_debug_spans.h - * \brief Interface of the InstallDebugSpans pass - */ - -#ifndef TVM_TIR_TRANSFORMS_INSTALL_DEBUG_SPANS_H_ -#define TVM_TIR_TRANSFORMS_INSTALL_DEBUG_SPANS_H_ - -#include -#include -#include -#include - -#include -#include - -#ifndef TVM_TIR_TRANSFORMS_INSTALL_DEBUG_SPANS_OPS_H_ -#define TVM_TIR_TRANSFORMS_INSTALL_DEBUG_SPANS_OPS_H_ - -#define TVM_TIR_TRANSFORMS_INSTALL_DEBUG_SPANS_SUPPORTED_EXPRS \ - X(Call) \ - X(Add) \ - X(Sub) \ - X(Mul) \ - X(Div) \ - X(Mod) \ - X(FloorDiv) \ - X(FloorMod) \ - X(Min) \ - X(Max) \ - X(EQ) \ - X(NE) \ - X(LT) \ - X(LE) \ - X(GT) \ - X(GE) \ - X(And) \ - X(Or) \ - X(Reduce) \ - X(Cast) \ - X(Not) \ - X(Select) \ - X(Ramp) \ - X(Broadcast) \ - X(Shuffle) \ - X(IntImm) \ - X(FloatImm) \ - X(StringImm) - -#define TVM_TIR_TRANSFORMS_INSTALL_DEBUG_SPANS_SUPPORTED_STMTS \ - X(AttrStmt) \ - X(IfThenElse) \ - X(LetStmt) \ - X(For) \ - X(While) \ - X(Allocate) \ - X(AllocateConst) \ - X(DeclBuffer) \ - X(BufferStore) \ - X(BufferRealize) \ - X(AssertStmt) \ - X(ProducerStore) \ - X(ProducerRealize) \ - X(Prefetch) \ - X(SeqStmt) \ - X(Evaluate) \ - X(BlockRealize) - -#endif // TVM_TIR_TRANSFORMS_INSTALL_DEBUG_SPANS_OPS_H_ - -namespace tvm { -namespace tir { - -/*! - * \brief This Pass prints out the provided 'stmt' through the TIR debug printer - while recording the statements and expressions printed on each line. Running - this pass uses the per-line information to change the Spans attached to each - statement and expression to the source location in the printed TIR. This pass - also writes to a file called '.tir' so the line information used is - saved to disk. - */ -class DebugInfoInstaller : public StmtExprMutator { - public: - static Stmt InstallInfo(const std::string& name, const Stmt& stmt); - - PrimExpr VisitExpr(const PrimExpr& expr) override; - Stmt VisitStmt(const Stmt& stmt) override; - - protected: - DebugInfoInstaller(const Stmt& stmt, const std::string& filename); - -#define X(TypeName) PrimExpr VisitExpr_(const TypeName##Node* op) override; - TVM_TIR_TRANSFORMS_INSTALL_DEBUG_SPANS_SUPPORTED_EXPRS -#undef X - -#define X(TypeName) Stmt VisitStmt_(const TypeName##Node* op) override; - TVM_TIR_TRANSFORMS_INSTALL_DEBUG_SPANS_SUPPORTED_STMTS -#undef X - - private: - std::unordered_map stmt_lines_; - std::unordered_map expr_lines_; - std::string filename_; - - Span MaybeSpan(const StmtNode* op); - Span MaybeSpan(const PrimExprNode* op); -}; - -} // namespace tir -} // namespace tvm - -#endif // TVM_TIR_TRANSFORMS_INSTALL_DEBUG_SPANS_H_ diff --git a/src/tir/transforms/primfunc_utils.cc b/src/tir/transforms/primfunc_utils.cc index 8a5317a3c84a..7f45fee9a26c 100644 --- a/src/tir/transforms/primfunc_utils.cc +++ b/src/tir/transforms/primfunc_utils.cc @@ -23,7 +23,6 @@ */ #include -#include #include namespace tvm { @@ -63,13 +62,6 @@ transform::Pass BindTarget(Target target) { transform::Pass AnnotateEntryFunc() { auto fpass = [](IRModule mod, transform::PassContext ctx) -> IRModule { - // AOT tracks the entry function, no annotation required - auto executor = mod->GetAttr("executor"); - const bool is_aot_executor = executor.defined() && executor.value()->name == "aot"; - if (is_aot_executor) { - return mod; - } - // If only a single function exists, that function must be the entry if (mod->functions.size() == 1) { auto [gvar, base_func] = *mod->functions.begin(); diff --git a/src/tir/transforms/using_assume_to_reduce_branches.cc b/src/tir/transforms/using_assume_to_reduce_branches.cc index 2e45bb0ff8fb..3cd33b85905b 100644 --- a/src/tir/transforms/using_assume_to_reduce_branches.cc +++ b/src/tir/transforms/using_assume_to_reduce_branches.cc @@ -36,19 +36,15 @@ */ #include -#include +#include #include #include #include #include #include -#include - #include "../../arith/constraint_extract.h" #include "../../arith/ir_mutator_with_analyzer.h" -#include "../../arith/unwrap_vector_expr.h" -#include "simplify.h" #include "tvm/ir/expr.h" namespace tvm { namespace tir { @@ -363,11 +359,11 @@ Pass UseAssumeToReduceBranches() { if (n->attrs.GetAttr("op_pattern").defined()) { Optional opt_pattern = f->GetAttr("op_pattern"); if (opt_pattern.defined()) { - relay::OpPatternKind pattern; - pattern = static_cast(Downcast(opt_pattern)->value); + relax::OpPatternKind pattern; + pattern = static_cast(Downcast(opt_pattern)->value); - if (pattern == relay::OpPatternKind::kElemWise || - pattern == relay::OpPatternKind::kBroadcast) { + if (pattern == relax::OpPatternKind::kElemWise || + pattern == relax::OpPatternKind::kBroadcast) { // If the primfunc contains assume statement then, run the mutator pass. AssumeChecker assume_checker; assume_checker(std::move(n->body)); diff --git a/src/tir/usmp/algo/greedy.cc b/src/tir/usmp/algo/greedy.cc deleted file mode 100644 index ec4f5a5d7215..000000000000 --- a/src/tir/usmp/algo/greedy.cc +++ /dev/null @@ -1,231 +0,0 @@ -/* - * 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 tir/analysis/usmp/algo/greedy.cc - * \brief This source contains greedy algorithms for planning - * memory for USMP. There are two algorithms present here : - * 1) greedy_by_size and 2) greedy_by_conflicts. - * - * greedy_by_size : this algorithm prioritizes placing the - * largest size buffer to the given pools. The BufferInfo objects - * are sorted based on the size and placed on each pool adhering - * to size_hint constraint. - * - * greedy_by_conflicts : this algorithm prioritizes placing the - * the most liveness conflicted buffer to the given pools. The - * BufferInfo objects are sorted based on the number of conflicts - * and placed on each pool adhering to size_hint constraint. - */ - -#include -#include -#include -#include -#include -#include -#include -#include - -namespace tvm { -namespace tir { -namespace usmp { -namespace algo { - -/*! - * \brief Rounds up the offset to satisfy the alignement requirement - */ -size_t GreedyBase::round_up_to_byte_alignment(const size_t& non_aligned_byte_offset, - const int& byte_alignment) { - return ((non_aligned_byte_offset + byte_alignment - 1) / byte_alignment) * byte_alignment; -} - -/*! - * \brief A helper function check whether a offset is valid given the constraints - */ -bool GreedyBase::IsValidPlacement(const PoolInfo& candidate_pool, const size_t& next_offset, - const size_t& size_bytes) { - Integer size_hint_bytes = -1; - if (const auto* p = candidate_pool.as()) { - size_hint_bytes = p->size_hint_bytes; - } else if (const auto* p = candidate_pool.as()) { - size_hint_bytes = p->size_hint_bytes; - } else { - LOG(FATAL) << "Pool '" << candidate_pool->GetTypeKey() << "' is not supported"; - } - - if (size_hint_bytes == kUnrestrictedPoolSizeHint) { - // this means pool is not bounded - return true; - } - auto pool_size = static_cast(size_hint_bytes.IntValue()); - auto max_address = next_offset + size_bytes; - if (max_address <= pool_size) { - return true; - } - return false; -} - -/*! - * \brief Selects a pool for placement in the given set of ordered pool candidates - */ -PoolInfo GreedyBase::SelectPlacementPool( - const BufferInfo& buf_info, - const std::unordered_map& pool_offsets) { - // Here the pool candidates are ordered when it is consumed by the algorithm. - // This could be from order the user has specified. However, schedulers are - // welcome to change the order for performance reasons. - for (const auto& pool_info : buf_info->pool_candidates) { - if (pool_offsets.count(pool_info)) { - return pool_info; - } - } - CHECK(false) << "TVM USMP Error: the space available in the provided pools exceeded when " - "trying to allocate the buffer : " - << buf_info << "\n. Please increase the size_hints for memory pools."; - return PoolInfo(); -} - -/*! - * \brief This is the base allocation function that works on sorted BufferInfo objects based - * on the greedy heuristic. The sorting algorithm has to be called before calling this. - */ -Map GreedyBase::PostSortAllocation( - const std::vector& buffer_info_vec) { - Map pool_allocations; - for (const auto& buf_info : buffer_info_vec) { - std::unordered_map pool_offset_candidates; - for (const auto& pool_info : buf_info->pool_candidates) { - // Mark pool candidates that satisfy the size constraints. - if (IsValidPlacement(pool_info, 0, buf_info->size_bytes->value)) { - pool_offset_candidates[pool_info] = 0; - } - } - - for (const auto& conflict_buf_info_obj : buf_info->conflicts) { - auto conflict_buf_info = Downcast(conflict_buf_info_obj); - size_t next_offset = 0; - // We only look at already allocated BufferInfo in-terms of conflicts. - if (pool_allocations.count(conflict_buf_info)) { - auto pool_allocation = pool_allocations[conflict_buf_info]; - next_offset = - pool_allocation->byte_offset.IntValue() + conflict_buf_info->size_bytes.IntValue(); - next_offset = round_up_to_byte_alignment(next_offset, conflict_buf_info->alignment->value); - // Checks whether the next offset in the same pool as the conflicting BufferInfo is valid. - if (IsValidPlacement(pool_allocation->pool_info, next_offset, - buf_info->size_bytes->value)) { - // There could be multiple conflicting BufferInfo in the same pool. - // Thus, we need to make sure we pick the largest offset of them all. - if (next_offset > pool_offset_candidates[pool_allocation->pool_info]) { - pool_offset_candidates[pool_allocation->pool_info] = next_offset; - } - } else { - pool_offset_candidates.erase(pool_allocation->pool_info); - } - } - } - auto selected_pool = SelectPlacementPool(buf_info, pool_offset_candidates); - pool_allocations.Set( - buf_info, PoolAllocation(selected_pool, Integer(pool_offset_candidates[selected_pool]))); - } - return pool_allocations; -} - -/*! - * \brief This class implements Greedy by the size of BufferInfo - * greedy algorithm. Please refer to main documentation of the file - * for more details. - */ -class GreedySize : public GreedyBase { - public: - GreedySize() {} - Map PlanMemory(const Array& buffer_info_arr) { - std::vector buffer_info_vec; - Map pool_allocations; - for (const auto& buffer_info : buffer_info_arr) { - buffer_info_vec.push_back(std::move(buffer_info)); - } - std::sort(buffer_info_vec.begin(), buffer_info_vec.end(), - [](const BufferInfo& a, const BufferInfo& b) { - if (a->size_bytes->value == b->size_bytes->value) { - if (a->conflicts.size() == b->conflicts.size()) { - return std::string(a->name_hint->data) > std::string(b->name_hint->data); - } else { - return a->conflicts.size() > b->conflicts.size(); - } - } - return a->size_bytes.IntValue() > b->size_bytes.IntValue(); - }); - return PostSortAllocation(buffer_info_vec); - } -}; - -/*! - * \brief This class implements Greedy by the number of conflicts of - * BufferInfo greedy algorithm. Please refer to main documentation - * of the file for more details. - */ -class GreedyConflicts : public GreedyBase { - public: - GreedyConflicts() {} - Map PlanMemory(const Array& buffer_info_arr) { - std::vector buffer_info_vec; - Map pool_allocations; - for (const auto& buffer_info : buffer_info_arr) { - buffer_info_vec.push_back(std::move(buffer_info)); - } - std::sort(buffer_info_vec.begin(), buffer_info_vec.end(), - [](const BufferInfo& a, const BufferInfo& b) { - if (a->conflicts.size() == b->conflicts.size()) { - if (a->size_bytes->value == b->size_bytes->value) { - return std::string(a->name_hint->data) > std::string(b->name_hint->data); - } else { - return a->size_bytes->value > b->size_bytes->value; - } - } - return a->conflicts.size() > b->conflicts.size(); - }); - return PostSortAllocation(buffer_info_vec); - } -}; - -Map GreedyBySize(const Array& buffer_info_arr, - const Integer& memory_pressure) { - return GreedySize().PlanMemory(buffer_info_arr); -} - -Map GreedyByConflicts(const Array& buffer_info_arr, - const Integer& memory_pressure) { - return GreedyConflicts().PlanMemory(buffer_info_arr); -} - -TVM_REGISTER_GLOBAL("tir.usmp.algo.greedy_by_size") - .set_body_typed([](Array buffer_info_arr, Integer memory_pressure) { - return GreedyBySize(buffer_info_arr, memory_pressure); - }); - -TVM_REGISTER_GLOBAL("tir.usmp.algo.greedy_by_conflicts") - .set_body_typed([](Array buffer_info_arr, Integer memory_pressure) { - return GreedyByConflicts(buffer_info_arr, memory_pressure); - }); - -} // namespace algo -} // namespace usmp -} // namespace tir -} // namespace tvm diff --git a/src/tir/usmp/algo/hill_climb.cc b/src/tir/usmp/algo/hill_climb.cc deleted file mode 100644 index 6e1de1e43cd3..000000000000 --- a/src/tir/usmp/algo/hill_climb.cc +++ /dev/null @@ -1,375 +0,0 @@ -/* - * 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 tir/analysis/usmp/algo/hill_climb.cc - * \brief Implement greedy by size memory planning algorithm - */ -#include -#include -#include -#include -#include -#include -#include - -#include -#include -#include - -namespace tvm { -namespace tir { -namespace usmp { -namespace algo { - -/* - * Simulated annealing / Hill climb - * - * Works by continiously invoking 'greedy-by-size' allocation, - * assessing the result, and introducing permutations to the allocation - * order which hopefully will led to more 'compact' memory allocation. - * Do not forget to use srand for repeatable results - */ -class HillClimbAllocator : public GreedyBase { - private: - size_t memory_pressure_ = 0; - - public: - explicit HillClimbAllocator(size_t memory_pressure) - : GreedyBase(), memory_pressure_(memory_pressure) {} - - protected: - using alloc_map_t = std::unordered_map; - - /* - * Initial sorting routine - */ - template - void sort_vector(std::vector* buffer_info_vec) { - std::sort(buffer_info_vec->begin(), buffer_info_vec->end(), [](const T& a, const T& b) { - if (a->size_bytes->value == b->size_bytes->value) { - if (a->conflicts.size() == b->conflicts.size()) { - return std::string(a->name_hint->data) > std::string(b->name_hint->data); - } else { - return a->conflicts.size() > b->conflicts.size(); - } - } - return a->size_bytes->value > b->size_bytes->value; - }); - } - - /* - * HillClimb's version of greedy allocation - * \param buffer_info_vec - buffers in specific order for allocation - */ - alloc_map_t greedy(const std::vector& buffer_info_vec, bool* could_not_fit) { - alloc_map_t pool_allocations(buffer_info_vec.size()); - for (const auto& buf_info : buffer_info_vec) { - std::unordered_map pool_offset_candidates; - - // check whether we can fit the buffer into the empty pool candidate - for (const auto& pool_info : buf_info->pool_candidates) { - if (IsValidPlacement(pool_info, 0, buf_info->size_bytes->value)) { - pool_offset_candidates[pool_info] = 0; - } - } - // select conflicting buffers which have already been allocated - std::vector buf_conf; - for (const auto& conflict_buf_info_obj : buf_info->conflicts) { - const BufferInfoNode* conflict_buf_info = conflict_buf_info_obj.as(); - if (pool_allocations.end() != pool_allocations.find(conflict_buf_info)) { - buf_conf.push_back(conflict_buf_info); - } - } - - // extra sorting for pool offsets - std::sort(buf_conf.begin(), buf_conf.end(), - [&pool_allocations](const auto* a, const auto* b) { - return pool_allocations[a]->byte_offset->value < - pool_allocations[b]->byte_offset->value; - }); - - for (const auto* conflict_buf_info : buf_conf) { - size_t next_offset = 0; - auto pool_allocation = pool_allocations[conflict_buf_info]; - if (!pool_offset_candidates.count(pool_allocation->pool_info)) { - continue; - } - - next_offset = - pool_allocation->byte_offset.IntValue() + conflict_buf_info->size_bytes.IntValue(); - next_offset = round_up_to_byte_alignment(next_offset, conflict_buf_info->alignment->value); - - if (IsValidPlacement(pool_allocation->pool_info, next_offset, - buf_info->size_bytes->value)) { - // extra check whether the previous attempt to fit the buffer is clashing with the current - // conflict - if (next_offset > pool_offset_candidates[pool_allocation->pool_info] && - pool_offset_candidates[pool_allocation->pool_info] + - static_cast(buf_info->size_bytes.IntValue()) > - static_cast(pool_allocation->byte_offset.IntValue())) { - pool_offset_candidates[pool_allocation->pool_info] = next_offset; - } - } else { - pool_offset_candidates.erase(pool_allocation->pool_info); - } - } - auto selected_pool = NullValue(); - for (const auto& pi : buf_info->pool_candidates) { - if (pool_offset_candidates.count(pi)) { - selected_pool = pi; - break; - } - } - - if (selected_pool.same_as(NullValue())) { - *could_not_fit = true; - } - - pool_allocations[buf_info.as()] = - PoolAllocation(selected_pool, Integer(pool_offset_candidates[selected_pool])); - } - return pool_allocations; - } - - /* - * Finds highest allocated memory address for each pool - */ - std::unordered_map find_highest( - alloc_map_t* pool_allocations) { - std::unordered_map pool_sizes; - for (const auto& it : *pool_allocations) { - const BufferInfoNode* buf = it.first; - const PoolAllocation& pa = it.second; - if (pa->pool_info.same_as(NullValue())) { - continue; - } - size_t high_sz = pa->byte_offset.IntValue() + buf->size_bytes.IntValue(); - if (pool_sizes[pa->pool_info] <= high_sz) { - pool_sizes[pa->pool_info] = high_sz; - } - } - return pool_sizes; - } - - /* - * Collects lists of first and secind level neigbors for provided buf. - * First level are the immediate neighbors of the buf and - * second level are the immediate neighbors of the first level nodes - */ - template - void collect_neighbor_lists(const BufferInfoNode* buf, - std::vector* first_level, - std::vector* second_level, const TPos& _pos) { - auto buf_pos = _pos(buf); - for (const auto& c1 : buf->conflicts) { - const auto* c1_buf = c1.as(); - int c1_pos = _pos(c1_buf); - if (buf_pos > c1_pos) { - first_level->push_back(c1_buf); - } - int c2_pos = -1; - for (const auto& c2 : c1_buf->conflicts) { - const auto c2_buf = c2.as(); - if (c1_pos > (c2_pos = _pos(c2_buf))) { - second_level->push_back(c2_buf); - } - } - } - } - - public: - Map PlanMemory(const Array& buffer_info_arr) { -// rand_r does not exist on Windows platform -#if defined(__linux__) || defined(__ANDROID__) - unsigned int _seedp = 0; -#define rnd_func() rand_r(&_seedp) -#else -#define rnd_func() rand() -#endif - Map result; - if (!buffer_info_arr.size()) { - return result; - } - std::vector buffer_info_vec; - for (const auto& buffer_info : buffer_info_arr) { - ICHECK(buffer_info->pool_candidates.size()) - << "Cannot process buffer \"" << buffer_info->name_hint << "\" with no pool candidates"; - buffer_info_vec.push_back(std::move(buffer_info)); - } - sort_vector(&buffer_info_vec); - - // populate positional index map - std::unordered_map _pos_map; - for (size_t index = 0; index < buffer_info_vec.size(); ++index) { - _pos_map[buffer_info_vec[index].as()] = index; - } - - size_t total_size = 0; - int attempts = 0; - - int swap_i1 = -1; - int swap_i2 = -1; - size_t desired_bytes_ = memory_pressure_; - constexpr auto _max_attempts = 500; - alloc_map_t rollback_pool_allocations; - alloc_map_t result_pool_allocations; - alloc_map_t pool_allocations; - - auto swap_buffers = [&buffer_info_vec, &_pos_map](int i1, int i2) { - if (i1 == i2) return; - auto b1 = buffer_info_vec[i1]; - auto b2 = buffer_info_vec[i2]; - buffer_info_vec[i1] = b2; - buffer_info_vec[i2] = b1; - - _pos_map[b1.as()] = i2; - _pos_map[b2.as()] = i1; - }; - - auto _pos = [&_pos_map](const auto* e) { - auto it = _pos_map.find(e); - if (it != _pos_map.end()) { - return it->second; - } - LOG(FATAL) << "node is not indexed in the _pos_map"; - }; - - for (; attempts < _max_attempts; ++attempts) { - rollback_pool_allocations = std::move(pool_allocations); - bool could_not_fit = false; - pool_allocations = std::move(greedy(buffer_info_vec, &could_not_fit)); - - // estimate result buffers - std::unordered_map pool_sizes = - find_highest(&pool_allocations); - if (!pool_sizes.size()) { - CHECK(false) << "TVM USMP Error: Please increase the size_hints for memory pools."; - } - - // calculate summary - size_t total = 0; - for (const auto& el : pool_sizes) { - total += el.second; - } - // accept/reject result heuristic - if (!total_size || /* first run */ - (!could_not_fit && - (total_size > total || /* always accept if better or with some probability */ - rnd_func() % 100 < static_cast(50 * (total - total_size) / total / attempts)))) { - // remember winning combination - result_pool_allocations = pool_allocations; - if (!could_not_fit) { - total_size = total; - // reached desired size - if (total_size <= desired_bytes_) { - break; - } - } - - } else { - // rollback - swap_buffers(swap_i2, swap_i1); - pool_allocations = std::move(rollback_pool_allocations); - pool_sizes = find_highest(&pool_allocations); - } - - std::vector max_pool_buf; - - for (const auto& it : pool_allocations) { - const auto* buf = it.first; - const auto pa = it.second; - if (pa->pool_info.same_as(NullValue())) { - continue; - } - size_t high_sz = pa->byte_offset.IntValue() + buf->size_bytes.IntValue(); - if (pool_sizes[pa->pool_info] == high_sz) { - max_pool_buf.push_back(buf); - } - } - if (!max_pool_buf.size()) { - CHECK(false) << "TVM USMP Error: Please increase the size_hints for memory pools."; - } - sort(max_pool_buf.begin(), max_pool_buf.end(), - [&_pos](const auto* a, const auto* b) { return _pos(a) < _pos(b); }); - // pick highest - const BufferInfoNode* node = max_pool_buf[rnd_func() % max_pool_buf.size()]; - std::vector first_level; - std::vector second_level; - collect_neighbor_lists(node, &first_level, &second_level, _pos); - sort(first_level.begin(), first_level.end(), - [&_pos](const auto* a, const auto* b) { return _pos(a) < _pos(b); }); - sort(second_level.begin(), second_level.end(), - [&_pos](const auto* a, const auto* b) { return _pos(a) < _pos(b); }); - - // retry if no first level neightbors were collected - if (!first_level.size()) { - continue; - } - - // pick the buffers - const BufferInfoNode* swap_buf1 = first_level[rnd_func() % first_level.size()]; - const BufferInfoNode* swap_buf2 = swap_buf1; - while (swap_buf2 == swap_buf1) { - swap_buf2 = second_level.size() && (!first_level.size() || (rnd_func() % 100 > 25)) - ? second_level[rnd_func() % second_level.size()] - : first_level[rnd_func() % first_level.size()]; - - if (second_level.size() < 2 && first_level.size() < 2) break; - } - if (swap_buf1 == swap_buf2) { - continue; - } - - swap_i1 = _pos(swap_buf1); - swap_i2 = _pos(swap_buf2); - // do swap - swap_buffers(swap_i1, swap_i2); - } - - // return winning combination - for (auto it : result_pool_allocations) { - // post-check that everything was fit - const BufferInfoNode* buf = it.first; - const PoolAllocation& pa = it.second; - if (NullValue().same_as(pa->pool_info) || - !IsValidPlacement(pa->pool_info, pa->byte_offset->value, buf->size_bytes->value)) { - std::unordered_map m = {}; - SelectPlacementPool(GetRef(buf), m); - } - result.Set(GetRef(it.first), it.second); - } - return result; - } -}; - -Map HillClimb(const Array& buffer_info_arr, - const Integer& memory_pressure) { - return HillClimbAllocator(memory_pressure.IntValue()).PlanMemory(buffer_info_arr); -} - -TVM_REGISTER_GLOBAL("tir.usmp.algo.hill_climb") - .set_body_typed([](Array buffer_info_arr, Integer memory_pressure) { - return HillClimb(buffer_info_arr, memory_pressure); - }); - -} // namespace algo -} // namespace usmp -} // namespace tir -} // namespace tvm diff --git a/src/tir/usmp/analysis/extract_buffer_info.cc b/src/tir/usmp/analysis/extract_buffer_info.cc deleted file mode 100644 index 5abfe24f434d..000000000000 --- a/src/tir/usmp/analysis/extract_buffer_info.cc +++ /dev/null @@ -1,627 +0,0 @@ -/* - * 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 tir/analysis/usmp/extract_buffer_info.cc - * - * \brief This analysis pass consumes a TIR IRModule with a main function - * that defines a ordering in the callees to operators and produces BufferInfo - * objects that contains information about tir.allocate nodes and liveness - * conflicts between other tir.allocate nodes. - */ -#include -#include -#include -#include -#include -#include -#include -#include - -#include - -#include "../../../runtime/thread_storage_scope.h" - -namespace tvm { -namespace tir { -namespace usmp { - -/*! - * \brief The visitor class to obtain buffer information - * - * The visitor would initiate the traversal from the main - * function and visits into the operator PrimFuncs. It will - * crate unique BufferInfo objects for each Allocate node. - * - * Every time the buffer variable of the allocate node is referenced - * it will be recorded using the stmt index. However, note that - * the same buffer variable could be references multiple times - * from different calls. Thereafter, a sweep is done on all the - * BufferInfo objects using the per-call liveness events. In the sweep, - * The BufferInfo objects that are live together will be recorded as - * mutual conflicts of each other. - */ -class BufferInfoExtractor : public StmtExprVisitor { - public: - explicit BufferInfoExtractor(const IRModule& module) : module_(module) { - for (const auto& gv_func : module_->functions) { - if (gv_func.second->IsInstance()) { - functions_.Set(gv_func.first->name_hint, Downcast(gv_func.second)); - } - } - // Pushing a scope info for the initial body of the main function - scope_stack_.push(ScopeInfo()); - } - BufferInfoAnalysis operator()(const PrimFunc& func); - - private: - void VisitStmt(const Stmt& n) override; - void VisitStmt_(const AllocateNode* op) override; - void VisitStmt_(const AllocateConstNode* op) override; - void VisitExpr_(const CallNode* op) override; - void VisitExpr_(const VarNode* op) override; - void VisitExpr_(const BufferLoadNode* op) override; - void VisitStmt_(const BufferStoreNode* op) override; - void VisitStmt_(const ForNode* op) override; - - void UpdateAliases(const Array& args, const PrimFunc& func); - void RecordAllocateNodeInfo(const AllocateNode* op); - void RecordAllocateConstNodeInfo(const AllocateConstNode* op); - void VisitPrimFunc(const PrimFunc& func, const Call& call); - - /*! - * \brief Maintains the mapping of BufferInfo to their associated TIR Statements. - */ - Map buffer_info_map_; - /*! - * \brief Records the order of calls in the main for stability. - */ - std::vector call_order_; - /*! - * \brief Lookup to avoid adding duplicates to `call_order_`. - */ - std::unordered_set call_order_contents_; - /*! - * \brief Records first access in-terms of Stmts to each buffer per call - * - * This is because multiple calls could happen to the same PrimFunc. - */ - std::unordered_map, ObjectPtrHash, ObjectPtrEqual> - buffer_info_start_stmt_idx_; - /*! - * \brief Records last access in-terms of Stmts to each buffer per call - * - * This is because multiple calls could happen to the same PrimFunc. - */ - std::unordered_map, ObjectPtrHash, ObjectPtrEqual> - buffer_info_end_stmt_idx_; - - /*! - * \brief This structure contains information regarding a Allocate node. - */ - struct AllocateInfo { - tir::Stmt Allocate; - PrimFunc prim_func; - Call call; - }; - - /*! - * \brief Maintains the mapping of buffer variable to their allocate nodes to ensure - * that only one BufferInfo object is created. - */ - std::unordered_map allocate_infos; - /*! - * \brief Indicates a count of stmts visited so far to use as a metric of liveness - */ - int current_stmt_idx_ = 0; - /*! - * \brief This structure is supposed to contain information around the scope - * the visitor is currently in. - */ - struct ScopeInfo { - /*! - * \brief We need to record access per call - */ - Call call; - /*! - * \brief Having access to PrimFunc metadata is useful - */ - PrimFunc func; - /*! - * \brief We currently support only serial for loops. Therefore - * need to know what kind of for loop the visitor is in. - */ - For for_loop; - /*! - * \brief We record the live allocate_nodes because once in loops - * the liveness range has to be extended to the whole of the nested - * loops structure. - */ - std::unordered_set allocate_nodes; - /* - * \brief We record the live allocate_const_nodes because once in loops - * the liveness range has to be extended to the whole of the nested - * loops structure. - */ - std::unordered_set allocate_const_nodes; - /*! - * \brief This is recorded to extend the liveness of all allocates within - * nested loop structure. - */ - Integer initial_stmt_of_the_nested_loops; - }; - std::stack scope_stack_; - - /*! - * \brief A liveness event is an event that when - * traversing the tir.Stmts where tir.allocate node - * begins or ceases to be Live. This particular struct - * is used to solve interval overlap problem using - * a sweep-line algorithm. For that, we need to record - * where the liveness event occurred in a chronological - * order. - */ - enum LivenessEventType { START = 0, END = 1 }; - struct LivenessEvent { - size_t tick; - LivenessEventType le_type; - BufferInfo buffer_info; - bool operator==(const LivenessEvent& other) { - if (tick == other.tick && le_type == other.le_type && buffer_info == other.buffer_info) { - return true; - } - return false; - } - }; - /*! - * \brief We need to create unique buffer name is the same name is used in - * two allocate nodes for clarity for memory planning algorithms. - */ - std::string GetUniqueBufferName(std::string name); - - /*! - * \brief This is per buffer name counter to aid the generating the above - * unique name. - */ - std::unordered_map buffer_names; - /*! - * \brief The TIR main function calls by name to PrimFuncs to be able to - * support BYOC. Therefore, this Map records functions that are present - * in the IRModule by name/ - */ - Map functions_; - /*! - * \brief The IRModule being analyzed. - */ - IRModule module_; -}; - -std::string BufferInfoExtractor::GetUniqueBufferName(std::string name) { - if (buffer_names.find(name) == buffer_names.end()) { - buffer_names[name] = 1; - return name; - } else { - buffer_names[name] = buffer_names[name] + 1; - return name + std::to_string(buffer_names[name]); - } -} - -void BufferInfoExtractor::VisitStmt(const Stmt& n) { - current_stmt_idx_ += 1; - StmtExprVisitor::VisitStmt(n); -} - -void BufferInfoExtractor::RecordAllocateNodeInfo(const AllocateNode* op) { - auto size_bytes = CalculateExtentsSize(op); - // We only statically memory plan only allocates with known - // compile time sizes. - if (size_bytes.defined()) { - if (allocate_infos.find(op->buffer_var) == allocate_infos.end()) { - // By default, the core compiler is assumed to attach the a default pool to each allocate. - ICHECK(op->annotations.count(kPoolCandidatesAllocateAttr)) - << "Every statically sized allocate node needs an pool candidate attribute"; - auto pool_candidates = - Downcast>(op->annotations[kPoolCandidatesAllocateAttr]); - - ICHECK(pool_candidates.size() > 0) - << "The AssignPoolInfo pass should at least attach a single PoolInfo. If there were no " - "user-given arguments for memory pools, the default behaviour is a single size " - "un-restricted pool is assigned"; - PrimFunc func = scope_stack_.top().func; - Optional executor_config = - module_->GetAttr(tvm::attr::kExecutor); - Integer workspace_alignment = 16; - if (executor_config) { - workspace_alignment = - executor_config.value()->GetAttr("workspace-byte-alignment").value_or(16); - } - - BufferInfoKind bi_kind = BufferInfoKind::kIntermediate; - String buffer_info_name = op->buffer_var->name_hint; - if (op->annotations.find(kInputTensorAllocate) != op->annotations.end()) { - bi_kind = BufferInfoKind::kInput; - // using original input name instead of the buffer_var name - // because this name will be used in the lowering to convey - // the pool allocation. - buffer_info_name = Downcast(op->annotations[kInputTensorAllocate]); - } else if (op->annotations.find(kOutputTensorAllocate) != op->annotations.end()) { - bi_kind = BufferInfoKind::kOutput; - // using original output name instead of the buffer_var name - // because this name will be used in the lowering to convey - // the pool allocation. - buffer_info_name = Downcast(op->annotations[kOutputTensorAllocate]); - } - auto buffer_info = BufferInfo(GetUniqueBufferName(buffer_info_name), size_bytes, - pool_candidates, workspace_alignment, bi_kind); - auto allocate = GetRef(op); - allocate_infos[op->buffer_var] = - AllocateInfo{allocate, scope_stack_.top().func, scope_stack_.top().call}; - buffer_info_map_.Set(buffer_info, allocate); - } else { - // Update the allocate info with the latest call - AllocateInfo ai = allocate_infos[op->buffer_var]; - ai.call = scope_stack_.top().call; - allocate_infos[op->buffer_var] = ai; - } - } -} - -void BufferInfoExtractor::VisitStmt_(const AllocateNode* op) { - ScopeInfo& current_scope_info = scope_stack_.top(); - const auto& type = Downcast(op->buffer_var->type_annotation); - const auto& storage_scope = runtime::StorageScope::Create(type->storage_scope); - - // If the allocate is in a for loop, USMP currently only looks at serial for loops. - // If its not a serial for loop, then memory planner will omit them in the current memory planning - // process leaving them to as tir.allocate nodes for codegen. Additionally, the USMP can only work - // with buffers that have global storage_scope - - if (storage_scope.rank == runtime::StorageRank::kGlobal) { - if (!current_scope_info.for_loop.defined()) { - RecordAllocateNodeInfo(op); - } else if (current_scope_info.for_loop.defined() && - current_scope_info.for_loop->kind == ForKind::kSerial) { - RecordAllocateNodeInfo(op); - } - } - StmtExprVisitor::VisitStmt(op->body); - current_scope_info.allocate_nodes.erase(GetRef(op)); -} - -void BufferInfoExtractor::VisitStmt_(const AllocateConstNode* op) { - ScopeInfo& current_scope_info = scope_stack_.top(); - RecordAllocateConstNodeInfo(op); - StmtExprVisitor::VisitStmt(op->body); - current_scope_info.allocate_const_nodes.erase(GetRef(op)); -} - -void BufferInfoExtractor::RecordAllocateConstNodeInfo(const AllocateConstNode* op) { - if (!op->annotations.count(kPoolCandidatesAllocateAttr)) { - return; - } - Integer size_bytes = CalculateExtentsSize(op); - ICHECK(size_bytes.defined()) << "constant node size should be defined"; - const auto& buffer_var = op->buffer_var; - if (allocate_infos.find(buffer_var) == allocate_infos.end()) { - // By default, the core compiler is assumed to attach the a default pool to each allocate. - ICHECK(op->annotations.count(kPoolCandidatesAllocateAttr)) - << "Every statically sized allocate node needs an pool candidate attribute"; - auto pool_candidates = Downcast>(op->annotations[kPoolCandidatesAllocateAttr]); - ICHECK(pool_candidates.size() > 0) - << "The core compiler should at least attach a single PoolInfo. If there were no " - "user-given arguments for memory pools, the default behaviour is a single size " - "un-restricted pool is assigned"; - PrimFunc func = scope_stack_.top().func; - Optional executor_config = - module_->GetAttr(tvm::attr::kExecutor); - Integer alignment = 16; - if (executor_config) { - alignment = - executor_config.value()->GetAttr("constant-byte-alignment").value_or(alignment); - } - auto buffer_info = BufferInfo(GetUniqueBufferName(buffer_var->name_hint), size_bytes, - pool_candidates, alignment); - auto allocate = GetRef(op); - allocate_infos[buffer_var] = - AllocateInfo{allocate, scope_stack_.top().func, scope_stack_.top().call}; - buffer_info_map_.Set(buffer_info, allocate); - } else { - // Update the allocate info with the latest call - AllocateInfo ai = allocate_infos[buffer_var]; - ai.call = scope_stack_.top().call; - allocate_infos[buffer_var] = ai; - } -} - -void BufferInfoExtractor::VisitStmt_(const ForNode* op) { - ScopeInfo si{scope_stack_.top().call, - scope_stack_.top().func, - GetRef(op), - scope_stack_.top().allocate_nodes, - scope_stack_.top().allocate_const_nodes, - scope_stack_.top().initial_stmt_of_the_nested_loops}; - if (!scope_stack_.top().initial_stmt_of_the_nested_loops.defined()) { - si.initial_stmt_of_the_nested_loops = Integer(current_stmt_idx_); - } - Call current_call = scope_stack_.top().call; - PrimFunc current_primfunc = scope_stack_.top().func; - scope_stack_.push(si); - StmtExprVisitor::VisitStmt_(op); - // Extending the liveness to beginning of for-loop next and end of the current for-loop - for (const Allocate& allocate : scope_stack_.top().allocate_nodes) { - AllocateInfo ai = allocate_infos[allocate->buffer_var]; - Call update_call = current_call; - // If the allocate does not belong to current prim func - // We need to update the call to which the allocate belong to - if (ai.prim_func != current_primfunc) { - update_call = ai.call; - } - if (scope_stack_.top().initial_stmt_of_the_nested_loops->value < - buffer_info_start_stmt_idx_[update_call][allocate].IntValue()) { - buffer_info_start_stmt_idx_[update_call].Set( - allocate, scope_stack_.top().initial_stmt_of_the_nested_loops->value); - } - if (current_stmt_idx_ > buffer_info_end_stmt_idx_[update_call][allocate].IntValue()) { - buffer_info_end_stmt_idx_[update_call].Set(allocate, current_stmt_idx_); - } - } - scope_stack_.pop(); -} - -void BufferInfoExtractor::VisitExpr_(const BufferLoadNode* op) { - this->VisitExpr(op->buffer->data); - StmtExprVisitor::VisitExpr_(op); -} - -void BufferInfoExtractor::VisitStmt_(const BufferStoreNode* op) { - this->VisitExpr(op->buffer->data); - StmtExprVisitor::VisitStmt_(op); -} - -void BufferInfoExtractor::VisitExpr_(const VarNode* op) { - auto var = GetRef(op); - Call current_call = scope_stack_.top().call; - PrimFunc current_primfunc = scope_stack_.top().func; - if (allocate_infos.count(var)) { - auto allocate = allocate_infos[var].Allocate; - auto allocate_primfunc = allocate_infos[var].prim_func; - Call update_call = current_call; - if (allocate_primfunc != current_primfunc) { - // If the allocate node does not belong to the current primfunc. - // It's access should be reported to the call to PrimFunc that - // Allocate belong to. - update_call = allocate_infos[var].call; - } - if (buffer_info_start_stmt_idx_[update_call].count(allocate) == 0) { - buffer_info_start_stmt_idx_[update_call].Set(allocate, current_stmt_idx_); - } - buffer_info_end_stmt_idx_[update_call].Set(allocate, current_stmt_idx_); - - ScopeInfo& currect_scope_info = scope_stack_.top(); - if (currect_scope_info.for_loop.defined()) { - if (allocate->IsInstance()) { - currect_scope_info.allocate_nodes.insert(Downcast(allocate)); - } else if (allocate->IsInstance()) { - currect_scope_info.allocate_const_nodes.insert(Downcast(allocate)); - } else { - LOG(FATAL) << "Handling of " << allocate->GetTypeKey() << " is not implemented"; - } - } - } - StmtExprVisitor::VisitExpr_(op); -} - -Array static GetMatchedBuffers(const PrimFunc& func) { - Array buffer_vars; - if (func->params.size() > 0) { - for (unsigned int i = 0; i < func->params.size() - 1; i++) { - Var param = func->params[i]; - buffer_vars.push_back(func->buffer_map[param]->data); - } - Var last_param = func->params.back(); - // Checks whether last var is present in the buffer map - // because it could be the resource handle - if (func->buffer_map.find(last_param) != func->buffer_map.end()) { - buffer_vars.push_back(func->buffer_map[last_param]->data); - } - } - return buffer_vars; -} - -void BufferInfoExtractor::UpdateAliases(const Array& args, const PrimFunc& func) { - auto param_buffers = GetMatchedBuffers(func); - // Last var could be a resource handle that does not have a Buffer - ICHECK(args.size() == param_buffers.size() || args.size() - 1 == param_buffers.size()); - for (size_t i = 0; i < param_buffers.size(); i++) { - auto arg = args[i]; - auto param_buf = param_buffers[i]; - // If tir.allocates are passed in to functions - // The function params are re-directed to point - // to the original allocate - if (arg->IsInstance()) { - auto var = Downcast(arg); - if (allocate_infos.count(var)) { - allocate_infos[param_buf] = allocate_infos[var]; - } - } - } -} - -void BufferInfoExtractor::VisitPrimFunc(const PrimFunc& func, const Call& call) { - ScopeInfo si{call, - func, - scope_stack_.top().for_loop, - scope_stack_.top().allocate_nodes, - scope_stack_.top().allocate_const_nodes, - scope_stack_.top().initial_stmt_of_the_nested_loops}; - if (call_order_contents_.count(call) == 0) { - call_order_contents_.insert(call); - call_order_.push_back(call); - } - scope_stack_.push(si); - this->VisitStmt(func->body); - scope_stack_.pop(); -} - -void BufferInfoExtractor::VisitExpr_(const CallNode* op) { - if (op->op.same_as(builtin::call_extern()) || op->op.same_as(builtin::tvm_call_cpacked())) { - StringImm func_name = Downcast(op->args[0])->value; - if (functions_.find(func_name->value) != functions_.end()) { - auto func = functions_.at(func_name->value); - auto actual_args = Array(op->args.begin() + 1, op->args.end()); - this->UpdateAliases(actual_args, func); - VisitPrimFunc(func, GetRef(op)); - return; - } - } - if (op->op->IsInstance()) { - auto func = Downcast(op->op); - this->UpdateAliases(op->args, func); - VisitPrimFunc(func, GetRef(op)); - return; - } - StmtExprVisitor::VisitExpr_(op); -} - -BufferInfoAnalysis BufferInfoExtractor::operator()(const PrimFunc& main_func) { - VisitPrimFunc(main_func, Call()); - - // Create a vector of liveness events - // associated with each BufferNodes. - std::vector le_events_timeline; - for (const auto& kv1 : buffer_info_map_) { - if (!kv1.second->IsInstance() && !kv1.second->IsInstance()) { - continue; - } - - auto allocate = Downcast(kv1.second); - auto buffer_info = Downcast(kv1.first); - - ICHECK(call_order_.size() >= buffer_info_end_stmt_idx_.size()); - ICHECK(call_order_.size() >= buffer_info_end_stmt_idx_.size()); - - for (const Call& call : call_order_) { - Map buffer_info_starts = buffer_info_start_stmt_idx_[call]; - if (buffer_info_starts.find(allocate) != buffer_info_starts.end()) { - LivenessEvent le_event_start; - le_event_start.buffer_info = buffer_info; - le_event_start.le_type = START; - le_event_start.tick = buffer_info_starts[allocate].IntValue(); - le_events_timeline.push_back(le_event_start); - } - } - - for (const Call& call : call_order_) { - Map buffer_info_ends = buffer_info_end_stmt_idx_[call]; - if (buffer_info_ends.find(allocate) != buffer_info_ends.end()) { - LivenessEvent le_event_end; - le_event_end.buffer_info = buffer_info; - le_event_end.le_type = END; - le_event_end.tick = buffer_info_ends[allocate].IntValue(); - le_events_timeline.push_back(le_event_end); - } - } - } - - // Sort the liveness events based on the chronological - // ordering. For events that are simultaneous, START event - // takes precedence. - std::sort(le_events_timeline.begin(), le_events_timeline.end(), - [](const LivenessEvent& lhs, const LivenessEvent& rhs) { - if (lhs.tick < rhs.tick) { - return true; - } else if (lhs.tick == rhs.tick && lhs.le_type == START && rhs.le_type == END) { - return true; - } - return false; - }); - - // Traverse the liveness events using a open set to track what - // is live while updating the conflicts through out the linear traversal - - int open_set_size = 0; - int max_open_set_size = 0; - std::unordered_set open_set; - for (const auto& le_event : le_events_timeline) { - if (le_event.le_type == START) { - for (const BufferInfo& open_buffer_info : open_set) { - open_buffer_info->conflicts.push_back(le_event.buffer_info); - if (le_event.buffer_info != open_buffer_info) { - le_event.buffer_info->conflicts.push_back(open_buffer_info); - } - } - open_set_size += le_event.buffer_info->size_bytes.IntValue(); - if (open_set_size > max_open_set_size) { - max_open_set_size = open_set_size; - } - open_set.insert(le_event.buffer_info); - } else { - open_set_size -= le_event.buffer_info->size_bytes.IntValue(); - open_set.erase(le_event.buffer_info); - } - } - - // All ConstantPoolInfo items should have conflicts with each other - // as they will be placed in RO segment and pre-initialized. To achieve this - // first, split buffers to vars (WorkspacePoolInfo items) and constants (ConstantPoolInfo items): - Array buffer_info_vars; - Array buffer_info_constants; - for (const auto& kv : this->buffer_info_map_) { - const auto& stmt = kv.second; - if (stmt->IsInstance()) { - buffer_info_constants.push_back(kv.first); - } else { - buffer_info_vars.push_back(kv.first); - } - } - ICHECK(buffer_info_map_.size() == buffer_info_vars.size() + buffer_info_constants.size()) - << "missing value"; - - Map srch; - // Then intersect constants with each other, as all constants should exist at the same time: - for (const auto& buf : buffer_info_constants) { - srch.Set(buf, buf); - Array conflicts; - std::copy_if(buffer_info_constants.begin(), buffer_info_constants.end(), - std::back_inserter(conflicts), [buf](const auto& b) { return b != buf; }); - buf->conflicts.Assign(conflicts.begin(), conflicts.end()); - } - - // And third, remove all conflicts between constants and vars: - for (const auto& buf : buffer_info_vars) { - Array conflicts; - std::copy_if(buf->conflicts.begin(), buf->conflicts.end(), std::back_inserter(conflicts), - [&srch](const auto& c) { return srch.end() == srch.find(c); }); - buf->conflicts.Assign(conflicts.begin(), conflicts.end()); - } - return BufferInfoAnalysis(this->buffer_info_map_, max_open_set_size); -} - -BufferInfoAnalysis ExtractBufferInfo(const PrimFunc& main_func, const IRModule& mod) { - return BufferInfoExtractor(mod)(main_func); -} - -TVM_REGISTER_GLOBAL("tir.usmp.analysis.extract_buffer_info") - .set_body_typed([](PrimFunc main_func, IRModule mod) { - return (ExtractBufferInfo(main_func, mod)); - }); - -} // namespace usmp -} // namespace tir -} // namespace tvm diff --git a/src/tir/usmp/transform/assign_pool_info.cc b/src/tir/usmp/transform/assign_pool_info.cc deleted file mode 100644 index 3acceab6e31b..000000000000 --- a/src/tir/usmp/transform/assign_pool_info.cc +++ /dev/null @@ -1,192 +0,0 @@ -/* - * 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 - -namespace tvm { -namespace tir { -namespace usmp { - -/*! \brief Assign PoolInfo objects to allocate that does not have any. - * The schedulers have the oppurtunity to assign PoolInfo objects to - * allocate nodes. However, each allocate node is expected to have - * at least one PoolInfo node assigned to it. If it was not the case, - * this Pass will assign all PoolInfo objects that the target could - * access.*/ -class PoolInfoAssigner : public StmtExprMutator { - public: - explicit PoolInfoAssigner(const IRModule& module) { - PrimFunc main_func = - Downcast(module->Lookup(::tvm::runtime::symbol::tvm_module_main)); - ICHECK(main_func.defined()) << "main function is not in the module"; - Optional target_host = main_func->GetAttr(tvm::attr::kTarget); - ICHECK(target_host) << "main function does not have a target attr"; - WorkspaceMemoryPools workspace_pools = - module->GetAttr(tvm::attr::kWorkspaceMemoryPools) - .value_or(WorkspaceMemoryPools({CreateDefaultWorkspaceMemoryPool(module)})); - // make default ConstantPoolInfo if no constant and no workspace pool infos supplied - ConstantMemoryPools constant_pools = - module->GetAttr(tvm::attr::kConstantMemoryPools) - .value_or( - module->GetAttr(tvm::attr::kWorkspaceMemoryPools).defined() - ? ConstantMemoryPools() - : ConstantMemoryPools({CreateDefaultConstantMemoryPool(module)})); - auto to_map = [](auto pool_infos) { - Map> pool_map; - for (const PoolInfo& pool_info : pool_infos) { - for (const auto& tgt : pool_info->targets) { - if (pool_map.find(tgt->str()) == pool_map.end()) { - pool_map.Set(tgt->str(), Array()); - } - Array pool_info_arr = pool_map[tgt->str()]; - pool_info_arr.push_back(pool_info); - pool_map.Set(tgt->str(), pool_info_arr); - } - } - return pool_map; - }; - - target_pool_infos_ = to_map(workspace_pools->pools); - if (constant_pools.defined()) { - target_const_pool_infos_ = to_map(constant_pools->pools); - } - mod_ = module->ShallowCopy(); - } - - IRModule operator()(); - - private: - Stmt VisitStmt_(const AllocateNode* op) override; - Stmt VisitStmt_(const AllocateConstNode* op) override; - - IRModule mod_; - Map> target_pool_infos_; - Map> target_const_pool_infos_; - PrimFunc func_; - WorkspacePoolInfo CreateDefaultWorkspaceMemoryPool(const IRModule& module); - ConstantPoolInfo CreateDefaultConstantMemoryPool(const IRModule& module) { - auto p = CreateDefaultWorkspaceMemoryPool(module); - return ConstantPoolInfo( - "global_const_workspace", {p->targets}, {}, - PoolInfoProperties(kUnrestrictedPoolSizeHint, kUnknownClockFrequency, kUnknownReadBandwidth, - kUnknownWriteBandwidth, 0, 0, {p->target_burst_bytes}, Bool(true))); - } -}; - -WorkspacePoolInfo PoolInfoAssigner::CreateDefaultWorkspaceMemoryPool(const tvm::IRModule& module) { - VLOG(1) << "Creating default memory pool for:" << std::endl << module; - Map target_access; - tir::PrimFunc tir_main_func = - Downcast(module->Lookup(::tvm::runtime::symbol::tvm_module_main)); - Target target_host = tir_main_func->GetAttr(tvm::attr::kTarget).value(); - for (const auto& kv : module->functions) { - BaseFunc func = kv.second; - Optional target = func->GetAttr(tvm::attr::kTarget); - target_access.Set(target.value_or(target_host), kTargetPoolReadWriteAccess); - } - Array targets; - for (const auto& kv : target_access) { - bool exist = false; - // Exclude targets with the same string representation - for (const auto& t : targets) { - if (t->str() == kv.first->str()) { - exist = true; - } - } - if (!exist) { - targets.push_back(kv.first); - } - } - return WorkspacePoolInfo( - "global_workspace", targets, - PoolInfoProperties(kUnrestrictedPoolSizeHint, kUnknownClockFrequency, kUnknownReadBandwidth, - kUnknownWriteBandwidth, 0, 0, {{target_host, 1}}, Bool(true))); -} - -Stmt PoolInfoAssigner::VisitStmt_(const AllocateNode* op) { - Optional tgt = func_->GetAttr(tvm::attr::kTarget).value(); - ICHECK(tgt) << "The following PrimFunc does not have a target attr: \n" << func_; - Map annotations = Map(op->annotations); - if (op->annotations.find(kPoolCandidatesAllocateAttr) == op->annotations.end()) { - ICHECK(target_pool_infos_.count(tgt.value()->str()) > 0) - << "Target " << tgt << " not found among " << target_pool_infos_; - annotations.Set(kPoolCandidatesAllocateAttr, target_pool_infos_[tgt.value()->str()]); - } - Stmt body = VisitStmt(op->body); - auto allocate = - Allocate(op->buffer_var, op->dtype, op->extents, op->condition, body, annotations); - return std::move(allocate); -} - -Stmt PoolInfoAssigner::VisitStmt_(const AllocateConstNode* op) { - if (!target_const_pool_infos_.size()) { - return StmtExprMutator::VisitStmt_(op); - } - Optional tgt = func_->GetAttr(tvm::attr::kTarget).value(); - ICHECK(tgt) << "The following PrimFunc does not have a target attr: \n" << func_; - Map annotations = Map(op->annotations); - if (op->annotations.find(kPoolCandidatesAllocateAttr) == op->annotations.end()) { - annotations.Set(kPoolCandidatesAllocateAttr, target_const_pool_infos_[tgt.value()->str()]); - annotations.Set(kTargetPoolReadOnlyAccess, Integer(1)); - } - Stmt body = VisitStmt(op->body); - auto allocate_const = - AllocateConst(op->buffer_var, op->dtype, op->extents, op->data, body, annotations); - return std::move(allocate_const); -} - -IRModule PoolInfoAssigner::operator()() { - for (const auto& kv : mod_->functions) { - GlobalVar gv = kv.first; - if (kv.second->IsInstance()) { - func_ = Downcast(kv.second); - Stmt body = this->VisitStmt(func_->body); - PrimFunc new_prim_func = - PrimFunc(func_->params, body, func_->ret_type, func_->buffer_map, func_->attrs); - mod_->Update(gv, new_prim_func); - } - } - return mod_; -} - -namespace transform { - -tvm::transform::Pass AssignPoolInfo() { - auto pass_func = [=](IRModule m, tvm::transform::PassContext ctx) { - return PoolInfoAssigner(m)(); - }; - return tvm::transform::CreateModulePass(pass_func, 0, "tir.usmp.AssignPoolInfo", {}); -} - -TVM_REGISTER_GLOBAL("tir.usmp.transform.AssignPoolInfo").set_body_typed(AssignPoolInfo); - -} // namespace transform - -} // namespace usmp -} // namespace tir -} // namespace tvm diff --git a/src/tir/usmp/transform/convert_pool_allocations_to_offsets.cc b/src/tir/usmp/transform/convert_pool_allocations_to_offsets.cc deleted file mode 100644 index 0426c0cb1e7d..000000000000 --- a/src/tir/usmp/transform/convert_pool_allocations_to_offsets.cc +++ /dev/null @@ -1,506 +0,0 @@ -/* - * 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 tir/analysis/usmp/transform/convert_pool_allocations_to_offsets.cc - * \brief This pass would convert the pool allocations to offsets from pools - */ - -#include -#include -#include -#include -#include -#include -#include -#include -#include - -#include - -namespace tvm { -namespace tir { -namespace usmp { - -/*! - * \brief The StmtExpr mutator class to replace allocate nodes - * with offsets within memory pools - * - * This mutator class will add Pool variables recursively to every PrimFunc - * starting from the main PrimFunc. For all allocate nodes, that have been - * memory planned, will be mutated into an offset using a Let binding. - */ -class PoolAllocationToOffsetConverter : public StmtExprMutator { - public: - PoolAllocationToOffsetConverter(const IRModule& module, - const Map& pool_allocations, - bool emit_tvmscript_printable = false) - : pool_allocations_(pool_allocations), emit_tvmscript_printable_(emit_tvmscript_printable) { - module_ = module->ShallowCopy(); - for (const auto& kv : pool_allocations) { - size_t extent_size = -1; - if (kv.first->IsInstance()) { - Allocate allocate_node = Downcast(kv.first); - extent_size = CalculateExtentsSize(allocate_node.operator->()).IntValue(); - } else if (kv.first->IsInstance()) { - AllocateConst allocate_const_node = Downcast(kv.first); - extent_size = CalculateExtentsSize(allocate_const_node.operator->()).IntValue(); - } else { - ICHECK(false) << "Not supported node type " << kv.first->GetTypeKey(); - } - PoolAllocation pool_allocation = kv.second; - PoolInfo pool_info = pool_allocation->pool_info; - int byte_pool_offset = pool_allocation->byte_offset->value; - int required_pool_size_for_allocation = byte_pool_offset + extent_size; - if (all_pools_sizes_.find(pool_info) == all_pools_sizes_.end()) { - all_pools_sizes_[pool_info] = required_pool_size_for_allocation; - } else { - int prev_required_pool_size = all_pools_sizes_[pool_info]; - if (prev_required_pool_size < required_pool_size_for_allocation) { - all_pools_sizes_[pool_info] = required_pool_size_for_allocation; - } - } - } - - for (const auto& kv : all_pools_sizes_) { - PoolInfo pi = kv.first; - int allocated_size = kv.second; - allocated_pool_ordering_.push_back(AllocatedPoolInfo(pi, allocated_size)); - } - std::sort(allocated_pool_ordering_.begin(), allocated_pool_ordering_.end(), - [](const AllocatedPoolInfo& lhs, const AllocatedPoolInfo& rhs) { - if (lhs->pool_info->pool_name < rhs->pool_info->pool_name) { - return true; - } - return false; - }); - } - IRModule operator()(); - - private: - PrimExpr VisitExpr_(const CallNode* op) override; - Stmt VisitStmt_(const AllocateNode* op) override; - PrimExpr VisitExpr_(const VarNode* op) override; - PrimExpr VisitExpr_(const BufferLoadNode* op) override; - Stmt VisitStmt_(const BufferStoreNode* op) override; - Stmt VisitStmt_(const DeclBufferNode* op) override; - - Stmt VisitStmt_(const AllocateConstNode* op) override; - LetStmt ToLetStmt(const PoolAllocation& pool_allocation, const Var& buffer_var, const Stmt& body); - /*! \brief This is a structure where the modified function - * signature is kept while body of the function is mutated - */ - struct ScopeInfo { - Array params; - Map pools_to_params; - Array allocated_pool_params; - Map buffer_map; - }; - - /*! \brief The function scope information that are needed - * in the mutation of the function need to be stacked and - * popped when each function is entered/exited in the - * mutation process. - */ - std::stack scope_stack; - /*! \brief Each PrimFunc signature needs to be updated - * with pool variables. This is a helper function to - * capture the updated information to ScopeInfo object. - */ - ScopeInfo UpdateFunctionScopeInfo(const PrimFunc& original_func); - /*! \brief This is a helper to create the PrimFunc with - * pool variables that calls the UpdateFunctionScopeInfo - * inside of it. - */ - PrimFunc CreatePrimFuncWithPoolParams(const PrimFunc& original_primfunc); - /*! \brief This is a helper to append the pool args to - * the callsite of the function. - */ - Array AppendPoolParamsToArgs(Array args, bool has_device_context); - /*! \brief Some arguments that used to be Allocate nodes - * should be replaced by Let nodes in the pass that loads - * the space from a pool variable. - */ - Array ReplaceAllocateArgsWithLetArgs(const Array& args); - /*! \brief Obtain a resource handle if its there - */ - Optional GetResourceHandle(const PrimFunc& func); - /*! \brief Get the Buffer object representing the mapped access into - * the pool. - */ - Buffer GetRemappedBuffer(Buffer buf); - - /*! \brief The tir::Var map to PoolInfo objects */ - Map primfunc_args_to_pool_info_map_; - /*! \brief The buffer var map to their allocate nodes */ - Map allocate_var_to_stmt_map_; - /*! \brief The IRModule being constructed/mutated */ - IRModule module_; - /*! \brief The input allocate node to PoolAllocation map */ - Map pool_allocations_; - /*! \brief The set of ordered pools to ensure an unique order of args for functions */ - std::vector allocated_pool_ordering_; - /*! \brief The storage of calculated pool size at init */ - std::unordered_map all_pools_sizes_; - /*! \brief After mutation, each allocate buffer is replaced with tir::Var that is let bounded - * to position from a pool as designated by a PoolAllocation - */ - Map allocate_var_to_let_var_; - /*! \brief A map from the original buffer object - * - * Each key-value pair in this map satisfies - * ``allocate_buf_to_let_var[key->data] = value->data``. However, - * since more than one `tir::Buffer` may use the same Var, they must - * be tracked separately. - */ - Map original_buf_to_let_buf_; - - Map signature_has_device_context_; - /*! \brief A counter to give references to pools a reproducible unique set of names */ - int pool_var_count_ = 0; - /*! \brief This toggles to remove non tvmscript printable items for IRModule for unit tests */ - bool emit_tvmscript_printable_ = false; - /*! \brief A counter to give references to pools a reproducible unique set of names */ - std::unordered_set visited_primfuncs; - - Map> pool_initializations_; - void AppdendConstInitializationData(ScopeInfo si); -}; - -Optional PoolAllocationToOffsetConverter::GetResourceHandle(const PrimFunc& func) { - if (!func->params.empty() && - func->buffer_map.find(func->params.back()) == func->buffer_map.end()) { - return func->params.back(); - } - return Optional(); -} - -PoolAllocationToOffsetConverter::ScopeInfo PoolAllocationToOffsetConverter::UpdateFunctionScopeInfo( - const PrimFunc& original_func) { - ScopeInfo si; - - Optional resource_handle = GetResourceHandle(original_func); - si.params = original_func->params; - if (resource_handle) { - si.params.pop_back(); - ICHECK(si.params.size() == original_func->params.size() - 1); - } - si.buffer_map = original_func->buffer_map; - Map ret; - for (const AllocatedPoolInfo& allocated_pool_info : allocated_pool_ordering_) { - PoolInfo pool_info = allocated_pool_info->pool_info; - String pool_ref_name = pool_info->pool_name + "_" + std::to_string(pool_var_count_++); - String var_name = pool_ref_name + "_var"; - DataType elem_dtype = DataType::UInt(8); - Var buffer_var(var_name, PointerType(PrimType(elem_dtype), "global")); - Var pool_var = Var(var_name, PointerType(PrimType(elem_dtype), "global")); - si.params.push_back(pool_var); - si.pools_to_params.Set(pool_info, pool_var); - si.allocated_pool_params.push_back(AllocatedPoolInfo( - allocated_pool_info->pool_info, allocated_pool_info->allocated_size, si.params.size() - 1)); - - int pool_size = all_pools_sizes_[pool_info]; - String buffer_var_name = pool_ref_name + "_buffer_var"; - si.buffer_map.Set(pool_var, - Buffer(buffer_var /* data */, elem_dtype /* dtype */, {pool_size} /* shape */, - {1} /* strides */, 0 /* elem_offset */, buffer_var_name /* name */, - 16 /* data_alignment */, 1 /* offset_factor */, - BufferType::kDefault /* buffer-type */)); - } - if (resource_handle) { - si.params.push_back(resource_handle.value()); - } - return si; -} - -PrimFunc PoolAllocationToOffsetConverter::CreatePrimFuncWithPoolParams( - const PrimFunc& original_primfunc) { - // Only create the new function if it was not modified with pool params - if (visited_primfuncs.find(original_primfunc) == visited_primfuncs.end()) { - ScopeInfo si = UpdateFunctionScopeInfo(original_primfunc); - this->scope_stack.push(si); - Stmt new_body = this->VisitStmt(original_primfunc->body); - this->scope_stack.pop(); - DictAttrs original_attrs = original_primfunc->attrs; - // We dont need attrs of PrimFunc that might include non printable attrs such as target - // for unit tests where emit_tvmscript_printable_ is to be used. - if (emit_tvmscript_printable_) { - // keep global symbol if it's there because it determines if the private attribute is printed - if (original_attrs->dict.count(tvm::attr::kGlobalSymbol)) { - original_attrs = DictAttrs( - {{tvm::attr::kGlobalSymbol, original_attrs->dict.at(tvm::attr::kGlobalSymbol)}}); - } else { - original_attrs = DictAttrs(); - } - } - PrimFunc ret = - PrimFunc(si.params, new_body, original_primfunc->ret_type, si.buffer_map, original_attrs); - if (!emit_tvmscript_printable_) { - ret = WithAttr(ret, tvm::attr::kPoolArgs, si.allocated_pool_params); - } - visited_primfuncs.insert(ret); - return ret; - } - return original_primfunc; -} - -Array PoolAllocationToOffsetConverter::AppendPoolParamsToArgs(Array args, - bool has_device_context) { - Array new_args; - PrimExpr resource_handle_arg; - // name, params...params[, context] - if (has_device_context) { - resource_handle_arg = args.back(); - args.pop_back(); - } - for (const auto& arg : args) { - new_args.push_back(VisitExpr(arg)); - } - ScopeInfo top_scope = this->scope_stack.top(); - for (const auto& pools_vars : top_scope.pools_to_params) { - tir::Var pool_var = pools_vars.second; - Buffer buffer_var = top_scope.buffer_map[pool_var]; - new_args.push_back(buffer_var->data); - } - if (resource_handle_arg.defined()) { - new_args.push_back(resource_handle_arg); - } - return new_args; -} - -Array PoolAllocationToOffsetConverter::ReplaceAllocateArgsWithLetArgs( - const Array& args) { - Array ret; - for (const PrimExpr& arg : args) { - if (arg->IsInstance() && - allocate_var_to_let_var_.find(Downcast(arg)) != allocate_var_to_let_var_.end()) { - ret.push_back(allocate_var_to_let_var_[Downcast(arg)]); - } else { - ret.push_back(VisitExpr(arg)); - } - } - return ret; -} - -PrimExpr PoolAllocationToOffsetConverter::VisitExpr_(const CallNode* op) { - if (op->op.same_as(builtin::call_extern()) || op->op.same_as(builtin::tvm_call_cpacked())) { - String func_name = Downcast(op->args[0])->value; - Array new_args; - if (module_->ContainGlobalVar(func_name) && - module_->Lookup(func_name)->IsInstance()) { - GlobalVar gv = module_->GetGlobalVar(func_name); - PrimFunc func = Downcast(module_->Lookup(gv)); - - if (!signature_has_device_context_.count(func_name)) { - if (op->args.size() == func->params.size() + 2) { - signature_has_device_context_.Set(func_name, Bool(true)); - } else { - signature_has_device_context_.Set(func_name, Bool(false)); - } - } - - PrimFunc prim_func = CreatePrimFuncWithPoolParams(func); - module_->Update(gv, prim_func); - new_args = AppendPoolParamsToArgs(op->args, signature_has_device_context_[func_name]); - new_args = ReplaceAllocateArgsWithLetArgs(new_args); - } else { - new_args = ReplaceAllocateArgsWithLetArgs(op->args); - } - return Call(op->dtype, op->op, new_args); - } - if (op->op->IsInstance()) { - String func_name = Downcast(op->args[0])->value; - PrimFunc func = Downcast(op->op); - PrimFunc prim_func = CreatePrimFuncWithPoolParams(func); - Array new_args = - AppendPoolParamsToArgs(op->args, signature_has_device_context_[func_name]); - new_args = ReplaceAllocateArgsWithLetArgs(new_args); - return Call(op->dtype, prim_func, new_args); - } - return StmtExprMutator::VisitExpr_(op); -} - -LetStmt PoolAllocationToOffsetConverter::ToLetStmt(const PoolAllocation& pool_allocation, - const Var& buffer_var, const Stmt& body) { - ScopeInfo scope_info = scope_stack.top(); - Var param = scope_info.pools_to_params[pool_allocation->pool_info]; - BufferLoad load_node = BufferLoad(scope_info.buffer_map[param], {pool_allocation->byte_offset}); - Call address_of_load = Call(DataType::Handle(), builtin::address_of(), {load_node}); - - Type let_var_type = buffer_var->type_annotation; - if (emit_tvmscript_printable_) { - // Strip the storage_scope from the variable type, as TVMScript - // doesn't parsethe scoped pointers (e.g. ``T.Ptr[global T.int32]``) - // correctly. - let_var_type = PointerType(Downcast(let_var_type)->element_type); - } - Var let_var(buffer_var->name_hint + "_let", let_var_type); - allocate_var_to_let_var_.Set(buffer_var, let_var); - Stmt new_body = VisitStmt(body); - allocate_var_to_let_var_.erase(buffer_var); - return LetStmt(let_var, address_of_load, new_body); -} - -Stmt PoolAllocationToOffsetConverter::VisitStmt_(const AllocateNode* op) { - if (pool_allocations_.count(GetRef(op))) { - return ToLetStmt(pool_allocations_[GetRef(op)], op->buffer_var, op->body); - } - return StmtExprMutator::VisitStmt_(op); -} - -Stmt PoolAllocationToOffsetConverter::VisitStmt_(const AllocateConstNode* op) { - if (pool_allocations_.count(GetRef(op))) { - const auto& result = ToLetStmt(pool_allocations_[GetRef(op)], op->buffer_var, op->body); - - PoolInfo pool_info = pool_allocations_[GetRef(op)]->pool_info; - if (pool_initializations_.find(pool_info) == pool_initializations_.end()) { - pool_initializations_.Set(pool_info, {}); - } - - auto consts = pool_initializations_[pool_info]; - consts.push_back({result->var->name_hint, pool_allocations_[GetRef(op)]->byte_offset, - op->data.value()}); - - pool_initializations_.Set(pool_info, consts); - return result; - } - return StmtExprMutator::VisitStmt_(op); -} - -Stmt PoolAllocationToOffsetConverter::VisitStmt_(const BufferStoreNode* op) { - BufferStore store = Downcast(StmtExprMutator::VisitStmt_(op)); - - Buffer remapped = GetRemappedBuffer(store->buffer); - if (!op->buffer.same_as(remapped)) { - store.CopyOnWrite()->buffer = remapped; - } - return std::move(store); -} - -Stmt PoolAllocationToOffsetConverter::VisitStmt_(const DeclBufferNode* op) { - auto decl = Downcast(StmtExprMutator::VisitStmt_(op)); - - Buffer remapped = GetRemappedBuffer(decl->buffer); - if (!op->buffer.same_as(remapped)) { - decl.CopyOnWrite()->buffer = remapped; - } - return std::move(decl); -} - -PrimExpr PoolAllocationToOffsetConverter::VisitExpr_(const BufferLoadNode* op) { - BufferLoad load = Downcast(StmtExprMutator::VisitExpr_(op)); - - Buffer remapped = GetRemappedBuffer(load->buffer); - if (!op->buffer.same_as(remapped)) { - load.CopyOnWrite()->buffer = remapped; - } - return std::move(load); -} - -PrimExpr PoolAllocationToOffsetConverter::VisitExpr_(const VarNode* op) { - auto it = allocate_var_to_let_var_.find(GetRef(op)); - if (it != allocate_var_to_let_var_.end()) { - return (*it).second; - } - - return StmtExprMutator::VisitExpr_(op); -} - -Buffer PoolAllocationToOffsetConverter::GetRemappedBuffer(Buffer original) { - { - auto it = original_buf_to_let_buf_.find(original); - if (it != original_buf_to_let_buf_.end()) { - return (*it).second; - } - } - - Buffer remapped = original; - - auto it = allocate_var_to_let_var_.find(original->data); - if (it != allocate_var_to_let_var_.end()) { - remapped = Buffer((*it).second, original->dtype, original->shape, original->strides, - original->elem_offset, original->name, original->data_alignment, - original->offset_factor, original->buffer_type, original->axis_separators, - original->span); - } - - original_buf_to_let_buf_.Set(original, remapped); - return remapped; -} - -void PoolAllocationToOffsetConverter::AppdendConstInitializationData( - PoolAllocationToOffsetConverter::ScopeInfo si) { - for (AllocatedPoolInfo api : si.allocated_pool_params) { - const auto& it = pool_initializations_.find(api->pool_info); - if (it != pool_initializations_.end()) { - auto* pi = const_cast(api->pool_info.as()); - pi->constant_info_array = (*it).second; - } - } -} - -IRModule PoolAllocationToOffsetConverter::operator()() { - GlobalVar gv = module_->GetGlobalVar(::tvm::runtime::symbol::tvm_module_main); - PrimFunc main_func = Downcast(module_->Lookup(gv)); - ScopeInfo si = UpdateFunctionScopeInfo(main_func); - this->scope_stack.push(si); - Stmt main_func_body = this->VisitStmt(main_func->body); - this->scope_stack.pop(); - AppdendConstInitializationData(si); - // We dont need attrs of PrimFunc that might include non printable attrs such as target - // for unit tests where emit_tvmscript_printable_ is to be used. - if (!emit_tvmscript_printable_) { - main_func = - PrimFunc(si.params, main_func_body, main_func->ret_type, si.buffer_map, main_func->attrs); - main_func = WithAttr(main_func, tvm::attr::kPoolArgs, si.allocated_pool_params); - } else { - auto new_attrs = DictAttrs(); - if (main_func->attrs->dict.count(tvm::attr::kGlobalSymbol)) { - new_attrs = DictAttrs( - {{tvm::attr::kGlobalSymbol, main_func->attrs->dict.at(tvm::attr::kGlobalSymbol)}}); - } - main_func = PrimFunc(si.params, main_func_body, main_func->ret_type, si.buffer_map, new_attrs); - } - module_->Update(gv, main_func); - if (!emit_tvmscript_printable_) { - return WithAttr(this->module_, tvm::attr::kPoolArgs, si.allocated_pool_params); - } - return this->module_; -} - -namespace transform { - -tvm::transform::Pass ConvertPoolAllocationsToOffsets( - const Map& pool_allocations, Bool emit_tvmscript_printable) { - auto pass_func = [=](IRModule m, tvm::transform::PassContext ctx) { - return Downcast(PoolAllocationToOffsetConverter( - m, pool_allocations, emit_tvmscript_printable->value != 0)()); - }; - return tvm::transform::CreateModulePass(pass_func, 0, "tir.usmp.ConvertPoolAllocationsToOffsets", - {}); -} - -TVM_REGISTER_GLOBAL("tir.usmp.transform.ConvertPoolAllocationsToOffsets") - .set_body_typed(ConvertPoolAllocationsToOffsets); - -} // namespace transform - -} // namespace usmp -} // namespace tir -} // namespace tvm diff --git a/src/tir/usmp/transform/create_io_allocates.cc b/src/tir/usmp/transform/create_io_allocates.cc deleted file mode 100644 index ca06095f0bdc..000000000000 --- a/src/tir/usmp/transform/create_io_allocates.cc +++ /dev/null @@ -1,212 +0,0 @@ -/* - * 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 - -namespace tvm { -namespace tir { -namespace usmp { - -/*! \brief Creates Allocate nodes with special annotations - * for I/O tensors in the graph to be memory planned.*/ -class IOAllocateCreator : public StmtExprVisitor { - public: - explicit IOAllocateCreator(const IRModule& module) { - main_func_ = Downcast(module->Lookup(::tvm::runtime::symbol::tvm_module_main)); - ICHECK(main_func_.defined()) << "main function is not in the module"; - for (const auto& gv_func : module->functions) { - if (gv_func.second->IsInstance()) { - functions_.Set(gv_func.first->name_hint, Downcast(gv_func.second)); - } - } - mod_ = module->ShallowCopy(); - } - IRModule operator()(); - - private: - void VisitExpr_(const BufferLoadNode* op) override; - void VisitExpr_(const CallNode* op) override; - void VisitStmt_(const BufferStoreNode* op) override; - - /*! \brief Updates aliases that buffer vars inside the primfunc refer - * to in terms call arguments they get bound to.*/ - void UpdateAliases(const Array& args, const PrimFunc& func); - - /*! \brief The IRModule that is being mutated */ - IRModule mod_; - /*! \brief The main function that calls into operator subgraphs */ - PrimFunc main_func_; - /*! \brief The input Vars of the main function */ - std::unordered_set inputs_; - /*! \brief The output Vars of the main function */ - std::unordered_set outputs_; - /*! \brief The buffer vars associated with the I/O Vars */ - std::unordered_set io_buffer_vars_; - /*! \brief The aliases that buffer vars inside the primfunc refer - * to in terms call arguments */ - std::unordered_map aliases_; - /*! - * \brief The TIR main function calls by name to PrimFuncs to be able to - * support BYOC. Therefore, this Map records functions that are present - * in the IRModule by name/ - */ - Map functions_; -}; - -/*! - * \brief The function obtains the matched buffer vars for - * the params of the PrimFunc. - */ -Array static GetMatchedBuffers(const PrimFunc& func) { - Array buffer_vars; - for (unsigned int i = 0; i < func->params.size() - 1; i++) { - Var param = func->params[i]; - buffer_vars.push_back(func->buffer_map[param]->data); - } - Var last_param = func->params.back(); - // Checks whether last var is present in the buffer map - // because it could be the resource handle - if (func->buffer_map.find(last_param) != func->buffer_map.end()) { - buffer_vars.push_back(func->buffer_map[last_param]->data); - } - return buffer_vars; -} - -/*! - * \brief The function updates aliases that each buffer var with its - * associated argument in the callsite. - */ -void IOAllocateCreator::UpdateAliases(const Array& args, const PrimFunc& func) { - auto param_buffers = GetMatchedBuffers(func); - // Last var could be a resource handle that does not have a Buffer - ICHECK(args.size() == param_buffers.size() || args.size() - 1 == param_buffers.size()); - for (size_t i = 0; i < param_buffers.size(); i++) { - auto arg = args[i]; - if (arg->IsInstance()) { - auto param_buf = param_buffers[i]; - aliases_[param_buf] = Downcast(arg); - } - } -} - -void IOAllocateCreator::VisitExpr_(const CallNode* op) { - if (op->op.same_as(builtin::call_extern()) || op->op.same_as(builtin::tvm_call_cpacked())) { - StringImm func_name = Downcast(op->args[0])->value; - if (functions_.find(func_name->value) != functions_.end()) { - auto func = functions_.at(func_name->value); - auto actual_args = Array(op->args.begin() + 1, op->args.end()); - this->UpdateAliases(actual_args, func); - VisitStmt(func->body); - return; - } - } - if (op->op->IsInstance()) { - auto func = Downcast(op->op); - this->UpdateAliases(op->args, func); - VisitStmt(func->body); - return; - } - StmtExprVisitor::VisitExpr_(op); -} - -void IOAllocateCreator::VisitExpr_(const BufferLoadNode* op) { - if (aliases_.find(op->buffer->data) != aliases_.end()) { - Var aliased_var = aliases_[op->buffer->data]; - if (io_buffer_vars_.find(aliased_var) != io_buffer_vars_.end()) { - ICHECK(outputs_.find(aliased_var) == outputs_.end()) - << "BufferLoad nodes should not be reading from output buffer vars."; - inputs_.insert(aliased_var); - } - } - StmtExprVisitor::VisitExpr_(op); -} - -void IOAllocateCreator::VisitStmt_(const BufferStoreNode* op) { - if (aliases_.find(op->buffer->data) != aliases_.end()) { - Var aliased_var = aliases_[op->buffer->data]; - if (io_buffer_vars_.find(aliased_var) != io_buffer_vars_.end()) { - ICHECK(inputs_.find(aliased_var) == inputs_.end()) - << "BufferStore nodes should not be writing to input buffer vars."; - outputs_.insert(aliased_var); - } - } - StmtExprVisitor::VisitStmt_(op); -} - -IRModule IOAllocateCreator::operator()() { - Array new_main_params; - Stmt main_body = main_func_->body; - for (const Var& param : main_func_->params) { - if (main_func_->buffer_map.find(param) != main_func_->buffer_map.end()) { - Var buffer_var = main_func_->buffer_map[param]->data; - io_buffer_vars_.insert(buffer_var); - aliases_[buffer_var] = buffer_var; - } - } - VisitStmt(main_body); - ICHECK(io_buffer_vars_.size() == inputs_.size() + outputs_.size()) - << "Every IO Buffer var should be categorized either to be input or output"; - for (const Var& param : main_func_->params) { - if (main_func_->buffer_map.find(param) != main_func_->buffer_map.end()) { - Buffer param_buffer = main_func_->buffer_map[param]; - String io_annotation; - if (inputs_.find(param_buffer->data) != inputs_.end()) { - io_annotation = String(kInputTensorAllocate); - } else { - io_annotation = String(kOutputTensorAllocate); - } - main_body = Allocate(param_buffer->data, param_buffer->dtype, param_buffer->shape, - const_true(), main_body, {{io_annotation, param->name_hint}}); - } else { - new_main_params.push_back(param); - } - } - const GlobalVar& gv = mod_->GetGlobalVar(::tvm::runtime::symbol::tvm_module_main); - mod_->Update(gv, PrimFunc(new_main_params, main_body, main_func_->ret_type, - main_func_->buffer_map, main_func_->attrs, main_func_->span)); - return mod_; -} - -namespace transform { - -tvm::transform::Pass CreateAllocatesForIO() { - auto pass_func = [=](IRModule m, tvm::transform::PassContext ctx) { - return IOAllocateCreator(m)(); - }; - return tvm::transform::CreateModulePass(pass_func, 0, "tir.usmp.CreateAllocatesForIO", {}); -} - -TVM_REGISTER_GLOBAL("tir.usmp.transform.CreateAllocatesForIO").set_body_typed(CreateAllocatesForIO); - -} // namespace transform - -} // namespace usmp -} // namespace tir -} // namespace tvm diff --git a/src/tir/usmp/unified_static_memory_planner.cc b/src/tir/usmp/unified_static_memory_planner.cc deleted file mode 100644 index 60030c1595d9..000000000000 --- a/src/tir/usmp/unified_static_memory_planner.cc +++ /dev/null @@ -1,140 +0,0 @@ -/* - * 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 tir/analysis/usmp/unified_static_memory_planner.cc - * \brief This is the pass that integrates the USMP passes to - * a single composite pass. - */ - -#include -#include -#include -#include -#include -#include -#include -#include -#include - -#include -#include - -namespace tvm { - -TVM_REGISTER_PASS_CONFIG_OPTION(kUSMPEnableOption, Bool); -TVM_REGISTER_PASS_CONFIG_OPTION(kUSMPAlgorithmOption, String); -TVM_REGISTER_PASS_CONFIG_OPTION(kUSMPUseWorkspaceIO, Bool); -TVM_REGISTER_PASS_CONFIG_OPTION(kUSMPCustomAlgorithmOption, String); - -namespace tir { -namespace usmp { - -static constexpr const char* kDefaultAlgo = "greedy_by_size"; - -static std::unordered_map( - const Array&, const Integer&)>> - algorithms{{"greedy_by_size", algo::GreedyBySize}, - {"greedy_by_conflicts", algo::GreedyByConflicts}, - {"hill_climb", algo::HillClimb}}; - -IRModule PlanMemory(const IRModule& mod, String algo, bool use_workspace_io, - Optional opt_custom_algo) { - VLOG(1) << "workspace required = " << CalculateModuleWorkspaceSize(mod); - IRModule module = mod->ShallowCopy(); - if (use_workspace_io) { - module = transform::CreateAllocatesForIO()(module); - } - module = transform::AssignPoolInfo()(module); - PrimFunc main_func = Downcast(module->Lookup(::tvm::runtime::symbol::tvm_module_main)); - BufferInfoAnalysis buffer_info_analysis = ExtractBufferInfo(main_func, module); - Array buffer_info_arr = - ConvertToArrayOfBufferInfo(buffer_info_analysis->buffer_info_stmts); - decltype(algorithms)::mapped_type algorithm; - if (opt_custom_algo) { - String algo_func_name = "tir.usmp.algo." + opt_custom_algo.value(); - const runtime::PackedFunc* pfAlgo = runtime::Registry::Get(algo_func_name); - CHECK(pfAlgo) << "The selected custom USMP algorithm : " << opt_custom_algo.value() - << " is not defined. Please register it as " << algo_func_name; - algorithm = *pfAlgo; - } else { - CHECK(algorithms.count(algo)) - << "The selected USMP algorithm : " << algo - << " is not defined. Please define it in the above algorithms map."; - algorithm = algorithms[algo]; - } - Map buffer_info_pool_allocations = - algorithm(buffer_info_arr, buffer_info_analysis->memory_pressure); - - Map stmt_pool_allocations = AssignStmtPoolAllocations( - buffer_info_analysis->buffer_info_stmts, buffer_info_pool_allocations); - - module = transform::ConvertPoolAllocationsToOffsets(stmt_pool_allocations)(module); - if (use_workspace_io) { - Map io_pool_allocations = - GetIOPoolAllocations(buffer_info_pool_allocations); - module = WithAttr(module, tvm::attr::kIOTensorPoolAllocations, io_pool_allocations); - } - tir::PrimFunc tir_main_func = - Downcast(module->Lookup(::tvm::runtime::symbol::tvm_module_main)); - Optional> allocated_pool_infos = - tir_main_func->GetAttr>(tvm::attr::kPoolArgs); - if (allocated_pool_infos) { - for (const tir::usmp::AllocatedPoolInfo& allocated_pool_info : allocated_pool_infos.value()) { - VLOG(1) << "pool_size = " << allocated_pool_info->allocated_size; - } - } - return module; -} - -} // namespace usmp - -namespace transform { - -tvm::transform::Pass UnifiedStaticMemoryPlanner() { - auto usmp_main_pass_func = [=](IRModule m, tvm::transform::PassContext ctx) { - auto algorithm_str = ctx->GetConfig(kUSMPAlgorithmOption, String(usmp::kDefaultAlgo)); - auto use_workspace_io = ctx->GetConfig(kUSMPUseWorkspaceIO, Bool(false)); - auto custom_algorithm_str = ctx->GetConfig(kUSMPCustomAlgorithmOption); - tvm::relay::Executor executor_config = - m->GetAttr(tvm::attr::kExecutor).value(); - String interface_api = executor_config->GetAttr("interface-api").value_or("packed"); - tvm::relay::Runtime runtime_config = - m->GetAttr(tvm::attr::kRuntime).value(); - if (use_workspace_io.value()) { - CHECK(interface_api == "c") << kUSMPUseWorkspaceIO - << " option is only compatible with interface_api c.\n" - << "Please use interface_api c to be able to enable " - << kUSMPUseWorkspaceIO << "\n"; - } - return Downcast( - usmp::PlanMemory(m, algorithm_str.value_or(String(usmp::kDefaultAlgo)), - use_workspace_io.value_or(Bool(false)), custom_algorithm_str)); - }; - - return tvm::transform::CreateModulePass(usmp_main_pass_func, 0, - "tir.transform.UnifiedStaticMemoryPlanner", {}); -} - -TVM_REGISTER_GLOBAL("tir.transform.UnifiedStaticMemoryPlanner") - .set_body_typed(UnifiedStaticMemoryPlanner); - -} // namespace transform -} // namespace tir -} // namespace tvm diff --git a/src/tir/usmp/utils.cc b/src/tir/usmp/utils.cc deleted file mode 100644 index d640e9fa073e..000000000000 --- a/src/tir/usmp/utils.cc +++ /dev/null @@ -1,280 +0,0 @@ -/* - * 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 tir/usmp/utils.cc - * \brief Utilities for Unified Static Memory Planner - */ - -#include -#include -#include -#include -#include -#include -#include -#include -#include - -namespace tvm { -namespace tir { -namespace usmp { - -BufferInfo::BufferInfo(String name_hint, Integer size_bytes, Array pool_candidates, - Integer alignment, BufferInfoKind kind) { - auto bufinfo_node = make_object(); - bufinfo_node->name_hint = name_hint; - bufinfo_node->size_bytes = size_bytes; - bufinfo_node->pool_candidates = pool_candidates; - bufinfo_node->alignment = alignment; - bufinfo_node->kind = kind; - data_ = std::move(bufinfo_node); -} - -void BufferInfoNode::SetConflicts(Array conflicting_buffer_info_objs) { - this->conflicts = conflicting_buffer_info_objs; -} - -TVM_REGISTER_NODE_TYPE(BufferInfoNode); -TVM_REGISTER_GLOBAL("tir.usmp.BufferInfo") - .set_body_typed([](String name_hint, Integer size_bytes, Array pool_candidates, - Integer alignment) { - if (!alignment.defined()) { - return BufferInfo(name_hint, size_bytes, pool_candidates); - } - return BufferInfo(name_hint, size_bytes, pool_candidates, alignment); - }); -TVM_REGISTER_GLOBAL("tir.usmp.BufferInfoSetConflicts") - .set_body_method(&BufferInfoNode::SetConflicts); - -TVM_STATIC_IR_FUNCTOR(ReprPrinter, vtable) - .set_dispatch([](const ObjectRef& ref, ReprPrinter* p) { - auto* node = static_cast(ref.get()); - std::unordered_map toString = { - {BufferInfoKind::kIntermediate, "kIntermediate"}, - {BufferInfoKind::kInput, "kInput"}, - {BufferInfoKind::kOutput, "kOutput"}}; - p->stream << "BufferInfoNode(\n" - << "name_hint=" << node->name_hint << ",\n size_bytes=" << node->size_bytes - << ",\n pool_candidates=" << node->pool_candidates - << ",\n alignment=" << node->alignment << ",\n kind=" << toString[node->kind] - << ",\n conflicts=" << node->conflicts.size() << ")"; - }); - -BufferInfoAnalysis::BufferInfoAnalysis(Map buffer_info_stmts, - Integer memory_pressure) { - auto bufinfo_analysis_node = make_object(); - bufinfo_analysis_node->buffer_info_stmts = buffer_info_stmts; - bufinfo_analysis_node->memory_pressure = memory_pressure; - data_ = std::move(bufinfo_analysis_node); -} - -TVM_REGISTER_NODE_TYPE(BufferInfoAnalysisNode); -TVM_REGISTER_GLOBAL("tir.usmp.BufferInfoAnalysis") - .set_body_typed([](Map buffer_info_stmts, Integer memory_pressure) { - return BufferInfoAnalysis(buffer_info_stmts, memory_pressure); - }); - -TVM_STATIC_IR_FUNCTOR(ReprPrinter, vtable) - .set_dispatch([](const ObjectRef& ref, ReprPrinter* p) { - auto* node = static_cast(ref.get()); - p->stream << "BufferInfoAnalysisNode(\n" - << "buffer_info_stmts=" << node->buffer_info_stmts - << ",\n memory_pressure=" << node->memory_pressure << ")"; - }); - -PoolAllocation::PoolAllocation(PoolInfo pool_info, Integer byte_offset) { - auto pool_allocation_node = make_object(); - pool_allocation_node->pool_info = pool_info; - pool_allocation_node->byte_offset = byte_offset; - data_ = std::move(pool_allocation_node); -} - -TVM_REGISTER_NODE_TYPE(PoolAllocationNode); -TVM_REGISTER_GLOBAL("tir.usmp.PoolAllocation") - .set_body_typed([](PoolInfo pool_info, Integer byte_offset) { - return PoolAllocation(pool_info, byte_offset); - }); - -TVM_STATIC_IR_FUNCTOR(ReprPrinter, vtable) - .set_dispatch([](const ObjectRef& ref, ReprPrinter* p) { - auto* node = static_cast(ref.get()); - p->stream << "PoolAllocationNode(\n" - << "pool_info=" << node->pool_info << ",\n byte_offset=" << node->byte_offset - << ")"; - }); - -AllocatedPoolInfo::AllocatedPoolInfo(PoolInfo pool_info, Integer allocated_size, - Integer pool_var_idx) { - auto allocated_poolinfo_node = make_object(); - allocated_poolinfo_node->pool_info = pool_info; - allocated_poolinfo_node->allocated_size = allocated_size; - if (pool_var_idx.defined()) { - allocated_poolinfo_node->pool_var_idx = pool_var_idx; - } - data_ = std::move(allocated_poolinfo_node); -} - -TVM_REGISTER_NODE_TYPE(AllocatedPoolInfoNode); -TVM_REGISTER_GLOBAL("ir.AllocatedPoolInfo") - .set_body_typed([](PoolInfo pool_info, Integer allocated_size, Integer pool_var_idx) { - return AllocatedPoolInfo(pool_info, allocated_size, pool_var_idx); - }); - -TVM_STATIC_IR_FUNCTOR(ReprPrinter, vtable) - .set_dispatch([](const ObjectRef& ref, ReprPrinter* p) { - auto* node = static_cast(ref.get()); - p->stream << "AllocatedPoolInfoNode(\n" - << "pool_info=" << node->pool_info << ",\n allocated_size=" << node->allocated_size - << ")"; - }); - -Array ConvertToArrayOfBufferInfo(const Map& buffer_info_map) { - Array ret; - for (const auto& kv : buffer_info_map) { - auto buffer_info = kv.first; - ret.push_back(buffer_info); - } - return ret; -} - -Map AssignStmtPoolAllocations( - const Map& buffer_info_to_stmt, - const Map& buffer_info_to_pool_allocation) { - Map ret; - for (const auto& kv : buffer_info_to_pool_allocation) { - BufferInfo bi = kv.first; - Stmt stmt_ = buffer_info_to_stmt[bi]; - PoolAllocation pa = kv.second; - ret.Set(stmt_, pa); - } - return ret; -} - -Map GetIOPoolAllocations( - const Map& buffer_info_to_pool_allocation) { - Map io_tensor_name_to_pool_allocation; - for (const auto& kv : buffer_info_to_pool_allocation) { - BufferInfo buffer_info = kv.first; - PoolAllocation pool_allocation = kv.second; - if (buffer_info->kind != BufferInfoKind::kIntermediate) { - io_tensor_name_to_pool_allocation.Set(buffer_info->name_hint, pool_allocation); - } - } - return io_tensor_name_to_pool_allocation; -} - -static Integer CalculateExtentsSize(const DataType& dtype, const Array& extents) { - if (dtype.is_scalable_vector()) { - // We cannot statically calculate workspace for scalable types - return Integer(); - } - size_t element_size_bytes = dtype.bytes() * dtype.lanes(); - size_t num_elements = 1; - for (const auto& ext : extents) { - if (ext->IsInstance()) { - num_elements *= Downcast(ext)->value; - } else { - // We can't statically calculate workspace for dynamic shapes - return Integer(); - } - } - return Integer(num_elements * element_size_bytes); -} - -Integer CalculateExtentsSize(const AllocateNode* op) { - return CalculateExtentsSize(op->dtype, op->extents); -} - -Integer CalculateExtentsSize(const AllocateConstNode* op) { - return CalculateExtentsSize(op->dtype, op->extents); -} - -class ModuleWorkspaceSizeCalculator : public StmtExprVisitor { - public: - explicit ModuleWorkspaceSizeCalculator(const IRModule& module) : mod_(module) { - for (const auto& gv_func : mod_->functions) { - if ((gv_func.second)->IsInstance()) { - functions_.Set(gv_func.first->name_hint, Downcast(gv_func.second)); - } - } - main_func_ = Downcast(module->Lookup(::tvm::runtime::symbol::tvm_module_main)); - ICHECK(main_func_.defined()) << "main function is not in the module"; - Optional target_host = main_func_->GetAttr(tvm::attr::kTarget); - ICHECK(target_host) << "main function does not have a target attr"; - target_host_ = target_host.value(); - } - - Integer operator()() { - UpdateWorkspaceData(main_func_); - return Integer(max_workspace_size); - } - - private: - void UpdateWorkspaceData(const PrimFunc& func) { - Target tgt = func->GetAttr(tvm::attr::kTarget).value_or(target_host_); - Integer workspace_byte_alignment = - tgt->GetAttr("workspace-byte-alignment").value_or(16); - Integer workspace_req = CalculateWorkspaceBytes(func, workspace_byte_alignment); - if (workspace_req.IntValue() != 0) { - current_workspace_size_ += workspace_req->value; - } - if (max_workspace_size < current_workspace_size_) { - max_workspace_size = current_workspace_size_; - } - this->VisitStmt(func->body); - if (workspace_req.IntValue() != 0) { - current_workspace_size_ -= workspace_req->value; - } - } - - void VisitExpr_(const CallNode* op) override { - if (op->op.same_as(builtin::call_extern())) { - PrimFunc func = functions_.at(Downcast(op->args[0])->value); - UpdateWorkspaceData(func); - } else if (op->op->IsInstance()) { - PrimFunc func = Downcast(op->op); - UpdateWorkspaceData(func); - } else { - StmtExprVisitor::VisitExpr_(op); - } - } - - IRModule mod_; - Target target_host_; - PrimFunc main_func_; - Map functions_; - size_t current_workspace_size_ = 0; - size_t max_workspace_size = 0; -}; - -Integer CalculateModuleWorkspaceSize(const IRModule& mod) { - return ModuleWorkspaceSizeCalculator(mod)(); -} - -TVM_REGISTER_GLOBAL("tir.usmp.CreateArrayBufferInfo") - .set_body_typed([](Map buffer_info_map) { - return (ConvertToArrayOfBufferInfo(buffer_info_map)); - }); - -TVM_REGISTER_GLOBAL("tir.usmp.AssignStmtPoolAllocations").set_body_typed(AssignStmtPoolAllocations); - -} // namespace usmp -} // namespace tir -} // namespace tvm diff --git a/tests/cpp-runtime/opencl/texture_copy_test.cc b/tests/cpp-runtime/opencl/texture_copy_test.cc index 23b490f695e2..701fec4d8baf 100644 --- a/tests/cpp-runtime/opencl/texture_copy_test.cc +++ b/tests/cpp-runtime/opencl/texture_copy_test.cc @@ -17,8 +17,8 @@ * under the License. */ -#include #include +#include #include #include diff --git a/src/relay/collage/README.md b/tests/cpp/README.md similarity index 70% rename from src/relay/collage/README.md rename to tests/cpp/README.md index 945a775e383d..884054809e90 100644 --- a/src/relay/collage/README.md +++ b/tests/cpp/README.md @@ -14,13 +14,10 @@ +# tests/cpp -The `CollagePartition` pass for finding optimal partitionings of Relay models. +This folder contains some unit tests for C++ utilities in the codebase. -See the [RFC](https://github.com/mbs-octoml/mbs-tvm-rfcs/blob/mbs-rfcs-collage/rfcs/xxxx-collage.md). - -Based on: -> *Collage: Automated Integration of Deep Learning Backends* -> Byungsoo Jeon, Sunghyun Park, Peiyuan Liao, Sheng Xu, Tianqi Chen, Zhihao Jia - -CAUTION: This is a prototype, do not use in prod. +In principle we aim to do most compiler related tests in the python to +bring more development velocity, and only use this folder for low-level unit-tests. +All tests should finish fast and not dependent on a presence of an accelerator device. diff --git a/tests/cpp/aot_metadata_test.cc b/tests/cpp/aot_metadata_test.cc deleted file mode 100644 index 7d280538e0a3..000000000000 --- a/tests/cpp/aot_metadata_test.cc +++ /dev/null @@ -1,425 +0,0 @@ -/* - * 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 "../src/target/metadata.h" -#include "../src/target/metadata_utils.h" - -namespace { - -const int64_t kNormalInput1Shape[4] = {1, 5, 5, 3}; -const struct TVMTensorInfo kNormalInputs[2] = { - {"input1", kNormalInput1Shape, 4, DLDataType{1, 2, 3}}, - {"input2", kNormalInput1Shape, 4, DLDataType{2, 3, 4}}}; - -const int64_t kNormalOutput1Shape[3] = {3, 8, 8}; -const struct TVMTensorInfo kNormalOutputs[1] = { - {"output1", kNormalOutput1Shape, 3, DLDataType{3, 4, 5}}}; - -const int64_t kNormalPool1Shape[3] = {3, 8, 8}; -const struct TVMTensorInfo kNormalWorkspacePools[1] = { - {"workspace_pool1", kNormalPool1Shape, 3, DLDataType{3, 4, 7}}}; -const struct TVMConstantInfo kNormalConstantPools[1] = {{"constant_pool1", 0, 0, {}}}; - -const struct TVMMetadata kNormal = { - TVM_METADATA_VERSION, - kNormalInputs, - 2, - kNormalOutputs, - 1, - kNormalWorkspacePools, - 1, - kNormalConstantPools, - 1, - "default", -}; -} // namespace - -using ::testing::ElementsAre; -using ::testing::ElementsAreArray; -using ::testing::Eq; -using ::testing::Matcher; -using ::testing::MatcherInterface; -using ::testing::MatchResultListener; -using ::testing::StrEq; - -using ::tvm::codegen::metadata::DiscoverArraysVisitor; -using ::tvm::codegen::metadata::DiscoverComplexTypesVisitor; -using ::tvm::codegen::metadata::kMetadataGlobalSymbol; - -using ::tvm::runtime::Array; -using ::tvm::runtime::Downcast; -using ::tvm::runtime::ObjectRef; - -using ::tvm::runtime::metadata::ConstantInfoMetadata; -using ::tvm::runtime::metadata::Metadata; -using ::tvm::runtime::metadata::MetadataArray; -using ::tvm::runtime::metadata::MetadataKind; -using ::tvm::runtime::metadata::TensorInfo; - -TEST(Metadata, ParseStruct) { - Metadata md = Metadata(&kNormal); - EXPECT_THAT(md->version(), Eq(TVM_METADATA_VERSION)); - EXPECT_THAT(md->num_inputs(), Eq(2)); - - auto inputs = md->inputs(); - EXPECT_THAT(inputs.size(), Eq(2)); - - auto input1 = inputs[0]; - EXPECT_THAT(input1->name(), Eq("input1")); - EXPECT_THAT(input1->shape(), ElementsAre(1, 5, 5, 3)); - EXPECT_THAT(input1->dtype(), Eq(tvm::runtime::DataType(DLDataType{1, 2, 3}))); - - auto input2 = inputs[1]; - EXPECT_THAT(input2->name(), Eq("input2")); - EXPECT_THAT(input2->shape(), ElementsAre(1, 5, 5, 3)); - EXPECT_THAT(input2->dtype(), Eq(tvm::runtime::DataType(DLDataType{2, 3, 4}))); - - EXPECT_THAT(md->num_outputs(), Eq(1)); - auto outputs = md->outputs(); - EXPECT_THAT(outputs.size(), Eq(1)); - - auto output1 = outputs[0]; - EXPECT_THAT(output1->name(), Eq("output1")); - EXPECT_THAT(output1->shape(), ElementsAre(3, 8, 8)); - EXPECT_THAT(output1->dtype(), Eq(tvm::runtime::DataType(DLDataType{3, 4, 5}))); - - auto pools = md->workspace_pools(); - EXPECT_THAT(pools.size(), Eq(1)); - - auto workspace_pool1 = pools[0]; - EXPECT_THAT(workspace_pool1->name(), Eq("workspace_pool1")); - EXPECT_THAT(workspace_pool1->shape(), ElementsAre(3, 8, 8)); - EXPECT_THAT(workspace_pool1->dtype(), Eq(tvm::runtime::DataType(DLDataType{3, 4, 7}))); - - EXPECT_THAT(md->mod_name(), Eq("default")); -} - -class TestVisitor : public tvm::AttrVisitor { - public: - using Element = ::std::tuple<::std::string, ::tvm::runtime::ObjectRef>; - void Visit(const char* key, double* value) final { - keys.push_back(key); - values.push_back(::tvm::FloatImm(::tvm::runtime::DataType(kDLFloat, 64, 1), *value)); - } - void Visit(const char* key, int64_t* value) final { - keys.push_back(key); - values.push_back(::tvm::IntImm(::tvm::runtime::DataType(kDLInt, 64, 1), *value)); - } - void Visit(const char* key, uint64_t* value) final { - keys.push_back(key); - int64_t v; - *(reinterpret_cast(&v)) = *value; - values.push_back(::tvm::IntImm(::tvm::runtime::DataType(kDLUInt, 64, 1), v)); - } - void Visit(const char* key, int* value) final { - keys.push_back(key); - values.push_back(::tvm::IntImm(::tvm::runtime::DataType(kDLInt, 64, 1), *value)); - } - void Visit(const char* key, bool* value) final { - keys.push_back(key); - values.push_back(::tvm::Bool(*value)); - } - void Visit(const char* key, std::string* value) final { - keys.push_back(key); - values.push_back(::tvm::runtime::String(*value)); - } - void Visit(const char* key, tvm::runtime::DataType* value) final { - keys.push_back(key); - values.push_back(::tvm::PrimType(*value)); - } - void Visit(const char* key, tvm::runtime::NDArray* value) final { - keys.push_back(key); - values.push_back(*value); - } - void Visit(const char* key, void** value) final { CHECK(false) << "Do not expect this type"; } - - void Visit(const char* key, ::tvm::runtime::ObjectRef* value) final { - keys.push_back(key); - values.push_back(*value); - } - - std::vector keys; - std::vector<::tvm::runtime::ObjectRef> values; -}; - -TEST(Metadata, Visitor) { - Metadata md = Metadata(&kNormal); - TestVisitor v; - ::tvm::ReflectionVTable::Global()->VisitAttrs(md.operator->(), &v); - - EXPECT_THAT(v.keys, ElementsAre(StrEq("version"), StrEq("inputs"), StrEq("num_inputs"), - StrEq("outputs"), StrEq("num_outputs"), StrEq("workspace_pools"), - StrEq("num_workspace_pools"), StrEq("constant_pools"), - StrEq("num_constant_pools"), StrEq("mod_name"))); - EXPECT_THAT(Downcast(v.values[0])->value, Eq(TVM_METADATA_VERSION)); - - EXPECT_THAT(Downcast(v.values[0])->value, Eq(TVM_METADATA_VERSION)); - - // Just identify the tensor. - auto input_array = Downcast(v.values[1]); - EXPECT_THAT(input_array->kind, Eq(MetadataKind::kMetadata)); - EXPECT_THAT(input_array->type_key, StrEq("metadata.TensorInfoNode")); - EXPECT_THAT(input_array->array.size(), Eq(2)); - - auto input1 = Downcast(input_array->array[0]); - EXPECT_THAT(input1->name(), StrEq("input1")); - EXPECT_THAT(input1->shape(), ElementsAre(1, 5, 5, 3)); - EXPECT_THAT(input1->dtype(), tvm::runtime::DataType(DLDataType{1, 2, 3})); - - auto input2 = Downcast(input_array->array[1]); - EXPECT_THAT(input1->name(), StrEq("input1")); - EXPECT_THAT(input1->shape(), ElementsAre(1, 5, 5, 3)); - EXPECT_THAT(input1->dtype(), tvm::runtime::DataType(DLDataType{1, 2, 3})); - - auto num_inputs = Downcast(v.values[2]); - EXPECT_THAT(num_inputs->value, Eq(2)); - - auto output_array = Downcast(v.values[3]); - EXPECT_THAT(output_array->kind, Eq(MetadataKind::kMetadata)); - EXPECT_THAT(output_array->type_key, StrEq("metadata.TensorInfoNode")); - auto output1 = Downcast(output_array->array[0]); - - EXPECT_THAT(output1->name(), Eq("output1")); - - auto num_outputs = Downcast(v.values[4]); - EXPECT_THAT(num_outputs->value, Eq(1)); - - auto pool_array = Downcast(v.values[5]); - EXPECT_THAT(pool_array->kind, Eq(MetadataKind::kMetadata)); - EXPECT_THAT(pool_array->type_key, StrEq("metadata.TensorInfoNode")); - auto workspace_pool1 = Downcast(pool_array->array[0]); - - EXPECT_THAT(workspace_pool1->name(), Eq("workspace_pool1")); - - auto num_workspace_pools = Downcast(v.values[6]); - EXPECT_THAT(num_workspace_pools->value, Eq(1)); - - auto consts_array = Downcast(v.values[7]); - EXPECT_THAT(consts_array->kind, Eq(MetadataKind::kMetadata)); - EXPECT_THAT(consts_array->type_key, StrEq("metadata.ConstantInfoNode")); - auto consts1 = Downcast(consts_array->array[0]); - - EXPECT_THAT(consts1->name_hint(), Eq("constant_pool1")); - - auto num_consts = Downcast(v.values[8]); - EXPECT_THAT(num_consts->value, Eq(1)); - - auto mod_name = Downcast(v.values[9]); - EXPECT_THAT(mod_name, Eq("default")); -} - -using ::tvm::runtime::make_object; -TEST(Metadata, InMemory) { - Metadata md = Metadata(make_object( - TVM_METADATA_VERSION, - std::vector( - {TensorInfo(make_object( - tvm::String("Input1"), std::vector{1, 5, 5, 3}, - tvm::runtime::DataType(DLDataType{1, 2, 3}))), - TensorInfo(make_object( - tvm::String("Input2"), std::vector{1, 5, 5, 3}, - tvm::runtime::DataType(DLDataType{2, 3, 4})))}), - std::vector( - {TensorInfo(make_object( - tvm::String("Output1"), std::vector{3, 8, 8}, - tvm::runtime::DataType(DLDataType{3, 4, 5})))}), - std::vector( - {TensorInfo(make_object( - tvm::String("Workspace_Pool1"), std::vector{5, 10, 10}, - tvm::runtime::DataType(DLDataType{3, 4, 7})))}), - std::vector({tvm::ConstantInfo( - "Constant_Pool1", 64, - tvm::runtime::NDArray::Empty({64}, tvm::runtime::DataType::Int(64), {kDLCPU}))}), - "default")); - - auto md_data = md->data(); - EXPECT_THAT(md_data->version, Eq(TVM_METADATA_VERSION)); - EXPECT_THAT(md_data->num_inputs, Eq(2)); - - auto input0 = &md_data->inputs[0]; - EXPECT_THAT(input0->name, StrEq("Input1")); - EXPECT_THAT(std::vector(input0->shape, input0->shape + input0->num_shape), - ElementsAre(1, 5, 5, 3)); - EXPECT_THAT(tvm::runtime::DataType(input0->dtype), - Eq(tvm::runtime::DataType(DLDataType({1, 2, 3})))); - - auto input1 = &md_data->inputs[1]; - EXPECT_THAT(input1->name, StrEq("Input2")); - EXPECT_THAT(std::vector(input1->shape, input1->shape + input1->num_shape), - ElementsAre(1, 5, 5, 3)); - EXPECT_THAT(tvm::runtime::DataType(input1->dtype), - Eq(tvm::runtime::DataType(DLDataType({2, 3, 4})))); - - auto output0 = &md_data->outputs[0]; - EXPECT_THAT(output0->name, StrEq("Output1")); - EXPECT_THAT(std::vector(output0->shape, output0->shape + output0->num_shape), - ElementsAre(3, 8, 8)); - EXPECT_THAT(tvm::runtime::DataType(output0->dtype), - Eq(tvm::runtime::DataType(DLDataType({3, 4, 5})))); - - auto workspace_pool0 = &md_data->workspace_pools[0]; - EXPECT_THAT(workspace_pool0->name, StrEq("Workspace_Pool1")); - EXPECT_THAT(std::vector(workspace_pool0->shape, - workspace_pool0->shape + workspace_pool0->num_shape), - ElementsAre(5, 10, 10)); - EXPECT_THAT(tvm::runtime::DataType(workspace_pool0->dtype), - Eq(tvm::runtime::DataType(DLDataType({3, 4, 7})))); - - auto constant_pool0 = &md_data->constant_pools[0]; - EXPECT_THAT(constant_pool0->name_hint, StrEq("Constant_Pool1")); - - EXPECT_THAT(md_data->mod_name, StrEq("default")); -} - -TEST(Metadata, ZeroElementLists) { - Metadata md = Metadata(make_object( - TVM_METADATA_VERSION, std::vector({}), - std::vector( - {TensorInfo(make_object( - tvm::String("Output1"), std::vector{}, - tvm::runtime::DataType(DLDataType{3, 4, 5})))}), - std::vector({}), std::vector({}), "default")); - - EXPECT_THAT(md->data()->num_inputs, Eq(0)); - EXPECT_THAT(md->inputs().size(), Eq(0)); - EXPECT_THAT(md->num_inputs(), Eq(0)); - EXPECT_THAT(md->inputs(), ElementsAre()); - - auto output0 = md->data()->outputs[0]; - EXPECT_THAT(output0.num_shape, Eq(0)); - EXPECT_THAT(md->outputs()[0]->shape().size(), Eq(0)); - EXPECT_THAT(md->outputs()[0]->shape(), ElementsAre()); - - EXPECT_THAT(md->workspace_pools().size(), Eq(0)); - EXPECT_THAT(md->num_workspace_pools(), Eq(0)); - EXPECT_THAT(md->workspace_pools(), ElementsAre()); -} - -TEST(MetadataArray, GetElementCStructName) { - MetadataArray arr_struct{make_object( - Array(), MetadataKind::kMetadata, "metadata.FooMetadataNode")}; - EXPECT_THAT(arr_struct->kind, Eq(MetadataKind::kMetadata)); - EXPECT_THAT(arr_struct->get_element_c_struct_name(), StrEq("TVMFooMetadata")); - - MetadataArray arr_int{make_object( - Array(), MetadataKind::kInt64, nullptr)}; - EXPECT_THROW(arr_int->get_element_c_struct_name(), std::runtime_error); -} - -namespace { -std::string ExplainDiscoveredNameEq(bool negation, std::string expected_name) { - std::stringstream ss; - ss << "std::get<0>(discovered_array) " << (negation ? "isn't" : "is") << " equal to " - << expected_name; - return ss.str(); -} -} // namespace - -MATCHER_P(DiscoveredNameEq, expected_name, ExplainDiscoveredNameEq(negation, expected_name)) { - return std::string(std::get<0>(arg)) == expected_name; -} - -TEST(DiscoverArraysVisitor, DiscoverArrays) { - std::vector q; - DiscoverArraysVisitor visitor(&q); - - Metadata md = Metadata(&kNormal); - visitor.Visit(kMetadataGlobalSymbol, &md); - - EXPECT_THAT(q, ElementsAreArray({DiscoveredNameEq("kTvmgenMetadata_inputs_0_shape"), - DiscoveredNameEq("kTvmgenMetadata_inputs_1_shape"), - DiscoveredNameEq("kTvmgenMetadata_inputs"), - DiscoveredNameEq("kTvmgenMetadata_outputs_0_shape"), - DiscoveredNameEq("kTvmgenMetadata_outputs"), - DiscoveredNameEq("kTvmgenMetadata_workspace_pools_0_shape"), - DiscoveredNameEq("kTvmgenMetadata_workspace_pools"), - DiscoveredNameEq("kTvmgenMetadata_constant_pools")})); -} - -// In Debug builds the _type_key is no longer inlined but also has no -// link-time definition. -#define WITH_TYPE_KEY 0 - -template ::value, bool> = - true> -class TVMObjectIsInstanceMatcher : public MatcherInterface { - public: - using is_gtest_matcher = void; - - bool MatchAndExplain(tvm::runtime::metadata::MetadataBase arg, - MatchResultListener* os) const override { - bool result = arg->IsInstance(); - if (!result) { -#if WITH_TYPE_KEY - (*os) << "is an instance of type " << T::ContainerType::_type_key; -#else - (*os) << "is not of expected instance type"; -#endif - } - - return result; - } - - void DescribeTo(std::ostream* os) const override { -#if WITH_TYPE_KEY - (*os) << "is an instance of type " << T::ContainerType::_type_key; -#else - (*os) << "is not of expected instance type"; -#endif - } - - void DescribeNegationTo(std::ostream* os) const override { -#if WITH_TYPE_KEY - (*os) << "is not an instance of type " << T::ContainerType::_type_key; -#else - (*os) << "is not of expected instance type"; -#endif - } -}; - -template -Matcher TVMObjectIsInstance() { - return Matcher(new TVMObjectIsInstanceMatcher()); -} - -TEST(DiscoverComplexTypesVisitor, DiscoverComplexTypes) { - std::vector q; - DiscoverComplexTypesVisitor visitor(&q); - - Metadata md = Metadata(&kNormal); - visitor.Discover(md); - - EXPECT_THAT( - q, ElementsAre(TVMObjectIsInstance(), TVMObjectIsInstance(), - TVMObjectIsInstance())); -} - -TEST(Metadata, TVMConstantInfo) { - std::vector q; - auto ci = std::make_unique(10); - EXPECT_TRUE(ci.get() != nullptr); -} diff --git a/tests/cpp/arith_integer_set_test.cc b/tests/cpp/arith_integer_set_test.cc index 04546abba9a6..4454598c2049 100644 --- a/tests/cpp/arith_integer_set_test.cc +++ b/tests/cpp/arith_integer_set_test.cc @@ -18,9 +18,9 @@ */ #if TVM_MLIR_VERSION >= 150 -#include #include #include +#include #include #include "../src/arith/presburger_set.h" diff --git a/tests/cpp/arith_simplify_test.cc b/tests/cpp/arith_simplify_test.cc index 073b4269eb6f..23bcd8a7a7e5 100644 --- a/tests/cpp/arith_simplify_test.cc +++ b/tests/cpp/arith_simplify_test.cc @@ -17,9 +17,9 @@ * under the License. */ -#include #include #include +#include #include TEST(Simplify, MinMax) { diff --git a/tests/cpp/attrs_test.cc b/tests/cpp/attrs_test.cc index d836639043d1..5a6f03088929 100644 --- a/tests/cpp/attrs_test.cc +++ b/tests/cpp/attrs_test.cc @@ -17,9 +17,9 @@ * under the License. */ -#include #include #include +#include #include #include diff --git a/tests/cpp/auto_scheduler_test.cc b/tests/cpp/auto_scheduler_test.cc deleted file mode 100644 index 0a753bc9a740..000000000000 --- a/tests/cpp/auto_scheduler_test.cc +++ /dev/null @@ -1,172 +0,0 @@ -/* - * 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 - -// Compute declaration for test -tvm::Array conv2d_nchw_bn_relu_func(int N, int H, int W, int CI, int CO, - int kernel_size, int strides, int padding, - int dilation = 1) { - using namespace tvm; - using namespace tvm::te; - - Tensor data = placeholder({N, CI, H, W}, DataType::Float(32), "Data"); - Tensor kernel = placeholder({CO, CI, kernel_size, kernel_size}, DataType::Float(32), "Kernel"); - Tensor bias = placeholder({CO, 1, 1}, DataType::Float(32), "Bias"); - Tensor bn_scale = placeholder({CO, 1, 1}, DataType::Float(32), "Bn_scale"); - Tensor bn_offset = placeholder({CO, 1, 1}, DataType::Float(32), "Bn_offset"); - - int OH = (H + 2 * padding - (kernel_size - 1) * dilation - 1) / strides + 1; - int OW = (W + 2 * padding - (kernel_size - 1) * dilation - 1) / strides + 1; - - const auto& conv = topi::conv2d_nchw(data, kernel, padding, padding, strides, strides); - ICHECK(conv->shape[2].as()->value == OH); - ICHECK(conv->shape[3].as()->value == OW); - - const auto& bias_add = compute( - {N, CO, OH, OW}, [&](Var i, Var j, Var k, Var l) { return conv[i][j][k][l] + bias[j][0][0]; }, - "Bias_add"); - const auto& bn_mul = compute( - {N, CO, OH, OW}, - [&](Var i, Var j, Var k, Var l) { return bias_add[i][j][k][l] * bn_scale[j][0][0]; }, - "Bn_mul"); - const auto& bn_add = compute( - {N, CO, OH, OW}, - [&](Var i, Var j, Var k, Var l) { return bn_mul[i][j][k][l] + bn_offset[j][0][0]; }, - "Bn_add"); - const auto& out = topi::relu(bn_add); - - return {data, kernel, bias, bn_scale, bn_offset, out}; -} - -using namespace tvm::auto_scheduler; - -// Test Access Analyzer -TEST(ComputeDAG, AccessAnalyzer) { - const auto& tensors = conv2d_nchw_bn_relu_func(1, 224, 224, 3, 64, 7, 2, 3); - const auto& dag = tvm::auto_scheduler::ComputeDAG(tensors); - State s0 = dag->init_state; - - int data = 0, padding = 1, kernel = 2, conv = 3, bias = 4, bias_add = 5; - int bn_scale = 6, bn_mul = 7, bn_offset = 8, bn_add = 9, relu = 10; - - std::set needs_multi_level_tiling = {conv}; - for (size_t stage_id = 0; stage_id < dag->ops.size(); stage_id++) { - if (needs_multi_level_tiling.count(stage_id)) { - ICHECK(dag->access_analyzer.NeedsMultiLevelTiling(dag->ops[stage_id])); - } else { - ICHECK(!dag->access_analyzer.NeedsMultiLevelTiling(dag->ops[stage_id])); - } - } - - std::set is_simple_access = {data, padding, kernel, bias, bias_add, - bn_scale, bn_mul, bn_offset, bn_add, relu}; - for (size_t stage_id = 0; stage_id < dag->ops.size(); stage_id++) { - if (is_simple_access.count(stage_id)) { - ICHECK(dag->access_analyzer.IsSimpleAccess(dag->ops[stage_id])); - } else { - ICHECK(!dag->access_analyzer.IsSimpleAccess(dag->ops[stage_id])); - } - } - - std::set is_strictly_inlinable = {bias_add, bn_mul, bn_add, relu}; - for (size_t stage_id = 0; stage_id < dag->ops.size(); stage_id++) { - if (is_strictly_inlinable.count(stage_id)) { - ICHECK(dag->access_analyzer.IsStrictlyInlineable(dag->ops[stage_id])); - } else { - ICHECK(!dag->access_analyzer.IsStrictlyInlineable(dag->ops[stage_id])); - } - } - - std::set is_output = {relu}; - for (size_t stage_id = 0; stage_id < dag->ops.size(); stage_id++) { - if (is_output.count(stage_id)) { - ICHECK(dag->access_analyzer.IsOutput(dag->ops[stage_id])); - } else { - ICHECK(!dag->access_analyzer.IsOutput(dag->ops[stage_id])); - } - } - - ICHECK_EQ(dag->access_analyzer.GetNumCommonOuterIterator(dag->ops[conv], dag->ops[bias_add]), 4); - ICHECK_EQ(dag->access_analyzer.GetNumCommonOuterIterator(dag->ops[conv], dag->ops[relu]), 4); - ICHECK_EQ(dag->access_analyzer.GetNumCommonOuterIterator(dag->ops[data], dag->ops[relu]), 1); - - ICHECK(dag->access_analyzer.ElementWiseMatch(dag->ops[conv], dag->ops[bias_add])); - ICHECK(dag->access_analyzer.ElementWiseMatch(dag->ops[conv], dag->ops[relu])); - ICHECK(!dag->access_analyzer.ElementWiseMatch(dag->ops[data], dag->ops[padding])); - - std::unordered_set op_set; - { - std::vector> consumer_list = { - {data, padding}, {padding, conv}, {kernel, conv}, {conv, bias_add}, - {bias, bias_add}, {bias_add, bn_mul}, {bn_scale, bn_mul}, {bn_mul, bn_add}, - {bn_offset, bn_add}, {bn_add, relu}}; - for (const auto& pair : consumer_list) { - op_set = dag->access_analyzer.GetConsumers(s0, s0->stages[pair.first]->op); - ICHECK_EQ(op_set.size(), 1); - ICHECK_EQ((*op_set.begin()), s0->stages[pair.second]->op); - } - std::vector>> producer_list = {{padding, {data}}, - {conv, {padding, kernel}}, - {bias_add, {conv, bias}}, - {bn_mul, {bias_add, bn_scale}}, - {bn_add, {bn_mul, bn_offset}}, - {relu, {bn_add}}}; - for (const auto& pair : producer_list) { - op_set = dag->access_analyzer.GetProducers(s0, s0->stages[pair.first]->op); - ICHECK_EQ(op_set.size(), pair.second.size()); - for (const auto& target : pair.second) { - ICHECK(op_set.count(s0->stages[target]->op)); - } - } - } - - s0.compute_inline(bn_add); - s0.compute_inline(bn_mul); - s0.compute_inline(bias_add); - s0.compute_inline(padding); - { - std::vector> consumer_list = {{data, conv}, {kernel, conv}, {conv, relu}}; - for (const auto& pair : consumer_list) { - op_set = dag->access_analyzer.GetConsumers(s0, s0->stages[pair.first]->op); - ICHECK_EQ(op_set.size(), 1); - ICHECK_EQ((*op_set.begin()), s0->stages[pair.second]->op); - } - std::vector>> producer_list = {{padding, {data}}, - {conv, {padding, kernel}}, - {bias_add, {conv, bias}}, - {bn_mul, {bias_add, bn_scale}}, - {bn_add, {bn_mul, bn_offset}}, - {relu, {bn_add}}}; - for (const auto& pair : producer_list) { - op_set = dag->access_analyzer.GetDirectProducers(s0->stages[pair.first]->op); - ICHECK_EQ(op_set.size(), pair.second.size()); - for (const auto& target : pair.second) { - ICHECK(op_set.count(s0->stages[target]->op)); - } - } - } -} diff --git a/tests/cpp/build_module_test.cc b/tests/cpp/build_module_test.cc index 181a1fa3de4c..cedc9b62701d 100644 --- a/tests/cpp/build_module_test.cc +++ b/tests/cpp/build_module_test.cc @@ -65,143 +65,3 @@ TEST(BuildModule, Basic) { ICHECK_EQ(mali_target->GetAttr("model").value(), "Mali-T860MP4@800Mhz"); ICHECK_EQ(mali_target->GetAttr("max_num_threads").value(), 256); } - -TEST(BuildModule, Heterogeneous) { - /* The testing network is like following, where the element-wise add and sub - * ops are allocated to GPU and CPU, respectively: - * - * A B - * \ / - * elemwise_add (gpu) - * \ - * copy C - * \ / - * elemwise_sub (cpu) - */ - - using namespace tvm; - using namespace tvm::te; - bool enabled = tvm::runtime::RuntimeEnabled("cuda"); - if (!enabled) { - LOG(INFO) << "Skip heterogeneous test because cuda is not enabled." - << "\n"; - return; - } - - auto target_llvm = Target("llvm"); - auto target_cuda = Target("cuda"); - - // The shape of input tensors. - const int n = 4; - Array shape{n}; - - auto A = placeholder(shape, DataType::Float(32), "A"); - auto B = placeholder(shape, DataType::Float(32), "B"); - auto C = placeholder(shape, DataType::Float(32), "C"); - - auto elemwise_add = compute( - A->shape, [&A, &B](PrimExpr i) { return A[i] + B[i]; }, "elemwise_add"); - - // TODO(mbs): device_copy cleanup. - auto copy = placeholder(shape, DataType::Float(32), "__copy"); - auto elemwise_sub = compute( - C->shape, [©, &C](PrimExpr i) { return copy[i] - C[i]; }, "elemwise_sub"); - - auto fcreate_s1 = [=]() { - With cuda_scope(target_cuda); - return topi::cuda::schedule_injective(target_cuda, {elemwise_add}); - }; - - auto fcreate_s2 = [=]() { - With llvm_scope(target_llvm); - return create_schedule({elemwise_sub->op}); - }; - - auto args1 = Array({A, B, elemwise_add}); - auto args2 = Array({copy, C, elemwise_sub}); - - std::unordered_map binds; - GlobalVarSupply global_var_supply = GlobalVarSupply(); - auto lowered_s1 = LowerSchedule(fcreate_s1(), args1, "elemwise_add", binds, global_var_supply); - auto lowered_s2 = LowerSchedule(fcreate_s2(), args2, "elemwise_sub", binds, global_var_supply); - Map inputs = {{target_cuda, lowered_s1}, {target_llvm, lowered_s2}}; - auto module = build(inputs, Target()); - - // Assertion for build. - ICHECK_EQ(module->imports().size(), 1); - - // Execute the graph and check the correctness. - // Setup graph json. - std::string json = - "{\"nodes\": [{\"op\": \"null\", \"name\": \"A\", \"inputs\": []}, " - "{\"op\": \"null\", \"name\": \"B\", \"inputs\": []}, {\"op\": " - "\"tvm_op\", \"name\": \"elemwise_add\", \"attrs\": {\"flatten_data\": " - "\"1\", \"func_name\": \"elemwise_add\", \"num_inputs\": \"2\", " - "\"num_outputs\": \"1\"}, \"inputs\": [[0, 0, 0], [1, 0, 0]]}, {\"op\": " - "\"tvm_op\", \"name\": \"__copy_add_to_sub\", \"attrs\": " - "{\"flatten_data\": \"0\", \"func_name\": \"__copy\", \"num_inputs\": " - "\"1\", \"num_outputs\": \"1\"}, \"inputs\": [[2, 0, 0]]}, {\"op\": " - "\"null\", \"name\": \"C\", \"inputs\": []}, {\"op\": \"tvm_op\", " - "\"name\": \"elemwise_sub\", \"attrs\": {\"flatten_data\": \"0\", " - "\"func_name\": \"elemwise_sub\", \"num_inputs\": \"2\", " - "\"num_outputs\": \"1\"}, \"inputs\": [[3, 0, 0], [4, 0, 0]]}], " - "\"arg_nodes\": [0, 1, 4], \"node_row_ptr\": [0, 1, 2, 3, 4, 5, 6], " - "\"heads\": [[5, 0, 0]], \"attrs\": {\"storage_id\": [\"list_int\", [3, " - "4, 0, 1, 5, 2]], \"shape\": [\"list_shape\", [[4], [4], [4], [4], [4], " - "[4]]], \"device_index\": [\"list_int\", [2, 2, 2, 1, 1, 1]], \"dtype\": " - "[\"list_int\", [0, 0, 0, 0, 0, 0]], \"dltype\": [\"list_str\", " - "[\"float32\", \"float32\", \"float32\", \"float32\", \"float32\", " - "\"float32\"]]}}"; - - // Setup inputs. - auto a_val = runtime::NDArray::Empty({n}, {kDLFloat, 32, 1}, {kDLCPU, 0}); - auto b_val = runtime::NDArray::Empty({n}, {kDLFloat, 32, 1}, {kDLCPU, 0}); - auto c_val = runtime::NDArray::Empty({n}, {kDLFloat, 32, 1}, {kDLCPU, 0}); - - auto pa = static_cast(a_val->data); - auto pb = static_cast(b_val->data); - auto pc = static_cast(c_val->data); - - // Assign values. - for (int i = 0; i < n; i++) { - pa[i] = i; - pb[i] = i + 1.0; - pc[i] = i - 1.0; - } - - // Initialize graph executor. - int cpu_dev_ty = static_cast(kDLCPU); - int cpu_dev_id = 0; - int gpu_dev_ty = static_cast(kDLCUDA); - int gpu_dev_id = 0; - - const runtime::PackedFunc* graph_executor = - tvm::runtime::Registry::Get("tvm.graph_executor.create"); - runtime::Module mod = - (*graph_executor)(json, module, cpu_dev_ty, cpu_dev_id, gpu_dev_ty, gpu_dev_id); - - // test FFI for module. - auto test_ffi = PackedFunc([](TVMArgs args, TVMRetValue* rv) { - int tcode = args[1]; - ICHECK_EQ(args[0].type_code(), tcode); - }); - - test_ffi(runtime::Module(mod), static_cast(kTVMModuleHandle)); - test_ffi(Optional(mod), static_cast(kTVMModuleHandle)); - - PackedFunc set_input = mod.GetFunction("set_input", false); - PackedFunc run = mod.GetFunction("run", false); - PackedFunc get_output = mod.GetFunction("get_output", false); - set_input("A", a_val); - set_input("B", b_val); - set_input("C", c_val); - - run(); - tvm::runtime::NDArray out = get_output(0); - float* p_out = static_cast(out->data); - - // Check correctness. - for (int i = 0; i < n; ++i) { - ICHECK_LT(std::fabs(p_out[i] - (i + (i + 1.0) - (i - 1.0))), 1e-5); - } -} diff --git a/tests/cpp/c_codegen_test.cc b/tests/cpp/c_codegen_test.cc deleted file mode 100644 index 5f783830495e..000000000000 --- a/tests/cpp/c_codegen_test.cc +++ /dev/null @@ -1,129 +0,0 @@ -/* - * 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 - -TEST(CCodegen, MainFunctionOrder) { - using namespace tvm; - using namespace tvm::te; - - std::string tvm_module_main = std::string(runtime::symbol::tvm_module_main); - - tvm::Target target_c = tvm::Target("c -keys=cpu"); - - const int n = 4; - Array shape{n}; - - auto A = placeholder(shape, DataType::Float(32), "A"); - auto B = placeholder(shape, DataType::Float(32), "B"); - - auto elemwise_add = compute( - A->shape, [&A, &B](PrimExpr i) { return A[i] + B[i]; }, "elemwise_add"); - - auto fcreate = [=]() { - With llvm_scope(target_c); - return create_schedule({elemwise_add->op}); - }; - - auto args = Array({A, B, elemwise_add}); - - std::unordered_map binds; - auto lowered = LowerSchedule(fcreate(), args, "elemwise_add", binds, GlobalVarSupply()); - Map inputs = {{target_c, lowered}}; - runtime::Module module = build(inputs, Target()); - Array functions = module->GetFunction("get_func_names", false)(); - - ICHECK(functions.back().compare(tvm_module_main) == 0); -} - -auto BuildLowered(std::string op_name, tvm::Target target) { - using namespace tvm; - using namespace tvm::te; - - // The shape of input tensors. - const int n = 4; - Array shape{n}; - - auto A = placeholder(shape, DataType::Float(32), "A"); - auto B = placeholder(shape, DataType::Float(32), "B"); - - auto op = compute( - A->shape, [&A, &B](PrimExpr i) { return A[i] + B[i]; }, op_name); - - auto fcreate_s = [=]() { - With llvm_scope(target); - return create_schedule({op->op}); - }; - - auto args = Array({A, B, op}); - std::unordered_map binds; - auto lowered_s = LowerSchedule(fcreate_s(), args, op_name, binds, GlobalVarSupply()); - return lowered_s; -} - -bool IsSorted(tvm::Map inputs) { - std::vector schedule_names; - for (auto const& module : inputs) { - for (auto const& func : module.second->functions) { - schedule_names.push_back(func.first->name_hint); - } - } - return std::is_sorted(schedule_names.begin(), schedule_names.end()); -} - -TEST(CCodegen, FunctionOrder) { - using testing::_; - using testing::ElementsAre; - using testing::StrEq; - using namespace tvm; - using namespace tvm::te; - - Target target = Target("c -keys=cpu"); - - // add schedules in reverse order - Map inputs; - inputs.Set(Target("c -keys=cpu"), BuildLowered("op_2", target)); - inputs.Set(Target("c -keys=cpu"), BuildLowered("op_1", target)); - - for (uint32_t counter = 99; IsSorted(inputs) && counter > 0; counter--) { - std::string op_name = "op_" + std::to_string(counter); - inputs.Set(Target("c -keys=cpu"), BuildLowered(op_name, target)); - } - - EXPECT_FALSE(IsSorted(inputs)); - - auto module = build(inputs, Target()); - Array func_array = module->GetFunction("get_func_names", false)(); - std::vector functions{func_array.begin(), func_array.end()}; - // The entry point is handled separately from the other functions. - functions.erase(std::remove_if(functions.begin(), functions.end(), - [](const std::string& name) { - return name == tvm::runtime::symbol::tvm_module_main; - }), - functions.end()); - EXPECT_TRUE(std::is_sorted(functions.begin(), functions.end())); -} diff --git a/tests/cpp/container_test.cc b/tests/cpp/container_test.cc index 9d2f1437b9ab..0a089eaebde3 100644 --- a/tests/cpp/container_test.cc +++ b/tests/cpp/container_test.cc @@ -17,13 +17,13 @@ * under the License. */ -#include #include #include #include #include #include #include +#include #include #include diff --git a/tests/cpp/dataflow_pattern_test.cc b/tests/cpp/dataflow_pattern_test.cc deleted file mode 100644 index 0452d0047b05..000000000000 --- a/tests/cpp/dataflow_pattern_test.cc +++ /dev/null @@ -1,213 +0,0 @@ -/* - * 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 - -TEST(DFPattern, IsVar) { - using namespace tvm; - using namespace tvm::relay; - auto pattern = IsVar("add"); - auto* node = pattern.as(); - ICHECK(node); - ICHECK(node->name == String("add")); -} - -TEST(DFPattern, IsConstant) { - using namespace tvm; - using namespace tvm::relay; - auto pattern = IsConstant(); - auto* node = pattern.as(); - ICHECK(node); -} - -TEST(DFPattern, IsOp) { - using namespace tvm; - using namespace tvm::relay; - auto pattern = IsOp("add"); - auto* node = pattern.as(); - ICHECK(node); - ICHECK(node->expr == Op::Get("add")); -} - -TEST(DFPattern, IsTuple) { - using namespace tvm; - using namespace tvm::relay; - auto a = WildcardPattern(); - auto b = WildcardPattern(); - auto pattern = IsTuple({a, b}); - auto* node = pattern.as(); - ICHECK(node); - ICHECK(node->fields[0] == a); - ICHECK(node->fields[1] == b); -} - -TEST(DFPattern, IsTupleGetItem) { - using namespace tvm; - using namespace tvm::relay; - auto a = WildcardPattern(); - auto b = WildcardPattern(); - auto tuple = IsTuple({a, b}); - auto pattern = IsTupleGetItem(tuple, 1); - auto* node = pattern.as(); - ICHECK(node); - ICHECK(node->tuple == tuple); - ICHECK(node->index == 1); -} - -TEST(DFPattern, ADD) { - using namespace tvm; - using namespace tvm::relay; - auto a = WildcardPattern(); - auto b = WildcardPattern(); - auto pattern = a + b; - auto* node = pattern.as(); - ICHECK(node); - ICHECK(node->args[0] == a); - ICHECK(node->args[1] == b); - auto* expr_pattern = node->op.as(); - ICHECK(expr_pattern); - ICHECK(expr_pattern->expr == Op::Get("add")); -} - -TEST(DFPattern, SUB) { - using namespace tvm; - using namespace tvm::relay; - auto a = WildcardPattern(); - auto b = WildcardPattern(); - auto pattern = a - b; - auto* node = pattern.as(); - ICHECK(node); - ICHECK(node->args[0] == a); - ICHECK(node->args[1] == b); - auto* expr_pattern = node->op.as(); - ICHECK(expr_pattern); - ICHECK(expr_pattern->expr == Op::Get("subtract")); -} - -TEST(DFPattern, MUL) { - using namespace tvm; - using namespace tvm::relay; - auto a = WildcardPattern(); - auto b = WildcardPattern(); - auto pattern = a * b; - auto* node = pattern.as(); - ICHECK(node); - ICHECK(node->args[0] == a); - ICHECK(node->args[1] == b); - auto* expr_pattern = node->op.as(); - ICHECK(expr_pattern); - ICHECK(expr_pattern->expr == Op::Get("multiply")); -} - -TEST(DFPattern, DIV) { - using namespace tvm; - using namespace tvm::relay; - auto a = WildcardPattern(); - auto b = WildcardPattern(); - auto pattern = a / b; - auto* node = pattern.as(); - ICHECK(node); - ICHECK(node->args[0] == a); - ICHECK(node->args[1] == b); - auto* expr_pattern = node->op.as(); - ICHECK(expr_pattern); - ICHECK(expr_pattern->expr == Op::Get("divide")); -} - -TEST(DFPattern, OR) { - using namespace tvm; - using namespace tvm::relay; - auto a = WildcardPattern(); - auto b = WildcardPattern(); - auto pattern = a || b; - auto* node = pattern.as(); - ICHECK(node); - ICHECK(node->left == a); - ICHECK(node->right == b); -} - -TEST(DFPattern, Optional) { - using namespace tvm; - using namespace tvm::relay; - DFPattern a = WildcardPattern(); - DFPattern b = WildcardPattern(); - auto pattern = a.Optional([b](const DFPattern& other) { return other + b; }); - auto* node = pattern.as(); - ICHECK(node); - ICHECK(node->left == a); - auto* right_node = node->right.as(); - ICHECK(right_node); - ICHECK(right_node->args.size() == 2); - ICHECK(right_node->args[0] == a); - ICHECK(right_node->args[1] == b); - auto* expr_pattern = right_node->op.as(); - ICHECK(expr_pattern); - ICHECK(expr_pattern->expr == Op::Get("add")); -} - -TEST(DFPattern, HasAttr) { - using namespace tvm; - using namespace tvm::relay; - auto a = WildcardPattern(); - Map attrs; - auto b = String("b"); - attrs.Set("a", b); - auto pattern = a.HasAttr(attrs); - auto* node = pattern.as(); - ICHECK(node); - ICHECK(node->pattern == a); - ICHECK(node->attrs->dict.at("a") == b); -} - -TEST(DFPattern, HasType) { - using namespace tvm; - using namespace tvm::relay; - auto a = WildcardPattern(); - TensorType type({1, 2, 3}, DataType(runtime::String2DLDataType("float32"))); - auto pattern = a.HasType(type); - auto* node = pattern.as(); - ICHECK(node); - ICHECK(node->pattern == a); - ICHECK(node->type == type); -} - -TEST(DFPattern, HasDtype) { - using namespace tvm; - using namespace tvm::relay; - auto a = WildcardPattern(); - auto pattern = a.HasDtype("float32"); - auto* node = pattern.as(); - ICHECK(node); - ICHECK(node->pattern == a); - ICHECK(runtime::DLDataType2String(node->dtype.operator DLDataType()) == "float32"); -} - -TEST(DFPattern, HasShape) { - using namespace tvm; - using namespace tvm::relay; - auto a = WildcardPattern(); - Array shape{1, 2, 3}; - auto pattern = a.HasShape(shape); - auto* node = pattern.as(); - ICHECK(node); - ICHECK(node->pattern == a); - ICHECK(node->shape == shape); -} diff --git a/tests/cpp/expr_test.cc b/tests/cpp/expr_test.cc index 82de46616cb4..579479ccc0e5 100644 --- a/tests/cpp/expr_test.cc +++ b/tests/cpp/expr_test.cc @@ -17,9 +17,9 @@ * under the License. */ -#include #include #include +#include #include TEST(Expr, Basic) { diff --git a/tests/cpp/ir_functor_test.cc b/tests/cpp/ir_functor_test.cc index 30b1bc78247a..9449787218cc 100644 --- a/tests/cpp/ir_functor_test.cc +++ b/tests/cpp/ir_functor_test.cc @@ -17,11 +17,10 @@ * under the License. */ -#include #include #include #include -#include +#include #include #include #include @@ -56,22 +55,6 @@ TEST(IRF, CountVar) { ICHECK_EQ(n_var, 2); } -TEST(IRF, VisitPrimFuncs) { - using namespace tvm; - using namespace tvm::tir; - PrimFunc prim_func(/*params=*/{}, /*body=*/Evaluate(Integer(0))); - auto c_data = tvm::runtime::NDArray::Empty({1, 2, 3}, {kDLFloat, 32, 1}, {kDLCPU, 0}); - relay::Function relay_func(/*params=*/{}, /*body=*/relay::Expr(relay::Constant(c_data)), - /*ret_type=*/relay::Type(), /*ty_params=*/{}); - IRModule mod({ - {GlobalVar("main"), prim_func}, - {GlobalVar("main2"), relay_func}, - }); - int n_visited = 0; - VisitPrimFuncs(mod, [&](const PrimFuncNode* func) { ++n_visited; }); - ASSERT_EQ(n_visited, 1); -} - TEST(IRF, PreOrderVisit) { using namespace tvm; using namespace tvm::tir; diff --git a/tests/cpp/llvm_codegen_test.cc b/tests/cpp/llvm_codegen_registry_test.cc similarity index 100% rename from tests/cpp/llvm_codegen_test.cc rename to tests/cpp/llvm_codegen_registry_test.cc diff --git a/tests/cpp/name_supply_test.cc b/tests/cpp/name_supply_test.cc deleted file mode 100644 index 023d2e903aba..000000000000 --- a/tests/cpp/name_supply_test.cc +++ /dev/null @@ -1,129 +0,0 @@ -/* - * 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 - -using namespace tvm; - -NameSupply preambleNameSupply() { - NameSupply name_supply("prefix"); - name_supply->FreshName("test"); - return name_supply; -} - -TEST(NameSupply, FreshName) { - NameSupply name_supply = preambleNameSupply(); - String fresh = name_supply->FreshName("test"); - - EXPECT_EQ(fresh.compare("prefix_test_1"), 0); -} - -TEST(NameSupply, FreshNameNoConflict) { - NameSupply name_supply = preambleNameSupply(); - String fresh = name_supply->FreshName("name_2"); - EXPECT_EQ(fresh.compare("prefix_name_2"), 0); - - fresh = name_supply->FreshName("name"); - EXPECT_EQ(fresh.compare("prefix_name"), 0); - - fresh = name_supply->FreshName("name"); - EXPECT_EQ(fresh.compare("prefix_name_1"), 0); - - fresh = name_supply->FreshName("name"); - EXPECT_EQ(fresh.compare("prefix_name_3"), 0); -} - -TEST(NameSupply, ContainsName) { - NameSupply name_supply = preambleNameSupply(); - - EXPECT_TRUE(name_supply->ContainsName("test")); - EXPECT_FALSE(name_supply->ContainsName("test_1")); -} - -TEST(NameSupply, ReserveName) { - NameSupply name_supply = preambleNameSupply(); - name_supply->ReserveName("otherTest", false); - - EXPECT_TRUE(name_supply->ContainsName("otherTest", false)); - EXPECT_FALSE(name_supply->ContainsName("otherTest")); - - name_supply->ReserveName("otherTest"); - EXPECT_TRUE(name_supply->ContainsName("prefix_otherTest", false)); - EXPECT_TRUE(name_supply->ContainsName("otherTest")); -} - -GlobalVarSupply preambleVarSupply() { - GlobalVarSupply global_var_supply; - global_var_supply->FreshGlobal("test"); - return global_var_supply; -} - -TEST(GlobalVarSupply, FreshGlobal) { - GlobalVarSupply global_var_supply = preambleVarSupply(); - GlobalVar first_var = global_var_supply->FreshGlobal("test"); - GlobalVar second_var = global_var_supply->FreshGlobal("test"); - - EXPECT_FALSE(tvm::StructuralEqual()(first_var, second_var)); - EXPECT_EQ(first_var->name_hint.compare("test_1"), 0); - EXPECT_EQ(second_var->name_hint.compare("test_2"), 0); -} - -TEST(GlobalVarSupply, UniqueGlobalFor) { - GlobalVarSupply global_var_supply = preambleVarSupply(); - GlobalVar first_var = global_var_supply->UniqueGlobalFor("someName"); - GlobalVar second_var = global_var_supply->UniqueGlobalFor("someName"); - - EXPECT_TRUE(tvm::StructuralEqual()(first_var, second_var)); - EXPECT_EQ(first_var->name_hint.compare("someName"), 0); - EXPECT_EQ(second_var->name_hint.compare("someName"), 0); -} - -TEST(GlobalVarSupply, ReserveGlobal) { - GlobalVarSupply global_var_supply = preambleVarSupply(); - GlobalVar var = GlobalVar("someName"); - global_var_supply->ReserveGlobalVar(var); - GlobalVar second_var = global_var_supply->UniqueGlobalFor("someName"); - GlobalVar third_var = global_var_supply->FreshGlobal("someName"); - - EXPECT_TRUE(tvm::StructuralEqual()(var, second_var)); - EXPECT_FALSE(tvm::StructuralEqual()(var, third_var)); - EXPECT_EQ(second_var->name_hint.compare("someName"), 0); - EXPECT_EQ(third_var->name_hint.compare("someName_1"), 0); -} - -TEST(GlobalVarSupply, BuildIRModule) { - auto x = relay::Var("x", relay::Type()); - auto f = relay::Function(tvm::Array{x}, x, relay::Type(), {}); - GlobalVar var = GlobalVar("test"); - IRModule module = IRModule({{var, f}}); - - GlobalVarSupply global_var_supply = GlobalVarSupply(module); - GlobalVar second_var = global_var_supply->UniqueGlobalFor("test", false); - GlobalVar third_var = global_var_supply->FreshGlobal("test", false); - - EXPECT_TRUE(tvm::StructuralEqual()(var, second_var)); - EXPECT_FALSE(tvm::StructuralEqual()(var, third_var)); - EXPECT_EQ(second_var->name_hint.compare("test"), 0); - EXPECT_EQ(third_var->name_hint.compare("test_1"), 0); -} diff --git a/tests/cpp/name_transforms_test.cc b/tests/cpp/name_transforms_test.cc deleted file mode 100644 index 7e3cfe1d779c..000000000000 --- a/tests/cpp/name_transforms_test.cc +++ /dev/null @@ -1,142 +0,0 @@ -/* - * 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 "../src/relay/backend/name_transforms.h" - -#include -#include -#include - -namespace tvm { -namespace relay { -namespace backend { - -using namespace tvm::runtime; - -std::string ToCamel(const std::string& original_name); - -TEST(NameTransforms, ToCFunctionStyle) { - ASSERT_EQ(ToCFunctionStyle("TVM_Woof"), "TVMWoof"); - ASSERT_EQ(ToCFunctionStyle("TVM_woof"), "TVMWoof"); - ASSERT_EQ(ToCFunctionStyle("TVM_woof_woof"), "TVMWoofWoof"); - ASSERT_EQ(ToCFunctionStyle("TVMGen_woof_woof"), "TVMGenWoofWoof"); - EXPECT_THROW(ToCFunctionStyle("Cake_Bakery"), InternalError); // Incorrect prefix - EXPECT_THROW(ToCFunctionStyle(""), InternalError); -} - -TEST(NameTransforms, ToCVariableStyle) { - ASSERT_EQ(ToCVariableStyle("TVM_Woof"), "tvm_woof"); - ASSERT_EQ(ToCVariableStyle("TVM_woof"), "tvm_woof"); - ASSERT_EQ(ToCVariableStyle("TVM_woof_Woof"), "tvm_woof_woof"); - EXPECT_THROW(ToCVariableStyle("Cake_Bakery"), InternalError); // Incorrect prefix - EXPECT_THROW(ToCVariableStyle(""), InternalError); -} - -TEST(NameTransforms, ToCConstantStyle) { - ASSERT_EQ(ToCConstantStyle("TVM_Woof"), "TVM_WOOF"); - ASSERT_EQ(ToCConstantStyle("TVM_woof"), "TVM_WOOF"); - ASSERT_EQ(ToCConstantStyle("TVM_woof_Woof"), "TVM_WOOF_WOOF"); - EXPECT_THROW(ToCConstantStyle("Cake_Bakery"), InternalError); // Incorrect prefix - EXPECT_THROW(ToCConstantStyle(""), InternalError); -} - -TEST(NameTransforms, ToRustStructStyle) { - ASSERT_EQ(ToRustStructStyle("Woof"), "Woof"); - ASSERT_EQ(ToRustStructStyle("woof"), "Woof"); - ASSERT_EQ(ToRustStructStyle("woof_woof"), "WoofWoof"); - EXPECT_THROW(ToRustStructStyle(""), InternalError); -} - -TEST(NameTransforms, ToRustMacroStyle) { - ASSERT_EQ(ToRustMacroStyle("Woof"), "woof"); - ASSERT_EQ(ToRustMacroStyle("woof"), "woof"); - ASSERT_EQ(ToRustMacroStyle("woof_Woof"), "woof_woof"); - EXPECT_THROW(ToRustMacroStyle(""), InternalError); -} - -TEST(NameTransforms, ToRustConstantStyle) { - ASSERT_EQ(ToRustConstantStyle("Woof"), "WOOF"); - ASSERT_EQ(ToRustConstantStyle("woof"), "WOOF"); - ASSERT_EQ(ToRustConstantStyle("woof_Woof"), "WOOF_WOOF"); - EXPECT_THROW(ToRustConstantStyle(""), InternalError); -} - -TEST(NameTransforms, PrefixName) { - ASSERT_EQ(PrefixName({"Woof"}), "TVM_Woof"); - ASSERT_EQ(PrefixName({"woof"}), "TVM_woof"); - ASSERT_EQ(PrefixName({"woof", "moo"}), "TVM_woof_moo"); - EXPECT_THROW(PrefixName({}), InternalError); - EXPECT_THROW(PrefixName({""}), InternalError); -} - -TEST(NameTransforms, PrefixGeneratedName) { - ASSERT_EQ(PrefixGeneratedName({"Woof"}), "TVMGen_Woof"); - ASSERT_EQ(PrefixGeneratedName({"woof"}), "TVMGen_woof"); - ASSERT_EQ(PrefixGeneratedName({"woof", "moo"}), "TVMGen_woof_moo"); - EXPECT_THROW(PrefixGeneratedName({}), InternalError); - EXPECT_THROW(PrefixGeneratedName({""}), InternalError); -} - -TEST(NameTransforms, CombineNames) { - ASSERT_EQ(CombineNames({"woof"}), "woof"); - ASSERT_EQ(CombineNames({"Woof", "woof"}), "Woof_woof"); - ASSERT_EQ(CombineNames({"Woof", "woof", "woof"}), "Woof_woof_woof"); - ASSERT_EQ(CombineNames({"Woof", "moo", "t"}), "Woof_moo_t"); - - EXPECT_THROW(CombineNames({}), InternalError); - EXPECT_THROW(CombineNames({""}), InternalError); - EXPECT_THROW(CombineNames({"Woof", ""}), InternalError); - EXPECT_THROW(CombineNames({"", "Woof"}), InternalError); -} - -TEST(NameTransforms, SanitizeName) { - ASSERT_EQ(SanitizeName("+_+ "), "____"); - ASSERT_EQ(SanitizeName("input+"), "input_"); - ASSERT_EQ(SanitizeName("input-"), "input_"); - ASSERT_EQ(SanitizeName("input++"), "input__"); - ASSERT_EQ(SanitizeName("woof:1"), "woof_1"); - EXPECT_THROW(SanitizeName(""), InternalError); -} - -TEST(NameTransforms, CombinedLogic) { - ASSERT_EQ(ToCFunctionStyle(PrefixName({"Device", "target", "Invoke"})), "TVMDeviceTargetInvoke"); - ASSERT_EQ(ToCFunctionStyle(PrefixGeneratedName({"model", "Run"})), "TVMGenModelRun"); - ASSERT_EQ(ToCVariableStyle(PrefixName({"Device", "target", "t"})), "tvm_device_target_t"); - ASSERT_EQ(ToCVariableStyle(PrefixGeneratedName({"model", "Devices"})), "tvmgen_model_devices"); -} - -TEST(NameTransforms, Internal_ToCamel) { - ASSERT_EQ(ToCamel("Woof"), "Woof"); - ASSERT_EQ(ToCamel("woof"), "Woof"); - ASSERT_EQ(ToCamel("woof_woof"), "WoofWoof"); -} - -TEST(NameTransforms, Internal_ToCamel_Allocation) { - std::string woof = "Woof_woof_woof_woof"; - std::string camel = ToCamel(woof); - std::string check; - check.reserve(woof.size()); - - // Check that the pre-allocation happens - ASSERT_EQ(camel.capacity(), check.capacity()); -} - -} // namespace backend -} // namespace relay -} // namespace tvm diff --git a/tests/cpp/ndarray_test.cc b/tests/cpp/ndarray_test.cc index cd5c75410aae..57ad3ba90b40 100644 --- a/tests/cpp/ndarray_test.cc +++ b/tests/cpp/ndarray_test.cc @@ -17,8 +17,8 @@ * under the License. */ -#include #include +#include #include using namespace tvm; diff --git a/tests/cpp/nested_msg_test.cc b/tests/cpp/nested_msg_test.cc index 9ddae05e59e3..784ae2ab415a 100644 --- a/tests/cpp/nested_msg_test.cc +++ b/tests/cpp/nested_msg_test.cc @@ -17,11 +17,11 @@ * under the License. */ -#include #include #include #include #include +#include #include #include diff --git a/tests/cpp/object_protocol_test.cc b/tests/cpp/object_protocol_test.cc index 42928b484da9..c4c83dcd95c2 100644 --- a/tests/cpp/object_protocol_test.cc +++ b/tests/cpp/object_protocol_test.cc @@ -17,8 +17,8 @@ * under the License. */ -#include #include +#include #include #include diff --git a/tests/cpp/packed_func_test.cc b/tests/cpp/packed_func_test.cc index 183aca1385a7..001ef3310c75 100644 --- a/tests/cpp/packed_func_test.cc +++ b/tests/cpp/packed_func_test.cc @@ -17,8 +17,8 @@ * under the License. */ -#include #include +#include #include #include #include diff --git a/tests/cpp/parallel_for_test.cc b/tests/cpp/parallel_for_test.cc index e32fd32012a6..2057044cc13f 100644 --- a/tests/cpp/parallel_for_test.cc +++ b/tests/cpp/parallel_for_test.cc @@ -17,7 +17,6 @@ * under the License. */ -#include #include #include #include diff --git a/tests/cpp/pass_immutable_module_test.cc b/tests/cpp/pass_immutable_module_test.cc deleted file mode 100644 index b90f1deee737..000000000000 --- a/tests/cpp/pass_immutable_module_test.cc +++ /dev/null @@ -1,86 +0,0 @@ -/* - * 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 - -using namespace tvm; -using namespace transform; - -Pass MutateModulePass() { - auto pass_func = [=](IRModule mod, PassContext pc) -> IRModule { - GlobalVar var = mod->GetGlobalVar("dummyFunction"); - mod->Remove(var); - return mod; - }; - return tvm::transform::CreateModulePass(pass_func, 1, "ImmutableModulev1", {}); -} - -Pass DoNotMutateModulePass() { - auto pass_func = [=](IRModule mod, PassContext pc) -> IRModule { - IRModule result(mod->functions, mod->type_definitions, mod->Imports(), mod->source_map, - mod->attrs); - GlobalVar var = result->GetGlobalVar("dummyFunction"); - result->Remove(var); - return result; - }; - return tvm::transform::CreateModulePass(pass_func, 1, "ImmutableModulev2", {}); -} - -IRModule preamble() { - auto x = relay::Var("x", relay::Type()); - auto f = relay::Function(tvm::Array{x}, x, relay::Type(), {}); - ICHECK(f->IsInstance()); - - auto global_var = GlobalVar("dummyFunction"); - auto mod = IRModule::FromExpr(f, {{global_var, f}}, {}); - return mod; -} - -TEST(Relay, ModuleIsMutated) { - IRModule mod = preamble(); - - EXPECT_THROW( - { - auto pass_ctx = relay::transform::PassContext::Create(); - pass_ctx->config.Set("testing.immutable_module", Bool(true)); - { - tvm::With ctx_scope(pass_ctx); - mod = MutateModulePass()(mod); - } - }, - runtime::InternalError); -} - -TEST(Relay, ModuleIsNotMutated) { - IRModule mod = preamble(); - - auto pass_ctx = relay::transform::PassContext::Create(); - pass_ctx->config.Set("testing.immutable_module", Bool(true)); - { - tvm::With ctx_scope(pass_ctx); - mod = DoNotMutateModulePass()(mod); - } -} diff --git a/tests/cpp/profiling_test.cc b/tests/cpp/profiling_test.cc deleted file mode 100644 index d2fc0e95db2c..000000000000 --- a/tests/cpp/profiling_test.cc +++ /dev/null @@ -1,41 +0,0 @@ -/* - * 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 - -namespace tvm { -namespace runtime { -TEST(DefaultTimer, Basic) { - using namespace tvm::runtime; - Device dev; - dev.device_type = kDLCPU; - dev.device_id = 0; - - Timer t = Timer::Start(dev); - std::this_thread::sleep_for(std::chrono::milliseconds(10)); - t->Stop(); - int64_t elapsed = t->SyncAndGetElapsedNanos(); - CHECK_GT(elapsed, 9 * 1e6); -} -} // namespace runtime -} // namespace tvm diff --git a/tests/cpp/random_engine_test.cc b/tests/cpp/random_engine_test.cc index bc835dede4ee..078f99bd6e90 100644 --- a/tests/cpp/random_engine_test.cc +++ b/tests/cpp/random_engine_test.cc @@ -17,8 +17,8 @@ * under the License. */ -#include #include +#include #include TEST(RandomEngine, Randomness) { diff --git a/tests/cpp/si_builder_test.cc b/tests/cpp/si_builder_test.cc deleted file mode 100644 index f65debaa6b17..000000000000 --- a/tests/cpp/si_builder_test.cc +++ /dev/null @@ -1,399 +0,0 @@ -/* - * 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 - -tvm::Span _CreateSpan(std::string text) { - return tvm::Span(tvm::SourceName::Get(text), 0, 0, 0, 0); -} - -class RelayCheckSpan : public tvm::relay::ExprVisitor { - public: - std::vector tmp_result_; - std::vector lhs_spans_; - std::vector rhs_spans_; - - std::vector CollectSpan(tvm::relay::Expr expr) { - tmp_result_.clear(); - VisitExpr(expr); - return tmp_result_; - } - - void Check(tvm::relay::Expr lhs, tvm::relay::Expr rhs) { - tvm::relay::Function lhs_f = - tvm::relay::Function(tvm::relay::FreeVars(lhs), lhs, tvm::relay::Type(), {}); - tvm::relay::Function rhs_f = - tvm::relay::Function(tvm::relay::FreeVars(rhs), rhs, tvm::relay::Type(), {}); - EXPECT_TRUE(tvm::StructuralEqual()(lhs_f, rhs_f)); - lhs_spans_ = CollectSpan(lhs); - rhs_spans_ = CollectSpan(rhs); - - EXPECT_EQ(lhs_spans_.size(), rhs_spans_.size()); - for (std::size_t i = 0; i != lhs_spans_.size(); i++) { - EXPECT_TRUE(tvm::StructuralEqual()(lhs_spans_[i], rhs_spans_[i])); - } - } - - void VisitExpr(const tvm::relay::Expr& expr) { - if (expr->span.defined()) { - tmp_result_.push_back(expr->span); - } - using TParent = ExprFunctor; - TParent::VisitExpr(expr); - visit_counter_.emplace(expr.get(), 1); - } -}; - -TEST(SIBuilder, SequentialSpan) { - using namespace tvm; - Array ingredients = {_CreateSpan("first"), _CreateSpan("second"), _CreateSpan("third")}; - - SequentialSpan seq_span_1{ingredients[0], ingredients[1]}; - EXPECT_EQ(seq_span_1->spans.size(), 2); - for (std::size_t i = 0; i != seq_span_1->spans.size(); i++) { - EXPECT_EQ(seq_span_1->spans[i], ingredients[i]); - } - - // nested SequentialSpan test - SequentialSpan seq_span_2{seq_span_1, ingredients[2]}; - EXPECT_EQ(seq_span_2->spans.size(), 3); - for (std::size_t i = 0; i != seq_span_2->spans.size(); i++) { - EXPECT_EQ(seq_span_2->spans[i], ingredients[i]); - } - - // Array constructor test - Array tvm_array(ingredients); - SequentialSpan seq_span_3(tvm_array); - EXPECT_EQ(seq_span_3->spans.size(), 3); - for (std::size_t i = 0; i != seq_span_3->spans.size(); i++) { - EXPECT_EQ(seq_span_3->spans[i], ingredients[i]); - } -} - -TEST(SIBuilder, CreateSapn) { - using namespace tvm; - auto pass_ctx = transform::PassContext::Create(); - pass_ctx->config.Set("ir.enable_si_builder", Bool(true)); - tvm::With ctx_scope(pass_ctx); - Span span_1 = _CreateSpan("first"); - { - SIBuilder si_builder(span_1); - EXPECT_EQ(span_1, si_builder.Build()); - } - - Span span_2 = _CreateSpan("second"); - Array ingredients = {span_1, span_2}; - SequentialSpan seq_span_1{ingredients[0], ingredients[1]}; - { - SIBuilder si_builder_1(seq_span_1); - SIBuilder si_builder_2({span_1, span_2}); - SIBuilder si_builder_3{span_1, span_2}; - - Span created_span_1 = si_builder_1.Build(); - Span created_span_2 = si_builder_2.Build(); - Span created_span_3 = si_builder_3.Build(); - - auto created_seq_span_1 = created_span_1.as(); - auto created_seq_span_2 = created_span_2.as(); - auto created_seq_span_3 = created_span_3.as(); - EXPECT_EQ(created_seq_span_1->spans.size(), 2); - EXPECT_EQ(created_seq_span_2->spans.size(), 2); - EXPECT_EQ(created_seq_span_3->spans.size(), 2); - for (std::size_t i = 0; i != 2; i++) { - EXPECT_EQ(created_seq_span_1->spans[i], ingredients[i]); - EXPECT_EQ(created_seq_span_2->spans[i], ingredients[i]); - EXPECT_EQ(created_seq_span_3->spans[i], ingredients[i]); - } - } -} - -TEST(SIBuilder, DisableSIBuilder) { - using namespace tvm; - auto pass_ctx = transform::PassContext::Create(); - pass_ctx->config.Set("ir.enable_si_builder", Bool(false)); - tvm::With ctx_scope(pass_ctx); - Span span_1 = _CreateSpan("first"); - { - SIBuilder si_builder(span_1); - EXPECT_NE(span_1, si_builder.Build()); - } -} - -TEST(SIBuilder, RelayRecursivelyFill) { - using namespace tvm; - auto pass_ctx = transform::PassContext::Create(); - pass_ctx->config.Set("ir.enable_si_builder", Bool(true)); - tvm::With ctx_scope(pass_ctx); - Span test_span = _CreateSpan("test_span"); - Span a_node_span = _CreateSpan("a_node"); - - auto tensor_type = relay::TensorType({2, 3}, tvm::DataType::Float(32)); - relay::Expr add_op = relay::Op::Get("add"); - relay::Expr relu_op = relay::Op::Get("nn.relu"); - relay::Expr leaky_relu_op = relay::Op::Get("nn.leaky_relu"); - // Reset span of OpNode. Because a relay Op Node is a static reference, any change on it will - // be assigned the original object. - add_op->span = Span(); - relu_op->span = Span(); - leaky_relu_op->span = Span(); - - relay::Expr a = relay::Var("a", tensor_type, a_node_span); - relay::Expr x = relay::Call(relu_op, {a}, tvm::Attrs(), {}); - relay::Expr y = relay::Call(leaky_relu_op, {x}, tvm::Attrs(), {}); - relay::Expr z = relay::Call(add_op, {y, x}, tvm::Attrs(), {}); - - relay::Expr expected_a = relay::Var("a", tensor_type, a_node_span); - relay::Expr expected_x = relay::Call(relu_op, {expected_a}, tvm::Attrs(), {}, test_span); - relay::Expr expected_y = relay::Call(leaky_relu_op, {expected_x}, tvm::Attrs(), {}, test_span); - relay::Expr expected_z = - relay::Call(add_op, {expected_y, expected_x}, tvm::Attrs(), {}, test_span); - - SIBuilder si_builder(test_span); - si_builder.RecursivelyFillSpan(z, {a}); - RelayCheckSpan checker; - checker.Check(z, expected_z); -} - -TEST(SIBuilder, RelayCollectSpans) { - using namespace tvm; - auto pass_ctx = transform::PassContext::Create(); - pass_ctx->config.Set("ir.enable_si_builder", Bool(true)); - tvm::With ctx_scope(pass_ctx); - Span a_node_span = _CreateSpan("a_node"); - Span x_node_span = _CreateSpan("x_node"); - Span y_node_span = _CreateSpan("y_node"); - Span z_node_span = _CreateSpan("z_node"); - std::vector target = {z_node_span, y_node_span, x_node_span, a_node_span}; - - auto tensor_type = relay::TensorType({2, 3}, tvm::DataType::Float(32)); - relay::Expr add_op = relay::Op::Get("add"); - relay::Expr relu_op = relay::Op::Get("nn.relu"); - relay::Expr leaky_relu_op = relay::Op::Get("nn.leaky_relu"); - // Reset span of OpNode. Because a relay Op Node is a static reference, any change on it will - // be assigned the original object. - add_op->span = Span(); - relu_op->span = Span(); - leaky_relu_op->span = Span(); - - relay::Expr a = relay::Var("a", tensor_type, a_node_span); - relay::Expr x = relay::Call(relu_op, {a}, tvm::Attrs(), {}, x_node_span); - relay::Expr y = relay::Call(leaky_relu_op, {x}, tvm::Attrs(), {}, y_node_span); - relay::Expr z = relay::Call(add_op, {y, x}, tvm::Attrs(), {}, z_node_span); - - SIBuilder si_builder(z, {a}); - Span created_span = si_builder.Build(); - auto created_seq_span = created_span.as(); - EXPECT_EQ(created_seq_span->spans.size(), 4); - for (std::size_t i = 0; i != created_seq_span->spans.size(); i++) { - EXPECT_TRUE(StructuralEqual()(created_seq_span->spans[i], target[i])); - } -} - -TEST(SIBuilder, TirCollectSpansPrimExpr) { - using namespace tvm; - auto pass_ctx = transform::PassContext::Create(); - pass_ctx->config.Set("ir.enable_si_builder", Bool(true)); - tvm::With ctx_scope(pass_ctx); - Span a_node_span = _CreateSpan("a_node"); - Span b_node_span = _CreateSpan("b_node"); - Span x_node_span = _CreateSpan("x_node"); - Span add_1_node_span = _CreateSpan("add_1_node"); - Span add_2_node_span = _CreateSpan("add_2_node"); - Span z_node_span = _CreateSpan("z_node"); - std::vector target = {z_node_span, add_2_node_span, add_1_node_span, x_node_span, - a_node_span}; - tir::Var a("a"); - tir::Var b("b"); - auto x = a + b; - auto add_1 = x + 1; - auto add_2 = add_1 + 2; - auto z = max(add_2, 100); - x->span = x_node_span; - a->span = a_node_span; - b->span = b_node_span; - add_1->span = add_1_node_span; - add_2->span = add_2_node_span; - z->span = z_node_span; - - SIBuilder si_builder(z, {x}); - Span created_span = si_builder.Build(); - auto created_seq_span = created_span.as(); - - EXPECT_EQ(created_seq_span->spans.size(), 4); - for (std::size_t i = 0; i != created_seq_span->spans.size(); i++) { - EXPECT_TRUE(StructuralEqual()(created_seq_span->spans[i], target[i])); - } -} - -TEST(SIBuilder, TirCollectSpansStmtWithPrimInput) { - using namespace tvm; - auto pass_ctx = transform::PassContext::Create(); - pass_ctx->config.Set("ir.enable_si_builder", Bool(true)); - tvm::With ctx_scope(pass_ctx); - Span a_node_span = _CreateSpan("a_node"); - Span b_node_span = _CreateSpan("b_node"); - Span x_node_span = _CreateSpan("x_node"); - Span z_node_span = _CreateSpan("z_plus_1"); - Span stmt_node_span = _CreateSpan("stmt_node"); - std::vector target = {stmt_node_span, z_node_span, x_node_span}; - tir::Var a("a"); - tir::Var b("b"); - auto x = a + b; - x->span = x_node_span; - auto fmaketest = [&]() { - auto z = x + 1; - z->span = z_node_span; - tir::Stmt ret = te::Evaluate(z); - return ret; - }; - auto stmt = fmaketest(); - stmt->span = stmt_node_span; - SIBuilder si_builder(stmt, {x}); - Span created_span = si_builder.Build(); - auto created_seq_span = created_span.as(); - - EXPECT_EQ(created_seq_span->spans.size(), 3); - for (std::size_t i = 0; i != created_seq_span->spans.size(); i++) { - EXPECT_TRUE(StructuralEqual()(created_seq_span->spans[i], target[i])); - } -} - -TEST(SIBuilder, TirCollectSpansStmtWithStmtInput) { - using namespace tvm; - auto pass_ctx = transform::PassContext::Create(); - pass_ctx->config.Set("ir.enable_si_builder", Bool(true)); - tvm::With ctx_scope(pass_ctx); - Span zero_node_span = _CreateSpan("zero_node"); - Span body_node_span = _CreateSpan("body_node"); - Span init_node_span = _CreateSpan("init_node"); - Span block_node_span = _CreateSpan("block_node"); - std::vector target = {block_node_span, init_node_span, body_node_span}; - - tir::Stmt zero = tir::Evaluate(Integer(0), zero_node_span); - tir::Stmt body = tir::Evaluate(Integer(1), body_node_span); - tir::Stmt init = tir::IfThenElse(tir::const_true(), zero, zero, init_node_span); - tir::Block block({}, {}, {}, "block", body, init, Array(), - Array(), Map(), block_node_span); - SIBuilder si_builder(block, {init}); - Span created_span = si_builder.Build(); - auto created_seq_span = created_span.as(); - - EXPECT_EQ(created_seq_span->spans.size(), 3); - for (std::size_t i = 0; i != created_seq_span->spans.size(); i++) { - EXPECT_TRUE(StructuralEqual()(created_seq_span->spans[i], target[i])); - } -} - -TEST(SIBuilder, TirRecursivelyFillPrimExpr) { - using namespace tvm; - auto pass_ctx = transform::PassContext::Create(); - pass_ctx->config.Set("ir.enable_si_builder", Bool(true)); - tvm::With ctx_scope(pass_ctx); - Span test_span = _CreateSpan("test_span"); - tir::Var a("a"); - tir::Var b("b"); - auto x = a + b; - auto add_1 = x + 1; - auto add_2 = add_1 + 2; - auto z = max(add_2, 100); - - SIBuilder si_builder(test_span); - si_builder.RecursivelyFillSpan(z, {a, b}); - EXPECT_TRUE(!a->span.defined()); - EXPECT_TRUE(!b->span.defined()); - EXPECT_TRUE(StructuralEqual()(x->span, test_span)); - EXPECT_TRUE(StructuralEqual()(add_1->span, test_span)); - EXPECT_TRUE(StructuralEqual()(add_2->span, test_span)); - EXPECT_TRUE(StructuralEqual()(z->span, test_span)); - - ObjectRef tmp = z; - PrimExpr zz = Downcast(tmp); - std::ostringstream os; - os << z; - EXPECT_TRUE(zz.same_as(z)); - EXPECT_EQ(os.str(), "T.max(a + b + 1 + 2, 100)"); -} - -TEST(SIBuilder, TirRecursivelyFillStmtWithPrimInput) { - using namespace tvm; - auto pass_ctx = transform::PassContext::Create(); - pass_ctx->config.Set("ir.enable_si_builder", Bool(true)); - tvm::With ctx_scope(pass_ctx); - Span test_span = _CreateSpan("test_span"); - tir::Var a("a"); - tir::Var b("b"); - auto x = a + b; - auto z = x + 1; - tir::Stmt stmt = te::Evaluate(z); - SIBuilder si_builder(test_span); - const std::unordered_set inputs = {a, b}; - si_builder.RecursivelyFillSpan(stmt, inputs); - - EXPECT_TRUE(!a->span.defined()); - EXPECT_TRUE(!b->span.defined()); - EXPECT_TRUE(StructuralEqual()(x->span, test_span)); - EXPECT_TRUE(StructuralEqual()(z->span, test_span)); - EXPECT_TRUE(StructuralEqual()(stmt->span, test_span)); - - ObjectRef tmp = z; - PrimExpr zz = Downcast(tmp); - std::ostringstream os; - os << z; - EXPECT_TRUE(zz.same_as(z)); - EXPECT_EQ(os.str(), "a + b + 1"); -} - -TEST(SIBuilder, TirRecursivelyFillStmtWithStmtInput) { - using namespace tvm; - auto pass_ctx = transform::PassContext::Create(); - pass_ctx->config.Set("ir.enable_si_builder", Bool(true)); - tvm::With ctx_scope(pass_ctx); - tir::Stmt zero = tir::Evaluate(Integer(0)); - tir::Stmt init = tir::IfThenElse(tir::const_true(), zero, zero); - tir::Stmt body = tir::Evaluate(Integer(1)); - tir::Block block(/*iter_vars=*/{}, /*reads=*/{}, - /*writes=*/{}, /*name_hint=*/"block", /*body=*/body, - /*init=*/init); - - Span test_span = _CreateSpan("test_span"); - const std::unordered_set inputs = {init}; - SIBuilder si_builder(test_span); - si_builder.RecursivelyFillSpan(block, {init}); - EXPECT_TRUE(!zero->span.defined()); - EXPECT_TRUE(!init->span.defined()); - EXPECT_TRUE(StructuralEqual()(body->span, test_span)); - EXPECT_TRUE(StructuralEqual()(block->span, test_span)); - - tir::Stmt expected_zero = tir::Evaluate(Integer(0)); - tir::Stmt expected_init = tir::IfThenElse(tir::const_true(), zero, zero); - tir::Stmt expected_body = tir::Evaluate(Integer(1)); - tir::Block expected_block(/*iter_vars=*/{}, /*reads=*/{}, - /*writes=*/{}, /*name_hint=*/"block", /*body=*/expected_body, - /*init=*/expected_init); - EXPECT_TRUE(tvm::StructuralEqual()(block, expected_block)); -} diff --git a/tests/cpp/support/scalars_test.cc b/tests/cpp/support/scalars_test.cc index d55f0541fa40..52bd2dc148c8 100644 --- a/tests/cpp/support/scalars_test.cc +++ b/tests/cpp/support/scalars_test.cc @@ -20,7 +20,6 @@ #include "../../../src/support/scalars.h" #include -#include namespace tvm { namespace support { diff --git a/tests/cpp/support_test.cc b/tests/cpp/support_test.cc index 01111d910246..20272d67c42f 100644 --- a/tests/cpp/support_test.cc +++ b/tests/cpp/support_test.cc @@ -17,8 +17,8 @@ * under the License. */ -#include #include +#include #include "../../src/support/hexdump.h" #include "../../src/support/utils.h" diff --git a/tests/cpp/target/compilation_config_test.cc b/tests/cpp/target/compilation_config_test.cc deleted file mode 100644 index 88e30f9dd766..000000000000 --- a/tests/cpp/target/compilation_config_test.cc +++ /dev/null @@ -1,362 +0,0 @@ -/* - * 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 - -namespace tvm { -namespace { - -Target TestCpuTarget() { return Target("llvm -mcpu arm64"); } - -Target TestCudaTarget() { return Target("nvidia/tesla-p40"); } - -Target TestDefaultCpuTarget() { return Target("llvm"); } - -Target TestExtDevTarget() { return Target("ext_dev"); } - -TVM_REGISTER_TARGET_KIND("test_ext_codegen_1", kDLCUDA) - .set_attr(tvm::attr::kIsExternalCodegen, Bool(true)); - -TVM_REGISTER_TARGET_KIND("test_ext_codegen_2", kDLCUDA) - .set_attr(tvm::attr::kIsExternalCodegen, Bool(true)); - -Target TestExtCodegenTarget1() { return Target("test_ext_codegen_1"); } -Target TestExtCodegenTarget2() { return Target("test_ext_codegen_2"); } - -CompilationConfig TestCompilationConfig() { - transform::PassContext pass_ctx = transform::PassContext::Create(); - Target host_target = TestDefaultCpuTarget(); - Target cuda_target = Target::WithHost(TestCudaTarget(), host_target); - Target cpu_target = Target::WithHost(TestCpuTarget(), host_target); - return CompilationConfig(pass_ctx, {cuda_target, cpu_target}); -} - -TEST(CompilationConfig, Constructor_Heterogeneous_RuleA_RuleF_ReplaceHost) { - transform::PassContext pass_ctx = transform::PassContext::Create(); - - Target host_target = TestDefaultCpuTarget(); - Target cuda_target = Target::WithHost(TestCudaTarget(), host_target); - Target ignored_target = TestExtDevTarget(); - Target raw_cpu_target = Target::WithHost(TestCpuTarget(), ignored_target); - CompilationConfig config(pass_ctx, {cuda_target, raw_cpu_target}); - - Target cpu_target = Target::WithHost(TestCpuTarget(), host_target); - VirtualDevice expected_default_primitive_virtual_device(kDLCPU, 0, cpu_target); - VirtualDevice expected_host_virtual_device(kDLCPU, 0, host_target); - - // Host is chosen as per Rule A. - EXPECT_TRUE(config->host_target.defined()); - EXPECT_TRUE(StructuralEqual()(config->host_target, host_target)); - EXPECT_TRUE(StructuralEqual()(config->host_virtual_device, expected_host_virtual_device)); - - ASSERT_EQ(config->primitive_targets.size(), 2); - EXPECT_TRUE(StructuralEqual()(config->primitive_targets[0], cuda_target)); - // The host is taken from first raw target and overwritten in second. - EXPECT_TRUE(StructuralEqual()(config->primitive_targets[1], cpu_target)); - - // Default primitive virtual device chosen as per Rule F - EXPECT_TRUE(StructuralEqual()(config->default_primitive_virtual_device, - expected_default_primitive_virtual_device)); - - // Heterogeneous case. - ASSERT_FALSE(config->optional_homogeneous_target.defined()); -} - -TEST(CompilationConfig, Constructor_Homogeneous_RuleA_RuleE) { - transform::PassContext pass_ctx = transform::PassContext::Create(); - - Target host_target = TestDefaultCpuTarget(); - Target cuda_target = Target::WithHost(TestCudaTarget(), host_target); - CompilationConfig config(pass_ctx, {cuda_target}); - - VirtualDevice expected_default_primitive_virtual_device(kDLCUDA, 0, cuda_target); - VirtualDevice expected_host_virtual_device(kDLCPU, 0, host_target); - - // Host is chose as per Rule A. - EXPECT_TRUE(config->host_target.defined()); - EXPECT_TRUE(StructuralEqual()(config->host_target, host_target)); - EXPECT_TRUE(StructuralEqual()(config->host_virtual_device, expected_host_virtual_device)); - - ASSERT_EQ(config->primitive_targets.size(), 1); - EXPECT_TRUE(StructuralEqual()(config->primitive_targets[0], cuda_target)); - - // Default primitive virtual device chose as per rule E. - EXPECT_TRUE(StructuralEqual()(config->default_primitive_virtual_device, - expected_default_primitive_virtual_device)); - - // Homogeneous case. - ASSERT_TRUE(config->optional_homogeneous_target.defined()); - EXPECT_TRUE(StructuralEqual()(config->optional_homogeneous_target, cuda_target)); -} - -TEST(CompilationConfig, Constructor_Heterogeneous_RuleB_RuleD) { - transform::PassContext pass_ctx = transform::PassContext::Create(); - pass_ctx->config.Set("relay.fallback_device_type", Integer(static_cast(kDLCUDA))); - - Target raw_cuda_target = TestCudaTarget(); - Target raw_cpu_target = TestCpuTarget(); - CompilationConfig config(pass_ctx, {raw_cuda_target, raw_cpu_target}); - - Target host_target = TestCpuTarget(); - Target cuda_target = Target::WithHost(TestCudaTarget(), host_target); - Target cpu_target = Target::WithHost(TestCpuTarget(), host_target); - - VirtualDevice expected_default_primitive_virtual_device(kDLCUDA, 0, cuda_target); - VirtualDevice expected_host_virtual_device(kDLCPU, 0, host_target); - - // Host is chosen as per Rule B. - EXPECT_TRUE(config->host_target.defined()); - EXPECT_TRUE(StructuralEqual()(config->host_target, host_target)); - EXPECT_TRUE(StructuralEqual()(config->host_virtual_device, expected_host_virtual_device)); - - ASSERT_EQ(config->primitive_targets.size(), 2); - EXPECT_TRUE(StructuralEqual()(config->primitive_targets[0], cuda_target)); - EXPECT_TRUE(StructuralEqual()(config->primitive_targets[1], cpu_target)); - - // Default primitive virtual device chosen as per Rule D - EXPECT_TRUE(StructuralEqual()(config->default_primitive_virtual_device, - expected_default_primitive_virtual_device)); - - // Heterogeneous case. - ASSERT_FALSE(config->optional_homogeneous_target.defined()); -} - -TEST(CompilationConfig, Constructor_Homogeneous_RuleC_RuleE) { - transform::PassContext pass_ctx = transform::PassContext::Create(); - - Target raw_cuda_target = TestCudaTarget(); - CompilationConfig config(pass_ctx, {raw_cuda_target}); - - Target host_target = TestDefaultCpuTarget(); - Target cuda_target = Target::WithHost(TestCudaTarget(), host_target); - Target cpu_target = Target::WithHost(TestDefaultCpuTarget(), host_target); - - VirtualDevice expected_default_primitive_virtual_device(kDLCUDA, 0, cuda_target); - VirtualDevice expected_host_virtual_device(kDLCPU, 0, host_target); - - // Host is chosen as per Rule C. - EXPECT_TRUE(config->host_target.defined()); - EXPECT_TRUE(StructuralEqual()(config->host_target, host_target)); - EXPECT_TRUE(StructuralEqual()(config->host_virtual_device, expected_host_virtual_device)); - - ASSERT_EQ(config->primitive_targets.size(), 1); - EXPECT_TRUE(StructuralEqual()(config->primitive_targets[0], cuda_target)); - - // Default primitive virtual device chosen as per Rule E - EXPECT_TRUE(StructuralEqual()(config->default_primitive_virtual_device, - expected_default_primitive_virtual_device)); - - // Homogeneous case. - ASSERT_TRUE(config->optional_homogeneous_target.defined()); - EXPECT_TRUE(StructuralEqual()(config->optional_homogeneous_target, cuda_target)); -} - -TEST(CompilationConfig, Constructor_Heterogeneous_CorrectOrdering) { - transform::PassContext pass_ctx = transform::PassContext::Create(); - - Target host_target = TestDefaultCpuTarget(); - Target cuda_target = Target::WithHost(TestCudaTarget(), host_target); - Target ext_codegen1_target = Target::WithHost(TestExtCodegenTarget1(), host_target); - Target ext_codegen2_target = Target::WithHost(TestExtCodegenTarget2(), host_target); - CompilationConfig config(pass_ctx, {cuda_target, ext_codegen1_target, ext_codegen2_target}); - - ASSERT_EQ(config->primitive_targets.size(), 3); - EXPECT_TRUE(StructuralEqual()(config->primitive_targets[0], cuda_target)); - EXPECT_TRUE(StructuralEqual()(config->primitive_targets[1], ext_codegen1_target)); - EXPECT_TRUE(StructuralEqual()(config->primitive_targets[2], ext_codegen2_target)); -} - -TEST(CompilationConfig, Constructor_Heterogeneous_InvalidOrdering) { - transform::PassContext pass_ctx = transform::PassContext::Create(); - - Target host_target = TestDefaultCpuTarget(); - Target ext_codegen1_target = Target::WithHost(TestExtCodegenTarget1(), host_target); - Target cuda_target = Target::WithHost(TestCudaTarget(), host_target); - Target ext_codegen2_target = Target::WithHost(TestExtCodegenTarget2(), host_target); - - EXPECT_ANY_THROW( - CompilationConfig(pass_ctx, {ext_codegen1_target, cuda_target, ext_codegen2_target})); -} - -TEST(CompilationConfig, Constructor_Homogenous_JustExternalCodegen) { - transform::PassContext pass_ctx = transform::PassContext::Create(); - - Target host_target = TestDefaultCpuTarget(); - Target ext_codegen1_target = Target::WithHost(TestExtCodegenTarget1(), host_target); - - CompilationConfig config(pass_ctx, {ext_codegen1_target}); - ASSERT_EQ(config->primitive_targets.size(), 1); - EXPECT_TRUE(StructuralEqual()(config->primitive_targets[0], ext_codegen1_target)); -} - -TEST(CompliationConfig, Constructor_DuplicateKinds) { - transform::PassContext pass_ctx = transform::PassContext::Create(); - - Target host_target = TestDefaultCpuTarget(); - Target cuda_target_1 = Target::WithHost(TestCudaTarget(), host_target); - Target cuda_target_2 = Target::WithHost(TestCudaTarget(), host_target); - - EXPECT_ANY_THROW(CompilationConfig(pass_ctx, {cuda_target_1, cuda_target_2})); -} - -TEST(CompilationConfig, Constructor_NoTargets) { - transform::PassContext pass_ctx = transform::PassContext::Create(); - EXPECT_ANY_THROW(CompilationConfig(pass_ctx, {})); -} - -TEST(CompilationConfig, Constructor_InvalidAttribute) { - transform::PassContext pass_ctx = transform::PassContext::Create(); - pass_ctx->config.Set("relay.fallback_device_type", Integer(static_cast(kInvalidDeviceType))); - - Target cuda_target = Target::WithHost(TestCudaTarget(), TestDefaultCpuTarget()); - EXPECT_ANY_THROW(CompilationConfig(pass_ctx, {cuda_target})); -} - -TEST(CompilationConfig, Constructor_NoMatchingPrimitiveTarget) { - transform::PassContext pass_ctx = transform::PassContext::Create(); - pass_ctx->config.Set("relay.fallback_device_type", Integer(static_cast(kDLMetal))); - Target host_target = TestDefaultCpuTarget(); - Target cuda_target = Target::WithHost(TestCudaTarget(), host_target); - EXPECT_ANY_THROW(CompilationConfig(pass_ctx, {cuda_target})); -} - -TEST(CompilationConfig, Constructor_DefaultNoMatchingPrimitiveTarget) { - transform::PassContext pass_ctx = transform::PassContext::Create(); - Target host_target = TestDefaultCpuTarget(); - Target cuda_target = Target::WithHost(TestCudaTarget(), host_target); - Target ext_target = Target::WithHost(TestExtDevTarget(), host_target); - EXPECT_ANY_THROW(CompilationConfig config(pass_ctx, {cuda_target, ext_target})); -} - -TEST(CompilationConfig, Constructor_Idempotent) { - transform::PassContext pass_ctx = transform::PassContext::Create(); - - Target host_target = TestDefaultCpuTarget(); - Target cuda_target = Target::WithHost(TestCudaTarget(), host_target); - Target ignored_target = TestExtDevTarget(); - Target raw_cpu_target = Target::WithHost(TestCpuTarget(), ignored_target); - CompilationConfig orig_config(pass_ctx, {cuda_target, raw_cpu_target}); - - CompilationConfig reconstructed_config(pass_ctx, orig_config->primitive_targets); - - ASSERT_EQ(orig_config->primitive_targets.size(), reconstructed_config->primitive_targets.size()); - ASSERT_TRUE(StructuralEqual()(orig_config->primitive_targets[0], - reconstructed_config->primitive_targets[0])); - ASSERT_TRUE(StructuralEqual()(orig_config->primitive_targets[1], - reconstructed_config->primitive_targets[1])); -} - -TEST(CompilationConfig, FindPrimitiveTargetForDeviceOrFail_Valid) { - CompilationConfig config = TestCompilationConfig(); - Target cpu_target = Target::WithHost(TestCpuTarget(), TestDefaultCpuTarget()); - ASSERT_TRUE(StructuralEqual()(config->FindPrimitiveTargetForDeviceOrFail(kDLCPU), cpu_target)); -} - -TEST(CompilationConfig, FindPrimitiveTargetForDeviceOrFail_Invalid) { - CompilationConfig config = TestCompilationConfig(); - EXPECT_ANY_THROW(config->FindPrimitiveTargetForDeviceOrFail(kDLMetal)); -} - -TEST(CompilationConfig, FindPrimitiveTargetForKind_Found) { - CompilationConfig config = TestCompilationConfig(); - Target cuda_target = Target::WithHost(TestCudaTarget(), TestDefaultCpuTarget()); - ASSERT_TRUE(StructuralEqual()(config->FindPrimitiveTargetForKind("cuda").value(), cuda_target)); -} - -TEST(CompilationConfig, FindPrimitiveTargetForKind_NotFound) { - CompilationConfig config = TestCompilationConfig(); - ASSERT_FALSE(config->FindPrimitiveTargetForKind("cutlass").defined()); -} - -TEST(CompilationConfig, CanonicalTarget) { - Target host_target = TestDefaultCpuTarget(); - Target cuda_target = TestCudaTarget(); - Target cpu_target = TestCpuTarget(); - CompilationConfig config = TestCompilationConfig(); - - { - Target other_cuda_target = Target::WithHost(TestCudaTarget(), TestDefaultCpuTarget()); - ASSERT_NE(other_cuda_target, cuda_target); - ASSERT_EQ(config->CanonicalTarget(other_cuda_target), - config->FindPrimitiveTargetForKind("cuda")); - } - { - Target other_host_target = TestDefaultCpuTarget(); - ASSERT_NE(other_host_target, cuda_target); - ASSERT_EQ(config->CanonicalTarget(other_host_target), config->host_target); - } - { - Target other_target("cuda -max_num_threads=7"); - ASSERT_EQ(config->CanonicalTarget(other_target), other_target); - } -} - -TEST(CompilationConfig, CanonicalVirtualDevice) { - Target host_target = TestDefaultCpuTarget(); - Target cuda_target = TestCudaTarget(); - Target cpu_target = TestCpuTarget(); - CompilationConfig config = TestCompilationConfig(); - - { - VirtualDevice in = VirtualDevice(kDLCPU); - VirtualDevice actual = config->CanonicalVirtualDevice(in); - ASSERT_TRUE(actual->target.defined()); - EXPECT_TRUE(StructuralEqual()(actual->target, Target::WithHost(cpu_target, host_target))); - EXPECT_EQ(config->CanonicalVirtualDevice(in), actual); - } - { - VirtualDevice in = VirtualDevice(kDLCUDA); - VirtualDevice actual = config->CanonicalVirtualDevice(in); - ASSERT_TRUE(actual->target.defined()); - EXPECT_TRUE(StructuralEqual()(actual->target, Target::WithHost(cuda_target, host_target))); - EXPECT_EQ(config->CanonicalVirtualDevice(in), actual); - } - { - Target other_cuda_target = Target::WithHost(TestCudaTarget(), TestDefaultCpuTarget()); - VirtualDevice in = VirtualDevice(kDLCUDA, -1, other_cuda_target); - VirtualDevice actual = config->CanonicalVirtualDevice(in); - ASSERT_EQ(actual->target, config->FindPrimitiveTargetForKind("cuda")); - } - { - VirtualDevice in = VirtualDevice::ForMemoryScope("scope"); - VirtualDevice actual = config->CanonicalVirtualDevice(in); - EXPECT_EQ(config->CanonicalVirtualDevice(in), actual); - } - { - VirtualDevice in = VirtualDevice::FullyUnconstrained(); - VirtualDevice actual = config->CanonicalVirtualDevice(in); - EXPECT_EQ(config->CanonicalVirtualDevice(in), actual); - } - { - VirtualDevice in = VirtualDevice(); // ie structurally equal to FullyUnconstrained. - VirtualDevice actual = config->CanonicalVirtualDevice(in); - EXPECT_EQ(config->CanonicalVirtualDevice(in), VirtualDevice::FullyUnconstrained()); - } -} - -TEST(CompilationConfig, CanonicalVirtualDevice_NoMatchingTarget) { - CompilationConfig config = TestCompilationConfig(); - VirtualDevice no_such_target(kDLMetal); - EXPECT_ANY_THROW(config->CanonicalVirtualDevice(no_such_target)); -} - -} // namespace -} // namespace tvm diff --git a/tests/cpp/target/source/interface_c_test.cc b/tests/cpp/target/source/interface_c_test.cc deleted file mode 100644 index d9d9d80bbe31..000000000000 --- a/tests/cpp/target/source/interface_c_test.cc +++ /dev/null @@ -1,762 +0,0 @@ -/* - * 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 - -using ::testing::ContainsRegex; -using ::testing::HasSubstr; - -namespace tvm { -namespace codegen { - -runtime::Module InterfaceCCreate(std::string module_name, Array inputs, - Array outputs, Array pools, - Map io_pool_allocations, - Array devices, int workspace_size, - Map input_sizes, Map output_sizes); - -namespace { - -TEST(InterfaceAPI, ContainsHeaderGuards) { - std::stringstream upper_header_guard; - std::stringstream lower_header_guard; - - upper_header_guard << "#ifndef TVMGEN_ULTIMATE_CAT_SPOTTER_H_\n" - << "#define TVMGEN_ULTIMATE_CAT_SPOTTER_H_\n" - << "#include \n\n" - << "#ifdef __cplusplus\n" - << "extern \"C\" {\n" - << "#endif\n\n"; - - lower_header_guard << "\n#ifdef __cplusplus\n" - << "}\n" - << "#endif\n\n" - << "#endif // TVMGEN_ULTIMATE_CAT_SPOTTER_H_\n"; - - Map input_sizes; - input_sizes.Set("input", IntImm(DataType::Int(32), 0)); - Map output_sizes; - output_sizes.Set("output", IntImm(DataType::Int(32), 0)); - - runtime::Module test_module = InterfaceCCreate("ultimate_cat_spotter", {"input"}, {"output"}, {}, - {}, {}, 0, input_sizes, output_sizes); - std::string header_source = test_module->GetSource(); - - ASSERT_THAT(header_source, HasSubstr(upper_header_guard.str())); - ASSERT_THAT(header_source, HasSubstr(lower_header_guard.str())); -} - -TEST(InterfaceAPI, ContainsRunFunction) { - std::stringstream run_function; - - run_function << "/*!\n" - << " * \\brief entrypoint function for TVM module \"ultimate_cat_spotter\"\n" - << " * \\param inputs Input tensors for the module \n" - << " * \\param outputs Output tensors for the module \n" - << " */\n" - << "int32_t tvmgen_ultimate_cat_spotter_run(\n" - << " struct tvmgen_ultimate_cat_spotter_inputs* inputs,\n" - << " struct tvmgen_ultimate_cat_spotter_outputs* outputs\n" - << ");\n"; - - Map input_sizes; - input_sizes.Set("input", IntImm(DataType::Int(32), 0)); - Map output_sizes; - output_sizes.Set("output", IntImm(DataType::Int(32), 0)); - - runtime::Module test_module = InterfaceCCreate("ultimate_cat_spotter", {"input"}, {"output"}, {}, - {}, {}, 0, input_sizes, output_sizes); - std::string header_source = test_module->GetSource(); - ASSERT_THAT(header_source, HasSubstr(run_function.str())); -} - -TEST(InterfaceAPI, ContainsRunFunctionWithDevices) { - std::stringstream run_function; - - run_function << "/*!\n" - << " * \\brief entrypoint function for TVM module \"ultimate_cat_spotter\"\n" - << " * \\param inputs Input tensors for the module \n" - << " * \\param outputs Output tensors for the module \n" - << " * \\param devices Device context pointers for the module \n" - << " */\n" - << "int32_t tvmgen_ultimate_cat_spotter_run(\n" - << " struct tvmgen_ultimate_cat_spotter_inputs* inputs,\n" - << " struct tvmgen_ultimate_cat_spotter_outputs* outputs,\n" - << " struct tvmgen_ultimate_cat_spotter_devices* devices\n" - << ");\n"; - - Map input_sizes; - input_sizes.Set("input", IntImm(DataType::Int(32), 0)); - Map output_sizes; - output_sizes.Set("output", IntImm(DataType::Int(32), 0)); - - runtime::Module test_module = InterfaceCCreate("ultimate_cat_spotter", {"input"}, {"output"}, {}, - {}, {"device"}, 0, input_sizes, output_sizes); - std::string header_source = test_module->GetSource(); - - ASSERT_THAT(header_source, HasSubstr(run_function.str())); -} - -TEST(InterfaceAPI, ContainsRunFunctionWithWorkspacePools) { - std::stringstream run_function; - - run_function << "/*!\n" - << " * \\brief entrypoint function for TVM module \"ultimate_cat_spotter\"\n" - << " * \\param inputs Input tensors for the module \n" - << " * \\param outputs Output tensors for the module \n" - << " * \\param workspace_pools Workspace memory pool pointers for the module \n" - << " */\n" - << "int32_t tvmgen_ultimate_cat_spotter_run(\n" - << " struct tvmgen_ultimate_cat_spotter_inputs* inputs,\n" - << " struct tvmgen_ultimate_cat_spotter_outputs* outputs,\n" - << " struct tvmgen_ultimate_cat_spotter_workspace_pools* workspace_pools\n" - << ");\n"; - - Map input_sizes; - input_sizes.Set("input", IntImm(DataType::Int(32), 0)); - Map output_sizes; - output_sizes.Set("output", IntImm(DataType::Int(32), 0)); - - PoolInfo pool_info = WorkspacePoolInfo("my_memory_pool", {}); - tir::usmp::AllocatedPoolInfo allocated_pool_info = - tir::usmp::AllocatedPoolInfo(pool_info, 100000); - runtime::Module test_module = - InterfaceCCreate("ultimate_cat_spotter", {"input"}, {"output"}, {allocated_pool_info}, {}, {}, - 0, input_sizes, output_sizes); - std::string header_source = test_module->GetSource(); - - ASSERT_THAT(header_source, HasSubstr(run_function.str())); -} - -TEST(InterfaceAPI, ContainsRunFunctionWithWorkspaceAndConstantPools) { - std::stringstream run_function; - - run_function << "/*!\n" - << " * \\brief entrypoint function for TVM module \"ultimate_cat_spotter\"\n" - << " * \\param inputs Input tensors for the module \n" - << " * \\param outputs Output tensors for the module \n" - << " * \\param workspace_pools Workspace memory pool pointers for the module \n" - << " */\n" - << "int32_t tvmgen_ultimate_cat_spotter_run(\n" - << " struct tvmgen_ultimate_cat_spotter_inputs* inputs,\n" - << " struct tvmgen_ultimate_cat_spotter_outputs* outputs,\n" - << " struct tvmgen_ultimate_cat_spotter_workspace_pools* workspace_pools\n" - << ");\n"; - - Map input_sizes; - input_sizes.Set("input", IntImm(DataType::Int(32), 0)); - Map output_sizes; - output_sizes.Set("output", IntImm(DataType::Int(32), 0)); - - PoolInfo pool_info = WorkspacePoolInfo("my_memory_pool", {}); - PoolInfo const_info = ConstantPoolInfo( - "my_constant_pool", {}, - {{"const1", 0, runtime::NDArray::Empty({1}, DataType::Int(32), {kDLCPU, 0})}, - {"const2", 16, runtime::NDArray::Empty({1}, DataType::Float(64), {kDLCPU, 0})}}); - tir::usmp::AllocatedPoolInfo allocated_pool_info = - tir::usmp::AllocatedPoolInfo(pool_info, 100000); - tir::usmp::AllocatedPoolInfo allocated_const_info = - tir::usmp::AllocatedPoolInfo(const_info, 100000); - runtime::Module test_module = InterfaceCCreate("ultimate_cat_spotter", {"input"}, {"output"}, - {allocated_pool_info, allocated_const_info}, {}, - {}, 0, input_sizes, output_sizes); - std::string header_source = test_module->GetSource(); - ASSERT_THAT(header_source, HasSubstr(run_function.str())); - ASSERT_THAT( - header_source, - HasSubstr("#define TVMGEN_ULTIMATE_CAT_SPOTTER_MY_CONSTANT_POOL_CONSTANT_POOL_SIZE 24")); - ASSERT_THAT( - header_source, - ContainsRegex( - "#define TVMGEN_ULTIMATE_CAT_SPOTTER_MY_CONSTANT_POOL_CONSTANT_POOL_DATA \\\\\\\n " - "0x\\w\\w, 0x\\w\\w, 0x\\w\\w, 0x\\w\\w, 0x\\w\\w, 0x\\w\\w, 0x\\w\\w, 0x\\w\\w, " - "0x\\w\\w, 0x\\w\\w, 0x\\w\\w, 0x\\w\\w, 0x\\w\\w, " - "0x\\w\\w, 0x\\w\\w, 0x\\w\\w, \\\\\\\n 0x\\w\\w, 0x\\w\\w, 0x\\w\\w, 0x\\w\\w, " - "0x\\w\\w, 0x\\w\\w, 0x\\w\\w, 0x\\w\\w\\\\\\\n")); -} - -TEST(InterfaceAPI, ContainsRunFunctionWithWorkspacePoolsAndDevices) { - std::stringstream run_function; - - run_function << "/*!\n" - << " * \\brief entrypoint function for TVM module \"ultimate_cat_spotter\"\n" - << " * \\param inputs Input tensors for the module \n" - << " * \\param outputs Output tensors for the module \n" - << " * \\param workspace_pools Workspace memory pool pointers for the module \n" - << " * \\param devices Device context pointers for the module \n" - << " */\n" - << "int32_t tvmgen_ultimate_cat_spotter_run(\n" - << " struct tvmgen_ultimate_cat_spotter_inputs* inputs,\n" - << " struct tvmgen_ultimate_cat_spotter_outputs* outputs,\n" - << " struct tvmgen_ultimate_cat_spotter_workspace_pools* workspace_pools,\n" - << " struct tvmgen_ultimate_cat_spotter_devices* devices\n" - << ");\n"; - - Map input_sizes; - input_sizes.Set("input", IntImm(DataType::Int(32), 0)); - Map output_sizes; - output_sizes.Set("output", IntImm(DataType::Int(32), 0)); - - PoolInfo pool_info = WorkspacePoolInfo("my_memory_pool", {}); - tir::usmp::AllocatedPoolInfo allocated_pool_info = - tir::usmp::AllocatedPoolInfo(pool_info, 100000); - runtime::Module test_module = - InterfaceCCreate("ultimate_cat_spotter", {"input"}, {"output"}, {allocated_pool_info}, {}, - {"device"}, 0, input_sizes, output_sizes); - std::string header_source = test_module->GetSource(); - - ASSERT_THAT(header_source, HasSubstr(run_function.str())); -} - -TEST(InterfaceAPI, ContainsRunFunctionWithWorkspaceIO) { - std::stringstream run_function_with_map_functions; - - run_function_with_map_functions - << "/*!\n" - << " * \\brief Maps I/O inside the workspace pools for TVM module \"ultimate_cat_spotter\"\n" - << " * \\param workspace_pools Workspace memory pool struct for the module \n" - << " * \\return I/O tensor struct for the module \n" - << " */\n" - << "struct tvmgen_ultimate_cat_spotter_inputs tvmgen_ultimate_cat_spotter_map_inputs(\n" - << " struct tvmgen_ultimate_cat_spotter_workspace_pools* workspace_pools\n" - << ");\n" - << "\n" - << "/*!\n" - << " * \\brief Maps I/O inside the workspace pools for TVM module \"ultimate_cat_spotter\"\n" - << " * \\param workspace_pools Workspace memory pool struct for the module \n" - << " * \\return I/O tensor struct for the module \n" - << " */\n" - << "struct tvmgen_ultimate_cat_spotter_outputs tvmgen_ultimate_cat_spotter_map_outputs(\n" - << " struct tvmgen_ultimate_cat_spotter_workspace_pools* workspace_pools\n" - << ");\n" - << "\n" - << "/*!\n" - << " * \\brief entrypoint function for TVM module \"ultimate_cat_spotter\"\n" - << " * \\param workspace_pools Workspace memory pool pointers for the module \n" - << " */\n" - << "int32_t tvmgen_ultimate_cat_spotter_run(\n" - << " struct tvmgen_ultimate_cat_spotter_workspace_pools* workspace_pools\n" - << ");\n"; - - Map input_sizes; - input_sizes.Set("input", IntImm(DataType::Int(32), 0)); - Map output_sizes; - output_sizes.Set("output", IntImm(DataType::Int(32), 0)); - - PoolInfo pool_info = WorkspacePoolInfo("my_memory_pool", {}); - tir::usmp::AllocatedPoolInfo allocated_pool_info = - tir::usmp::AllocatedPoolInfo(pool_info, 100000); - tir::usmp::PoolAllocation pool_allocation_input{pool_info, 1000}; - tir::usmp::PoolAllocation pool_allocation_output{pool_info, 2000}; - runtime::Module test_module = - InterfaceCCreate("ultimate_cat_spotter", {"input"}, {"output"}, {allocated_pool_info}, - {{"input", pool_allocation_input}, {"output", pool_allocation_output}}, {}, - 0, input_sizes, output_sizes); - std::string header_source = test_module->GetSource(); - std::cout << header_source << "\n"; - ASSERT_THAT(header_source, HasSubstr(run_function_with_map_functions.str())); -} - -TEST(InterfaceAPI, ContainsInputStructSingle) { - std::stringstream input_struct; - std::stringstream input_size_macro; - - input_size_macro - << "/*!\n" - << " * \\brief Input tensor input size (in bytes) for TVM module \"ultimate_cat_spotter\" \n" - << " */\n" - << "#define TVMGEN_ULTIMATE_CAT_SPOTTER_INPUT_SIZE 537\n"; - - input_struct << "/*!\n" - << " * \\brief Input tensor pointers for TVM module \"ultimate_cat_spotter\" \n" - << " */\n" - << "struct tvmgen_ultimate_cat_spotter_inputs {\n" - << " void* input;\n" - << "};\n\n"; - - Map input_sizes; - input_sizes.Set("input", IntImm(DataType::Int(32), 537)); - Map output_sizes; - output_sizes.Set("output", IntImm(DataType::Int(32), 0)); - - runtime::Module test_module = InterfaceCCreate("ultimate_cat_spotter", {"input"}, {"output"}, {}, - {}, {}, 0, input_sizes, output_sizes); - std::string header_source = test_module->GetSource(); - - ASSERT_THAT(header_source, HasSubstr(input_struct.str())); - - ASSERT_THAT(header_source, HasSubstr(input_size_macro.str())); -} - -TEST(InterfaceAPI, ContainsInputStructMany) { - std::stringstream input_struct; - std::stringstream input1_size_macro; - std::stringstream input2_size_macro; - - input1_size_macro - << "/*!\n" - << " * \\brief Input tensor input1 size (in bytes) for TVM module \"ultimate_cat_spotter\" \n" - << " */\n" - << "#define TVMGEN_ULTIMATE_CAT_SPOTTER_INPUT1_SIZE 765\n"; - - input2_size_macro - << "/*!\n" - << " * \\brief Input tensor input2 size (in bytes) for TVM module \"ultimate_cat_spotter\" \n" - << " */\n" - << "#define TVMGEN_ULTIMATE_CAT_SPOTTER_INPUT2_SIZE 127\n"; - - input_struct << "struct tvmgen_ultimate_cat_spotter_inputs {\n" - << " void* input1;\n" - << " void* input2;\n" - << "};\n\n"; - - Map input_sizes; - input_sizes.Set("input1", IntImm(DataType::Int(32), 765)); - input_sizes.Set("input2", IntImm(DataType::Int(32), 127)); - Map output_sizes; - output_sizes.Set("output", IntImm(DataType::Int(32), 0)); - - runtime::Module test_module = - InterfaceCCreate("ultimate_cat_spotter", {"input1", "input2"}, {"output"}, {}, {}, {}, 0, - input_sizes, output_sizes); - std::string header_source = test_module->GetSource(); - - ASSERT_THAT(header_source, HasSubstr(input_struct.str())); - ASSERT_THAT(header_source, HasSubstr(input1_size_macro.str())); - ASSERT_THAT(header_source, HasSubstr(input2_size_macro.str())); -} - -TEST(InterfaceAPI, ContainsInputStructSanitised) { - std::stringstream input_struct; - std::stringstream input1_size_macro; - std::stringstream input2_size_macro; - - input1_size_macro << "/*!\n" - << " * \\brief Input tensor input_1 size (in bytes) for TVM module " - "\"ultimate_cat_spotter\" \n" - << " */\n" - << "#define TVMGEN_ULTIMATE_CAT_SPOTTER_INPUT_1_SIZE 765\n"; - - input2_size_macro << "/*!\n" - << " * \\brief Input tensor input_2 size (in bytes) for TVM module " - "\"ultimate_cat_spotter\" \n" - << " */\n" - << "#define TVMGEN_ULTIMATE_CAT_SPOTTER_INPUT_2_SIZE 127\n"; - - input_struct << "struct tvmgen_ultimate_cat_spotter_inputs {\n" - << " void* input_1;\n" - << " void* input_2;\n" - << "};\n\n"; - - Map input_sizes; - input_sizes.Set("input+1", IntImm(DataType::Int(32), 765)); - input_sizes.Set("input+2", IntImm(DataType::Int(32), 127)); - Map output_sizes; - output_sizes.Set("output", IntImm(DataType::Int(32), 0)); - - runtime::Module test_module = - InterfaceCCreate("ultimate_cat_spotter", {"input+1", "input+2"}, {"output"}, {}, {}, {}, 0, - input_sizes, output_sizes); - std::string header_source = test_module->GetSource(); - - std::cout << header_source << std::endl; - - ASSERT_THAT(header_source, HasSubstr(input_struct.str())); - ASSERT_THAT(header_source, HasSubstr(input1_size_macro.str())); - ASSERT_THAT(header_source, HasSubstr(input2_size_macro.str())); -} - -TEST(InterfaceAPI, ContainsInputStructClash) { - Map input_sizes; - input_sizes.Set("input+", IntImm(DataType::Int(32), 0)); - input_sizes.Set("input-", IntImm(DataType::Int(32), 0)); - Map output_sizes; - output_sizes.Set("output", IntImm(DataType::Int(32), 0)); - - runtime::Module test_module = - InterfaceCCreate("ultimate_cat_spotter", {"input+", "input-"}, {"output"}, {}, {}, {}, 0, - input_sizes, output_sizes); - ASSERT_THROW(test_module->GetSource(), InternalError); -} - -TEST(InterfaceAPI, ContainsOutputStructSingle) { - std::stringstream output_struct; - std::stringstream output_size_macro; - - output_size_macro << "/*!\n" - << " * \\brief Output tensor output size (in bytes) for TVM module " - "\"ultimate_cat_spotter\" \n" - << " */\n" - << "#define TVMGEN_ULTIMATE_CAT_SPOTTER_OUTPUT_SIZE 543\n"; - - output_struct << "/*!\n" - << " * \\brief Output tensor pointers for TVM module \"ultimate_cat_spotter\" \n" - << " */\n" - << "struct tvmgen_ultimate_cat_spotter_outputs {\n" - << " void* output;\n" - << "};\n\n"; - - Map input_sizes; - input_sizes.Set("input", IntImm(DataType::Int(32), 0)); - Map output_sizes; - output_sizes.Set("output", IntImm(DataType::Int(32), 543)); - - runtime::Module test_module = InterfaceCCreate("ultimate_cat_spotter", {"input"}, {"output"}, {}, - {}, {}, 0, input_sizes, output_sizes); - std::string header_source = test_module->GetSource(); - - ASSERT_THAT(header_source, HasSubstr(output_struct.str())); - ASSERT_THAT(header_source, HasSubstr(output_size_macro.str())); -} - -TEST(InterfaceAPI, ContainsOutputStructMany) { - std::stringstream output_struct; - std::stringstream output1_size_macro; - std::stringstream output2_size_macro; - - output1_size_macro << "/*!\n" - << " * \\brief Output tensor output1 size (in bytes) for TVM module " - "\"ultimate_cat_spotter\" \n" - << " */\n" - << "#define TVMGEN_ULTIMATE_CAT_SPOTTER_OUTPUT1_SIZE 345\n"; - - output2_size_macro << "/*!\n" - << " * \\brief Output tensor output2 size (in bytes) for TVM module " - "\"ultimate_cat_spotter\" \n" - << " */\n" - << "#define TVMGEN_ULTIMATE_CAT_SPOTTER_OUTPUT2_SIZE 984\n"; - - output_struct << "struct tvmgen_ultimate_cat_spotter_outputs {\n" - << " void* output1;\n" - << " void* output2;\n" - << "};\n\n"; - - Map input_sizes; - input_sizes.Set("input", IntImm(DataType::Int(32), 0)); - Map output_sizes; - output_sizes.Set("output1", IntImm(DataType::Int(32), 345)); - output_sizes.Set("output2", IntImm(DataType::Int(32), 984)); - - runtime::Module test_module = - InterfaceCCreate("ultimate_cat_spotter", {"input"}, {"output1", "output2"}, {}, {}, {}, 0, - input_sizes, output_sizes); - std::string header_source = test_module->GetSource(); - - ASSERT_THAT(header_source, HasSubstr(output_struct.str())); - ASSERT_THAT(header_source, HasSubstr(output1_size_macro.str())); - ASSERT_THAT(header_source, HasSubstr(output2_size_macro.str())); -} - -TEST(InterfaceAPI, ContainsOutputStructSanitised) { - std::stringstream output_struct; - std::stringstream output1_size_macro; - std::stringstream output2_size_macro; - - output1_size_macro << "/*!\n" - << " * \\brief Output tensor output_1 size (in bytes) for TVM module " - "\"ultimate_cat_spotter\" \n" - << " */\n" - << "#define TVMGEN_ULTIMATE_CAT_SPOTTER_OUTPUT_1_SIZE 345\n"; - - output2_size_macro << "/*!\n" - << " * \\brief Output tensor output_2 size (in bytes) for TVM module " - "\"ultimate_cat_spotter\" \n" - << " */\n" - << "#define TVMGEN_ULTIMATE_CAT_SPOTTER_OUTPUT_2_SIZE 984\n"; - - output_struct << "struct tvmgen_ultimate_cat_spotter_outputs {\n" - << " void* output_1;\n" - << " void* output_2;\n" - << "};\n\n"; - - Map input_sizes; - input_sizes.Set("input", IntImm(DataType::Int(32), 0)); - Map output_sizes; - output_sizes.Set("output+1", IntImm(DataType::Int(32), 345)); - output_sizes.Set("output-2", IntImm(DataType::Int(32), 984)); - - runtime::Module test_module = - InterfaceCCreate("ultimate_cat_spotter", {"input"}, {"output+1", "output-2"}, {}, {}, {}, 0, - input_sizes, output_sizes); - std::string header_source = test_module->GetSource(); - - ASSERT_THAT(header_source, HasSubstr(output_struct.str())); - ASSERT_THAT(header_source, HasSubstr(output1_size_macro.str())); - ASSERT_THAT(header_source, HasSubstr(output2_size_macro.str())); -} - -TEST(InterfaceAPI, ContainsOutputStructClash) { - Map input_sizes; - input_sizes.Set("input", IntImm(DataType::Int(32), 0)); - Map output_sizes; - output_sizes.Set("output+", IntImm(DataType::Int(32), 0)); - output_sizes.Set("output-", IntImm(DataType::Int(32), 0)); - runtime::Module test_module = - InterfaceCCreate("ultimate_cat_spotter", {"input"}, {"output+", "output-"}, {}, {}, {}, 0, - input_sizes, output_sizes); - ASSERT_THROW(test_module->GetSource(), InternalError); -} - -TEST(InterfaceAPI, NoDeviceAPIStructIfNoDevices) { - std::stringstream device_struct; - - device_struct << "/*!\n" - << " * \\brief Device context pointers for TVM module \"ultimate_cat_spotter\" \n" - << " */\n" - << "struct tvmgen_ultimate_cat_spotter_devices {\n" - << "};\n\n"; - - Map input_sizes; - input_sizes.Set("input", IntImm(DataType::Int(32), 0)); - Map output_sizes; - output_sizes.Set("output", IntImm(DataType::Int(32), 0)); - runtime::Module test_module = InterfaceCCreate("ultimate_cat_spotter", {"input"}, {"output"}, {}, - {}, {}, 0, input_sizes, output_sizes); - std::string header_source = test_module->GetSource(); - - ASSERT_THAT(header_source, Not(HasSubstr(device_struct.str()))); -} - -TEST(InterfaceAPI, ContainsDeviceStructSingle) { - std::stringstream device_struct; - - device_struct << "/*!\n" - << " * \\brief Device context pointers for TVM module \"ultimate_cat_spotter\" \n" - << " */\n" - << "struct tvmgen_ultimate_cat_spotter_devices {\n" - << " void* device;\n" - << "};\n\n"; - - Map input_sizes; - input_sizes.Set("input", IntImm(DataType::Int(32), 0)); - Map output_sizes; - output_sizes.Set("output", IntImm(DataType::Int(32), 0)); - runtime::Module test_module = InterfaceCCreate("ultimate_cat_spotter", {"input"}, {"output"}, {}, - {}, {"device"}, 0, input_sizes, output_sizes); - std::string header_source = test_module->GetSource(); - - ASSERT_THAT(header_source, HasSubstr(device_struct.str())); -} - -TEST(InterfaceAPI, ContainsDeviceStructMany) { - std::stringstream device_struct; - - device_struct << "struct tvmgen_ultimate_cat_spotter_devices {\n" - << " void* device1;\n" - << " void* device2;\n" - << "};\n\n"; - - Map input_sizes; - input_sizes.Set("input", IntImm(DataType::Int(32), 0)); - Map output_sizes; - output_sizes.Set("output", IntImm(DataType::Int(32), 0)); - runtime::Module test_module = - InterfaceCCreate("ultimate_cat_spotter", {"input"}, {"output"}, {}, {}, - {"device1", "device2"}, 0, input_sizes, output_sizes); - std::string header_source = test_module->GetSource(); - - ASSERT_THAT(header_source, HasSubstr(device_struct.str())); -} - -TEST(InterfaceAPI, ContainsDeviceStructSanitised) { - std::stringstream device_struct; - - device_struct << "struct tvmgen_ultimate_cat_spotter_devices {\n" - << " void* device_1;\n" - << " void* device_2;\n" - << "};\n\n"; - - Map input_sizes; - input_sizes.Set("input", IntImm(DataType::Int(32), 0)); - Map output_sizes; - output_sizes.Set("output", IntImm(DataType::Int(32), 0)); - runtime::Module test_module = - InterfaceCCreate("ultimate_cat_spotter", {"input"}, {"output"}, {}, {}, - {"device+1", "device+2"}, 0, input_sizes, output_sizes); - std::string header_source = test_module->GetSource(); - - ASSERT_THAT(header_source, HasSubstr(device_struct.str())); -} - -TEST(InterfaceAPI, ContainsDeviceStructClash) { - Map input_sizes; - input_sizes.Set("input", IntImm(DataType::Int(32), 0)); - Map output_sizes; - output_sizes.Set("output", IntImm(DataType::Int(32), 0)); - runtime::Module test_module = - InterfaceCCreate("ultimate_cat_spotter", {"input"}, {"output"}, {}, {}, - {"device+", "device-"}, 0, input_sizes, output_sizes); - ASSERT_THROW(test_module->GetSource(), InternalError); -} - -TEST(InterfaceAPI, ContainsWorkspaceSize) { - Map input_sizes; - input_sizes.Set("input", IntImm(DataType::Int(32), 0)); - Map output_sizes; - output_sizes.Set("output", IntImm(DataType::Int(32), 0)); - runtime::Module test_module = InterfaceCCreate("ultimate_cat_spotter", {"input"}, {"output"}, {}, - {}, {}, 765432, input_sizes, output_sizes); - std::string header_source = test_module->GetSource(); - - ASSERT_THAT(header_source, - HasSubstr("* \\brief Workspace size for TVM module \"ultimate_cat_spotter\"")); - - ASSERT_THAT(header_source, - HasSubstr("#define TVMGEN_ULTIMATE_CAT_SPOTTER_WORKSPACE_SIZE 765432")); -} - -TEST(InterfaceAPI, ContainsWorkspacePoolStructSingle) { - PoolInfo pool_info = WorkspacePoolInfo("my_memory_pool", {}); - tir::usmp::AllocatedPoolInfo allocated_pool_info = - tir::usmp::AllocatedPoolInfo(pool_info, 100000); - - std::stringstream workspace_struct; - - workspace_struct - << "/*!\n" - << " * \\brief Workspace pool pointers for TVM module \"ultimate_cat_spotter\" \n" - << " */\n" - << "struct tvmgen_ultimate_cat_spotter_workspace_pools {\n" - << " void* my_memory_pool;\n" - << "};\n\n"; - - Map input_sizes; - input_sizes.Set("input", IntImm(DataType::Int(32), 0)); - Map output_sizes; - output_sizes.Set("output", IntImm(DataType::Int(32), 0)); - runtime::Module test_module = - InterfaceCCreate("ultimate_cat_spotter", {"input"}, {"output"}, {allocated_pool_info}, {}, {}, - 0, input_sizes, output_sizes); - std::string header_source = test_module->GetSource(); - - ASSERT_THAT(header_source, HasSubstr(workspace_struct.str())); - - ASSERT_THAT(header_source, - HasSubstr("* \\brief my_memory_pool size for TVM module \"ultimate_cat_spotter\"")); - - ASSERT_THAT( - header_source, - HasSubstr("#define TVMGEN_ULTIMATE_CAT_SPOTTER_MY_MEMORY_POOL_WORKSPACE_POOL_SIZE 100000")); -} - -TEST(InterfaceAPI, ContainsWorkspacePoolStructMany) { - PoolInfo pool_info1 = WorkspacePoolInfo("my_memory_pool_1", {}); - tir::usmp::AllocatedPoolInfo allocated_pool_info1 = - tir::usmp::AllocatedPoolInfo(pool_info1, 100000); - PoolInfo pool_info2 = WorkspacePoolInfo("my_memory_pool_2", {}); - tir::usmp::AllocatedPoolInfo allocated_pool_info2 = - tir::usmp::AllocatedPoolInfo(pool_info2, 200000); - - std::stringstream workspace_struct; - - workspace_struct - << "/*!\n" - << " * \\brief Workspace pool pointers for TVM module \"ultimate_cat_spotter\" \n" - << " */\n" - << "struct tvmgen_ultimate_cat_spotter_workspace_pools {\n" - << " void* my_memory_pool_1;\n" - << " void* my_memory_pool_2;\n" - << "};\n\n"; - - Map input_sizes; - input_sizes.Set("input", IntImm(DataType::Int(32), 0)); - Map output_sizes; - output_sizes.Set("output", IntImm(DataType::Int(32), 0)); - runtime::Module test_module = InterfaceCCreate("ultimate_cat_spotter", {"input"}, {"output"}, - {allocated_pool_info1, allocated_pool_info2}, {}, - {}, 0, input_sizes, output_sizes); - std::string header_source = test_module->GetSource(); - - ASSERT_THAT(header_source, HasSubstr(workspace_struct.str())); - - ASSERT_THAT(header_source, - HasSubstr("* \\brief my_memory_pool_1 size for TVM module \"ultimate_cat_spotter\"")); - - ASSERT_THAT( - header_source, - HasSubstr("#define TVMGEN_ULTIMATE_CAT_SPOTTER_MY_MEMORY_POOL_1_WORKSPACE_POOL_SIZE 100000")); - - ASSERT_THAT(header_source, - HasSubstr("* \\brief my_memory_pool_2 size for TVM module \"ultimate_cat_spotter\"")); - - ASSERT_THAT( - header_source, - HasSubstr("#define TVMGEN_ULTIMATE_CAT_SPOTTER_MY_MEMORY_POOL_2_WORKSPACE_POOL_SIZE 200000")); -} - -TEST(InterfaceAPI, ContainsWorkspacePoolStructSanitized) { - PoolInfo pool_info = WorkspacePoolInfo("my_memory_pool+1", {}); - tir::usmp::AllocatedPoolInfo allocated_pool_info = - tir::usmp::AllocatedPoolInfo(pool_info, 100000); - - std::stringstream workspace_struct; - - workspace_struct - << "/*!\n" - << " * \\brief Workspace pool pointers for TVM module \"ultimate_cat_spotter\" \n" - << " */\n" - << "struct tvmgen_ultimate_cat_spotter_workspace_pools {\n" - << " void* my_memory_pool_1;\n" - << "};\n\n"; - - Map input_sizes; - input_sizes.Set("input", IntImm(DataType::Int(32), 0)); - Map output_sizes; - output_sizes.Set("output", IntImm(DataType::Int(32), 0)); - runtime::Module test_module = - InterfaceCCreate("ultimate_cat_spotter", {"input"}, {"output"}, {allocated_pool_info}, {}, {}, - 0, input_sizes, output_sizes); - std::string header_source = test_module->GetSource(); - - ASSERT_THAT(header_source, HasSubstr(workspace_struct.str())); - - ASSERT_THAT(header_source, - HasSubstr("* \\brief my_memory_pool_1 size for TVM module \"ultimate_cat_spotter\"")); - - ASSERT_THAT( - header_source, - HasSubstr("#define TVMGEN_ULTIMATE_CAT_SPOTTER_MY_MEMORY_POOL_1_WORKSPACE_POOL_SIZE 100000")); -} - -TEST(InterfaceAPI, ContainsWorkspacePoolStructClash) { - PoolInfo pool_info1 = WorkspacePoolInfo("my_memory_pool+", {}); - tir::usmp::AllocatedPoolInfo allocated_pool_info1 = - tir::usmp::AllocatedPoolInfo(pool_info1, 100000); - PoolInfo pool_info2 = WorkspacePoolInfo("my_memory_pool-", {}); - tir::usmp::AllocatedPoolInfo allocated_pool_info2 = - tir::usmp::AllocatedPoolInfo(pool_info2, 200000); - - Map input_sizes; - input_sizes.Set("input", IntImm(DataType::Int(32), 0)); - Map output_sizes; - output_sizes.Set("output", IntImm(DataType::Int(32), 0)); - runtime::Module test_module = InterfaceCCreate("ultimate_cat_spotter", {"input"}, {"output"}, - {allocated_pool_info1, allocated_pool_info2}, {}, - {}, 0, input_sizes, output_sizes); - ASSERT_THROW(test_module->GetSource(), InternalError); -} - -} // namespace -} // namespace codegen -} // namespace tvm diff --git a/tests/cpp/target_test.cc b/tests/cpp/target_test.cc index 0a2b8206d322..2d4a5a487afd 100644 --- a/tests/cpp/target_test.cc +++ b/tests/cpp/target_test.cc @@ -17,10 +17,9 @@ * under the License. */ -#include #include #include -#include +#include #include #include @@ -470,40 +469,6 @@ TEST(TargetCreation, DetectSystemTriple) { #endif -TVM_REGISTER_TARGET_KIND("test_external_codegen_0", kDLCUDA) - .set_attr(tvm::attr::kIsExternalCodegen, runtime::Bool(true)); - -TVM_REGISTER_TARGET_KIND("test_external_codegen_1", kDLCUDA) - .set_attr(tvm::attr::kIsExternalCodegen, runtime::Bool(true)); - -TVM_REGISTER_TARGET_KIND("test_external_codegen_2", kDLMetal) - .set_attr(tvm::attr::kIsExternalCodegen, runtime::Bool(true)); - -TVM_REGISTER_TARGET_KIND("test_external_codegen_3", kDLCPU) - .set_attr(tvm::attr::kRelayToTIR, - tvm::relay::transform::InferType()); - -TEST(Target, ExternalCodegen) { - Target regular("cuda"); - Target external0("test_external_codegen_0"); - Target external1("test_external_codegen_1"); - Target external2("test_external_codegen_2"); - Target external3("test_external_codegen_3"); - - ASSERT_FALSE(regular.IsExternalCodegen()); - ASSERT_TRUE(external0.IsExternalCodegen()); - ASSERT_TRUE(external1.IsExternalCodegen()); - ASSERT_TRUE(external2.IsExternalCodegen()); - ASSERT_TRUE(external3.IsExternalCodegen()); - - ASSERT_TRUE(external0.IsExternalCodegenFor(regular)); - ASSERT_FALSE(regular.IsExternalCodegenFor(external0)); - ASSERT_TRUE(external1.IsExternalCodegenFor(regular)); - ASSERT_FALSE(regular.IsExternalCodegenFor(external1)); - ASSERT_FALSE(external2.IsExternalCodegenFor(regular)); - ASSERT_FALSE(regular.IsExternalCodegenFor(external2)); -} - TEST(TargetCreation, DeduplicateKeys) { Map config = { {"kind", String("llvm")}, diff --git a/tests/cpp/tensor_test.cc b/tests/cpp/te_compute_test.cc similarity index 89% rename from tests/cpp/tensor_test.cc rename to tests/cpp/te_compute_test.cc index e53f6d05a991..7d2360f22603 100644 --- a/tests/cpp/tensor_test.cc +++ b/tests/cpp/te_compute_test.cc @@ -17,8 +17,8 @@ * under the License. */ -#include #include +#include #include TEST(Tensor, Basic) { @@ -47,7 +47,6 @@ TEST(Tensor, Reduce) { auto C = te::compute( {m, n}, [&](Var i, Var j) { return sum(max(1 + A[i][rv] + 1, B[j][rv]), {rv}); }, "C"); - LOG(INFO) << C->op.as()->body; } TEST(Tensor, Indexing) { @@ -56,7 +55,4 @@ TEST(Tensor, Indexing) { Var x("x"), y("y"); te::Tensor A = te::placeholder({x, y}, DataType::Float(32), "A"); - LOG(INFO) << A(0, 0); - LOG(INFO) << A.IndexWithNegativeIndices(-1, -1); - LOG(INFO) << A.IndexWithNegativeIndices(0, -1); } diff --git a/tests/cpp/texture_copy_test.cc b/tests/cpp/texture_copy_test.cc deleted file mode 100644 index 63e2ac1a0af4..000000000000 --- a/tests/cpp/texture_copy_test.cc +++ /dev/null @@ -1,125 +0,0 @@ -/* - * 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 - -TEST(TextureCopy, HostDeviceRT) { - using namespace tvm; - bool enabled = tvm::runtime::RuntimeEnabled("opencl"); - if (!enabled) { - LOG(INFO) << "Skip texture copy test because opencl runtime is disabled.\n"; - return; - } - - std::vector shape{16, 16, 4}; - auto cpu_arr0 = runtime::NDArray::Empty(shape, {kDLFloat, 32, 1}, {kDLCPU, 0}); - auto cpu_arr1 = runtime::NDArray::Empty(shape, {kDLFloat, 32, 1}, {kDLCPU, 0}); - String mem_scope = "global.texture"; - auto opencl_txarr0 = runtime::NDArray::Empty(shape, {kDLFloat, 32, 1}, {kDLOpenCL, 0}, mem_scope); - - size_t size = 1; - for (size_t i = 0; i < shape.size(); ++i) { - size *= static_cast(shape[i]); - } - - std::random_device dev; - std::mt19937 mt(dev()); - std::uniform_real_distribution<> random(-10.0, 10.0); - - // Random initialize host ndarray - for (size_t i = 0; i < size; i++) { - static_cast(cpu_arr0->data)[i] = random(mt); - } - - // Do a roundtrip from host storage to opencl texture storage and back - cpu_arr0.CopyTo(opencl_txarr0); - opencl_txarr0.CopyTo(cpu_arr1); - for (size_t i = 0; i < size; ++i) { - ICHECK_LT( - std::fabs(static_cast(cpu_arr1->data)[i] - static_cast(cpu_arr0->data)[i]), - 1e-5); - } -} - -TEST(TextureCopy, OverwritePoolSubview) { - using namespace tvm; - bool enabled = tvm::runtime::RuntimeEnabled("opencl"); - if (!enabled) { - LOG(INFO) << "Skip texture copy test because opencl runtime is disabled.\n"; - return; - } - - std::vector shape{16, 16, 4}; - std::vector shape_pool{32, 32, 4}; - auto cpu_arr0 = runtime::NDArray::Empty(shape, {kDLFloat, 32, 1}, {kDLCPU, 0}); - auto cpu_arr1 = runtime::NDArray::Empty(shape, {kDLFloat, 32, 1}, {kDLCPU, 0}); - auto cpu_pool0 = runtime::NDArray::Empty(shape_pool, {kDLFloat, 32, 1}, {kDLCPU, 0}); - auto cpu_pool1 = runtime::NDArray::Empty(shape_pool, {kDLFloat, 32, 1}, {kDLCPU, 0}); - - String mem_scope = "global.texture"; - auto opencl_txpool = - runtime::NDArray::Empty(shape_pool, {kDLFloat, 32, 1}, {kDLOpenCL, 0}, mem_scope); - auto opencl_txarr0 = opencl_txpool.CreateView(shape, {kDLFloat, 32, 1}); - - std::random_device dev; - std::mt19937 mt(dev()); - std::uniform_real_distribution<> random(-10.0, 10.0); - - size_t size = 1; - size_t size_pool = 1; - for (size_t i = 0; i < shape_pool.size(); ++i) { - size *= static_cast(shape[i]); - size_pool *= static_cast(shape_pool[i]); - } - - // Random initialize host pool storage - for (size_t i = 0; i < size_pool; i++) { - static_cast(cpu_pool0->data)[i] = random(mt); - } - - // Random initialize host array storage - for (size_t i = 0; i < size; i++) { - static_cast(cpu_arr0->data)[i] = random(mt); - } - - // Loop through pool - cpu_pool0.CopyTo(opencl_txpool); - opencl_txpool.CopyTo(cpu_pool1); - - for (size_t i = 0; i < size_pool; i++) { - ICHECK_LT(std::fabs(static_cast(cpu_pool0->data)[i] - - static_cast(cpu_pool1->data)[i]), - 1e-5); - } - - // Loop through view - cpu_arr0.CopyTo(opencl_txarr0); - opencl_txarr0.CopyTo(cpu_arr1); - - for (size_t i = 0; i < size; i++) { - ICHECK_LT( - std::fabs(static_cast(cpu_arr0->data)[i] - static_cast(cpu_arr1->data)[i]), - 1e-5); - } -} diff --git a/tests/cpp/threading_backend_test.cc b/tests/cpp/threading_backend_test.cc index b156eec8ab3a..60149b0ac93a 100644 --- a/tests/cpp/threading_backend_test.cc +++ b/tests/cpp/threading_backend_test.cc @@ -17,9 +17,9 @@ * under the License. */ -#include #include #include +#include #include #include @@ -63,7 +63,6 @@ class AffinityCheck { str << i << ","; } } - LOG(INFO) << "id:" << id_ << " taskid:" << task_id << " affinity:" << str.str() << std::endl; #endif } @@ -169,7 +168,6 @@ TEST(ThreadingBackend, TVMBackendAffinityConfigure) { std::atomic acc(0); AffinityCheck ac(thread_pool_index, sys_max_concurrency, &acc); std::vector cpus; - LOG(INFO) << affinity_mode << std::endl; for (int k = 0; k < cpus_num_per_thread; k++) { cpus.push_back(thread_pool_index * cpus_num_per_thread + k); } diff --git a/tests/cpp/tir_analysis_side_effect.cc b/tests/cpp/tir_analysis_side_effect.cc index bd7d7805e7aa..12c011fd6abb 100644 --- a/tests/cpp/tir_analysis_side_effect.cc +++ b/tests/cpp/tir_analysis_side_effect.cc @@ -17,8 +17,8 @@ * under the License. */ -#include #include +#include #include #include #include diff --git a/tests/python/codegen/test_target_codegen_cuda_fp8.py b/tests/python/codegen/test_target_codegen_cuda_fp8.py index d04262a3701a..d94153003c6a 100644 --- a/tests/python/codegen/test_target_codegen_cuda_fp8.py +++ b/tests/python/codegen/test_target_codegen_cuda_fp8.py @@ -218,7 +218,7 @@ def test_half_broadcast(bcast_length): dtype = "float16" @T.prim_func - def vector_broadcast(a: T.Buffer[(), dtype], vec: T.Buffer[(bcast_length,), dtype]): + def vector_broadcast(a: T.Buffer((), dtype), vec: T.Buffer((bcast_length,), dtype)): for t in range(1): with T.block("broadcast"): vec[0:bcast_length] = T.broadcast(a[()], bcast_length) @@ -256,7 +256,7 @@ def test_half_misaligned_vector_load(vector_length): @T.prim_func def vector_load( - A: T.Buffer[(length,), dtype], B: T.Buffer[(length // vector_length,), vec_dtype] + A: T.Buffer((length,), dtype), B: T.Buffer((length // vector_length,), vec_dtype) ): for b in T.thread_binding(1, thread="blockIdx.x"): for i in T.thread_binding(length // vector_length, thread="threadIdx.x"): diff --git a/tests/python/contrib/test_hexagon/test_launcher.py b/tests/python/contrib/test_hexagon/test_launcher.py deleted file mode 100644 index c84e7a9d4a4c..000000000000 --- a/tests/python/contrib/test_hexagon/test_launcher.py +++ /dev/null @@ -1,722 +0,0 @@ -# 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,missing-function-docstring,redefined-outer-name -""" Test rpc based launcher for hexagon """ -import pytest - -import numpy as np - -import tvm.testing -from tvm import relay, te -from tvm.contrib.hexagon.session import Session -from tvm.relay.backend import Executor, Runtime -from tvm.contrib.hexagon.build import HexagonLauncherRPC -from tvm.contrib.hexagon.hexagon_profiler import HexagonProfiler - -from .infrastructure import get_hexagon_target - - -@tvm.testing.requires_hexagon -def test_add(hexagon_session: Session): - """Test simple add""" - dtype = "int8" - placeholder_a = tvm.te.placeholder((2,), dtype=dtype) - placeholder_b = tvm.te.placeholder((1,), dtype=dtype) - compute_c = tvm.te.compute( - placeholder_a.shape, lambda i: placeholder_a[i] + placeholder_b[0], name="C" - ) - - func = tvm.build( - te.create_prim_func([placeholder_a, placeholder_b, compute_c]), - get_hexagon_target("v68"), - name="add", - ) - - mod = hexagon_session.load_module(func) - - a_data = tvm.nd.array(np.array([2, 3], dtype=dtype), device=hexagon_session.device) - assert (a_data.numpy() == np.array([2, 3])).all() - b_data = tvm.nd.array(np.array([4], dtype=dtype), device=hexagon_session.device) - assert (b_data.numpy() == np.array([4])).all() - c_data = tvm.nd.array(np.array([0, 0], dtype=dtype), device=hexagon_session.device) - assert (c_data.numpy() == np.array([0, 0])).all() - mod["add"](a_data, b_data, c_data) - assert (c_data.numpy() == np.array([6, 7])).all() - - -@tvm.testing.requires_hexagon -def test_add_vtcm(hexagon_session: Session): - """Test add on VTCM""" - dtype = "int8" - placeholder_a = tvm.te.placeholder((2,), dtype=dtype) - placeholder_b = tvm.te.placeholder((1,), dtype=dtype) - compute_c = tvm.te.compute( - placeholder_a.shape, lambda i: placeholder_a[i] + placeholder_b[0], name="C" - ) - - func = tvm.build( - te.create_prim_func([placeholder_a, placeholder_b, compute_c]), - get_hexagon_target("v68"), - name="add", - ) - - mod = hexagon_session.load_module(func) - - a_data = tvm.nd.empty( - placeholder_a.shape, placeholder_a.dtype, hexagon_session.device, "global.vtcm" - ) - a_data.copyfrom(np.array([2, 3])) - - b_data = tvm.nd.empty( - placeholder_b.shape, placeholder_b.dtype, hexagon_session.device, "global.vtcm" - ) - b_data.copyfrom(np.array([4])) - - c_data = tvm.nd.empty(compute_c.shape, compute_c.dtype, hexagon_session.device, "global.vtcm") - c_data.copyfrom(np.array([0, 0])) - - mod["add"](a_data, b_data, c_data) - result = c_data.numpy() - assert (result == np.array([6, 7])).all() - - -class TestMatMul: - """Test matmul class""" - - size_m = tvm.testing.parameter(32) - size_n = tvm.testing.parameter(32) - size_k = tvm.testing.parameter(32) - - @tvm.testing.requires_hexagon - def test_matmul(self, hexagon_session, size_m, size_n, size_k): - """Test matmul""" - placeholder_x = te.placeholder((size_m, size_k), dtype="float32") - placeholder_y = te.placeholder((size_k, size_n), dtype="float32") - reduce_k1 = te.reduce_axis((0, size_k), name="k1") - compute_z = te.compute( - (size_m, size_n), - lambda i, j: te.sum( - placeholder_x[i, reduce_k1] * placeholder_y[reduce_k1, j], axis=[reduce_k1] - ), - ) - - func = tvm.build( - te.create_prim_func([placeholder_x, placeholder_y, compute_z]), - get_hexagon_target("v68"), - ) - - mod = hexagon_session.load_module(func) - - x_data = np.random.uniform(size=[i.value for i in placeholder_x.shape]).astype( - placeholder_x.dtype - ) - y_data = np.random.uniform(size=[i.value for i in placeholder_y.shape]).astype( - placeholder_y.dtype - ) - z_data = np.zeros([i.value for i in compute_z.shape], dtype=compute_z.dtype) - - x_array = tvm.nd.array(x_data, device=hexagon_session.device) - y_array = tvm.nd.array(y_data, device=hexagon_session.device) - z_array = tvm.nd.array(z_data, device=hexagon_session.device) - mod(x_array, y_array, z_array) - - target_llvm = tvm.target.Target("llvm") - mod = tvm.build( - schedule, - [placeholder_x, placeholder_y, compute_z], - tvm.target.Target(target_llvm, host=target_llvm), - ) - device = tvm.cpu(0) - xtcpu = tvm.nd.array(x_data, device) - ytcpu = tvm.nd.array(y_data, device) - ztcpu = tvm.nd.array(z_data, device) - mod(xtcpu, ytcpu, ztcpu) - - tvm.testing.assert_allclose(z_array.numpy(), ztcpu.numpy(), rtol=1e-4) - - -@tvm.testing.requires_hexagon -def test_graph_executor(hexagon_session: Session): - """Test graph executor""" - dtype = "float32" - data = relay.var("data", relay.TensorType((1, 64, 64, 3), dtype)) - weight = relay.var("weight", relay.TensorType((5, 5, 3, 8), dtype)) - conv2d_op = relay.nn.conv2d( - data, - weight, - padding=(2, 2), - kernel_size=(5, 5), - data_layout="NHWC", - kernel_layout="HWIO", - out_dtype="float32", - ) - f = relay.Function([data, weight], conv2d_op) - relay_mod = tvm.IRModule.from_expr(f) - relay_mod = relay.transform.InferType()(relay_mod) - - runtime = Runtime("cpp") - executor = Executor("graph") - - weight_in = np.random.rand(5, 5, 3, 8).astype(dtype=dtype) - data_in = np.random.rand(1, 64, 64, 3).astype(dtype=dtype) - params = {"weight": weight_in} - inputs = {"data": data_in} - - with tvm.transform.PassContext(opt_level=3): - lowered = tvm.relay.build( - relay_mod, - get_hexagon_target("v68"), - runtime=runtime, - executor=executor, - ) - - graph_mod = hexagon_session.get_executor_from_factory(lowered) - graph_mod.set_input(**params) - graph_mod.run(**inputs) - hexagon_output = graph_mod.get_output(0).numpy() - - target_llvm = tvm.target.Target("llvm") - with tvm.transform.PassContext(opt_level=3): - llvm_lowered = tvm.relay.build( - relay_mod, - tvm.target.Target(target_llvm, host=target_llvm), - runtime=runtime, - executor=executor, - ) - llvm_graph_mod = tvm.contrib.graph_executor.GraphModule(llvm_lowered["default"](tvm.cpu(0))) - llvm_graph_mod.set_input(**params) - llvm_graph_mod.run(**inputs) - expected_output = llvm_graph_mod.get_output(0).numpy() - - tvm.testing.assert_allclose(hexagon_output, expected_output, rtol=1e-4, atol=1e-5) - - -@tvm.testing.requires_hexagon -def test_graph_executor_multiple_conv2d(hexagon_session: Session): - """Test multiple conv2d nodes with graph_executor""" - dtype = "float32" - input_shape = (1, 8, 8, 3) - w1_shape = (5, 5, 3, 1) - w2_shape = (5, 5, 1, 3) - data = relay.var("data", relay.TensorType(input_shape, dtype)) - weight1 = relay.var("weight1", relay.TensorType(w1_shape, dtype)) - weight2 = relay.var("weight2", relay.TensorType(w2_shape, dtype)) - conv2d_op1 = relay.nn.conv2d( - data, - weight1, - padding=(2, 2), - kernel_size=(5, 5), - data_layout="NHWC", - kernel_layout="HWIO", - out_dtype="float32", - ) - conv2d_op2 = relay.nn.conv2d( - conv2d_op1, - weight2, - padding=(2, 2), - kernel_size=(5, 5), - data_layout="NHWC", - kernel_layout="HWIO", - out_dtype="float32", - ) - f = relay.Function([data, weight1, weight2], conv2d_op2) - relay_mod = tvm.IRModule.from_expr(f) - relay_mod = relay.transform.InferType()(relay_mod) - - runtime = Runtime("cpp") - executor = Executor("graph") - - with tvm.transform.PassContext(opt_level=3): - lowered = tvm.relay.build( - relay_mod, - get_hexagon_target("v68"), - runtime=runtime, - executor=executor, - ) - - weight1_data = np.random.rand(w1_shape[0], w1_shape[1], w1_shape[2], w1_shape[3]).astype( - dtype=dtype - ) - weight2_data = np.random.rand(w2_shape[0], w2_shape[1], w2_shape[2], w2_shape[3]).astype( - dtype=dtype - ) - input_data = np.random.rand( - input_shape[0], input_shape[1], input_shape[2], input_shape[3] - ).astype(dtype=dtype) - - params = {"weight1": weight1_data, "weight2": weight2_data} - inputs = {"data": input_data} - - graph_mod = hexagon_session.get_executor_from_factory(lowered) - graph_mod.set_input(**params) - graph_mod.run(**inputs) - hexagon_output = graph_mod.get_output(0).numpy() - - target_llvm = tvm.target.Target("llvm") - with tvm.transform.PassContext(opt_level=3): - llvm_lowered = tvm.relay.build( - relay_mod, - tvm.target.Target(target_llvm, host=target_llvm), - runtime=runtime, - executor=executor, - ) - llvm_graph_mod = tvm.contrib.graph_executor.GraphModule(llvm_lowered["default"](tvm.cpu(0))) - llvm_graph_mod.set_input(**params) - llvm_graph_mod.run(**inputs) - expected_output = llvm_graph_mod.get_output(0).numpy() - - tvm.testing.assert_allclose(hexagon_output, expected_output, rtol=1e-4, atol=1e-5) - - -@tvm.testing.requires_hexagon -def test_aot_executor(hexagon_session: Session, aot_host_target, aot_target): - """Test AOT executor""" - dtype = "float32" - input_shape = (1, 128, 128, 3) - w_shape = (5, 5, 3, 8) - data = relay.var("data", relay.TensorType(input_shape, dtype)) - weight = relay.var("weight", relay.TensorType(w_shape, dtype)) - y = relay.nn.conv2d( - data, - weight, - padding=(2, 2), - kernel_size=(5, 5), - data_layout="NHWC", - kernel_layout="HWIO", - out_dtype="float32", - ) - f = relay.Function([data, weight], y) - relay_mod = tvm.IRModule.from_expr(f) - relay_mod = relay.transform.InferType()(relay_mod) - - weight_data = np.random.rand(w_shape[0], w_shape[1], w_shape[2], w_shape[3]).astype(dtype=dtype) - input_data = np.random.rand( - input_shape[0], input_shape[1], input_shape[2], input_shape[3] - ).astype(dtype=dtype) - - params = {"weight": weight_data} - inputs = {"data": input_data} - - with tvm.transform.PassContext(opt_level=3): - lowered = tvm.relay.build( - relay_mod, - params=params, - target=tvm.target.Target(aot_target, host=aot_host_target), - runtime=Runtime("cpp"), - executor=Executor("aot", {"unpacked-api": False, "interface-api": "packed"}), - ) - - aot_mod = hexagon_session.get_executor_from_factory(lowered) - aot_mod.set_input(**inputs) - aot_mod.run() - hexagon_output = aot_mod.get_output(0).numpy() - - target_llvm = tvm.target.Target("llvm") - with tvm.transform.PassContext(opt_level=3): - llvm_lowered = tvm.relay.build( - relay_mod, - tvm.target.Target(target_llvm, host=target_llvm), - runtime=Runtime("cpp"), - executor=Executor("graph"), - ) - - llvm_graph_mod = tvm.contrib.graph_executor.GraphModule(llvm_lowered["default"](tvm.cpu(0))) - llvm_graph_mod.set_input(**params) - llvm_graph_mod.run(**inputs) - expected_output = llvm_graph_mod.get_output(0).numpy() - - tvm.testing.assert_allclose(hexagon_output, expected_output, rtol=1e-4, atol=1e-5) - - -@tvm.testing.requires_hexagon -def test_aot_executor_multiple_conv2d(hexagon_session: Session, aot_host_target, aot_target): - """Test multiple conv2d nodes with AOT executor""" - dtype = "float32" - input_shape = (1, 8, 8, 3) - w1_shape = (5, 5, 3, 1) - w2_shape = (5, 5, 1, 3) - data = relay.var("data", relay.TensorType(input_shape, dtype)) - weight1 = relay.var("weight1", relay.TensorType(w1_shape, dtype)) - weight2 = relay.var("weight2", relay.TensorType(w2_shape, dtype)) - conv2d_op1 = relay.nn.conv2d( - data, - weight1, - padding=(2, 2), - kernel_size=(5, 5), - data_layout="NHWC", - kernel_layout="HWIO", - out_dtype="float32", - ) - conv2d_op2 = relay.nn.conv2d( - conv2d_op1, - weight2, - padding=(2, 2), - kernel_size=(5, 5), - data_layout="NHWC", - kernel_layout="HWIO", - out_dtype="float32", - ) - f = relay.Function([data, weight1, weight2], conv2d_op2) - relay_mod = tvm.IRModule.from_expr(f) - relay_mod = relay.transform.InferType()(relay_mod) - - weight1_data = np.random.rand(w1_shape[0], w1_shape[1], w1_shape[2], w1_shape[3]).astype( - dtype=dtype - ) - weight2_data = np.random.rand(w2_shape[0], w2_shape[1], w2_shape[2], w2_shape[3]).astype( - dtype=dtype - ) - input_data = np.random.rand( - input_shape[0], input_shape[1], input_shape[2], input_shape[3] - ).astype(dtype=dtype) - - params = {"weight1": weight1_data, "weight2": weight2_data} - inputs = {"data": input_data} - - with tvm.transform.PassContext(opt_level=3): - lowered = tvm.relay.build( - relay_mod, - params=params, - target=tvm.target.Target(aot_target, host=aot_host_target), - runtime=Runtime("cpp"), - executor=Executor("aot", {"unpacked-api": False, "interface-api": "packed"}), - ) - - aot_mod = hexagon_session.get_executor_from_factory(lowered) - aot_mod.set_input(**inputs) - aot_mod.run() - hexagon_output = aot_mod.get_output(0).numpy() - - target_llvm = tvm.target.Target("llvm") - with tvm.transform.PassContext(opt_level=3): - llvm_lowered = tvm.relay.build( - relay_mod, - tvm.target.Target(target_llvm, host=target_llvm), - runtime=Runtime("cpp"), - executor=Executor("graph"), - ) - - llvm_graph_mod = tvm.contrib.graph_executor.GraphModule(llvm_lowered["default"](tvm.cpu(0))) - llvm_graph_mod.set_input(**params) - llvm_graph_mod.run(**inputs) - expected_output = llvm_graph_mod.get_output(0).numpy() - - tvm.testing.assert_allclose(hexagon_output, expected_output, rtol=1e-4, atol=1e-5) - - -data_dtype = tvm.testing.parameter("int8", "uint8") -weight_dtype = tvm.testing.parameter("int8", "uint8") - - -@tvm.testing.requires_hexagon -def test_conv2d_relay_vrmpy(hexagon_session, data_dtype, weight_dtype): - if data_dtype == "int8" and weight_dtype == "uint8": - pytest.skip("(i8, u8) input pair is not supported") - - def get_conv2d_nchw(d_shape, w_shape, padding, strides=(1, 1)): - out_dtype = "int32" - - data = relay.var("data", shape=d_shape, dtype=data_dtype) - weight = relay.var("weight", shape=w_shape, dtype=weight_dtype) - out_channel = w_shape[0] - return relay.nn.conv2d( - data=data, - weight=weight, - kernel_size=w_shape[2:], - channels=out_channel, - padding=padding, - strides=strides, - out_dtype=out_dtype, - ) - - target = get_hexagon_target("v68") - I, O, H, W = 64, 256, 56, 56 - kH = kW = 3 - padding = (1, 1) - strides = (1, 1) - - data_shape = (1, I, H, W) - weight_shape = (O, I, kH, kW) - bias_shape = (weight_shape[0],) - - bias = relay.var("bias", shape=bias_shape, dtype="int32") - - conv2d = get_conv2d_nchw( - data_shape, - weight_shape, - padding, - strides=strides, - ) - bias_add = relay.nn.bias_add(conv2d, bias) - mod = tvm.IRModule.from_expr(bias_add) - - if data_dtype == "uint8": - data_np = np.random.uniform(0, 255, size=data_shape).astype("uint8") - else: - data_np = np.random.uniform(-128, 127, size=data_shape).astype("int8") - - if weight_dtype == "uint8": - weight_np = np.random.uniform(0, 255, size=weight_shape).astype("uint8") - else: - weight_np = np.random.uniform(-128, 127, size=weight_shape).astype("int8") - - bias_np = np.random.randint(low=-127, high=128, size=bias_shape).astype("int32") - params = {"weight": weight_np, "bias": bias_np} - - ref = ( - relay.create_executor("graph", mod=mod, device=tvm.cpu(0), target="llvm") - .evaluate()(*[data_np, weight_np, bias_np]) - .numpy() - ) - - with tvm.transform.PassContext( - opt_level=3, - ): - executor = relay.backend.Executor("graph", {"link-params": True}) - lib = relay.build(mod, target=target, params=params, executor=executor) - - asm = lib.lib.get_source("asm") - assert "vrmpy" in asm - - rt_mod = hexagon_session.get_executor_from_factory(lib) - - rt_mod.set_input("data", data_np) - - rt_mod.run() - - out = rt_mod.get_output(0).numpy() - - np.testing.assert_equal(out, ref) - - -@tvm.testing.requires_hexagon -def test_dense_relay_vrmpy(hexagon_session, data_dtype, weight_dtype): - if data_dtype == "int8" and weight_dtype == "uint8": - pytest.skip("(i8, u8) input pair is not supported") - - target = get_hexagon_target("v68") - - M = 128 - N = 1000 - K = 2048 - data_shape = (M, K) - weight_shape = (N, K) - - data = relay.var("data", shape=data_shape, dtype=data_dtype) - weight = relay.var("weight", shape=weight_shape, dtype=weight_dtype) - - dense = relay.nn.dense(data, weight, out_dtype="int32") - - if data_dtype == "uint8": - data_np = np.random.uniform(0, 255, size=data_shape).astype("uint8") - else: - data_np = np.random.uniform(-128, 127, size=data_shape).astype("int8") - - if weight_dtype == "uint8": - weight_np = np.random.uniform(0, 255, size=weight_shape).astype("uint8") - else: - weight_np = np.random.uniform(-128, 127, size=weight_shape).astype("int8") - - bias_np = np.random.uniform(1, 10, size=(weight_shape[0],)).astype("int32") - - params = {"weight": weight_np, "bias": bias_np} - - bias = relay.var("bias", shape=(weight_shape[0],), dtype="int32") - bias_add = relay.nn.bias_add(dense, bias) - mod = tvm.IRModule.from_expr(bias_add) - - with tvm.transform.PassContext( - opt_level=3, - ): - executor = relay.backend.Executor("graph", {"link-params": True}) - lib = relay.build(mod, target=target, params=params, executor=executor) - - asm = lib.lib.get_source("asm") - assert "vrmpy" in asm - - rt_mod = hexagon_session.get_executor_from_factory(lib) - - rt_mod.set_input("data", data_np) - - rt_mod.run() - - out = rt_mod.get_output(0).numpy() - - ref = np.dot(data_np.astype("int32"), weight_np.transpose().astype("int32")) - ref += bias_np - - np.testing.assert_equal(out, ref) - - -@tvm.testing.requires_hexagon -def test_lwp( - hexagon_server_process, - hexagon_launcher: HexagonLauncherRPC, - hexagon_session: Session, - hexagon_debug, -): - dtype = "float32" - data = relay.var("data", relay.TensorType((1, 64, 64, 3), dtype)) - weight = relay.var("weight", relay.TensorType((5, 5, 3, 8), dtype)) - y = relay.nn.conv2d( - data, - weight, - padding=(2, 2), - kernel_size=(5, 5), - data_layout="NHWC", - kernel_layout="HWIO", - out_dtype="float32", - ) - - f = relay.Function([data, weight], y) - relay_mod = tvm.IRModule.from_expr(f) - relay_mod = relay.transform.InferType()(relay_mod) - - target_hexagon = tvm.target.hexagon("v68") - runtime = Runtime("cpp") - executor = Executor("graph") - - weight_in = np.random.rand(5, 5, 3, 8).astype(dtype=dtype) - data_in = np.random.rand(1, 64, 64, 3).astype(dtype=dtype) - params = {"weight": weight_in} - inputs = {"data": data_in} - - with tvm.transform.PassContext(opt_level=3, config={"tir.instrument_lwp": True}): - lowered = tvm.relay.build( - relay_mod, - tvm.target.Target(target_hexagon, host=target_hexagon), - runtime=runtime, - executor=executor, - ) - # Create HexagonProfiler object - dso_binary = "test_binary.so" - profiler = HexagonProfiler(dso_binary, lowered, hexagon_server_process, hexagon_debug) - - graph_mod = hexagon_session.get_executor_from_factory(lowered) - graph_mod.set_input(**params) - graph_mod.run(**inputs) - hexagon_output = graph_mod.get_output(0).numpy() - - # Get lightweight profiling output as a CSV file - profiler.get_profile_output(hexagon_launcher, hexagon_session) - - target_llvm = tvm.target.Target("llvm") - with tvm.transform.PassContext(opt_level=3): - llvm_lowered = tvm.relay.build( - relay_mod, - tvm.target.Target(target_llvm, host=target_llvm), - runtime=runtime, - executor=executor, - ) - llvm_graph_mod = tvm.contrib.graph_executor.GraphModule(llvm_lowered["default"](tvm.cpu(0))) - llvm_graph_mod.set_input(weight=weight_in) - llvm_graph_mod.run(data=data_in) - expected_output = llvm_graph_mod.get_output(0).numpy() - - tvm.testing.assert_allclose(hexagon_output, expected_output, rtol=1e-4, atol=1e-5) - - -@tvm.testing.requires_hexagon -def test_lwp_multiple_conv2d( - hexagon_server_process, - hexagon_launcher: HexagonLauncherRPC, - hexagon_session: Session, - hexagon_debug, -): - dtype = "float32" - input_shape = (1, 8, 8, 3) - w1_shape = (5, 5, 3, 1) - w2_shape = (5, 5, 1, 3) - data = relay.var("data", relay.TensorType(input_shape, dtype)) - weight1 = relay.var("weight1", relay.TensorType(w1_shape, dtype)) - weight2 = relay.var("weight2", relay.TensorType(w2_shape, dtype)) - y1 = relay.nn.conv2d( - data, - weight1, - padding=(2, 2), - kernel_size=(5, 5), - data_layout="NHWC", - kernel_layout="HWIO", - out_dtype="float32", - ) - y2 = relay.nn.conv2d( - y1, - weight2, - padding=(2, 2), - kernel_size=(5, 5), - data_layout="NHWC", - kernel_layout="HWIO", - out_dtype="float32", - ) - f = relay.Function([data, weight1, weight2], y2) - relay_mod = tvm.IRModule.from_expr(f) - relay_mod = relay.transform.InferType()(relay_mod) - - target_hexagon = tvm.target.hexagon("v68") - runtime = Runtime("cpp") - executor = Executor("graph") - - weight1_data = np.random.rand(w1_shape[0], w1_shape[1], w1_shape[2], w1_shape[3]).astype( - dtype=dtype - ) - weight2_data = np.random.rand(w2_shape[0], w2_shape[1], w2_shape[2], w2_shape[3]).astype( - dtype=dtype - ) - input_data = np.random.rand( - input_shape[0], input_shape[1], input_shape[2], input_shape[3] - ).astype(dtype=dtype) - - params = {"weight1": weight1_data, "weight2": weight2_data} - inputs = {"data": input_data} - - with tvm.transform.PassContext(opt_level=3, config={"tir.instrument_lwp": True}): - lowered = tvm.relay.build( - relay_mod, - tvm.target.Target(target_hexagon, host=target_hexagon), - runtime=runtime, - executor=executor, - ) - # Create HexagonProfiler object - dso_binary = "test_binary.so" - profiler = HexagonProfiler(dso_binary, lowered, hexagon_server_process, hexagon_debug) - - graph_mod = hexagon_session.get_executor_from_factory(lowered) - graph_mod.set_input(**params) - graph_mod.run(**inputs) - hexagon_output = graph_mod.get_output(0).numpy() - - # Get lightweight profiling output as a CSV file - profiler.get_profile_output(hexagon_launcher, hexagon_session) - - target_llvm = tvm.target.Target("llvm") - with tvm.transform.PassContext(opt_level=3): - llvm_lowered = tvm.relay.build( - relay_mod, - tvm.target.Target(target_llvm, host=target_llvm), - runtime=runtime, - executor=executor, - ) - llvm_graph_mod = tvm.contrib.graph_executor.GraphModule(llvm_lowered["default"](tvm.cpu(0))) - llvm_graph_mod.set_input(**params) - llvm_graph_mod.run(**inputs) - expected_output = llvm_graph_mod.get_output(0).numpy() - - tvm.testing.assert_allclose(hexagon_output, expected_output, rtol=1e-4, atol=1e-5) - - -if __name__ == "__main__": - tvm.testing.main() diff --git a/tests/python/contrib/test_hexagon/test_pass_fq2i_avg_pool2d.py b/tests/python/contrib/test_hexagon/test_pass_fq2i_avg_pool2d.py deleted file mode 100644 index e45f56ba171c..000000000000 --- a/tests/python/contrib/test_hexagon/test_pass_fq2i_avg_pool2d.py +++ /dev/null @@ -1,313 +0,0 @@ -# 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-outer-name - -""" Tests for avg_pool2d fake quantization to integer """ - -import numpy as np -import pytest - -import tvm -import tvm.testing -import tvm.topi.testing -from tvm import relay -from tvm.contrib.hexagon.session import Session -from tvm.contrib.hexagon.pytest_plugin import HEXAGON_AOT_LLVM_TARGET - -from .infrastructure import quantize_np, build_module, run_module - - -def _make_avgpool_conv2d(): - """Test case with avg_pool2d followed by a conv2d""" - dtype = "int8" - shape_x = [1, 2, 9, 9] - shape_w = [1, 2, 3, 3] - kernel = [3, 3] - stride = [1, 1] - dilation = [1, 1] - inp = relay.var("input", shape=shape_x, dtype=dtype) - wgt = relay.var("weight", shape=shape_w, dtype=dtype) - - x_np = np.random.random(shape_x) - w_np = np.random.random(shape_w) - - fp_avg = tvm.topi.testing.poolnd_python( - x_np, - kernel, - stride, - dilation, - padding_before=[0, 0], - padding_after=[0, 0], - pool_type="avg", - ) - fp_output = tvm.topi.testing.conv2d_nchw_python( - fp_avg, - w_np, - [1, 1], - [0, 0], - ) - - # Computing quantization parameters - input_quant, input_scale, input_zero_point = quantize_np(x_np, dtype) - weight_quant, weight_scale, weight_zero_point = quantize_np(w_np, dtype) - _, output_scale, output_zero_point = quantize_np(fp_output, dtype) - - inp_zp = relay.const(input_zero_point) - inp_sc = relay.const(input_scale) - wgt_zp = relay.const(weight_zero_point) - wgt_sc = relay.const(weight_scale) - out_zp = relay.const(output_zero_point) - out_sc = relay.const(output_scale) - - # Tested expression. - op0 = relay.qnn.op.dequantize(inp, inp_sc, inp_zp) - op1 = relay.op.nn.avg_pool2d(op0, kernel) - op2 = relay.qnn.op.dequantize(wgt, wgt_sc, wgt_zp) - op3 = relay.op.nn.conv2d(op1, op2, kernel_size=kernel) - expr = relay.qnn.op.quantize(op3, out_sc, out_zp, out_dtype=dtype) - expr = relay.qnn.op.dequantize(expr, out_sc, out_zp) - args = {"input": input_quant, "weight": weight_quant} - - # Expected graph - op0 = relay.qnn.op.avg_pool2d( - inp, - input_scale=inp_sc, - input_zero_point=inp_zp, - output_scale=inp_sc, - output_zero_point=inp_zp, - pool_size=kernel, - strides=stride, - dilation=dilation, - padding=[0, 0, 0, 0], - layout="NCHW", - count_include_pad=False, - ) - op1 = relay.qnn.op.conv2d( - op0, - wgt, - input_scale=inp_sc, - input_zero_point=inp_zp, - kernel_scale=wgt_sc, - kernel_zero_point=wgt_zp, - kernel_size=kernel, - channels=None, - ) - op2 = relay.qnn.op.requantize( - op1, - input_scale=relay.const(input_scale * weight_scale), - input_zero_point=relay.const(0), - output_scale=out_sc, - output_zero_point=out_zp, - axis=1, - out_dtype="int8", - ) - ref_expr = relay.qnn.op.dequantize(op2, out_sc, out_zp) - - return expr, args, ref_expr - - -def _make_avgpool_avgpool(): - """Test case with avg_pool2d followed by an avg_pool2d""" - dtype = "uint8" - shape_x = [1, 2, 9, 9] - kernel = [3, 3] - stride = [1, 1] - dilation = [1, 1] - inp = relay.var("input", shape=shape_x, dtype=dtype) - x_np = np.random.random(shape_x) - - fp_avg = tvm.topi.testing.poolnd_python( - x_np, - kernel, - stride, - dilation, - padding_before=[0, 0], - padding_after=[0, 0], - pool_type="avg", - ) - fp_output = tvm.topi.testing.poolnd_python( - fp_avg, - kernel, - stride, - dilation, - padding_before=[0, 0], - padding_after=[0, 0], - pool_type="avg", - ) - - # Computing quantization parameters - input_quant, input_scale, input_zero_point = quantize_np(x_np, dtype) - _, output_scale, output_zero_point = quantize_np(fp_output, dtype) - - inp_zp = relay.const(input_zero_point) - inp_sc = relay.const(input_scale) - out_zp = relay.const(output_zero_point) - out_sc = relay.const(output_scale) - - # Tested expression. - op0 = relay.qnn.op.dequantize(inp, inp_sc, inp_zp) - op1 = relay.op.nn.avg_pool2d(op0, kernel) - op2 = relay.op.nn.avg_pool2d(op1, kernel) - expr = relay.qnn.op.quantize(op2, out_sc, out_zp, out_dtype=dtype) - expr = relay.qnn.op.dequantize(expr, out_sc, out_zp) - args = {"input": input_quant} - - # Expected graph - op0 = relay.qnn.op.avg_pool2d( - inp, - input_scale=inp_sc, - input_zero_point=inp_zp, - output_scale=inp_sc, - output_zero_point=inp_zp, - pool_size=kernel, - strides=stride, - dilation=dilation, - padding=[0, 0, 0, 0], - layout="NCHW", - count_include_pad=False, - ) - op1 = relay.qnn.op.avg_pool2d( - op0, - input_scale=inp_sc, - input_zero_point=inp_zp, - output_scale=out_sc, - output_zero_point=out_zp, - pool_size=kernel, - strides=stride, - dilation=dilation, - padding=[0, 0, 0, 0], - layout="NCHW", - count_include_pad=False, - ) - ref_expr = relay.qnn.op.dequantize(op1, out_sc, out_zp) - - return expr, args, ref_expr - - -def _make_avgpool(): - dtype = "int8" - shape_x = [1, 2, 9, 9] - kernel = [3, 3] - stride = [1, 1] - dilation = [1, 1] - inp = relay.var("input", shape=shape_x, dtype=dtype) - x_np = np.random.random(shape_x) - - fp_output = tvm.topi.testing.poolnd_python( - x_np, - kernel, - stride, - dilation, - padding_before=[0, 0], - padding_after=[0, 0], - pool_type="avg", - ) - - # Computing quantization parameters - input_quant, input_scale, input_zero_point = quantize_np(x_np, dtype) - _, output_scale, output_zero_point = quantize_np(fp_output, dtype) - - inp_zp = relay.const(input_zero_point) - inp_sc = relay.const(input_scale) - out_zp = relay.const(output_zero_point) - out_sc = relay.const(output_scale) - - # Tested expression - op0 = relay.qnn.op.dequantize(inp, inp_sc, inp_zp) - op1 = relay.op.nn.avg_pool2d(op0, kernel) - expr = relay.qnn.op.quantize(op1, out_sc, out_zp, out_dtype=dtype) - expr = relay.qnn.op.dequantize(expr, out_sc, out_zp) - args = {"input": input_quant} - - # Expected graph - op = relay.qnn.op.avg_pool2d( - inp, - input_scale=inp_sc, - input_zero_point=inp_zp, - output_scale=out_sc, - output_zero_point=out_zp, - pool_size=kernel, - strides=stride, - dilation=dilation, - padding=[0, 0, 0, 0], - layout="NCHW", - count_include_pad=False, - ) - ref_expr = relay.qnn.op.dequantize(op, out_sc, out_zp) - - return expr, args, ref_expr - - -def compare_graphs(expr, ref_expr): - """Compares the given graph with the expected graph""" - mod = tvm.IRModule.from_expr(expr) - mod = tvm.relay.transform.InferType()(mod) - mod_int = tvm.relay.transform.FakeQuantizationToInteger()(mod) - ref_mod = tvm.IRModule.from_expr(ref_expr) - ref_mod = tvm.relay.transform.InferType()(ref_mod) - tvm.ir.assert_structural_equal(mod_int["main"], ref_mod["main"], map_free_vars=True) - - -def compare_fq_to_int(hexagon_session, expr, inputs): - """Compares the float module output with the integer module output""" - mod = tvm.IRModule.from_expr(expr) - mod = tvm.relay.transform.InferType()(mod) - mod_int = tvm.relay.transform.FakeQuantizationToInteger()(mod) - assert not tvm.ir.structural_equal(mod, mod_int) - - mod = build_module( - mod, tvm.target.Target(HEXAGON_AOT_LLVM_TARGET, host=HEXAGON_AOT_LLVM_TARGET) - ) - mod_int = build_module( - mod_int, tvm.target.Target(HEXAGON_AOT_LLVM_TARGET, host=HEXAGON_AOT_LLVM_TARGET) - ) - - hexagon_mod = hexagon_session.get_executor_from_factory(mod) - result = run_module(hexagon_mod, inputs) - - hexagon_mod = hexagon_session.get_executor_from_factory(mod_int) - result_int = run_module(hexagon_mod, inputs) - - tvm.testing.assert_allclose(result, result_int, rtol=1e-02, atol=1e-02) - - -avgpool_test_case = tvm.testing.parameter( - _make_avgpool, - _make_avgpool_avgpool, - pytest.param( - _make_avgpool_conv2d, - marks=pytest.mark.xfail( - reason="Rounding differences causing mismatch of Constant, difference around 10^-7" - ), - ), -) - - -@tvm.testing.requires_hexagon -def test_execution(hexagon_session: Session, avgpool_test_case): - expr, args, _ = avgpool_test_case() - compare_fq_to_int(hexagon_session, expr, args) - - -def test_quantization(avgpool_test_case): - expr, _, ref_expr = avgpool_test_case() - compare_graphs(expr, ref_expr) - - -if __name__ == "__main__": - tvm.testing.main() diff --git a/tests/python/contrib/test_hexagon/test_usmp.py b/tests/python/contrib/test_hexagon/test_usmp.py deleted file mode 100644 index adfebcd122b3..000000000000 --- a/tests/python/contrib/test_hexagon/test_usmp.py +++ /dev/null @@ -1,109 +0,0 @@ -# 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. - -"""USMP tests""" - -import numpy as np -import pytest -import tvm.testing -from tvm import relay -from tvm.contrib.hexagon.session import Session -from tvm.relay.backend import Executor, Runtime -from tvm.testing.usmp import is_tvm_backendallocworkspace_calls - - -@pytest.mark.parametrize("usmp_enabled", [False, True]) -@tvm.testing.requires_hexagon -def test_conv2d(hexagon_session: Session, aot_host_target, aot_target, usmp_enabled): - """Try conv2d on AOT target with usmp_enabled and check for TVMBackendAllocWorkspace calls""" - dtype = "float32" - input_shape = (1, 8, 8, 3) - w1_shape = (5, 5, 3, 1) - w2_shape = (5, 5, 1, 3) - data = relay.var("data", relay.TensorType(input_shape, dtype)) - weight1 = relay.var("weight1", relay.TensorType(w1_shape, dtype)) - weight2 = relay.var("weight2", relay.TensorType(w2_shape, dtype)) - outpu1 = relay.nn.conv2d( - data, - weight1, - padding=(2, 2), - kernel_size=(5, 5), - data_layout="NHWC", - kernel_layout="HWIO", - out_dtype="float32", - ) - output2 = relay.nn.conv2d( - outpu1, - weight2, - padding=(2, 2), - kernel_size=(5, 5), - data_layout="NHWC", - kernel_layout="HWIO", - out_dtype="float32", - ) - f = relay.Function([data, weight1, weight2], output2) - relay_mod = tvm.IRModule.from_expr(f) - relay_mod = relay.transform.InferType()(relay_mod) - - weight1_data = np.random.rand(w1_shape[0], w1_shape[1], w1_shape[2], w1_shape[3]).astype( - dtype=dtype - ) - weight2_data = np.random.rand(w2_shape[0], w2_shape[1], w2_shape[2], w2_shape[3]).astype( - dtype=dtype - ) - input_data = np.random.rand( - input_shape[0], input_shape[1], input_shape[2], input_shape[3] - ).astype(dtype=dtype) - - params = {"weight1": weight1_data, "weight2": weight2_data} - inputs = {"data": input_data} - - with tvm.transform.PassContext(opt_level=3, config={"tir.usmp.enable": usmp_enabled}): - lowered = tvm.relay.build( - relay_mod, - params=params, - target=tvm.target.Target(aot_target, host=aot_host_target), - runtime=Runtime("cpp"), - executor=Executor("aot", {"unpacked-api": False, "interface-api": "packed"}), - ) - - assert is_tvm_backendallocworkspace_calls(lowered.lib) != usmp_enabled - - aot_mod = hexagon_session.get_executor_from_factory(lowered) - aot_mod.set_input(**inputs) - aot_mod.run() - hexagon_output = aot_mod.get_output(0).numpy() - - target_llvm = tvm.target.Target("llvm") - with tvm.transform.PassContext(opt_level=3): - llvm_lowered = tvm.relay.build( - relay_mod, - tvm.target.Target(target_llvm, host=target_llvm), - runtime=Runtime("cpp"), - executor=Executor("aot"), - ) - - llvm_mod = tvm.runtime.executor.AotModule(llvm_lowered["default"](tvm.cpu(0))) - llvm_mod.set_input(**params) - llvm_mod.run(**inputs) - expected_output = llvm_mod.get_output(0).numpy() - - tvm.testing.assert_allclose(hexagon_output, expected_output, rtol=1e-4, atol=1e-5) - - -if __name__ == "__main__": - tvm.testing.main() diff --git a/tests/python/contrib/test_msc/test_translate_tensorflow.py b/tests/python/contrib/test_msc/test_translate_tensorflow.py index 913056b88098..857cffbbd87a 100644 --- a/tests/python/contrib/test_msc/test_translate_tensorflow.py +++ b/tests/python/contrib/test_msc/test_translate_tensorflow.py @@ -391,7 +391,6 @@ def test_sigmoid(): def _test_argx(func, data, **kwargs): - with tf.Graph().as_default(): inp = array_ops.placeholder(shape=data.shape, dtype=data.dtype, name="c0") func(inp, name="argx", **kwargs) @@ -444,7 +443,6 @@ def test_matmul(): def _test_batch_matmul(a_shape, b_shape, adjoint_a=False, adjoint_b=False): - with tf.Graph().as_default(): a_in = tf.placeholder(shape=a_shape, dtype="float32", name="A") b_in = tf.placeholder(shape=b_shape, dtype="float32", name="B") @@ -801,7 +799,6 @@ def test_fill(): def _test_pack(axis, shape, **kwargs): - a = np.arange(np.prod(shape), dtype=np.float32).reshape(shape) b = np.arange(np.prod(shape), dtype=np.float32).reshape(shape) diff --git a/tests/python/ir/test_ir_type.py b/tests/python/ir/test_ir_type.py index 238cb5965ba3..0c95bc4554aa 100644 --- a/tests/python/ir/test_ir_type.py +++ b/tests/python/ir/test_ir_type.py @@ -38,22 +38,10 @@ def test_tensor_type_bad_constructor(): pass -def test_type_param(): - tp = tvm.ir.TypeVar("name", tvm.ir.TypeKind.Type) - assert tp.kind == tvm.ir.TypeKind.Type - # assert tp.span # TODO allow us to set span - str(tp) - check_json_roundtrip(tp) - - def test_func_type(): - type_params = tvm.runtime.convert([]) - type_constraints = tvm.runtime.convert([]) # TODO: fill me in arg_types = tvm.runtime.convert([]) ret_type = tvm.ir.TensorType((1, 2, 3), "float32") - tf = tvm.ir.FuncType(arg_types, ret_type, type_params, type_constraints) - assert tf.type_params == type_params - assert tf.type_constraints == type_constraints + tf = tvm.ir.FuncType(arg_types, ret_type) assert tf.arg_types == arg_types assert tf.ret_type == ret_type assert tf.span == None @@ -63,10 +51,9 @@ def test_func_type(): def test_tuple_type(): - tp = tvm.ir.TypeVar("tp", tvm.ir.TypeKind.Type) - tf = tvm.ir.FuncType([], tvm.ir.TupleType([]), [], []) + tf = tvm.ir.FuncType([], tvm.ir.TupleType([])) tt = tvm.ir.TensorType(tvm.runtime.convert([1, 2, 3]), "float32") - fields = tvm.runtime.convert([tp, tf, tt]) + fields = tvm.runtime.convert([tf, tt]) tup_ty = tvm.ir.TupleType(fields) assert tup_ty.fields == fields @@ -74,27 +61,7 @@ def test_tuple_type(): check_json_roundtrip(tup_ty) -def test_type_relation(): - tp = tvm.ir.TypeVar("tp", tvm.ir.TypeKind.Type) - tf = tvm.ir.FuncType([], None, [], []) - tt = tvm.ir.TensorType(tvm.runtime.convert([1, 2, 3]), "float32") - args = tvm.runtime.convert([tp, tf, tt]) - - num_inputs = 2 - func = tvm.ir.EnvFunc.get("tvm.relay.type_relation.Broadcast") - attrs = tvm.ir.make_node("attrs.TestAttrs", name="attr", padding=(3, 4)) - - tr = tvm.ir.TypeRelation(func, args, num_inputs, attrs) - assert tr.args == args - assert tr.num_inputs == num_inputs - str(tr) - check_json_roundtrip(tr) - - if __name__ == "__main__": test_tensor_type_bad_constructor() - test_tensor_type() - test_type_param() test_func_type() test_tuple_type() - test_type_relation() diff --git a/tests/python/nightly/test_nnapi/infrastructure.py b/tests/python/nightly/test_nnapi/infrastructure.py index aa5580c375ae..a86c681f0bc0 100644 --- a/tests/python/nightly/test_nnapi/infrastructure.py +++ b/tests/python/nightly/test_nnapi/infrastructure.py @@ -20,7 +20,6 @@ import tvm import tvm.script.relax as R -# from tvm.contrib.debugger import debug_runtime as graph_executor from tvm.contrib import ndk, utils from tvm.relax.backend.contrib.nnapi import partition_for_nnapi @@ -96,7 +95,6 @@ def _build(mod, enable_nnapi): def _run(remote, tracker, ex, inputs): - tmp = utils.tempdir() so_name = "test_mod.so" so_path = tmp / so_name @@ -106,7 +104,6 @@ def _run(remote, tracker, ex, inputs): dev = remote.cpu(0) try: - # Execute the model on the remote. remote_ex = remote.load_module(so_name) vm = tvm.relax.VirtualMachine(remote_ex, device=dev) diff --git a/tests/python/relax/test_dataflow_pattern.py b/tests/python/relax/test_dataflow_pattern.py index 4b5da0d9e608..a534b7c0c7c9 100644 --- a/tests/python/relax/test_dataflow_pattern.py +++ b/tests/python/relax/test_dataflow_pattern.py @@ -283,8 +283,8 @@ def test_op_attr(): yp = is_var("y") # TODO(@yuchen): reenable the assert after figuring out why it fails # assert is_op("nn.conv2d")(xp, yp).has_attr({"strides": [3, 3]}).match(conv2d) - assert not is_op("nn.conv2d")(xp, yp).has_attr({"strides": [4, 3]}).match(conv2d) - assert not is_op("nn.conv2d")(xp, yp).has_attr({"strides": [3, 3]}).match(conv2d) + assert not is_op("relax.nn.conv2d")(xp, yp).has_attr({"strides": [4, 3]}).match(conv2d) + assert not is_op("relax.nn.conv2d")(xp, yp).has_attr({"strides": [3, 3]}).match(conv2d) def test_match_call_attr(): diff --git a/tests/python/target/test_target_target.py b/tests/python/target/test_target_target.py index b99834aef35a..c6908d23f000 100644 --- a/tests/python/target/test_target_target.py +++ b/tests/python/target/test_target_target.py @@ -22,39 +22,11 @@ from tvm.target import Target, arm_cpu, bifrost, cuda, intel_graphics, mali, rocm -@tvm.target.generic_func -def mygeneric(data): - # default generic function - return data + 1 - - -@mygeneric.register(["cuda", "gpu"]) -def cuda_func(data): - return data + 2 - - -@mygeneric.register("rocm") -def rocm_func(data): - return data + 3 - - -@mygeneric.register("cpu") -def rocm_func(data): - return data + 10 - - def test_all_targets_device_type_verify(): """Consistency verification for all targets' device type""" all_targets = [tvm.target.Target(t) for t in tvm.target.Target.list_kinds()] for tgt in all_targets: - # skip targets with hooks or otherwise intended to be used with external codegen - relay_to_tir = tgt.get_kind_attr("RelayToTIR") - tir_to_runtime = tgt.get_kind_attr("TIRToRuntime") - is_external_codegen = tgt.get_kind_attr("is_external_codegen") - if relay_to_tir is not None or tir_to_runtime is not None or is_external_codegen: - continue - if tgt.kind.name not in tvm._ffi.runtime_ctypes.Device.STR2MASK: raise KeyError("Cannot find target kind: %s in Device.STR2MASK" % tgt.kind.name) @@ -63,80 +35,6 @@ def test_all_targets_device_type_verify(): ) -def test_target_dispatch(): - with tvm.target.cuda(): - assert mygeneric(1) == 3 - assert mygeneric.get_packed_func()(1) == 3 - - with tvm.target.rocm(): - assert mygeneric(1) == 4 - assert mygeneric.get_packed_func()(1) == 4 - - with tvm.target.Target("cuda"): - assert mygeneric(1) == 3 - assert mygeneric.get_packed_func()(1) == 3 - - with tvm.target.arm_cpu(): - assert mygeneric(1) == 11 - assert mygeneric.get_packed_func()(1) == 11 - - with tvm.target.Target("metal"): - assert mygeneric(1) == 3 - assert mygeneric.get_packed_func()(1) == 3 - - assert tvm.target.Target.current() is None - - -@tvm.target.override_native_generic_func("test_target_temp_strategy") -def target_generic(data): - # default generic function - return data + 1 - - -@target_generic.register(["cuda", "gpu"]) -def target_cuda_func(data): - return data + 2 - - -def temp_target_cuda_func(data): - return data + 3 - - -def test_target_temp_strategy(): - class TempStrategy(object): - def __init__(self, name, target, fstrategy): - generic_fstrategy = tvm.target.get_native_generic_func(name) - self.target = target - self.name = name - self.origin_func = {} - with tvm.target.Target(target) as target_obj: - for tgt_key in target_obj.keys: - self.origin_func[tgt_key] = generic_fstrategy.get_packed_func() - generic_fstrategy.register(fstrategy, tgt_key, allow_override=True) - - def __enter__(self): - return self - - def __exit__(self, typ, value, traceback): - generic_fstrategy = tvm.target.get_native_generic_func(self.name) - with tvm.target.Target(self.target) as target_obj: - for tgt_key in target_obj.keys: - generic_fstrategy.register( - self.origin_func[tgt_key], tgt_key, allow_override=True - ) - - with tvm.target.Target("cuda"): - assert target_generic(1) == 3 - - # The strategy func change to temp_target_cuda_func. - with TempStrategy("test_target_temp_strategy", "cuda", temp_target_cuda_func): - with tvm.target.Target("cuda"): - assert target_generic(1) == 4 - - with tvm.target.Target("cuda"): - assert target_generic(1) == 3 - - def test_target_string_parse(): target = tvm.target.Target("cuda -model=unknown -libs=cublas,cudnn") diff --git a/tests/python/tir-analysis/test_tir_analysis_calculate_workspace.py b/tests/python/tir-analysis/test_tir_analysis_calculate_workspace.py deleted file mode 100644 index 29bfc5845870..000000000000 --- a/tests/python/tir-analysis/test_tir_analysis_calculate_workspace.py +++ /dev/null @@ -1,126 +0,0 @@ -# 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 pytest - -import tvm -from tvm import tir -from tvm.script import tir as T - - -# fmt: off -@T.prim_func -def primfunc_global_allocates(placeholder_144: T.handle, placeholder_145: T.handle, placeholder_146: T.handle, T_cast_48: T.handle) -> None: - # function attr dict - T.func_attr({"global_symbol": "fused_nn_conv2d_add_cast_fixed_point_multiply_clip_cast_cast_13", "tir.noalias": True}) - placeholder_147 = T.match_buffer(placeholder_144, [100352], dtype="int16", elem_offset=0, align=64, offset_factor=1) - placeholder_148 = T.match_buffer(placeholder_145, [4608], dtype="int16", elem_offset=0, align=64, offset_factor=1) - placeholder_149 = T.match_buffer(placeholder_146, [512], dtype="int32", elem_offset=0, align=64, offset_factor=1) - T_cast_49 = T.match_buffer(T_cast_48, [100352], dtype="int16", elem_offset=0, align=64, offset_factor=1) - # body - PaddedInput_22 = T.decl_buffer([131072], "int16") - DepthwiseConv2d_9 = T.decl_buffer([100352], "int32") - for i1_29, i2_39, i3_40 in T.grid(16, 16, 512): - PaddedInput_22[(((i1_29*8192) + (i2_39*512)) + i3_40)] = T.if_then_else(((((1 <= i1_29) and (i1_29 < 15)) and (1 <= i2_39)) and (i2_39 < 15)), placeholder_147[((((i1_29*7168) + (i2_39*512)) + i3_40) - 7680)], T.int16(0), dtype="int16") - for i_9, j_9, c_9 in T.grid(14, 14, 512): - DepthwiseConv2d_9[(((i_9*7168) + (j_9*512)) + c_9)] = 0 - for di_9, dj_9 in T.grid(3, 3): - DepthwiseConv2d_9[(((i_9*7168) + (j_9*512)) + c_9)] = (DepthwiseConv2d_9[(((i_9*7168) + (j_9*512)) + c_9)] + (PaddedInput_22[(((((i_9*8192) + (di_9*8192)) + (j_9*512)) + (dj_9*512)) + c_9)].astype("int32")*placeholder_148[(((di_9*1536) + (dj_9*512)) + c_9)].astype("int32"))) - for ax1_27, ax2_28, ax3_30 in T.grid(14, 14, 512): - DepthwiseConv2d_9[(((ax1_27*7168) + (ax2_28*512)) + ax3_30)] = (DepthwiseConv2d_9[(((ax1_27*7168) + (ax2_28*512)) + ax3_30)] + placeholder_149[ax3_30]) - for i1_30, i2_40, i3_41 in T.grid(14, 14, 512): - DepthwiseConv2d_9[(((i1_30*7168) + (i2_40*512)) + i3_41)] = T.q_multiply_shift(DepthwiseConv2d_9[(((i1_30*7168) + (i2_40*512)) + i3_41)], 1269068532, 31, -4, dtype="int32") - for i1_31, i2_41, i3_42 in T.grid(14, 14, 512): - DepthwiseConv2d_9[(((i1_31*7168) + (i2_41*512)) + i3_42)] = T.max(T.max(DepthwiseConv2d_9[(((i1_31*7168) + (i2_41*512)) + i3_42)], 255), 0) - for ax1_28, ax2_29, ax3_31 in T.grid(14, 14, 512): - PaddedInput_22[(((ax1_28*7168) + (ax2_29*512)) + ax3_31)] = DepthwiseConv2d_9[(((ax1_28*7168) + (ax2_29*512)) + ax3_31)].astype("uint8") - for ax1_29, ax2_30, ax3_32 in T.grid(14, 14, 512): - T_cast_49[(((ax1_29*7168) + (ax2_30*512)) + ax3_32)] = PaddedInput_22[(((ax1_29*7168) + (ax2_30*512)) + ax3_32)].astype("int16") -# fmt: on - - -# fmt: off -@T.prim_func -def primfunc_local_allocates(placeholder_162: T.handle, placeholder_163: T.handle, placeholder_164: T.handle, T_cast_76: T.handle) -> None: - # function attr dict - T.func_attr({"global_symbol": "fused_nn_conv2d_add_cast_fixed_point_multiply_clip_cast_cast_9", "tir.noalias": True}) - placeholder_165 = T.match_buffer(placeholder_162, [100352], dtype="int16", elem_offset=0, align=64, offset_factor=1) - placeholder_166 = T.match_buffer(placeholder_163, [4608], dtype="int16", elem_offset=0, align=64, offset_factor=1) - placeholder_167 = T.match_buffer(placeholder_164, [512], dtype="int32", elem_offset=0, align=64, offset_factor=1) - T_cast_77 = T.match_buffer(T_cast_76, [100352], dtype="int16", elem_offset=0, align=64, offset_factor=1) - sid_21 = T.allocate_const([0,1,2,3,4,5,6,7], "int8", [8]) - # body - PaddedInput_25 = T.decl_buffer([131072], "int16") - for i1_35, i2_46, i3_47 in T.grid(16, 16, 512): - PaddedInput_25[(((i1_35*8192) + (i2_46*512)) + i3_47)] = T.if_then_else(((((1 <= i1_35) and (i1_35 < 15)) and (1 <= i2_46)) and (i2_46 < 15)), placeholder_165[((((i1_35*7168) + (i2_46*512)) + i3_47) - 7680)], T.int16(0), dtype="int16") - T_add_11 = T.decl_buffer([100352], "int32") - with T.decl_buffer([100352], "int32") as DepthwiseConv2d_11: - for i_11, j_11, c_11 in T.grid(14, 14, 512): - DepthwiseConv2d_11[(((i_11*7168) + (j_11*512)) + c_11)] = 0 - for di_11, dj_11 in T.grid(3, 3): - DepthwiseConv2d_11[(((i_11*7168) + (j_11*512)) + c_11)] = (DepthwiseConv2d_11[(((i_11*7168) + (j_11*512)) + c_11)] + (PaddedInput_25[(((((i_11*8192) + (di_11*8192)) + (j_11*512)) + (dj_11*512)) + c_11)].astype("int32")*placeholder_166[(((di_11*1536) + (dj_11*512)) + c_11)].astype("int32"))) - for ax1_44, ax2_45, ax3_47 in T.grid(14, 14, 512): - T_add_11[(((ax1_44*7168) + (ax2_45*512)) + ax3_47)] = (DepthwiseConv2d_11[(((ax1_44*7168) + (ax2_45*512)) + ax3_47)] + placeholder_167[ax3_47]) - compute_22 = T.decl_buffer([100352], "int32") - with T.decl_buffer([100352], "int32") as T_cast_78: - for ax1_45, ax2_46, ax3_48 in T.grid(14, 14, 512): - T_cast_78[(((ax1_45*7168) + (ax2_46*512)) + ax3_48)] = T_add_11[(((ax1_45*7168) + (ax2_46*512)) + ax3_48)] - for i1_36, i2_47, i3_48 in T.grid(14, 14, 512): - compute_22[(((i1_36*7168) + (i2_47*512)) + i3_48)] = T.q_multiply_shift(T_cast_78[(((i1_36*7168) + (i2_47*512)) + i3_48)], 1948805937, 31, -5, dtype="int32") - T_cast_79 = T.decl_buffer([100352], "uint8") - with T.decl_buffer([100352], "int32") as compute_23: - for i1_37, i2_48, i3_49 in T.grid(14, 14, 512): - compute_23[(((i1_37*7168) + (i2_48*512)) + i3_49)] = T.max(T.max(compute_22[(((i1_37*7168) + (i2_48*512)) + i3_49)], 255), 0) - for ax1_46, ax2_47, ax3_49 in T.grid(14, 14, 512): - T_cast_79[(((ax1_46*7168) + (ax2_47*512)) + ax3_49)] = compute_23[(((ax1_46*7168) + (ax2_47*512)) + ax3_49)].astype("uint8") - for ax1_47, ax2_48, ax3_50 in T.grid(14, 14, 512): - T_cast_77[(((ax1_47*7168) + (ax2_48*512)) + ax3_50)] = T_cast_79[(((ax1_47*7168) + (ax2_48*512)) + ax3_50)].astype("int16") -# fmt: on - - -@T.prim_func -def prim_func_decl_vector_type(a: T.handle, b: T.handle): - T.func_attr({"tir.noalias": True}) - A = T.match_buffer(a, (4,), "float32x4") - B = T.match_buffer(b, (4,), "float32x4") - C = T.decl_buffer((4,), "float32x4") - for i in range(3): - with T.block("block"): - vi = T.axis.remap("S", [i]) - B[vi] = A[vi] + C[vi] - - -@pytest.mark.parametrize("alignment,size,consts", [(1, 663552, 0), (10, 663560, 0)]) -def test_global_allocates(alignment, size, consts): - primfunc = primfunc_global_allocates - assert tvm.tir.analysis.calculate_constant_bytes(primfunc, alignment) == consts - assert tvm.tir.analysis.calculate_workspace_bytes(primfunc, alignment) == size - - -@pytest.mark.parametrize("alignment,size,consts", [(1, 1566720, 8), (100, 1567100, 100)]) -def test_local_allocates(alignment, size, consts): - primfunc = primfunc_local_allocates - assert tvm.tir.analysis.calculate_constant_bytes(primfunc, alignment) == consts - assert tvm.tir.analysis.calculate_workspace_bytes(primfunc, alignment) == size - - -def test_vector_type(): - primfunc = prim_func_decl_vector_type - assert tvm.tir.analysis.calculate_workspace_bytes(primfunc, 1) == 64 - - -if __name__ == "__main__": - tvm.testing.main() diff --git a/tests/python/tir-analysis/test_tir_analysis_identify_memcpy.py b/tests/python/tir-analysis/test_tir_analysis_identify_memcpy.py index 8510a66d308d..5a44a43ae70b 100644 --- a/tests/python/tir-analysis/test_tir_analysis_identify_memcpy.py +++ b/tests/python/tir-analysis/test_tir_analysis_identify_memcpy.py @@ -56,7 +56,7 @@ def test_identify_memcpy(self, func, expected): class Test1D(BaseTest): """Simplest test case""" - def func(A: T.Buffer[1024, "float32"], B: T.Buffer[1024, "float32"]): + def func(A: T.Buffer(1024, "float32"), B: T.Buffer(1024, "float32")): for i in T.serial(1024): B[i] = A[i] @@ -68,7 +68,7 @@ def expected(self, func): class Test1DCompute(BaseTest): """Like Test1D, but a computation prevents this being a memcpy""" - def func(A: T.Buffer[1024, "float32"], B: T.Buffer[1024, "float32"]): + def func(A: T.Buffer(1024, "float32"), B: T.Buffer(1024, "float32")): for i in T.serial(1024): B[i] = A[i] + 1.0 @@ -79,7 +79,7 @@ def expected(self, func): class Test1DConditional(BaseTest): """Like Test1D, but a conditionals prevents this being a memcpy""" - def func(A: T.Buffer[1024, "float32"], B: T.Buffer[1024, "float32"]): + def func(A: T.Buffer(1024, "float32"), B: T.Buffer(1024, "float32")): for i in T.serial(1024): if i < 1024: B[i] = A[i] @@ -92,7 +92,7 @@ def expected(self, func): class Test1DStridedInput(BaseTest): """Like Test1D, but strided input prevents this being a memcpy""" - def func(A: T.Buffer[2048, "float32"], B: T.Buffer[1024, "float32"]): + def func(A: T.Buffer(2048, "float32"), B: T.Buffer(1024, "float32")): for i in T.serial(1024): B[i] = A[i * 2] @@ -103,7 +103,7 @@ def expected(self, func): class Test1DStridedOutput(BaseTest): """Like Test1D, but strided output prevents this being a memcpy""" - def func(A: T.Buffer[1024, "float32"], B: T.Buffer[2048, "float32"]): + def func(A: T.Buffer(1024, "float32"), B: T.Buffer(2048, "float32")): for i in T.serial(1024): B[i * 2] = A[i] @@ -114,7 +114,7 @@ def expected(self, func): class Test1DInput2DOutputFusedLoop(BaseTest): """Like Test1D, but the output is written as a 2-d buffer""" - def func(A: T.Buffer[1024, "float32"], B: T.Buffer[(32, 32), "float32"]): + def func(A: T.Buffer(1024, "float32"), B: T.Buffer((32, 32), "float32")): for i in T.serial(1024): B[i // 32, i % 32] = A[i] @@ -126,7 +126,7 @@ def expected(self, func): class Test2DInput1DOutputFusedLoop(BaseTest): """Like Test1D, but the input is written as a 2-d buffer""" - def func(A: T.Buffer[(32, 32), "float32"], B: T.Buffer[1024, "float32"]): + def func(A: T.Buffer((32, 32), "float32"), B: T.Buffer(1024, "float32")): for i in T.serial(1024): B[i] = A[i // 32, i % 32] @@ -144,7 +144,7 @@ class Test1DInput1DOutputNestedLoop(BaseTest): is more convenient to return the results for all loops. """ - def func(A: T.Buffer[1024, "float32"], B: T.Buffer[1024, "float32"]): + def func(A: T.Buffer(1024, "float32"), B: T.Buffer(1024, "float32")): for i, j in T.grid(32, 32): B[i * 32 + j] = A[i * 32 + j] @@ -165,7 +165,7 @@ class Test1DInput1DOutputNestedLoopEquivalentExpressions(BaseTest): equivalent. """ - def func(A: T.Buffer[1024, "float32"], B: T.Buffer[1024, "float32"]): + def func(A: T.Buffer(1024, "float32"), B: T.Buffer(1024, "float32")): for i, j in T.grid(32, 32): B[i * 32 + j] = A[j + i * 32] @@ -181,7 +181,7 @@ def expected(self, func): class Test1DInput2DOutputNestedLoop(BaseTest): """Like Test1DInput1DOutputNestedLoop, but with a 2-d output buffer""" - def func(A: T.Buffer[1024, "float32"], B: T.Buffer[(32, 32), "float32"]): + def func(A: T.Buffer(1024, "float32"), B: T.Buffer((32, 32), "float32")): for i, j in T.grid(32, 32): B[i, j] = A[i * 32 + j] @@ -197,7 +197,7 @@ def expected(self, func): class Test2DInput1DOutputNestedLoop(BaseTest): """Like Test1DInput1DOutputNestedLoop, but with a 2-d input buffer""" - def func(A: T.Buffer[(32, 32), "float32"], B: T.Buffer[1024, "float32"]): + def func(A: T.Buffer((32, 32), "float32"), B: T.Buffer(1024, "float32")): for i, j in T.grid(32, 32): B[i * 32 + j] = A[i, j] @@ -213,7 +213,7 @@ def expected(self, func): class Test2DInput2DOutputNestedLoop(BaseTest): """Like Test1DInput1DOutputNestedLoop, but with 2-d input/output buffers""" - def func(A: T.Buffer[(32, 32), "float32"], B: T.Buffer[(32, 32), "float32"]): + def func(A: T.Buffer((32, 32), "float32"), B: T.Buffer((32, 32), "float32")): for i, j in T.grid(32, 32): B[i, j] = A[i, j] @@ -232,7 +232,7 @@ class Test2DInput2DOutputTransposeOutput(BaseTest): This is not recognized as a memcpy, because it results in a transpose. """ - def func(A: T.Buffer[(32, 32), "float32"], B: T.Buffer[(32, 32), "float32"]): + def func(A: T.Buffer((32, 32), "float32"), B: T.Buffer((32, 32), "float32")): for i, j in T.grid(32, 32): B[j, i] = A[i, j] @@ -249,7 +249,7 @@ class Test2DInput2DOutputTransposeInput(BaseTest): This is not recognized as a memcpy, because it results in a transpose. """ - def func(A: T.Buffer[(32, 32), "float32"], B: T.Buffer[(32, 32), "float32"]): + def func(A: T.Buffer((32, 32), "float32"), B: T.Buffer((32, 32), "float32")): for i, j in T.grid(32, 32): B[i, j] = A[j, i] @@ -269,7 +269,7 @@ class Test2DInput2DOutputTransposeBoth(BaseTest): region has been copied over, even though it occurs out of order. """ - def func(A: T.Buffer[(32, 32), "float32"], B: T.Buffer[(32, 32), "float32"]): + def func(A: T.Buffer((32, 32), "float32"), B: T.Buffer((32, 32), "float32")): for i, j in T.grid(32, 32): B[j, i] = A[j, i] @@ -288,7 +288,7 @@ class TestCacheRead(BaseTest): pattern would appear when B is a read cache of A. """ - def func(A: T.Buffer[(32, 32), "float32"], B: T.Buffer[32, "float32"]): + def func(A: T.Buffer((32, 32), "float32"), B: T.Buffer(32, "float32")): for i, j in T.grid(32, 32): B[j] = A[i, j] @@ -308,7 +308,7 @@ class TestCacheWrite(BaseTest): pattern would appear when A is a write cache of B. """ - def func(A: T.Buffer[32, "float32"], B: T.Buffer[(32, 32), "float32"]): + def func(A: T.Buffer(32, "float32"), B: T.Buffer((32, 32), "float32")): for i, j in T.grid(32, 32): B[i, j] = A[j] diff --git a/tests/python/tir-base/test_debug_info.py b/tests/python/tir-base/test_debug_info.py deleted file mode 100644 index 2e799815252d..000000000000 --- a/tests/python/tir-base/test_debug_info.py +++ /dev/null @@ -1,189 +0,0 @@ -# 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 line-level debug info for TIR""" -import tvm -import tvm.testing -from tvm import tir -from tvm.script import tir as T, ir as I - -from typing import List, Dict -import re - - -def find_di_locations(source: str) -> Dict[int, int]: - """ - Parse out DILocation references in printed LLVM IR - """ - result = {} - - for line in source.splitlines(): - m = re.match(r"!(\d+) = !DILocation\(line: (\d+).*", line) - if m: - debug_id, line = m.groups() - result[debug_id] = line - - return result - - -def _module(): - @tvm.script.ir_module - class MyModule: - @T.prim_func - def main(a: T.handle, b: T.handle): - # We exchange data between function by handles, which are similar to pointer. - T.func_attr( - { - "tir.noalias": True, - "target": T.target("llvm"), - } - ) - # Create buffer from handles. - A = T.match_buffer(a, (8,), dtype="float32") - B = T.match_buffer(b, (8,), dtype="float32") - for i in range(8): - # A block is an abstraction for computation. - with T.block("B"): - # Define a spatial block iterator and bind it to value i. - vi = T.axis.spatial(8, i) - assert 1 == 0, "Some numbers" - B[vi] = A[vi] + 1.0 - - return MyModule - - -def test_tir_debug_info(): - """ - Test that Spans are correctly replaced with debug spans that reference - the printed TIR - """ - - def find_span(m): - func = next(m.functions.values()) - return func.body.block.body.span - - module_before = _module() - span_before = find_span(module_before) - assert span_before is None - - module_after = tir.transform.InstallDebugSpans()(module_before) - span_after = find_span(module_after) - - # Check that the module name has been added and a line number is present - assert span_after.source_name.name == "main.tir" - assert span_after.line == 4 - - -def test_tir_debug_info_with_subroutine(): - """Like test_tir_debug_info, but with a TIR subroutine - - The current InstallDebugSpans applies to a single PrimFunc. This - test verifies that the existence of device-side subroutines - - """ - - def find_span(m): - func = next(m.functions.values()) - return func.body.block.body.span - - @tvm.script.ir_module - class module_before: - @T.prim_func - def main(a: T.handle, b: T.handle): - T.func_attr({"global_symbol": "main", "tir.noalias": True, "target": T.target("llvm")}) - A = T.match_buffer(a, (8,), dtype="float32") - B = T.match_buffer(b, (8,), dtype="float32") - for i in range(8): - with T.block("B"): - vi = T.axis.spatial(8, i) - module_before.subroutine(T.address_of(A[vi]), T.address_of(B[vi])) - - @T.prim_func - def subroutine(a_ptr: T.handle("float32"), b_ptr: T.handle("float32")): - T.func_attr({"global_symbol": "main", "tir.noalias": True}) - A = T.decl_buffer(1, "float32", data=a_ptr) - B = T.decl_buffer(1, "float32", data=b_ptr) - B[0] = A[1] + 1.0 - - span_before = find_span(module_before) - assert span_before is None - - module_after = tir.transform.InstallDebugSpans()(module_before) - span_after = find_span(module_after) - - # Check that the module name has been added and a line number is present - assert span_after.source_name.name == "main.tir" - assert span_after.line == 4 - - -def test_llvm_ir_debug_info(): - """ - Check that the right amount of debug locations are present - """ - MyModule = _module() - with tvm.transform.PassContext(opt_level=3, config={"tir.enable_debug": True}): - runtime_module = tvm.build(MyModule, target="llvm") - - source = runtime_module.get_source() - - locations = find_di_locations(source) - assert len(locations) == 41 - - -def test_llvm_ir_debug_accuracy(): - """ - Check that the debug location on an assert is correct - """ - MyModule = _module() - with tvm.transform.PassContext(opt_level=3, config={"tir.enable_debug": True}): - runtime_module = tvm.build(MyModule, target="llvm") - source = runtime_module.get_source() - locations = find_di_locations(source) - - # Find the 'assert' from MyModule - debug_dir_match = re.search(r"tail call void %0\(.* !dbg !(\d+)\n", source) - - # Extract out the debug directive line - directive_idx = debug_dir_match.groups()[0] - - # Check that it matches the expected line number (in main.tir) - debug_line_no = int(locations[directive_idx]) - assert debug_line_no == 60 - - -def test_building_without_llvm_equivalent(): - """A TIR PrimFunc may contain non-LLVM types - - Types used in optimized kernels (e.g. "e4m3_float8") may not have - an equivalent in DWARF, or the mapping from TIR type to DWARF type - may not be defined. If this occurs, the function should still be - able to be built. - """ - - @I.ir_module - class Module: - @T.prim_func(private=True) - def main(A_data: T.handle("e4m3_float8"), B_data: T.handle("e4m3_float8")): - A = T.decl_buffer(128, "e4m3_float8", data=A_data) - B = T.decl_buffer(128, "e4m3_float8", data=B_data) - for i in range(128): - B[i] = A[i] - - tvm.target.codegen.build_module(Module, "llvm") - - -if __name__ == "__main__": - tvm.testing.main() diff --git a/tests/python/tir-schedule/test_tir_schedule_compute_at.py b/tests/python/tir-schedule/test_tir_schedule_compute_at.py index 2c44c9b29569..7e561cbf0e23 100644 --- a/tests/python/tir-schedule/test_tir_schedule_compute_at.py +++ b/tests/python/tir-schedule/test_tir_schedule_compute_at.py @@ -1227,7 +1227,7 @@ def test_compute_at_tiled_repeat_op(use_block_name): def test_compute_at_rev_iter(): @T.prim_func - def before(X: T.Buffer[(10, 10), "float32"], Z: T.Buffer[(10, 10), "float32"]): + def before(X: T.Buffer((10, 10), "float32"), Z: T.Buffer((10, 10), "float32")): Y = T.alloc_buffer([10, 10], "float32") for i, j in T.grid(10, 10): with T.block("b0"): @@ -1239,7 +1239,7 @@ def before(X: T.Buffer[(10, 10), "float32"], Z: T.Buffer[(10, 10), "float32"]): Z[vi, vj] = Y[vj, vi] + 2.0 @T.prim_func - def after(X: T.Buffer[(10, 10), "float32"], Z: T.Buffer[(10, 10), "float32"]): + def after(X: T.Buffer((10, 10), "float32"), Z: T.Buffer((10, 10), "float32")): Y = T.alloc_buffer([10, 10], "float32") for i in range(10): for j in range(10): diff --git a/tests/python/tir-transform/test_tir_transform_loop_partition.py b/tests/python/tir-transform/test_tir_transform_loop_partition.py index bec4129ffcbf..25660880e13f 100644 --- a/tests/python/tir-transform/test_tir_transform_loop_partition.py +++ b/tests/python/tir-transform/test_tir_transform_loop_partition.py @@ -714,7 +714,7 @@ def concat_five_buffers_with_equalities_expected( @T.prim_func -def nested_partition_with_single_points(A: T.Buffer[(25,), "int32"]): +def nested_partition_with_single_points(A: T.Buffer((25,), "int32")): for i in T.serial(5, annotations={"pragma_loop_partition_hint": 1}): if i == 1: for j in T.serial(5, annotations={"pragma_loop_partition_hint": 1}): @@ -727,7 +727,7 @@ def nested_partition_with_single_points(A: T.Buffer[(25,), "int32"]): @T.prim_func -def nested_partition_with_single_points_expected(A: T.Buffer[(25,), "int32"]): +def nested_partition_with_single_points_expected(A: T.Buffer((25,), "int32")): for j in range(2): A[j + 3] = j + 3 for j in range(2): @@ -764,7 +764,7 @@ def test_single_point_partition(origin, expected): def test_equation_on_floordiv(): @T.prim_func - def before(A: T.Buffer[(2, 2, 20), "int32"]): + def before(A: T.Buffer((2, 2, 20), "int32")): for i in T.serial(5, annotations={"pragma_loop_partition_hint": 1}): if i == 1: for vv in T.vectorized(640, annotations={"pragma_loop_partition_hint": 1}): @@ -772,7 +772,7 @@ def before(A: T.Buffer[(2, 2, 20), "int32"]): A[i - 1, i * 2 + vv // 320 - 3, vv % 320 // 16] = 1 @T.prim_func - def expected(A: T.Buffer[(2, 2, 20), "int32"]): + def expected(A: T.Buffer((2, 2, 20), "int32")): for vv in T.vectorized(320): A[0, 0, vv // 16] = 1 @@ -787,7 +787,7 @@ def test_ignore_loop_partition_hint(): """Skip unroll body and prologue for pipeline case""" @T.prim_func - def before(A: T.Buffer[(10), "float32"], D: T.Buffer[(10), "float32"]): + def before(A: T.Buffer((10), "float32"), D: T.Buffer((10), "float32")): B = T.decl_buffer([2], "float32") C = T.decl_buffer([2], "float32") for i in T.serial(12, annotations={"pragma_loop_partition_hint": 1}): @@ -799,7 +799,7 @@ def before(A: T.Buffer[(10), "float32"], D: T.Buffer[(10), "float32"]): D[i - 2] = C[i % 2] + 3.0 @T.prim_func - def expected(A: T.Buffer[(10), "float32"], D: T.Buffer[(10), "float32"]): + def expected(A: T.Buffer((10), "float32"), D: T.Buffer((10), "float32")): B = T.decl_buffer([2], "float32") C = T.decl_buffer([2], "float32") for i in range(2): diff --git a/tests/python/tir-usmp/test_tir_usmp_algo.py b/tests/python/tir-usmp/test_tir_usmp_algo.py deleted file mode 100644 index 80f7f6b999ce..000000000000 --- a/tests/python/tir-usmp/test_tir_usmp_algo.py +++ /dev/null @@ -1,683 +0,0 @@ -# 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 pytest - -import tvm -from tvm import tir, script -from tvm.script import tir as T -from tvm.tir import stmt_functor -from tvm.tir.usmp import utils as usmp_utils -from tvm.target import Target -from tvm import WorkspacePoolInfo, PoolInfoProperties - - -def _replace_stmt_with_buf_var_names(buffer_info_map): - """helper to replace tir.allocates with buffer names""" - new_buffer_info_map = dict() - for k, v in buffer_info_map.items(): - new_buffer_info_map[v.buffer_var.name] = k - return new_buffer_info_map - - -def _verify_conflicts(main_buf_name, conflicting_buf_names, buffer_info_map): - """helper to check expected liveness conflicts""" - buf_info = buffer_info_map[main_buf_name] - for conflict in buf_info.conflicts: - assert conflict.name_hint in conflicting_buf_names - - -def _get_allocates(primfunc): - """helper to extract all allocate nodes by name""" - allocates = dict() - - def get_allocate(stmt): - if isinstance(stmt, tvm.tir.Allocate): - allocates[str(stmt.buffer_var.name)] = stmt - - stmt_functor.post_order_visit(primfunc.body, get_allocate) - return allocates - - -def _assign_poolinfos_to_allocates_in_primfunc(primfunc, pool_infos): - """helper to assing poolinfos to allocate nodes in a tir.PrimFunc""" - - def set_poolinfos(stmt): - if isinstance(stmt, tvm.tir.Allocate): - return tvm.tir.Allocate( - buffer_var=stmt.buffer_var, - dtype=stmt.dtype, - extents=stmt.extents, - condition=stmt.condition, - body=stmt.body, - annotations={tvm.tir.usmp.utils.CANDIDATE_MEMORY_POOL_ATTR: pool_infos}, - ) - - return primfunc.with_body(stmt_functor.ir_transform(primfunc.body, None, set_poolinfos)) - - -def _assign_poolinfos_to_allocates_in_irmodule(mod, pool_infos): - """helper to assing poolinfos to allocate nodes in a IRModule""" - ret = tvm.IRModule() - for global_var, basefunc in mod.functions.items(): - if isinstance(basefunc, tvm.tir.PrimFunc): - ret[global_var] = _assign_poolinfos_to_allocates_in_primfunc(basefunc, pool_infos) - return ret - - -def _assign_targets_to_primfuncs_irmodule(mod, target): - """helper to assign target for PrimFunc in a IRModule""" - ret = tvm.IRModule() - for global_var, basefunc in mod.functions.items(): - if isinstance(basefunc, tvm.tir.PrimFunc): - ret[global_var] = basefunc.with_attr("target", target) - return ret - - -def _check_max_workspace_size(buffer_pool_allocations, pool_info, size): - max_workspace_size = 0 - for buffer_info, pool_allocation in buffer_pool_allocations.items(): - if pool_allocation.pool_info == pool_info: - size_candidate = pool_allocation.byte_offset + buffer_info.size_bytes - if size_candidate > max_workspace_size: - max_workspace_size = size_candidate - assert max_workspace_size == size - - -def test_no_pool_error(): - target = Target("c") - tiny_workspace_pool = WorkspacePoolInfo( - "tiny_workspace", - [target], - PoolInfoProperties(size_hint_bytes=10), - ) - bi_a = usmp_utils.BufferInfo( - name_hint="bi_a", size_bytes=10, pool_candidates=[tiny_workspace_pool] - ) - bi_b = usmp_utils.BufferInfo( - name_hint="bi_b", size_bytes=10, pool_candidates=[tiny_workspace_pool] - ) - bi_c = usmp_utils.BufferInfo( - name_hint="bi_c", size_bytes=10, pool_candidates=[tiny_workspace_pool] - ) - bi_a.set_conflicts([bi_b]) - bi_b.set_conflicts([bi_c]) - bi_c.set_conflicts([bi_a]) - buffer_info_arr = [bi_a, bi_b, bi_c] - fusmp_algo = tvm.get_global_func(f"tir.usmp.algo.greedy_by_size") - with pytest.raises( - tvm.TVMError, match="TVM USMP Error: the space available in the provided pools exceeded" - ): - buffer_pool_allocations = fusmp_algo(buffer_info_arr, 0) - - -@pytest.mark.parametrize("algorithm", ["greedy_by_size", "greedy_by_conflicts", "hill_climb"]) -def test_name_based_ordering(algorithm): - """This checks when the size and conlicts are same a stable result is generated""" - - def _test(): - target = Target("c") - global_workspace_pool = WorkspacePoolInfo( - "global_workspace", - [target], - ) - bi_a = usmp_utils.BufferInfo( - name_hint="bi_a", size_bytes=10, pool_candidates=[global_workspace_pool] - ) - bi_b = usmp_utils.BufferInfo( - name_hint="bi_b", size_bytes=10, pool_candidates=[global_workspace_pool] - ) - bi_c = usmp_utils.BufferInfo( - name_hint="bi_c", size_bytes=10, pool_candidates=[global_workspace_pool] - ) - bi_a.set_conflicts([bi_b, bi_c]) - bi_b.set_conflicts([bi_c, bi_a]) - bi_c.set_conflicts([bi_a, bi_b]) - - buffer_info_arr = [bi_a, bi_b, bi_c] - fusmp_algo = tvm.get_global_func(f"tir.usmp.algo.{algorithm}") - buffer_pool_allocations = fusmp_algo(buffer_info_arr, 0) - assert buffer_pool_allocations[bi_a].byte_offset == 20 - assert buffer_pool_allocations[bi_b].byte_offset == 10 - assert buffer_pool_allocations[bi_c].byte_offset == 0 - - # This is tested for several times to check stability - for x in range(0, 10): - _test() - - -@pytest.mark.parametrize( - ["algorithm", "workspace_size"], - [("greedy_by_size", 140), ("greedy_by_conflicts", 140), ("hill_climb", 140)], -) -def test_linear(algorithm, workspace_size): - """ - The test case here represent BufferInfo objects - that could get generated for a linear sequence - such as : - (Op A) - | - bi_a - | - (Op B) - | - bi_b - | - . - . - . - (Op F) - | - bi_f - """ - target = Target("c") - global_workspace_pool = WorkspacePoolInfo( - "global_workspace", - [target], - ) - bi_a = usmp_utils.BufferInfo( - name_hint="bi_a", size_bytes=10, pool_candidates=[global_workspace_pool] - ) - bi_b = usmp_utils.BufferInfo( - name_hint="bi_b", size_bytes=20, pool_candidates=[global_workspace_pool] - ) - bi_c = usmp_utils.BufferInfo( - name_hint="bi_c", size_bytes=100, pool_candidates=[global_workspace_pool] - ) - bi_d = usmp_utils.BufferInfo( - name_hint="bi_d", size_bytes=40, pool_candidates=[global_workspace_pool] - ) - bi_e = usmp_utils.BufferInfo( - name_hint="bi_e", size_bytes=50, pool_candidates=[global_workspace_pool] - ) - bi_f = usmp_utils.BufferInfo( - name_hint="bi_f", size_bytes=50, pool_candidates=[global_workspace_pool] - ) - - # Creating conflicts for a linear graph - bi_a.set_conflicts([bi_b]) - bi_b.set_conflicts([bi_a, bi_c]) - bi_c.set_conflicts([bi_b, bi_d]) - bi_d.set_conflicts([bi_c, bi_e]) - bi_e.set_conflicts([bi_d, bi_f]) - bi_f.set_conflicts([bi_e]) - - buffer_info_arr = [bi_a, bi_b, bi_c, bi_d, bi_e, bi_f] - fusmp_algo = tvm.get_global_func(f"tir.usmp.algo.{algorithm}") - buffer_pool_allocations = fusmp_algo(buffer_info_arr, 0) - _check_max_workspace_size(buffer_pool_allocations, global_workspace_pool, workspace_size) - - -@pytest.mark.parametrize( - ["algorithm", "workspace_size"], - [("greedy_by_size", 190), ("greedy_by_conflicts", 320), ("hill_climb", 190)], -) -def test_fanout(algorithm, workspace_size): - """ - The test case here represent BufferInfo objects - that could get generated for a fanout topology - such as : - (Op A) - | - bi_a --------- - | | - (Op B) (Op C) - | | - bi_b bi_c - | | - (Op D) (Op E) - | | - bi_d bi_e - | | - (Op F) ------ - | - bi_f - | - (Op G) - | - bi_g - """ - target = Target("c") - global_workspace_pool = WorkspacePoolInfo( - "global_workspace", - targets=[target], - ) - bi_a = usmp_utils.BufferInfo( - name_hint="bi_a", size_bytes=10, pool_candidates=[global_workspace_pool] - ) - bi_b = usmp_utils.BufferInfo( - name_hint="bi_b", size_bytes=20, pool_candidates=[global_workspace_pool] - ) - bi_c = usmp_utils.BufferInfo( - name_hint="bi_c", size_bytes=100, pool_candidates=[global_workspace_pool] - ) - bi_d = usmp_utils.BufferInfo( - name_hint="bi_d", size_bytes=40, pool_candidates=[global_workspace_pool] - ) - bi_e = usmp_utils.BufferInfo( - name_hint="bi_e", size_bytes=50, pool_candidates=[global_workspace_pool] - ) - bi_f = usmp_utils.BufferInfo( - name_hint="bi_f", size_bytes=60, pool_candidates=[global_workspace_pool] - ) - bi_g = usmp_utils.BufferInfo( - name_hint="bi_g", size_bytes=70, pool_candidates=[global_workspace_pool] - ) - - # Creating conflicts for a linear graph - bi_a.set_conflicts([bi_b, bi_c]) - bi_b.set_conflicts([bi_a, bi_c, bi_e]) - bi_c.set_conflicts([bi_e, bi_a, bi_b, bi_d]) - bi_d.set_conflicts([bi_b, bi_f, bi_c, bi_e]) - bi_e.set_conflicts([bi_c, bi_f, bi_b, bi_d]) - bi_f.set_conflicts([bi_d, bi_e, bi_f]) - bi_g.set_conflicts([bi_f]) - - buffer_info_arr = [bi_a, bi_b, bi_c, bi_d, bi_e, bi_f, bi_g] - fusmp_algo = tvm.get_global_func(f"tir.usmp.algo.{algorithm}") - buffer_pool_allocations = fusmp_algo(buffer_info_arr, 0) - _check_max_workspace_size(buffer_pool_allocations, global_workspace_pool, workspace_size) - - -# fmt: off -@tvm.script.ir_module -class MobilenetStructure: - @T.prim_func - def tvmgen_default_fused_cast_subtract(placeholder_2: T.handle, placeholder_3: T.handle, T_subtract: T.handle) -> None: - # function attr dict - T.func_attr({"global_symbol": "tvmgen_default_fused_cast_subtract", "tir.noalias": True}) - placeholder_4 = T.match_buffer(placeholder_2, [150528], dtype="uint8", elem_offset=0, align=64, offset_factor=1) - placeholder_5 = T.match_buffer(placeholder_3, [1], dtype="int16", elem_offset=0, align=64, offset_factor=1) - T_subtract_1 = T.match_buffer(T_subtract, [150528], dtype="int16", elem_offset=0, align=64, offset_factor=1) - # body - for ax0_ax1_fused_1 in T.serial(0, 224): - for ax2_1, ax3_inner_1 in T.grid(224, 3): - T_subtract_1[(((ax0_ax1_fused_1*672) + (ax2_1*3)) + ax3_inner_1)] = (T.cast(placeholder_4[(((ax0_ax1_fused_1*672) + (ax2_1*3)) + ax3_inner_1)], "int16") - placeholder_5[0]) - - @T.prim_func - def tvmgen_default_fused_nn_conv2d_add_fixed_point_multiply_clip_cast(placeholder_62: T.handle, placeholder_63: T.handle, placeholder_64: T.handle, T_cast_20: T.handle) -> None: - # function attr dict - T.func_attr({"global_symbol": "tvmgen_default_fused_nn_conv2d_add_fixed_point_multiply_clip_cast", "tir.noalias": True}) - placeholder_65 = T.match_buffer(placeholder_62, [150528], dtype="int16", elem_offset=0, align=64, offset_factor=1) - placeholder_66 = T.match_buffer(placeholder_63, [9408], dtype="int16", elem_offset=0, align=64, offset_factor=1) - placeholder_67 = T.match_buffer(placeholder_64, [64], dtype="int32", elem_offset=0, align=64, offset_factor=1) - T_cast_21 = T.match_buffer(T_cast_20, [802816], dtype="uint8", elem_offset=0, align=64, offset_factor=1) - # body - PaddedInput_7 = T.decl_buffer([157323], "int16") - for i0_i1_fused_7 in T.serial(0, 229): - for i2_7, i3_7 in T.grid(229, 3): - PaddedInput_7[(((i0_i1_fused_7*687) + (i2_7*3)) + i3_7)] = T.if_then_else(((((2 <= i0_i1_fused_7) and (i0_i1_fused_7 < 226)) and (2 <= i2_7)) and (i2_7 < 226)), placeholder_65[((((i0_i1_fused_7*672) + (i2_7*3)) + i3_7) - 1350)], T.int16(0), dtype="int16") - for ax0_ax1_fused_ax2_fused_7 in T.serial(0, 12544): - Conv2dOutput_7 = T.decl_buffer([64], "int32") - for ff_3 in T.serial(0, 64): - Conv2dOutput_7[ff_3] = 0 - for ry_2, rx_2, rc_7 in T.grid(7, 7, 3): - Conv2dOutput_7[ff_3] = (Conv2dOutput_7[ff_3] + (T.cast(PaddedInput_7[(((((T.floordiv(ax0_ax1_fused_ax2_fused_7, 112)*1374) + (ry_2*687)) + (T.floormod(ax0_ax1_fused_ax2_fused_7, 112)*6)) + (rx_2*3)) + rc_7)], "int32")*T.cast(placeholder_66[((((ry_2*1344) + (rx_2*192)) + (rc_7*64)) + ff_3)], "int32"))) - for ax3_inner_7 in T.serial(0, 64): - T_cast_21[((ax0_ax1_fused_ax2_fused_7*64) + ax3_inner_7)] = T.cast(T.max(T.min(T.q_multiply_shift((Conv2dOutput_7[ax3_inner_7] + placeholder_67[ax3_inner_7]), 1939887962, 31, -9, dtype="int32"), 255), 0), "uint8") - - @T.prim_func - def tvmgen_default_fused_nn_max_pool2d_cast(placeholder_28: T.handle, T_cast_6: T.handle) -> None: - # function attr dict - T.func_attr({"global_symbol": "tvmgen_default_fused_nn_max_pool2d_cast", "tir.noalias": True}) - placeholder_29 = T.match_buffer(placeholder_28, [802816], dtype="uint8", elem_offset=0, align=64, offset_factor=1) - T_cast_7 = T.match_buffer(T_cast_6, [200704], dtype="int16", elem_offset=0, align=64, offset_factor=1) - # body - tensor_2 = T.decl_buffer([200704], "uint8") - for ax0_ax1_fused_4 in T.serial(0, 56): - for ax2_4 in T.serial(0, 56): - for ax3_init in T.serial(0, 64): - tensor_2[(((ax0_ax1_fused_4*3584) + (ax2_4*64)) + ax3_init)] = T.uint8(0) - for rv0_rv1_fused_1, ax3_2 in T.grid(9, 64): - tensor_2[(((ax0_ax1_fused_4*3584) + (ax2_4*64)) + ax3_2)] = T.max(tensor_2[(((ax0_ax1_fused_4*3584) + (ax2_4*64)) + ax3_2)], T.if_then_else(((((ax0_ax1_fused_4*2) + T.floordiv(rv0_rv1_fused_1, 3)) < 112) and (((ax2_4*2) + T.floormod(rv0_rv1_fused_1, 3)) < 112)), placeholder_29[(((((ax0_ax1_fused_4*14336) + (T.floordiv(rv0_rv1_fused_1, 3)*7168)) + (ax2_4*128)) + (T.floormod(rv0_rv1_fused_1, 3)*64)) + ax3_2)], T.uint8(0), dtype="uint8")) - for ax0_ax1_fused_5 in T.serial(0, 56): - for ax2_5, ax3_3 in T.grid(56, 64): - T_cast_7[(((ax0_ax1_fused_5*3584) + (ax2_5*64)) + ax3_3)] = T.cast(tensor_2[(((ax0_ax1_fused_5*3584) + (ax2_5*64)) + ax3_3)], "int16") - - @T.prim_func - def run_model(input: T.handle, output: T.handle) -> None: - # function attr dict - T.func_attr({"global_symbol": "tvmgen_default_run_model", "runner_function": True}) - # body - T.attr("default", "device_id", 0) - T.attr("default", "device_type", 1) - sid_9 = T.allocate([301056], "int8", "global") - sid_8 = T.allocate([802816], "int8", "global") - T.evaluate(T.call_extern("tvmgen_default_fused_cast_subtract", input, T.lookup_param("p0", dtype="handle"), sid_9, dtype="int32")) - T.evaluate(T.call_extern("tvmgen_default_fused_nn_conv2d_add_fixed_point_multiply_clip_cast", sid_9, T.lookup_param("p1", dtype="handle"), T.lookup_param("p2", dtype="handle"), sid_8, dtype="int32")) - T.evaluate(T.call_extern("tvmgen_default_fused_nn_max_pool2d_cast", sid_8, output, dtype="int32")) -# fmt: on - - -@pytest.mark.parametrize( - ["algorithm", "fast_memory_size", "slow_memory_size"], - [ - ("greedy_by_size", 200704, 1418528), - ("greedy_by_conflicts", 200704, 1418528), - ("hill_climb", 200704, 1117462), - ], -) -def test_mobilenet_subgraph(algorithm, fast_memory_size, slow_memory_size): - target = Target("c") - fast_memory_pool = WorkspacePoolInfo( - "fast_memory", - [target], - PoolInfoProperties(size_hint_bytes=200704), - ) - slow_memory_pool = WorkspacePoolInfo( - "slow_memory", - [target], - ) - tir_mod = MobilenetStructure - tir_mod = _assign_targets_to_primfuncs_irmodule(tir_mod, target) - tir_mod = _assign_poolinfos_to_allocates_in_irmodule( - tir_mod, [fast_memory_pool, slow_memory_pool] - ) - main_func = tir_mod["run_model"] - buffer_info_analysis = tvm.tir.usmp.analysis.extract_buffer_info(main_func, tir_mod) - assert buffer_info_analysis.memory_pressure == 1117718 - - fcreate_array_bi = tvm.get_global_func("tir.usmp.CreateArrayBufferInfo") - buffer_info_arr = fcreate_array_bi(buffer_info_analysis.buffer_info_stmts) - fusmp_algo = tvm.get_global_func(f"tir.usmp.algo.{algorithm}") - buffer_pool_allocations = fusmp_algo(buffer_info_arr, buffer_info_analysis.memory_pressure) - - buffer_info_map_names = dict() - for buf_info in buffer_info_arr: - buffer_info_map_names[buf_info.name_hint] = buf_info - - # check conflicts - _verify_conflicts("PaddedInput_7", ["sid_9", "sid_8", "Conv2dOutput_7"], buffer_info_map_names) - _verify_conflicts("tensor_2", ["sid_8"], buffer_info_map_names) - _verify_conflicts("sid_9", ["PaddedInput_7"], buffer_info_map_names) - _verify_conflicts( - "sid_8", ["PaddedInput_7", "Conv2dOutput_7", "tensor_2"], buffer_info_map_names - ) - _verify_conflicts("Conv2dOutput_7", ["sid_8", "PaddedInput_7"], buffer_info_map_names) - - _check_max_workspace_size(buffer_pool_allocations, slow_memory_pool, slow_memory_size) - _check_max_workspace_size(buffer_pool_allocations, fast_memory_pool, fast_memory_size) - - -# fmt: off -@tvm.script.ir_module -class ResnetStructure: - @T.prim_func - def tvmgen_default_fused_cast_subtract_fixed_point_multiply_add_clip_cast_cast(placeholder: T.handle, placeholder_1: T.handle, T_cast: T.handle) -> None: - # function attr dict - T.func_attr({"global_symbol": "tvmgen_default_fused_cast_subtract_fixed_point_multiply_add_clip_cast_cast", "tir.noalias": True}) - placeholder_2 = T.match_buffer(placeholder, [360000], dtype="uint8") - placeholder_3 = T.match_buffer(placeholder_1, [64], dtype="int32") - T_cast_1 = T.match_buffer(T_cast, [360000], dtype="int16") - # body - for ax0_ax1_fused, ax2, ax3_outer, ax3_inner in T.grid(75, 75, 4, 16): - T_cast_1[ax0_ax1_fused * 4800 + ax2 * 64 + ax3_outer * 16 + ax3_inner] = T.cast(T.cast(T.max(T.min(T.q_multiply_shift(T.cast(placeholder_2[ax0_ax1_fused * 4800 + ax2 * 64 + ax3_outer * 16 + ax3_inner], "int32") - 94, 1843157232, 31, 1, dtype="int32") + placeholder_3[ax3_outer * 16 + ax3_inner], 255), 0), "uint8"), "int16") - - @T.prim_func - def tvmgen_default_fused_nn_conv2d_add_fixed_point_multiply_clip_cast_cast_1(placeholder_10: T.handle, placeholder_11: T.handle, placeholder_12: T.handle, T_cast_4: T.handle) -> None: - # function attr dict - T.func_attr({"global_symbol": "tvmgen_default_fused_nn_conv2d_add_fixed_point_multiply_clip_cast_cast_1", "tir.noalias": True}) - placeholder_13 = T.match_buffer(placeholder_10, [360000], dtype="int16") - placeholder_14 = T.match_buffer(placeholder_11, [36864], dtype="int16") - placeholder_15 = T.match_buffer(placeholder_12, [64], dtype="int32") - T_cast_5 = T.match_buffer(T_cast_4, [360000], dtype="int16") - # body - PaddedInput_1 = T.decl_buffer([379456], "int16") - for i0_i1_fused_1, i2_1, i3_1 in T.grid(77, 77, 64): - PaddedInput_1[i0_i1_fused_1 * 4928 + i2_1 * 64 + i3_1] = T.if_then_else(1 <= i0_i1_fused_1 and i0_i1_fused_1 < 76 and 1 <= i2_1 and i2_1 < 76, placeholder_13[i0_i1_fused_1 * 4800 + i2_1 * 64 + i3_1 - 4864], T.int16(0), dtype="int16") - for ax0_ax1_fused_ax2_fused_1 in T.serial(0, 5625): - Conv2dOutput_1 = T.decl_buffer([64], "int32") - for ff_1 in T.serial(0, 64): - Conv2dOutput_1[ff_1] = 0 - for ry, rx, rc_1 in T.grid(3, 3, 64): - Conv2dOutput_1[ff_1] = Conv2dOutput_1[ff_1] + T.cast(PaddedInput_1[T.floordiv(ax0_ax1_fused_ax2_fused_1, 75) * 4928 + ry * 4928 + rx * 64 + T.floormod(ax0_ax1_fused_ax2_fused_1, 75) * 64 + rc_1], "int32") * T.cast(placeholder_14[ry * 12288 + rx * 4096 + rc_1 * 64 + ff_1], "int32") - for ax3_inner_2 in T.serial(0, 64): - T_cast_5[ax0_ax1_fused_ax2_fused_1 * 64 + ax3_inner_2] = T.cast(T.cast(T.max(T.min(T.q_multiply_shift(Conv2dOutput_1[ax3_inner_2] + placeholder_15[ax3_inner_2], 1608879842, 31, -7, dtype="int32"), 255), 0), "uint8"), "int16") - - @T.prim_func - def tvmgen_default_fused_nn_conv2d_add_fixed_point_multiply_add_clip_cast_cast_subtract_fixed_point_15934180698220515269_(placeholder_16: T.handle, placeholder_17: T.handle, placeholder_18: T.handle, T_add: T.handle) -> None: - # function attr dict - T.func_attr({"global_symbol": "tvmgen_default_fused_nn_conv2d_add_fixed_point_multiply_add_clip_cast_cast_subtract_fixed_point_15934180698220515269_", "tir.noalias": True}) - placeholder_19 = T.match_buffer(placeholder_16, [360000], dtype="int16") - placeholder_20 = T.match_buffer(placeholder_17, [16384], dtype="int16") - placeholder_21 = T.match_buffer(placeholder_18, [256], dtype="int32") - T_add_1 = T.match_buffer(T_add, [1440000], dtype="int32") - # body - PaddedInput_2 = T.decl_buffer([360000], "int16") - for i0_i1_fused_2, i2_2, i3_2 in T.grid(75, 75, 64): - PaddedInput_2[i0_i1_fused_2 * 4800 + i2_2 * 64 + i3_2] = placeholder_19[i0_i1_fused_2 * 4800 + i2_2 * 64 + i3_2] - for ax0_ax1_fused_ax2_fused_2 in T.serial(0, 5625): - Conv2dOutput_2 = T.decl_buffer([64], "int32") - for ax3_outer_1 in T.serial(0, 4): - for ff_2 in T.serial(0, 64): - Conv2dOutput_2[ff_2] = 0 - for rc_2 in T.serial(0, 64): - Conv2dOutput_2[ff_2] = Conv2dOutput_2[ff_2] + T.cast(PaddedInput_2[ax0_ax1_fused_ax2_fused_2 * 64 + rc_2], "int32") * T.cast(placeholder_20[rc_2 * 256 + ax3_outer_1 * 64 + ff_2], "int32") - for ax3_inner_3 in T.serial(0, 64): - T_add_1[ax0_ax1_fused_ax2_fused_2 * 256 + ax3_outer_1 * 64 + ax3_inner_3] = T.q_multiply_shift(T.cast(T.cast(T.max(T.min(T.q_multiply_shift(Conv2dOutput_2[ax3_inner_3] + placeholder_21[ax3_outer_1 * 64 + ax3_inner_3], 1711626602, 31, -8, dtype="int32") + 132, 255), 0), "uint8"), "int32") - 132, 2094289803, 31, -2, dtype="int32") + 136 - - @T.prim_func - def tvmgen_default_fused_nn_conv2d_add_fixed_point_multiply_add_clip_cast_cast_subtract_fixed_point_4200876283395191415_(placeholder_22: T.handle, placeholder_23: T.handle, placeholder_24: T.handle, placeholder_25: T.handle, T_cast_6: T.handle) -> None: - # function attr dict - T.func_attr({"global_symbol": "tvmgen_default_fused_nn_conv2d_add_fixed_point_multiply_add_clip_cast_cast_subtract_fixed_point_4200876283395191415_", "tir.noalias": True}) - placeholder_29 = T.match_buffer(placeholder_22, [360000], dtype="int16") - placeholder_27 = T.match_buffer(placeholder_23, [16384], dtype="int16") - placeholder_26 = T.match_buffer(placeholder_24, [256], dtype="int32") - placeholder_28 = T.match_buffer(placeholder_25, [1440000], dtype="int32") - T_cast_7 = T.match_buffer(T_cast_6, [1440000], dtype="uint8") - # body - PaddedInput_3 = T.decl_buffer([360000], "int16") - for i0_i1_fused_3, i2_3, i3_3 in T.grid(75, 75, 64): - PaddedInput_3[i0_i1_fused_3 * 4800 + i2_3 * 64 + i3_3] = placeholder_29[i0_i1_fused_3 * 4800 + i2_3 * 64 + i3_3] - for ax0_ax1_fused_ax2_fused_3 in T.serial(0, 5625): - Conv2dOutput_3 = T.decl_buffer([64], "int32") - for ax3_outer_2 in T.serial(0, 4): - for ff_3 in T.serial(0, 64): - Conv2dOutput_3[ff_3] = 0 - for rc_3 in T.serial(0, 64): - Conv2dOutput_3[ff_3] = Conv2dOutput_3[ff_3] + T.cast(PaddedInput_3[ax0_ax1_fused_ax2_fused_3 * 64 + rc_3], "int32") * T.cast(placeholder_27[rc_3 * 256 + ax3_outer_2 * 64 + ff_3], "int32") - for ax3_inner_4 in T.serial(0, 64): - T_cast_7[ax0_ax1_fused_ax2_fused_3 * 256 + ax3_outer_2 * 64 + ax3_inner_4] = T.cast(T.max(T.min(T.q_multiply_shift(T.cast(T.cast(T.max(T.min(T.q_multiply_shift(Conv2dOutput_3[ax3_inner_4] + placeholder_26[ax3_outer_2 * 64 + ax3_inner_4], 1343014664, 31, -8, dtype="int32") + 136, 255), 0), "uint8"), "int32") - 136, 1073903788, 31, 1, dtype="int32") + placeholder_28[ax0_ax1_fused_ax2_fused_3 * 256 + ax3_outer_2 * 64 + ax3_inner_4], 255), 0), "uint8") - - @T.prim_func - def tvmgen_default_run_model(input: T.handle, output: T.handle) -> None: - # function attr dict - T.func_attr({"global_symbol": "tvmgen_default_run_model", "runner_function": True}) - # body - T.attr("default", "device_id", 0) - T.attr("default", "device_type", 1) - sid_2 = T.allocate([720000], "int8", "global") - sid_6 = T.allocate([5760000], "int8", "global") - sid_7 = T.allocate([720000], "int8", "global") - sid_8 = T.allocate([720000], "int8", "global") - T.evaluate(T.call_extern("tvmgen_default_fused_cast_subtract_fixed_point_multiply_add_clip_cast_cast", input, T.lookup_param("p0", dtype="handle"), sid_2, dtype="int32")) - T.evaluate(T.call_extern("tvmgen_default_fused_nn_conv2d_add_fixed_point_multiply_clip_cast_cast", sid_2, T.lookup_param("p3", dtype="handle"), T.lookup_param("p4", dtype="handle"), sid_8, dtype="int32")) - T.evaluate(T.call_extern("tvmgen_default_fused_nn_conv2d_add_fixed_point_multiply_clip_cast_cast_1", sid_8, T.lookup_param("p5", dtype="handle"), T.lookup_param("p6", dtype="handle"), sid_7, dtype="int32")) - T.evaluate(T.call_extern("tvmgen_default_fused_nn_conv2d_add_fixed_point_multiply_add_clip_cast_cast_subtract_fixed_point_15934180698220515269_", sid_7, T.lookup_param("p7", dtype="handle"), T.lookup_param("p8", dtype="handle"), sid_6, dtype="int32")) - T.evaluate(T.call_extern("tvmgen_default_fused_nn_conv2d_add_fixed_point_multiply_add_clip_cast_cast_subtract_fixed_point_4200876283395191415_", sid_2, T.lookup_param("p1", dtype="handle"), T.lookup_param("p2", dtype="handle"), sid_6, output, dtype="int32")) - - @T.prim_func - def tvmgen_default_fused_nn_conv2d_add_fixed_point_multiply_clip_cast_cast(placeholder_4: T.handle, placeholder_5: T.handle, placeholder_6: T.handle, T_cast_2: T.handle) -> None: - # function attr dict - T.func_attr({"global_symbol": "tvmgen_default_fused_nn_conv2d_add_fixed_point_multiply_clip_cast_cast", "tir.noalias": True}) - placeholder_7 = T.match_buffer(placeholder_4, [360000], dtype="int16") - placeholder_8 = T.match_buffer(placeholder_5, [4096], dtype="int16") - placeholder_9 = T.match_buffer(placeholder_6, [64], dtype="int32") - T_cast_3 = T.match_buffer(T_cast_2, [360000], dtype="int16") - # body - PaddedInput = T.decl_buffer([360000], "int16") - for i0_i1_fused, i2, i3 in T.grid(75, 75, 64): - PaddedInput[i0_i1_fused * 4800 + i2 * 64 + i3] = placeholder_7[i0_i1_fused * 4800 + i2 * 64 + i3] - for ax0_ax1_fused_ax2_fused in T.serial(0, 5625): - Conv2dOutput = T.decl_buffer([64], "int32") - for ff in T.serial(0, 64): - Conv2dOutput[ff] = 0 - for rc in T.serial(0, 64): - Conv2dOutput[ff] = Conv2dOutput[ff] + T.cast(PaddedInput[ax0_ax1_fused_ax2_fused * 64 + rc], "int32") * T.cast(placeholder_8[rc * 64 + ff], "int32") - for ax3_inner_1 in T.serial(0, 64): - T_cast_3[ax0_ax1_fused_ax2_fused * 64 + ax3_inner_1] = T.cast(T.cast(T.max(T.min(T.q_multiply_shift(Conv2dOutput[ax3_inner_1] + placeholder_9[ax3_inner_1], 1843106743, 31, -6, dtype="int32"), 255), 0), "uint8"), "int16") -# fmt: on - - -@pytest.mark.parametrize( - ["algorithm", "workspace_size"], - [("greedy_by_size", 7920256), ("greedy_by_conflicts", 7200256), ("hill_climb", 7200256)], -) -def test_resnet_subgraph(algorithm, workspace_size): - target = Target("c") - global_workspace_pool = WorkspacePoolInfo( - "global_workspace", - [target], - ) - tir_mod = ResnetStructure - tir_mod = _assign_targets_to_primfuncs_irmodule(tir_mod, target) - tir_mod = _assign_poolinfos_to_allocates_in_irmodule(tir_mod, [global_workspace_pool]) - main_func = tir_mod["tvmgen_default_run_model"] - buffer_info_analysis = tvm.tir.usmp.analysis.extract_buffer_info(main_func, tir_mod) - assert buffer_info_analysis.memory_pressure == 7200256 - - fcreate_array_bi = tvm.get_global_func("tir.usmp.CreateArrayBufferInfo") - buffer_info_arr = fcreate_array_bi(buffer_info_analysis.buffer_info_stmts) - fusmp_algo = tvm.get_global_func(f"tir.usmp.algo.{algorithm}") - buffer_pool_allocations = fusmp_algo(buffer_info_arr, buffer_info_analysis.memory_pressure) - - buffer_info_map_names = dict() - for buf_info in buffer_info_arr: - buffer_info_map_names[buf_info.name_hint] = buf_info - - # check conflicts - _verify_conflicts( - "sid_7", - [ - "PaddedInput_1", - "sid_2", - "Conv2dOutput_1", - "PaddedInput_2", - ], - buffer_info_map_names, - ) - _verify_conflicts( - "Conv2dOutput_3", - [ - "PaddedInput_3", - "sid_6", - ], - buffer_info_map_names, - ) - _verify_conflicts( - "sid_6", - [ - "Conv2dOutput_2", - "PaddedInput_2", - "sid_2", - "PaddedInput_3", - "Conv2dOutput_3", - ], - buffer_info_map_names, - ) - _verify_conflicts( - "Conv2dOutput", - [ - "sid_8", - "sid_2", - "PaddedInput", - ], - buffer_info_map_names, - ) - _verify_conflicts( - "PaddedInput_3", - [ - "sid_6", - "sid_2", - "Conv2dOutput_3", - ], - buffer_info_map_names, - ) - _verify_conflicts( - "Conv2dOutput_2", - [ - "PaddedInput_2", - "sid_2", - "sid_6", - ], - buffer_info_map_names, - ) - _verify_conflicts( - "PaddedInput_1", - [ - "sid_8", - "sid_2", - "sid_7", - "Conv2dOutput_1", - ], - buffer_info_map_names, - ) - _verify_conflicts( - "Conv2dOutput_1", - [ - "sid_7", - "PaddedInput_1", - "sid_2", - ], - buffer_info_map_names, - ) - _verify_conflicts( - "PaddedInput", - [ - "sid_2", - "sid_8", - "Conv2dOutput", - ], - buffer_info_map_names, - ) - _verify_conflicts( - "sid_8", - [ - "PaddedInput", - "sid_2", - "Conv2dOutput", - "PaddedInput_1", - ], - buffer_info_map_names, - ) - _verify_conflicts( - "sid_2", - [ - "PaddedInput", - "sid_8", - "Conv2dOutput", - "PaddedInput_1", - "sid_7", - "Conv2dOutput_1", - "PaddedInput_2", - "Conv2dOutput_2", - "sid_6", - "PaddedInput_3", - ], - buffer_info_map_names, - ) - _verify_conflicts( - "PaddedInput_2", - [ - "sid_7", - "sid_2", - "Conv2dOutput_2", - "sid_6", - ], - buffer_info_map_names, - ) - - _check_max_workspace_size(buffer_pool_allocations, global_workspace_pool, workspace_size) diff --git a/tests/python/tir-usmp/test_tir_usmp_algo_hill_climb.py b/tests/python/tir-usmp/test_tir_usmp_algo_hill_climb.py deleted file mode 100644 index 6450673e71dd..000000000000 --- a/tests/python/tir-usmp/test_tir_usmp_algo_hill_climb.py +++ /dev/null @@ -1,404 +0,0 @@ -# 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 sys -import pytest -import random -import tvm -import tvm.testing -from tvm.tir.usmp.utils import BufferInfo -from tvm import WorkspacePoolInfo, PoolInfoProperties - - -def _check_max_workspace_size(buffer_pool_allocations, pool_info, size, tolerance=0): - """Helper to check maximum allocated memory size""" - max_workspace_size = 0 - for buffer_info, pool_allocation in buffer_pool_allocations.items(): - if pool_allocation.pool_info == pool_info: - size_candidate = pool_allocation.byte_offset + buffer_info.size_bytes - if size_candidate > max_workspace_size: - max_workspace_size = size_candidate - _diff = max_workspace_size.value - size - return ( - (max_workspace_size.value == size if tolerance == 0 else tolerance > 100 * _diff / size), - "'{}': expected {} got {}, diff {:0.2f}% ({} bytes)".format( - pool_info.pool_name, size, max_workspace_size, 100 * _diff / size, _diff - ), - ) - - -def _verify_conflicts(buffer_info, pool_allocation, buffer_info_map): - """Helper to check expected liveness conflicts""" - for conflict in buffer_info.conflicts: - conflict_pool_allocation = buffer_info_map[conflict] - - if conflict_pool_allocation.pool_info == pool_allocation.pool_info: - assert conflict_pool_allocation.byte_offset != pool_allocation.byte_offset - l2 = max( - conflict_pool_allocation.byte_offset + conflict.size_bytes, - pool_allocation.byte_offset + buffer_info.size_bytes, - ) - min(conflict_pool_allocation.byte_offset, pool_allocation.byte_offset) - assert ( - conflict.size_bytes + buffer_info.size_bytes <= l2 - ), 'Conflicting: \n"{} @{}"\n"{} @{}"'.format( - conflict, conflict_pool_allocation, buffer_info, pool_allocation - ) - - -def _verify_all_conflicts(buffer_pool_allocations): - """Helper to verify liveness conflicts""" - for buffer_info, pool_allocation in buffer_pool_allocations.items(): - _verify_conflicts(buffer_info, pool_allocation, buffer_pool_allocations) - - -def test_bounded( - random_len=150, - pools=[ - WorkspacePoolInfo("default", [], PoolInfoProperties(65535)), - WorkspacePoolInfo("slow", []), - ], -): - """Tests two pools, one is bounded and one is not limited""" - random.seed(0) - mem_range = [BufferInfo(str(i), random.randrange(1, 65535), pools) for i in range(random_len)] - for mr in mem_range: - pr = random.choice(mem_range) - while pr in (*mr.conflicts, mr): - pr = random.choice(mem_range) - - mr.set_conflicts([*mr.conflicts, pr]) - pr.set_conflicts([*pr.conflicts, mr]) - - fusmp_algo = tvm.get_global_func("tir.usmp.algo.hill_climb") - result_map = fusmp_algo(mem_range, 0) - _verify_all_conflicts(result_map) - - -def __test_data_alloc_max(): - """Test data""" - intervals = [ - (0, 159, 2048), - (0, 13, 7904), - (4, 35, 16), - (12, 17, 32768), - (16, 21, 32768), - ] - return intervals - - -def __test_data_deep_speech(): - """Test data""" - intervals = [ - (0, 159, 2048), - (0, 151, 2048), - (0, 13, 7904), - (2, 49, 16), - (4, 35, 16), - (6, 21, 16), - (12, 17, 32768), - (16, 21, 32768), - (20, 27, 32768), - (26, 31, 32768), - (30, 35, 32768), - (34, 41, 32768), - (40, 45, 32768), - (44, 49, 32768), - (48, 145, 32768), - (54, 59, 2048), - (58, 483, 4096), - (60, 65, 2048), - (64, 461, 4096), - (66, 71, 2048), - (70, 439, 4096), - (72, 77, 2048), - (76, 417, 4096), - (78, 83, 2048), - (82, 395, 4096), - (84, 89, 2048), - (88, 373, 4096), - (90, 95, 2048), - (94, 351, 4096), - (96, 101, 2048), - (100, 329, 4096), - (102, 107, 2048), - (106, 307, 4096), - (108, 113, 2048), - (112, 285, 4096), - (114, 119, 2048), - (118, 263, 4096), - (120, 125, 2048), - (124, 241, 4096), - (126, 131, 2048), - (130, 219, 4096), - (132, 137, 2048), - (136, 197, 4096), - (138, 143, 2048), - (142, 175, 4096), - (144, 149, 2048), - (148, 153, 4096), - (152, 163, 8192), - (154, 171, 2048), - (156, 181, 2048), - (160, 167, 2048), - (162, 165, 2048), - (168, 171, 2048), - (170, 509, 2048), - (174, 185, 8192), - (176, 193, 2048), - (178, 203, 2048), - (182, 189, 2048), - (184, 187, 2048), - (190, 193, 2048), - (192, 511, 2048), - (196, 207, 8192), - (198, 215, 2048), - (200, 225, 2048), - (204, 211, 2048), - (206, 209, 2048), - (212, 215, 2048), - (214, 513, 2048), - (218, 229, 8192), - (220, 237, 2048), - (222, 247, 2048), - (226, 233, 2048), - (228, 231, 2048), - (234, 237, 2048), - (236, 515, 2048), - (240, 251, 8192), - (242, 259, 2048), - (244, 269, 2048), - (248, 255, 2048), - (250, 253, 2048), - (256, 259, 2048), - (258, 517, 2048), - (262, 273, 8192), - (264, 281, 2048), - (266, 291, 2048), - (270, 277, 2048), - (272, 275, 2048), - (278, 281, 2048), - (280, 519, 2048), - (284, 295, 8192), - (286, 303, 2048), - (288, 313, 2048), - (292, 299, 2048), - (294, 297, 2048), - (300, 303, 2048), - (302, 521, 2048), - (306, 317, 8192), - (308, 325, 2048), - (310, 335, 2048), - (314, 321, 2048), - (316, 319, 2048), - (322, 325, 2048), - (324, 523, 2048), - (328, 339, 8192), - (330, 347, 2048), - (332, 357, 2048), - (336, 343, 2048), - (338, 341, 2048), - (344, 347, 2048), - (346, 525, 2048), - (350, 361, 8192), - (352, 369, 2048), - (354, 379, 2048), - (358, 365, 2048), - (360, 363, 2048), - (366, 369, 2048), - (368, 527, 2048), - (372, 383, 8192), - (374, 391, 2048), - (376, 401, 2048), - (380, 387, 2048), - (382, 385, 2048), - (388, 391, 2048), - (390, 529, 2048), - (394, 405, 8192), - (396, 413, 2048), - (398, 423, 2048), - (402, 409, 2048), - (404, 407, 2048), - (410, 413, 2048), - (412, 531, 2048), - (416, 427, 8192), - (418, 435, 2048), - (420, 445, 2048), - (424, 431, 2048), - (426, 429, 2048), - (432, 435, 2048), - (434, 533, 2048), - (438, 449, 8192), - (440, 457, 2048), - (442, 467, 2048), - (446, 453, 2048), - (448, 451, 2048), - (454, 457, 2048), - (456, 535, 2048), - (460, 471, 8192), - (462, 479, 2048), - (464, 489, 2048), - (468, 475, 2048), - (470, 473, 2048), - (476, 479, 2048), - (478, 537, 2048), - (482, 493, 8192), - (484, 501, 2048), - (486, 497, 2048), - (490, 497, 2048), - (492, 495, 2048), - (496, 626, 2048), - (498, 501, 2048), - (500, 626, 2048), - (504, 549, 16), - (508, 543, 32768), - (542, 549, 32768), - (548, 555, 32768), - (554, 563, 464), - (560, 563, 256), - (562, 617, 2048), - (564, 567, 1856), - (566, 573, 1024), - (568, 619, 1024), - (570, 573, 1024), - (572, 577, 1024), - (576, 579, 1024), - (578, 605, 1024), - (580, 593, 1024), - (584, 587, 1024), - (586, 603, 1024), - (594, 597, 1024), - (596, 613, 1024), - (604, 607, 1024), - (606, 617, 1024), - (616, 621, 2048), - (618, 621, 1024), - (620, 626, 464), - ] - return intervals - - -def __test_data_five(): - """Test data""" - return [ - (4, 5, 95), - (1, 4, 52135), - (3, 4, 12136), - (3, 5, 62099), - (4, 5, 50458), - ] - - -def __test_data_simple(): - """Test data""" - return [ - (0, 23, 131072), # 0 - (4, 5, 65568), # 1 - (4, 9, 8192), # 2 - (8, 30, 15360), # 3 - (10, 11, 65568), # 4 - (10, 15, 4096), # 5 - (16, 17, 65552), # 6 - (16, 21, 2048), # 7 - (22, 23, 32784), # 8 - (22, 27, 1024), # 9 - ] - - -def find_maximum_from_intervals(intervals): - """Expected list of intervals of (start, end, size)""" - sorted_list = sorted(intervals, key=lambda _: _[0]) - max_mem = 0 - for t in range(sorted_list[0][0], sorted_list[-1][1] + 1): - max_mem = max( - max_mem, sum([size for (start, end, size) in sorted_list if t >= start and t <= end]) - ) - return max_mem - - -@pytest.mark.parametrize( - "intervals", - [__test_data_alloc_max(), __test_data_simple(), __test_data_deep_speech(), __test_data_five()], -) -def test_intervals(intervals): - """Tests supplied intervals""" - random.seed(0) - result = run_intervals(intervals, 5) - assert result["tir.usmp.algo.hill_climb"] == True, f" {result}" - - -def generate_range(sz, max_segment_sz=65535): - """Helper func to generate list of size sz of ranges of random size max_segment_sz""" - for i in range(0, sz): - start = random.randrange(i, sz) - stop = random.randrange(start + 1, start + 2 + ((sz - start) // 2)) - assert stop - start > 0 - yield (start, stop, random.randrange(1, max_segment_sz)) - - -def test_random_intervals(interval_len=16): - """Tests randomly generated interval of length interval_len""" - random.seed(0) - intervals = list(generate_range(interval_len)) - return run_intervals(intervals) - - -def run_intervals(intervals, tolerance=0): - """Helper to run intervals""" - expected_mem = find_maximum_from_intervals(intervals) - pools = [WorkspacePoolInfo("default", [])] - buffers = [] - # populate - for i, (start, stop, size) in enumerate(intervals): - buf = BufferInfo(str(i), size, pools) - # buf.set_pool_candidates( ["default"] ) - buffers.append(buf) - - # intersect - for i, (i_start, i_stop, _) in enumerate(intervals): - conflicts = set() - for j, (j_start, j_stop, _) in enumerate(intervals): - start = min(i_start, j_start) - stop = max(i_stop, j_stop) - i_dur = i_stop - i_start + 1 - j_dur = j_stop - j_start + 1 - - if i != j and (stop - start + 1 < i_dur + j_dur): - conflicts.add(buffers[j]) - - buffers[i].set_conflicts([c for c in sorted(conflicts, key=lambda c: c.name_hint)]) - - result = {} - for (alg, params) in [ - ("tir.usmp.algo.hill_climb", (expected_mem,)), - ("tir.usmp.algo.greedy_by_size", (expected_mem,)), - ]: - fusmp_algo = tvm.get_global_func(alg) - print("\n", "started", alg) - buffer_info_arr = fusmp_algo(buffers, *params) - print() - - _verify_all_conflicts(buffer_info_arr) - result[alg], msg = _check_max_workspace_size( - buffer_info_arr, pools[0], expected_mem, tolerance - ) - if not result[alg]: - print(alg, msg) - - return result - - -if __name__ == "__main__": - tvm.testing.main() diff --git a/tests/python/tir-usmp/test_tir_usmp_analysis_extract_bufferinfo.py b/tests/python/tir-usmp/test_tir_usmp_analysis_extract_bufferinfo.py deleted file mode 100644 index f8da0ef9f42d..000000000000 --- a/tests/python/tir-usmp/test_tir_usmp_analysis_extract_bufferinfo.py +++ /dev/null @@ -1,1690 +0,0 @@ -# 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 pytest -import sys - -import tvm -import tvm.testing -from tvm import tir, script -from tvm.ir import Range -from tvm.script import tir as T -from tvm.tir import stmt_functor -from tvm.tir import PrimFunc -from tvm.tir.usmp import utils as usmp_utils -from tvm.target import Target -from tvm import WorkspacePoolInfo, ConstantPoolInfo - - -def _replace_stmt_with_buf_var_names(buffer_info_map): - """helper to replace tir.allocates with buffer names""" - new_buffer_info_map = dict() - for k, v in buffer_info_map.items(): - new_buffer_info_map[k.name_hint] = k - return new_buffer_info_map - - -def _verify_conflicts(main_buf_name, conflicting_buf_names, buffer_info_map): - """helper to check expected liveness conflicts""" - buf_info = buffer_info_map[main_buf_name] - for conflict in buf_info.conflicts: - assert conflict.name_hint in conflicting_buf_names - - -def _get_allocates(primfunc): - """helper to extract all allocate nodes by name""" - allocates = dict() - - def get_allocate(stmt): - if isinstance(stmt, tvm.tir.Allocate): - allocates[str(stmt.buffer_var.name)] = stmt - - stmt_functor.post_order_visit(primfunc.body, get_allocate) - return allocates - - -def _assign_poolinfos_to_allocates_in_primfunc(primfunc, pool_infos, constant_pool_infos): - """helper to assing poolinfos to allocate nodes in a tir.PrimFunc""" - - def set_poolinfos(stmt): - if isinstance(stmt, tvm.tir.Allocate): - return tvm.tir.Allocate( - buffer_var=stmt.buffer_var, - dtype=stmt.dtype, - extents=stmt.extents, - condition=stmt.condition, - body=stmt.body, - annotations={tvm.tir.usmp.utils.CANDIDATE_MEMORY_POOL_ATTR: pool_infos}, - ) - elif isinstance(stmt, tvm.tir.AllocateConst): - return tvm.tir.AllocateConst( - buffer_var=stmt.buffer_var, - dtype=stmt.dtype, - extents=stmt.extents, - data_or_idx=stmt.data, - body=stmt.body, - annotations={tvm.tir.usmp.utils.CANDIDATE_MEMORY_POOL_ATTR: constant_pool_infos}, - ) - - return primfunc.with_body(stmt_functor.ir_transform(primfunc.body, None, set_poolinfos)) - - -def _assign_poolinfos_to_allocates_in_irmodule(mod, pool_infos, constant_pool_infos=None): - """helper to assign poolinfos to allocate nodes in a IRModule""" - ret = tvm.IRModule() - for global_var, basefunc in mod.functions.items(): - if isinstance(basefunc, tvm.tir.PrimFunc): - ret[global_var] = _assign_poolinfos_to_allocates_in_primfunc( - basefunc, pool_infos, constant_pool_infos - ) - return ret - - -def _assign_targets_to_primfuncs_irmodule(mod, target): - """helper to assign target for PrimFunc in a IRModule""" - ret = tvm.IRModule() - for global_var, basefunc in mod.functions.items(): - if isinstance(basefunc, tvm.tir.PrimFunc): - ret[global_var] = basefunc.with_attr("target", target) - return ret - - -# These are test IRModules that contains varied topologies of operator graphs -# that includes a main TIR function that includes call to such operators. - -# fmt: off -@tvm.script.ir_module -class LinearStructure: - @T.prim_func - def tvmgen_default_fused_cast_subtract(placeholder_2: T.handle, placeholder_3: T.handle, T_subtract: T.handle) -> None: - # function attr dict - T.func_attr({"global_symbol": "tvmgen_default_fused_cast_subtract", "tir.noalias": True}) - placeholder_4 = T.match_buffer(placeholder_2, [150528], dtype="uint8", elem_offset=0, align=64, offset_factor=1) - placeholder_5 = T.match_buffer(placeholder_3, [1], dtype="int16", elem_offset=0, align=64, offset_factor=1) - T_subtract_1 = T.match_buffer(T_subtract, [452], dtype="int16", elem_offset=0, align=64, offset_factor=1) - # body - for ax0_ax1_fused_1 in T.serial(0, 224): - for ax2_1, ax3_inner_1 in T.grid(224, 3): - T_subtract_1[(((ax0_ax1_fused_1*672) + (ax2_1*3)) + ax3_inner_1)] = (T.cast(placeholder_4[(((ax0_ax1_fused_1*672) + (ax2_1*3)) + ax3_inner_1)], "int16") - placeholder_5[0]) - - @T.prim_func - def tvmgen_default_fused_nn_conv2d_add_fixed_point_multiply_clip_cast(placeholder_62: T.handle, placeholder_63: T.handle, placeholder_64: T.handle, T_cast_20: T.handle) -> None: - # function attr dict - T.func_attr({"global_symbol": "tvmgen_default_fused_nn_conv2d_add_fixed_point_multiply_clip_cast", "tir.noalias": True}) - placeholder_65 = T.match_buffer(placeholder_62, [150528], dtype="int16", elem_offset=0, align=64, offset_factor=1) - placeholder_66 = T.match_buffer(placeholder_63, [9408], dtype="int16", elem_offset=0, align=64, offset_factor=1) - placeholder_67 = T.match_buffer(placeholder_64, [64], dtype="int32", elem_offset=0, align=64, offset_factor=1) - T_cast_21 = T.match_buffer(T_cast_20, [289], dtype="uint8", elem_offset=0, align=64, offset_factor=1) - # body - PaddedInput_7 = T.decl_buffer([157323], "int16") - for i0_i1_fused_7 in T.serial(0, 229): - for i2_7, i3_7 in T.grid(229, 3): - PaddedInput_7[(((i0_i1_fused_7*687) + (i2_7*3)) + i3_7)] = T.if_then_else(((((2 <= i0_i1_fused_7) and (i0_i1_fused_7 < 226)) and (2 <= i2_7)) and (i2_7 < 226)), placeholder_65[((((i0_i1_fused_7*672) + (i2_7*3)) + i3_7) - 1350)], T.int16(0), dtype="int16") - for ax0_ax1_fused_ax2_fused_7 in T.serial(0, 12544): - Conv2dOutput_7 = T.decl_buffer([64], "int32") - for ff_3 in T.serial(0, 64): - Conv2dOutput_7[ff_3] = 0 - for ry_2, rx_2, rc_7 in T.grid(7, 7, 3): - Conv2dOutput_7[ff_3] = (Conv2dOutput_7[ff_3] + (T.cast(PaddedInput_7[(((((T.floordiv(ax0_ax1_fused_ax2_fused_7, 112)*1374) + (ry_2*687)) + (T.floormod(ax0_ax1_fused_ax2_fused_7, 112)*6)) + (rx_2*3)) + rc_7)], "int32")*T.cast(placeholder_66[((((ry_2*1344) + (rx_2*192)) + (rc_7*64)) + ff_3)], "int32"))) - for ax3_inner_7 in T.serial(0, 64): - T_cast_21[((ax0_ax1_fused_ax2_fused_7*64) + ax3_inner_7)] = T.cast(T.max(T.min(T.q_multiply_shift((Conv2dOutput_7[ax3_inner_7] + placeholder_67[ax3_inner_7]), 1939887962, 31, -9, dtype="int32"), 255), 0), "uint8") - - @T.prim_func - def tvmgen_default_fused_nn_max_pool2d_cast(placeholder_28: T.handle, T_cast_6: T.handle) -> None: - # function attr dict - T.func_attr({"global_symbol": "tvmgen_default_fused_nn_max_pool2d_cast", "tir.noalias": True}) - placeholder_29 = T.match_buffer(placeholder_28, [802816], dtype="uint8", elem_offset=0, align=64, offset_factor=1) - T_cast_7 = T.match_buffer(T_cast_6, [177], dtype="int16", elem_offset=0, align=64, offset_factor=1) - # body - tensor_2 = T.decl_buffer([200704], "uint8") - for ax0_ax1_fused_4 in T.serial(0, 56): - for ax2_4 in T.serial(0, 56): - for ax3_init in T.serial(0, 64): - tensor_2[(((ax0_ax1_fused_4*3584) + (ax2_4*64)) + ax3_init)] = T.uint8(0) - for rv0_rv1_fused_1, ax3_2 in T.grid(9, 64): - tensor_2[(((ax0_ax1_fused_4*3584) + (ax2_4*64)) + ax3_2)] = T.max(tensor_2[(((ax0_ax1_fused_4*3584) + (ax2_4*64)) + ax3_2)], T.if_then_else(((((ax0_ax1_fused_4*2) + T.floordiv(rv0_rv1_fused_1, 3)) < 112) and (((ax2_4*2) + T.floormod(rv0_rv1_fused_1, 3)) < 112)), placeholder_29[(((((ax0_ax1_fused_4*14336) + (T.floordiv(rv0_rv1_fused_1, 3)*7168)) + (ax2_4*128)) + (T.floormod(rv0_rv1_fused_1, 3)*64)) + ax3_2)], T.uint8(0), dtype="uint8")) - for ax0_ax1_fused_5 in T.serial(0, 56): - for ax2_5, ax3_3 in T.grid(56, 64): - T_cast_7[(((ax0_ax1_fused_5*3584) + (ax2_5*64)) + ax3_3)] = T.cast(tensor_2[(((ax0_ax1_fused_5*3584) + (ax2_5*64)) + ax3_3)], "int16") - - @T.prim_func - def run_model(input: T.handle, output: T.handle) -> None: - # function attr dict - T.func_attr({"global_symbol": "tvmgen_default_run_model", "runner_function": True}) - # body - T.attr("default", "device_id", 0) - T.attr("default", "device_type", 1) - sid_9 = T.allocate([301056], "int8", "global") - sid_8 = T.allocate([802816], "int8", "global") - T.evaluate(T.call_extern("tvmgen_default_fused_cast_subtract", input, T.lookup_param("p0", dtype="handle"), sid_9, dtype="int32")) - T.evaluate(T.call_extern("tvmgen_default_fused_nn_conv2d_add_fixed_point_multiply_clip_cast", sid_9, T.lookup_param("p1", dtype="handle"), T.lookup_param("p2", dtype="handle"), sid_8, dtype="int32")) - T.evaluate(T.call_extern("tvmgen_default_fused_nn_max_pool2d_cast", sid_8, output, dtype="int32")) -# fmt: on - - -def test_linear(): - target = Target("c") - fast_memory_pool = WorkspacePoolInfo(pool_name="fast_memory", targets=[target]) - slow_memory_pool = WorkspacePoolInfo(pool_name="slow_memory", targets=[target]) - tir_mod = LinearStructure - tir_mod = _assign_targets_to_primfuncs_irmodule(tir_mod, target) - tir_mod = _assign_poolinfos_to_allocates_in_irmodule( - tir_mod, [fast_memory_pool, slow_memory_pool] - ) - buffer_info_analysis = tvm.tir.usmp.analysis.extract_buffer_info(tir_mod["run_model"], tir_mod) - assert buffer_info_analysis.memory_pressure == 1117718 - buffer_info_map = _replace_stmt_with_buf_var_names(buffer_info_analysis.buffer_info_stmts) - - # check conflicts - _verify_conflicts("PaddedInput_7", ["sid_9", "sid_8", "Conv2dOutput_7"], buffer_info_map) - _verify_conflicts("tensor_2", ["sid_8"], buffer_info_map) - _verify_conflicts("sid_9", ["PaddedInput_7"], buffer_info_map) - _verify_conflicts("sid_8", ["PaddedInput_7", "Conv2dOutput_7", "tensor_2"], buffer_info_map) - _verify_conflicts("Conv2dOutput_7", ["sid_8", "PaddedInput_7"], buffer_info_map) - - # check sizes - assert buffer_info_map["sid_8"].size_bytes == 802816 - assert buffer_info_map["Conv2dOutput_7"].size_bytes == 256 - assert buffer_info_map["PaddedInput_7"].size_bytes == 314646 - assert buffer_info_map["tensor_2"].size_bytes == 200704 - assert buffer_info_map["sid_9"].size_bytes == 301056 - - # check_pool_candidates - assert [ - pool_info.pool_name for pool_info in list(buffer_info_map["sid_8"].pool_candidates) - ] == ["fast_memory", "slow_memory"] - - -# fmt: off -@tvm.script.ir_module -class ParallelSerialMixedForLoops: - @T.prim_func - def tvmgen_default_fused_nn_conv2d_add_fixed_point_multiply_clip_cast_1(placeholder_68: T.handle, placeholder_69: T.handle, placeholder_70: T.handle, T_cast_22: T.handle) -> None: - # function attr dict - T.func_attr({"global_symbol": "tvmgen_default_fused_nn_conv2d_add_fixed_point_multiply_clip_cast_1", "tir.noalias": True}) - placeholder_71 = T.match_buffer(placeholder_68, [200704], dtype="int16", elem_offset=0, align=64, offset_factor=1) - placeholder_72 = T.match_buffer(placeholder_69, [110592], dtype="int16", elem_offset=0, align=64, offset_factor=1) - placeholder_73 = T.match_buffer(placeholder_70, [192], dtype="int32", elem_offset=0, align=64, offset_factor=1) - T_cast_23 = T.match_buffer(T_cast_22, [305], dtype="uint8", elem_offset=0, align=64, offset_factor=1) - # body - PaddedInput_8 = T.decl_buffer([215296], "int16") - for i0_i1_fused_8 in T.serial(0, 58): - for i2_8, i3_8 in T.grid(58, 64): - PaddedInput_8[(((i0_i1_fused_8*3712) + (i2_8*64)) + i3_8)] = T.if_then_else(((((1 <= i0_i1_fused_8) and (i0_i1_fused_8 < 57)) and (1 <= i2_8)) and (i2_8 < 57)), placeholder_71[((((i0_i1_fused_8*3584) + (i2_8*64)) + i3_8) - 3648)], T.int16(0), dtype="int16") - for ax0_ax1_fused_ax2_fused_8 in T.parallel(0, 3136): - dummy_allocate = T.decl_buffer([1], "int32") - for ax3_outer_4 in T.serial(0, 3): - Conv2dOutput_8 = T.decl_buffer([64], "int32") - for ff_4 in T.serial(0, 64): - Conv2dOutput_8[ff_4] = 0 - for ry_3, rx_3, rc_8 in T.grid(3, 3, 64): - Conv2dOutput_8[ff_4] = (Conv2dOutput_8[ff_4] + (T.cast(PaddedInput_8[(((((T.floordiv(ax0_ax1_fused_ax2_fused_8, 56)*3712) + (ry_3*3712)) + (rx_3*64)) + (T.floormod(ax0_ax1_fused_ax2_fused_8, 56)*64)) + rc_8)], "int32")*T.cast(placeholder_72[(((((ry_3*36864) + (rx_3*12288)) + (rc_8*192)) + (ax3_outer_4*64)) + ff_4)], "int32"))) - for ax3_inner_8 in T.serial(0, 64): - T_cast_23[(((ax0_ax1_fused_ax2_fused_8*192) + (ax3_outer_4*64)) + ax3_inner_8)] = T.cast(T.max(T.min(T.q_multiply_shift((Conv2dOutput_8[ax3_inner_8] + placeholder_73[((ax3_outer_4*64) + ax3_inner_8)]), 1139793473, 31, -6, dtype="int32"), 255), 0), "uint8") - - @T.prim_func - def run_model(input: T.handle, output: T.handle) -> None: - # function attr dict - T.func_attr({"global_symbol": "tvmgen_default_run_model", "runner_function": True}) - # body - T.attr("default", "device_id", 0) - T.attr("default", "device_type", 1) - T.evaluate(T.call_extern("tvmgen_default_fused_nn_conv2d_add_fixed_point_multiply_clip_cast_1", input, T.lookup_param("p5", dtype="handle"), T.lookup_param("p6", dtype="handle"), output, dtype="int32")) - - -# fmt: on - - -# fmt: off -@tvm.script.ir_module -class AllSerialForLoops: - @T.prim_func - def tvmgen_default_fused_nn_conv2d_add_fixed_point_multiply_clip_cast_1(placeholder_68: T.handle, placeholder_69: T.handle, placeholder_70: T.handle, T_cast_22: T.handle) -> None: - # function attr dict - T.func_attr({"global_symbol": "tvmgen_default_fused_nn_conv2d_add_fixed_point_multiply_clip_cast_1", "tir.noalias": True}) - placeholder_71 = T.match_buffer(placeholder_68, [200704], dtype="int16", elem_offset=0, align=64, offset_factor=1) - placeholder_72 = T.match_buffer(placeholder_69, [110592], dtype="int16", elem_offset=0, align=64, offset_factor=1) - placeholder_73 = T.match_buffer(placeholder_70, [192], dtype="int32", elem_offset=0, align=64, offset_factor=1) - T_cast_23 = T.match_buffer(T_cast_22, [305], dtype="uint8", elem_offset=0, align=64, offset_factor=1) - # body - PaddedInput_8 = T.decl_buffer([215296], "int16") - for i0_i1_fused_8 in T.serial(0, 58): - for i2_8, i3_8 in T.grid(58, 64): - PaddedInput_8[(((i0_i1_fused_8*3712) + (i2_8*64)) + i3_8)] = T.if_then_else(((((1 <= i0_i1_fused_8) and (i0_i1_fused_8 < 57)) and (1 <= i2_8)) and (i2_8 < 57)), placeholder_71[((((i0_i1_fused_8*3584) + (i2_8*64)) + i3_8) - 3648)], T.int16(0), dtype="int16") - for ax0_ax1_fused_ax2_fused_8 in T.serial(0, 3136): - dummy_allocate = T.decl_buffer([1], "int32") - for ax3_outer_4 in T.serial(0, 3): - Conv2dOutput_8 = T.decl_buffer([64], "int32") - for ff_4 in T.serial(0, 64): - Conv2dOutput_8[ff_4] = 0 - for ry_3, rx_3, rc_8 in T.grid(3, 3, 64): - Conv2dOutput_8[ff_4] = (Conv2dOutput_8[ff_4] + (T.cast(PaddedInput_8[(((((T.floordiv(ax0_ax1_fused_ax2_fused_8, 56)*3712) + (ry_3*3712)) + (rx_3*64)) + (T.floormod(ax0_ax1_fused_ax2_fused_8, 56)*64)) + rc_8)], "int32")*T.cast(placeholder_72[(((((ry_3*36864) + (rx_3*12288)) + (rc_8*192)) + (ax3_outer_4*64)) + ff_4)], "int32"))) - for ax3_inner_8 in T.serial(0, 64): - T_cast_23[(((ax0_ax1_fused_ax2_fused_8*192) + (ax3_outer_4*64)) + ax3_inner_8)] = T.cast(T.max(T.min(T.q_multiply_shift((Conv2dOutput_8[ax3_inner_8] + placeholder_73[((ax3_outer_4*64) + ax3_inner_8)]), 1139793473, 31, -6, dtype="int32"), 255), 0), "uint8") - - @T.prim_func - def run_model(input: T.handle, output: T.handle) -> None: - # function attr dict - T.func_attr({"global_symbol": "tvmgen_default_run_model", "runner_function": True}) - # body - T.attr("default", "device_id", 0) - T.attr("default", "device_type", 1) - T.evaluate(T.call_extern("tvmgen_default_fused_nn_conv2d_add_fixed_point_multiply_clip_cast_1", input, T.lookup_param("p5", dtype="handle"), T.lookup_param("p6", dtype="handle"), output, dtype="int32")) - - -# fmt: on - - -def test_parallel_serial_mixed_for_loops(): - target = Target("c") - global_ws_pool = WorkspacePoolInfo( - pool_name="global_workspace", - targets=[target], - ) - all_serial_tir_mod = AllSerialForLoops - all_serial_tir_mod = _assign_targets_to_primfuncs_irmodule(all_serial_tir_mod, target) - all_serial_tir_mod = _assign_poolinfos_to_allocates_in_irmodule( - all_serial_tir_mod, [global_ws_pool] - ) - main_func = all_serial_tir_mod["run_model"] - buffer_info_analysis = tvm.tir.usmp.analysis.extract_buffer_info(main_func, all_serial_tir_mod) - assert buffer_info_analysis.memory_pressure == 430848 - buffer_info_map = _replace_stmt_with_buf_var_names(buffer_info_analysis.buffer_info_stmts) - - # When all loops are serial all allocates are touched by USMP - assert len(buffer_info_map) == 3 - for name, _ in buffer_info_map.items(): - assert name in ["dummy_allocate", "Conv2dOutput_8", "PaddedInput_8"] - - parallel_serial_mixed_tir_mod = ParallelSerialMixedForLoops - parallel_serial_mixed_tir_mod = _assign_targets_to_primfuncs_irmodule( - parallel_serial_mixed_tir_mod, target - ) - parallel_serial_mixed_tir_mod = _assign_poolinfos_to_allocates_in_irmodule( - parallel_serial_mixed_tir_mod, [global_ws_pool] - ) - main_func = parallel_serial_mixed_tir_mod["run_model"] - buffer_info_analysis = tvm.tir.usmp.analysis.extract_buffer_info( - main_func, parallel_serial_mixed_tir_mod - ) - assert buffer_info_analysis.memory_pressure == 430848 - buffer_info_map = _replace_stmt_with_buf_var_names(buffer_info_analysis.buffer_info_stmts) - - # USMP will not touch (yet) the allocates inside parallel for loops - assert len(buffer_info_map) == 2 - for name, _ in buffer_info_map.items(): - assert name in ["Conv2dOutput_8", "PaddedInput_8"] - - -# fmt: off -@tvm.script.ir_module -class InceptionStructure: - @T.prim_func - def tvmgen_default_fused_nn_max_pool2d(placeholder: T.handle, tensor: T.handle) -> None: - # function attr dict - T.func_attr({"global_symbol": "tvmgen_default_fused_nn_max_pool2d", "tir.noalias": True}) - placeholder_1 = T.match_buffer(placeholder, [602112], dtype="uint8", elem_offset=0, align=64, offset_factor=1) - tensor_1 = T.match_buffer(tensor, [249], dtype="uint8", elem_offset=0, align=64, offset_factor=1) - # body - for ax0_ax1_fused in T.serial(0, 28): - for ax2 in T.serial(0, 28): - for ax3_outer_init, ax3_inner_init in T.grid(3, 64): - tensor_1[((((ax0_ax1_fused*5376) + (ax2*192)) + (ax3_outer_init*64)) + ax3_inner_init)] = T.uint8(0) - for rv0_rv1_fused, ax3_outer, ax3_inner in T.grid(9, 3, 64): - tensor_1[((((ax0_ax1_fused*5376) + (ax2*192)) + (ax3_outer*64)) + ax3_inner)] = T.max(tensor_1[((((ax0_ax1_fused*5376) + (ax2*192)) + (ax3_outer*64)) + ax3_inner)], T.if_then_else(((((ax0_ax1_fused*2) + T.floordiv(rv0_rv1_fused, 3)) < 56) and (((ax2*2) + T.floormod(rv0_rv1_fused, 3)) < 56)), placeholder_1[((((((ax0_ax1_fused*21504) + (T.floordiv(rv0_rv1_fused, 3)*10752)) + (ax2*384)) + (T.floormod(rv0_rv1_fused, 3)*192)) + (ax3_outer*64)) + ax3_inner)], T.uint8(0), dtype="uint8")) - - @T.prim_func - def tvmgen_default_fused_cast_subtract(placeholder_2: T.handle, placeholder_3: T.handle, T_subtract: T.handle) -> None: - # function attr dict - T.func_attr({"global_symbol": "tvmgen_default_fused_cast_subtract", "tir.noalias": True}) - placeholder_4 = T.match_buffer(placeholder_2, [150528], dtype="uint8", elem_offset=0, align=64, offset_factor=1) - placeholder_5 = T.match_buffer(placeholder_3, [1], dtype="int16", elem_offset=0, align=64, offset_factor=1) - T_subtract_1 = T.match_buffer(T_subtract, [452], dtype="int16", elem_offset=0, align=64, offset_factor=1) - # body - for ax0_ax1_fused_1 in T.serial(0, 224): - for ax2_1, ax3_inner_1 in T.grid(224, 3): - T_subtract_1[(((ax0_ax1_fused_1*672) + (ax2_1*3)) + ax3_inner_1)] = (T.cast(placeholder_4[(((ax0_ax1_fused_1*672) + (ax2_1*3)) + ax3_inner_1)], "int16") - placeholder_5[0]) - - @T.prim_func - def tvmgen_default_fused_cast(placeholder_6: T.handle, T_cast: T.handle) -> None: - # function attr dict - T.func_attr({"global_symbol": "tvmgen_default_fused_cast", "tir.noalias": True}) - placeholder_7 = T.match_buffer(placeholder_6, [150528], dtype="uint8", elem_offset=0, align=64, offset_factor=1) - T_cast_1 = T.match_buffer(T_cast, [249], dtype="int16", elem_offset=0, align=64, offset_factor=1) - # body - for ax0_ax1_fused_2 in T.serial(0, 28): - for ax2_2, ax3_outer_1, ax3_inner_2 in T.grid(28, 12, 16): - T_cast_1[((((ax0_ax1_fused_2*5376) + (ax2_2*192)) + (ax3_outer_1*16)) + ax3_inner_2)] = T.cast(placeholder_7[((((ax0_ax1_fused_2*5376) + (ax2_2*192)) + (ax3_outer_1*16)) + ax3_inner_2)], "int16") - - @T.prim_func - def tvmgen_default_fused_concatenate(placeholder_8: T.handle, placeholder_9: T.handle, placeholder_10: T.handle, placeholder_11: T.handle, T_concat: T.handle) -> None: - # function attr dict - T.func_attr({"global_symbol": "tvmgen_default_fused_concatenate", "tir.noalias": True}) - placeholder_12 = T.match_buffer(placeholder_8, [50176], dtype="uint8", elem_offset=0, align=64, offset_factor=1) - T_concat_1 = T.match_buffer(T_concat, [313], dtype="uint8", elem_offset=0, align=64, offset_factor=1) - placeholder_13 = T.match_buffer(placeholder_9, [100352], dtype="uint8", elem_offset=0, align=64, offset_factor=1) - placeholder_14 = T.match_buffer(placeholder_11, [25088], dtype="uint8", elem_offset=0, align=64, offset_factor=1) - placeholder_15 = T.match_buffer(placeholder_10, [25088], dtype="uint8", elem_offset=0, align=64, offset_factor=1) - # body - for ax0_ax1_fused_3 in T.serial(0, 28): - for ax2_3, ax3 in T.grid(28, 256): - T_concat_1[(((ax0_ax1_fused_3*7168) + (ax2_3*256)) + ax3)] = T.if_then_else((224 <= ax3), placeholder_14[((((ax0_ax1_fused_3*896) + (ax2_3*32)) + ax3) - 224)], T.if_then_else((192 <= ax3), placeholder_15[((((ax0_ax1_fused_3*896) + (ax2_3*32)) + ax3) - 192)], T.if_then_else((64 <= ax3), placeholder_13[((((ax0_ax1_fused_3*3584) + (ax2_3*128)) + ax3) - 64)], placeholder_12[(((ax0_ax1_fused_3*1792) + (ax2_3*64)) + ax3)], dtype="uint8"), dtype="uint8"), dtype="uint8") - - @T.prim_func - def tvmgen_default_fused_nn_conv2d_add_fixed_point_multiply_clip_cast_cast(placeholder_16: T.handle, placeholder_17: T.handle, placeholder_18: T.handle, T_cast_2: T.handle) -> None: - # function attr dict - T.func_attr({"global_symbol": "tvmgen_default_fused_nn_conv2d_add_fixed_point_multiply_clip_cast_cast", "tir.noalias": True}) - placeholder_19 = T.match_buffer(placeholder_16, [200704], dtype="int16", elem_offset=0, align=64, offset_factor=1) - placeholder_20 = T.match_buffer(placeholder_17, [4096], dtype="int16", elem_offset=0, align=64, offset_factor=1) - placeholder_21 = T.match_buffer(placeholder_18, [64], dtype="int32", elem_offset=0, align=64, offset_factor=1) - T_cast_3 = T.match_buffer(T_cast_2, [177], dtype="int16", elem_offset=0, align=64, offset_factor=1) - # body - PaddedInput = T.decl_buffer([200704], "int16") - for i0_i1_fused in T.serial(0, 56): - for i2, i3 in T.grid(56, 64): - PaddedInput[(((i0_i1_fused*3584) + (i2*64)) + i3)] = placeholder_19[(((i0_i1_fused*3584) + (i2*64)) + i3)] - for ax0_ax1_fused_ax2_fused in T.serial(0, 3136): - Conv2dOutput = T.decl_buffer([64], "int32") - for ff in T.serial(0, 64): - Conv2dOutput[ff] = 0 - for rc in T.serial(0, 64): - Conv2dOutput[ff] = (Conv2dOutput[ff] + (T.cast(PaddedInput[((ax0_ax1_fused_ax2_fused*64) + rc)], "int32")*T.cast(placeholder_20[((rc*64) + ff)], "int32"))) - for ax3_inner_3 in T.serial(0, 64): - T_cast_3[((ax0_ax1_fused_ax2_fused*64) + ax3_inner_3)] = T.cast(T.cast(T.max(T.min(T.q_multiply_shift((Conv2dOutput[ax3_inner_3] + placeholder_21[ax3_inner_3]), 1191576922, 31, -4, dtype="int32"), 255), 0), "uint8"), "int16") - - @T.prim_func - def tvmgen_default_fused_nn_conv2d_add_fixed_point_multiply_clip_cast_cast_1(placeholder_22: T.handle, placeholder_23: T.handle, placeholder_24: T.handle, T_cast_4: T.handle) -> None: - # function attr dict - T.func_attr({"global_symbol": "tvmgen_default_fused_nn_conv2d_add_fixed_point_multiply_clip_cast_cast_1", "tir.noalias": True}) - placeholder_25 = T.match_buffer(placeholder_22, [150528], dtype="int16", elem_offset=0, align=64, offset_factor=1) - placeholder_26 = T.match_buffer(placeholder_23, [18432], dtype="int16", elem_offset=0, align=64, offset_factor=1) - placeholder_27 = T.match_buffer(placeholder_24, [96], dtype="int32", elem_offset=0, align=64, offset_factor=1) - T_cast_5 = T.match_buffer(T_cast_4, [153], dtype="int16", elem_offset=0, align=64, offset_factor=1) - # body - PaddedInput_1 = T.decl_buffer([150528], "int16") - for i0_i1_fused_1 in T.serial(0, 28): - for i2_1, i3_1 in T.grid(28, 192): - PaddedInput_1[(((i0_i1_fused_1*5376) + (i2_1*192)) + i3_1)] = placeholder_25[(((i0_i1_fused_1*5376) + (i2_1*192)) + i3_1)] - for ax0_ax1_fused_ax2_fused_1 in T.serial(0, 784): - Conv2dOutput_1 = T.decl_buffer([1], "int32") - for ax3_1 in T.serial(0, 96): - Conv2dOutput_1[0] = 0 - for rc_1 in T.serial(0, 192): - Conv2dOutput_1[0] = (Conv2dOutput_1[0] + (T.cast(PaddedInput_1[((ax0_ax1_fused_ax2_fused_1*192) + rc_1)], "int32")*T.cast(placeholder_26[((rc_1*96) + ax3_1)], "int32"))) - T_cast_5[((ax0_ax1_fused_ax2_fused_1*96) + ax3_1)] = T.cast(T.cast(T.max(T.min(T.q_multiply_shift((Conv2dOutput_1[0] + placeholder_27[ax3_1]), 1201322342, 31, -6, dtype="int32"), 255), 0), "uint8"), "int16") - - @T.prim_func - def tvmgen_default_fused_nn_max_pool2d_cast(placeholder_28: T.handle, T_cast_6: T.handle) -> None: - # function attr dict - T.func_attr({"global_symbol": "tvmgen_default_fused_nn_max_pool2d_cast", "tir.noalias": True}) - placeholder_29 = T.match_buffer(placeholder_28, [802816], dtype="uint8", elem_offset=0, align=64, offset_factor=1) - T_cast_7 = T.match_buffer(T_cast_6, [177], dtype="int16", elem_offset=0, align=64, offset_factor=1) - # body - tensor_2 = T.decl_buffer([200704], "uint8") - for ax0_ax1_fused_4 in T.serial(0, 56): - for ax2_4 in T.serial(0, 56): - for ax3_init in T.serial(0, 64): - tensor_2[(((ax0_ax1_fused_4*3584) + (ax2_4*64)) + ax3_init)] = T.uint8(0) - for rv0_rv1_fused_1, ax3_2 in T.grid(9, 64): - tensor_2[(((ax0_ax1_fused_4*3584) + (ax2_4*64)) + ax3_2)] = T.max(tensor_2[(((ax0_ax1_fused_4*3584) + (ax2_4*64)) + ax3_2)], T.if_then_else(((((ax0_ax1_fused_4*2) + T.floordiv(rv0_rv1_fused_1, 3)) < 112) and (((ax2_4*2) + T.floormod(rv0_rv1_fused_1, 3)) < 112)), placeholder_29[(((((ax0_ax1_fused_4*14336) + (T.floordiv(rv0_rv1_fused_1, 3)*7168)) + (ax2_4*128)) + (T.floormod(rv0_rv1_fused_1, 3)*64)) + ax3_2)], T.uint8(0), dtype="uint8")) - for ax0_ax1_fused_5 in T.serial(0, 56): - for ax2_5, ax3_3 in T.grid(56, 64): - T_cast_7[(((ax0_ax1_fused_5*3584) + (ax2_5*64)) + ax3_3)] = T.cast(tensor_2[(((ax0_ax1_fused_5*3584) + (ax2_5*64)) + ax3_3)], "int16") - - @T.prim_func - def tvmgen_default_fused_nn_conv2d_add_fixed_point_multiply_clip_cast_2(placeholder_30: T.handle, placeholder_31: T.handle, placeholder_32: T.handle, T_cast_8: T.handle) -> None: - # function attr dict - T.func_attr({"global_symbol": "tvmgen_default_fused_nn_conv2d_add_fixed_point_multiply_clip_cast_2", "tir.noalias": True}) - placeholder_33 = T.match_buffer(placeholder_30, [150528], dtype="int16", elem_offset=0, align=64, offset_factor=1) - placeholder_34 = T.match_buffer(placeholder_31, [12288], dtype="int16", elem_offset=0, align=64, offset_factor=1) - placeholder_35 = T.match_buffer(placeholder_32, [64], dtype="int32", elem_offset=0, align=64, offset_factor=1) - T_cast_9 = T.match_buffer(T_cast_8, [121], dtype="uint8", elem_offset=0, align=64, offset_factor=1) - # body - PaddedInput_2 = T.decl_buffer([150528], "int16") - for i0_i1_fused_2 in T.serial(0, 28): - for i2_2, i3_2 in T.grid(28, 192): - PaddedInput_2[(((i0_i1_fused_2*5376) + (i2_2*192)) + i3_2)] = placeholder_33[(((i0_i1_fused_2*5376) + (i2_2*192)) + i3_2)] - for ax0_ax1_fused_ax2_fused_2 in T.serial(0, 784): - Conv2dOutput_2 = T.decl_buffer([64], "int32") - for ff_1 in T.serial(0, 64): - Conv2dOutput_2[ff_1] = 0 - for rc_2 in T.serial(0, 192): - Conv2dOutput_2[ff_1] = (Conv2dOutput_2[ff_1] + (T.cast(PaddedInput_2[((ax0_ax1_fused_ax2_fused_2*192) + rc_2)], "int32")*T.cast(placeholder_34[((rc_2*64) + ff_1)], "int32"))) - for ax3_inner_4 in T.serial(0, 64): - T_cast_9[((ax0_ax1_fused_ax2_fused_2*64) + ax3_inner_4)] = T.cast(T.max(T.min(T.q_multiply_shift((Conv2dOutput_2[ax3_inner_4] + placeholder_35[ax3_inner_4]), 1663316467, 31, -7, dtype="int32"), 255), 0), "uint8") - - @T.prim_func - def tvmgen_default_fused_nn_max_pool2d_cast_1(placeholder_36: T.handle, T_cast_10: T.handle) -> None: - # function attr dict - T.func_attr({"global_symbol": "tvmgen_default_fused_nn_max_pool2d_cast_1", "tir.noalias": True}) - placeholder_37 = T.match_buffer(placeholder_36, [150528], dtype="uint8", elem_offset=0, align=64, offset_factor=1) - T_cast_11 = T.match_buffer(T_cast_10, [249], dtype="int16", elem_offset=0, align=64, offset_factor=1) - # body - tensor_3 = T.decl_buffer([150528], "uint8") - for ax0_ax1_fused_6 in T.serial(0, 28): - for ax2_6 in T.serial(0, 28): - for ax3_outer_init_1, ax3_inner_init_1 in T.grid(3, 64): - tensor_3[((((ax0_ax1_fused_6*5376) + (ax2_6*192)) + (ax3_outer_init_1*64)) + ax3_inner_init_1)] = T.uint8(0) - for rv0_rv1_fused_2, ax3_outer_2, ax3_inner_5 in T.grid(9, 3, 64): - tensor_3[((((ax0_ax1_fused_6*5376) + (ax2_6*192)) + (ax3_outer_2*64)) + ax3_inner_5)] = T.max(tensor_3[((((ax0_ax1_fused_6*5376) + (ax2_6*192)) + (ax3_outer_2*64)) + ax3_inner_5)], T.if_then_else(((((1 <= (T.floordiv(rv0_rv1_fused_2, 3) + ax0_ax1_fused_6)) and ((T.floordiv(rv0_rv1_fused_2, 3) + ax0_ax1_fused_6) < 29)) and (1 <= (ax2_6 + T.floormod(rv0_rv1_fused_2, 3)))) and ((ax2_6 + T.floormod(rv0_rv1_fused_2, 3)) < 29)), placeholder_37[(((((((T.floordiv(rv0_rv1_fused_2, 3)*5376) + (ax0_ax1_fused_6*5376)) + (ax2_6*192)) + (T.floormod(rv0_rv1_fused_2, 3)*192)) + (ax3_outer_2*64)) + ax3_inner_5) - 5568)], T.uint8(0), dtype="uint8")) - for ax0_ax1_fused_7 in T.serial(0, 28): - for ax2_7, ax3_4 in T.grid(28, 192): - T_cast_11[(((ax0_ax1_fused_7*5376) + (ax2_7*192)) + ax3_4)] = T.cast(tensor_3[(((ax0_ax1_fused_7*5376) + (ax2_7*192)) + ax3_4)], "int16") - - @T.prim_func - def tvmgen_default_fused_nn_conv2d_add_fixed_point_multiply_clip_cast_cast_fixed_point_multiply_cli_4464294615199028320__2(placeholder_38: T.handle, placeholder_39: T.handle, placeholder_40: T.handle, T_cast_12: T.handle) -> None: - # function attr dict - T.func_attr({"global_symbol": "tvmgen_default_fused_nn_conv2d_add_fixed_point_multiply_clip_cast_cast_fixed_point_multiply_cli_4464294615199028320__2", "tir.noalias": True}) - placeholder_41 = T.match_buffer(placeholder_38, [150528], dtype="int16", elem_offset=0, align=64, offset_factor=1) - placeholder_42 = T.match_buffer(placeholder_39, [6144], dtype="int16", elem_offset=0, align=64, offset_factor=1) - placeholder_43 = T.match_buffer(placeholder_40, [32], dtype="int32", elem_offset=0, align=64, offset_factor=1) - T_cast_13 = T.match_buffer(T_cast_12, [89], dtype="uint8", elem_offset=0, align=64, offset_factor=1) - # body - PaddedInput_3 = T.decl_buffer([150528], "int16") - for i0_i1_fused_3 in T.serial(0, 28): - for i2_3, i3_3 in T.grid(28, 192): - PaddedInput_3[(((i0_i1_fused_3*5376) + (i2_3*192)) + i3_3)] = placeholder_41[(((i0_i1_fused_3*5376) + (i2_3*192)) + i3_3)] - for ax0_ax1_fused_ax2_fused_3 in T.serial(0, 784): - Conv2dOutput_3 = T.decl_buffer([1], "int32") - for ax3_5 in T.serial(0, 32): - Conv2dOutput_3[0] = 0 - for rc_3 in T.serial(0, 192): - Conv2dOutput_3[0] = (Conv2dOutput_3[0] + (T.cast(PaddedInput_3[((ax0_ax1_fused_ax2_fused_3*192) + rc_3)], "int32")*T.cast(placeholder_42[((rc_3*32) + ax3_5)], "int32"))) - T_cast_13[((ax0_ax1_fused_ax2_fused_3*32) + ax3_5)] = T.cast(T.max(T.min(T.q_multiply_shift(T.cast(T.cast(T.max(T.min(T.q_multiply_shift((Conv2dOutput_3[0] + placeholder_43[ax3_5]), 1811141736, 31, -6, dtype="int32"), 255), 0), "uint8"), "int32"), 1136333842, 31, 0, dtype="int32"), 255), 0), "uint8") - - @T.prim_func - def tvmgen_default_fused_nn_conv2d_add_fixed_point_multiply_clip_cast_cast_2(placeholder_44: T.handle, placeholder_45: T.handle, placeholder_46: T.handle, T_cast_14: T.handle) -> None: - # function attr dict - T.func_attr({"global_symbol": "tvmgen_default_fused_nn_conv2d_add_fixed_point_multiply_clip_cast_cast_2", "tir.noalias": True}) - placeholder_47 = T.match_buffer(placeholder_44, [150528], dtype="int16", elem_offset=0, align=64, offset_factor=1) - placeholder_48 = T.match_buffer(placeholder_45, [3072], dtype="int16", elem_offset=0, align=64, offset_factor=1) - placeholder_49 = T.match_buffer(placeholder_46, [16], dtype="int32", elem_offset=0, align=64, offset_factor=1) - T_cast_15 = T.match_buffer(T_cast_14, [73], dtype="int16", elem_offset=0, align=64, offset_factor=1) - # body - PaddedInput_4 = T.decl_buffer([150528], "int16") - for i0_i1_fused_4 in T.serial(0, 28): - for i2_4, i3_4 in T.grid(28, 192): - PaddedInput_4[(((i0_i1_fused_4*5376) + (i2_4*192)) + i3_4)] = placeholder_47[(((i0_i1_fused_4*5376) + (i2_4*192)) + i3_4)] - for ax0_ax1_fused_ax2_fused_4 in T.serial(0, 784): - Conv2dOutput_4 = T.decl_buffer([1], "int32") - for ax3_6 in T.serial(0, 16): - Conv2dOutput_4[0] = 0 - for rc_4 in T.serial(0, 192): - Conv2dOutput_4[0] = (Conv2dOutput_4[0] + (T.cast(PaddedInput_4[((ax0_ax1_fused_ax2_fused_4*192) + rc_4)], "int32")*T.cast(placeholder_48[((rc_4*16) + ax3_6)], "int32"))) - T_cast_15[((ax0_ax1_fused_ax2_fused_4*16) + ax3_6)] = T.cast(T.cast(T.max(T.min(T.q_multiply_shift((Conv2dOutput_4[0] + placeholder_49[ax3_6]), 1764006585, 31, -7, dtype="int32"), 255), 0), "uint8"), "int16") - - @T.prim_func - def tvmgen_default_fused_nn_conv2d_add_fixed_point_multiply_clip_cast_cast_fixed_point_multiply_cli_4464294615199028320__1(placeholder_50: T.handle, placeholder_51: T.handle, placeholder_52: T.handle, T_cast_16: T.handle) -> None: - # function attr dict - T.func_attr({"global_symbol": "tvmgen_default_fused_nn_conv2d_add_fixed_point_multiply_clip_cast_cast_fixed_point_multiply_cli_4464294615199028320__1", "tir.noalias": True}) - placeholder_53 = T.match_buffer(placeholder_50, [12544], dtype="int16", elem_offset=0, align=64, offset_factor=1) - placeholder_54 = T.match_buffer(placeholder_51, [4608], dtype="int16", elem_offset=0, align=64, offset_factor=1) - placeholder_55 = T.match_buffer(placeholder_52, [32], dtype="int32", elem_offset=0, align=64, offset_factor=1) - T_cast_17 = T.match_buffer(T_cast_16, [89], dtype="uint8", elem_offset=0, align=64, offset_factor=1) - # body - PaddedInput_5 = T.decl_buffer([14400], "int16") - for i0_i1_fused_5 in T.serial(0, 30): - for i2_5, i3_5 in T.grid(30, 16): - PaddedInput_5[(((i0_i1_fused_5*480) + (i2_5*16)) + i3_5)] = T.if_then_else(((((1 <= i0_i1_fused_5) and (i0_i1_fused_5 < 29)) and (1 <= i2_5)) and (i2_5 < 29)), placeholder_53[((((i0_i1_fused_5*448) + (i2_5*16)) + i3_5) - 464)], T.int16(0), dtype="int16") - for ax0_ax1_fused_ax2_fused_5 in T.serial(0, 784): - Conv2dOutput_5 = T.decl_buffer([1], "int32") - for ax3_7 in T.serial(0, 32): - Conv2dOutput_5[0] = 0 - for ry, rx, rc_5 in T.grid(3, 3, 16): - Conv2dOutput_5[0] = (Conv2dOutput_5[0] + (T.cast(PaddedInput_5[(((((T.floordiv(ax0_ax1_fused_ax2_fused_5, 28)*480) + (ry*480)) + (rx*16)) + (T.floormod(ax0_ax1_fused_ax2_fused_5, 28)*16)) + rc_5)], "int32")*T.cast(placeholder_54[((((ry*1536) + (rx*512)) + (rc_5*32)) + ax3_7)], "int32"))) - T_cast_17[((ax0_ax1_fused_ax2_fused_5*32) + ax3_7)] = T.cast(T.max(T.min(T.q_multiply_shift(T.cast(T.cast(T.max(T.min(T.q_multiply_shift((Conv2dOutput_5[0] + placeholder_55[ax3_7]), 1131968888, 31, -6, dtype="int32"), 255), 0), "uint8"), "int32"), 1900719667, 31, 0, dtype="int32"), 255), 0), "uint8") - - @T.prim_func - def tvmgen_default_fused_nn_conv2d_add_fixed_point_multiply_clip_cast_cast_fixed_point_multiply_cli_4464294615199028320_(placeholder_56: T.handle, placeholder_57: T.handle, placeholder_58: T.handle, T_cast_18: T.handle) -> None: - # function attr dict - T.func_attr({"global_symbol": "tvmgen_default_fused_nn_conv2d_add_fixed_point_multiply_clip_cast_cast_fixed_point_multiply_cli_4464294615199028320_", "tir.noalias": True}) - placeholder_59 = T.match_buffer(placeholder_56, [75264], dtype="int16", elem_offset=0, align=64, offset_factor=1) - placeholder_60 = T.match_buffer(placeholder_57, [110592], dtype="int16", elem_offset=0, align=64, offset_factor=1) - placeholder_61 = T.match_buffer(placeholder_58, [128], dtype="int32", elem_offset=0, align=64, offset_factor=1) - T_cast_19 = T.match_buffer(T_cast_18, [185], dtype="uint8", elem_offset=0, align=64, offset_factor=1) - # body - PaddedInput_6 = T.decl_buffer([86400], "int16") - for i0_i1_fused_6 in T.serial(0, 30): - for i2_6, i3_6 in T.grid(30, 96): - PaddedInput_6[(((i0_i1_fused_6*2880) + (i2_6*96)) + i3_6)] = T.if_then_else(((((1 <= i0_i1_fused_6) and (i0_i1_fused_6 < 29)) and (1 <= i2_6)) and (i2_6 < 29)), placeholder_59[((((i0_i1_fused_6*2688) + (i2_6*96)) + i3_6) - 2784)], T.int16(0), dtype="int16") - for ax0_ax1_fused_ax2_fused_6 in T.serial(0, 784): - Conv2dOutput_6 = T.decl_buffer([64], "int32") - for ax3_outer_3 in T.serial(0, 2): - for ff_2 in T.serial(0, 64): - Conv2dOutput_6[ff_2] = 0 - for ry_1, rx_1, rc_6 in T.grid(3, 3, 96): - Conv2dOutput_6[ff_2] = (Conv2dOutput_6[ff_2] + (T.cast(PaddedInput_6[(((((T.floordiv(ax0_ax1_fused_ax2_fused_6, 28)*2880) + (ry_1*2880)) + (rx_1*96)) + (T.floormod(ax0_ax1_fused_ax2_fused_6, 28)*96)) + rc_6)], "int32")*T.cast(placeholder_60[(((((ry_1*36864) + (rx_1*12288)) + (rc_6*128)) + (ax3_outer_3*64)) + ff_2)], "int32"))) - for ax3_inner_6 in T.serial(0, 64): - T_cast_19[(((ax0_ax1_fused_ax2_fused_6*128) + (ax3_outer_3*64)) + ax3_inner_6)] = T.cast(T.max(T.min(T.q_multiply_shift(T.cast(T.cast(T.max(T.min(T.q_multiply_shift((Conv2dOutput_6[ax3_inner_6] + placeholder_61[((ax3_outer_3*64) + ax3_inner_6)]), 1374050734, 31, -7, dtype="int32"), 255), 0), "uint8"), "int32"), 1544713713, 31, 0, dtype="int32"), 255), 0), "uint8") - - @T.prim_func - def tvmgen_default_fused_nn_conv2d_add_fixed_point_multiply_clip_cast(placeholder_62: T.handle, placeholder_63: T.handle, placeholder_64: T.handle, T_cast_20: T.handle) -> None: - # function attr dict - T.func_attr({"global_symbol": "tvmgen_default_fused_nn_conv2d_add_fixed_point_multiply_clip_cast", "T.noalias": True}) - placeholder_65 = T.match_buffer(placeholder_62, [150528], dtype="int16", elem_offset=0, align=64, offset_factor=1) - placeholder_66 = T.match_buffer(placeholder_63, [9408], dtype="int16", elem_offset=0, align=64, offset_factor=1) - placeholder_67 = T.match_buffer(placeholder_64, [64], dtype="int32", elem_offset=0, align=64, offset_factor=1) - T_cast_21 = T.match_buffer(T_cast_20, [289], dtype="uint8", elem_offset=0, align=64, offset_factor=1) - # body - PaddedInput_7 = T.decl_buffer([157323], "int16") - for i0_i1_fused_7 in T.serial(0, 229): - for i2_7, i3_7 in T.grid(229, 3): - PaddedInput_7[(((i0_i1_fused_7*687) + (i2_7*3)) + i3_7)] = T.if_then_else(((((2 <= i0_i1_fused_7) and (i0_i1_fused_7 < 226)) and (2 <= i2_7)) and (i2_7 < 226)), placeholder_65[((((i0_i1_fused_7*672) + (i2_7*3)) + i3_7) - 1350)], T.int16(0), dtype="int16") - for ax0_ax1_fused_ax2_fused_7 in T.serial(0, 12544): - Conv2dOutput_7 = T.decl_buffer([64], "int32") - for ff_3 in T.serial(0, 64): - Conv2dOutput_7[ff_3] = 0 - for ry_2, rx_2, rc_7 in T.grid(7, 7, 3): - Conv2dOutput_7[ff_3] = (Conv2dOutput_7[ff_3] + (T.cast(PaddedInput_7[(((((T.floordiv(ax0_ax1_fused_ax2_fused_7, 112)*1374) + (ry_2*687)) + (T.floormod(ax0_ax1_fused_ax2_fused_7, 112)*6)) + (rx_2*3)) + rc_7)], "int32")*T.cast(placeholder_66[((((ry_2*1344) + (rx_2*192)) + (rc_7*64)) + ff_3)], "int32"))) - for ax3_inner_7 in T.serial(0, 64): - T_cast_21[((ax0_ax1_fused_ax2_fused_7*64) + ax3_inner_7)] = T.cast(T.max(T.min(T.q_multiply_shift((Conv2dOutput_7[ax3_inner_7] + placeholder_67[ax3_inner_7]), 1939887962, 31, -9, dtype="int32"), 255), 0), "uint8") - - @T.prim_func - def tvmgen_default_fused_nn_conv2d_add_fixed_point_multiply_clip_cast_1(placeholder_68: T.handle, placeholder_69: T.handle, placeholder_70: T.handle, T_cast_22: T.handle) -> None: - # function attr dict - T.func_attr({"global_symbol": "tvmgen_default_fused_nn_conv2d_add_fixed_point_multiply_clip_cast_1", "tir.noalias": True}) - placeholder_71 = T.match_buffer(placeholder_68, [200704], dtype="int16", elem_offset=0, align=64, offset_factor=1) - placeholder_72 = T.match_buffer(placeholder_69, [110592], dtype="int16", elem_offset=0, align=64, offset_factor=1) - placeholder_73 = T.match_buffer(placeholder_70, [192], dtype="int32", elem_offset=0, align=64, offset_factor=1) - T_cast_23 = T.match_buffer(T_cast_22, [305], dtype="uint8", elem_offset=0, align=64, offset_factor=1) - # body - PaddedInput_8 = T.decl_buffer([215296], "int16") - for i0_i1_fused_8 in T.serial(0, 58): - for i2_8, i3_8 in T.grid(58, 64): - PaddedInput_8[(((i0_i1_fused_8*3712) + (i2_8*64)) + i3_8)] = T.if_then_else(((((1 <= i0_i1_fused_8) and (i0_i1_fused_8 < 57)) and (1 <= i2_8)) and (i2_8 < 57)), placeholder_71[((((i0_i1_fused_8*3584) + (i2_8*64)) + i3_8) - 3648)], T.int16(0), dtype="int16") - for ax0_ax1_fused_ax2_fused_8 in T.serial(0, 3136): - Conv2dOutput_8 = T.decl_buffer([64], "int32") - for ax3_outer_4 in T.serial(0, 3): - for ff_4 in T.serial(0, 64): - Conv2dOutput_8[ff_4] = 0 - for ry_3, rx_3, rc_8 in T.grid(3, 3, 64): - Conv2dOutput_8[ff_4] = (Conv2dOutput_8[ff_4] + (T.cast(PaddedInput_8[(((((T.floordiv(ax0_ax1_fused_ax2_fused_8, 56)*3712) + (ry_3*3712)) + (rx_3*64)) + (T.floormod(ax0_ax1_fused_ax2_fused_8, 56)*64)) + rc_8)], "int32")*T.cast(placeholder_72[(((((ry_3*36864) + (rx_3*12288)) + (rc_8*192)) + (ax3_outer_4*64)) + ff_4)], "int32"))) - for ax3_inner_8 in T.serial(0, 64): - T_cast_23[(((ax0_ax1_fused_ax2_fused_8*192) + (ax3_outer_4*64)) + ax3_inner_8)] = T.cast(T.max(T.min(T.q_multiply_shift((Conv2dOutput_8[ax3_inner_8] + placeholder_73[((ax3_outer_4*64) + ax3_inner_8)]), 1139793473, 31, -6, dtype="int32"), 255), 0), "uint8") - - @T.prim_func - def run_model(input: T.handle, output: T.handle) -> None: - # function attr dict - T.func_attr({"global_symbol": "tvmgen_default_run_model", "runner_function": True}) - # body - T.attr("default", "device_id", 0) - T.attr("default", "device_type", 1) - sid_32 = T.allocate([301056], "int8", "global") - sid_20 = T.allocate([150528], "int8", "global") - sid_6 = T.allocate([401408], "int8", "global") - sid_9 = T.allocate([301056], "int8", "global") - sid_7 = T.allocate([401408], "int8", "global") - sid_8 = T.allocate([802816], "int8", "global") - sid_2 = T.allocate([50176], "int8", "global") - sid_3 = T.allocate([301056], "int8", "global") - sid_19 = T.allocate([100352], "int8", "global") - sid_4 = T.allocate([150528], "int8", "global") - sid_5 = T.allocate([602112], "int8", "global") - sid_25 = T.allocate([25088], "int8", "global") - sid_26 = T.allocate([25088], "int8", "global") - sid_31 = T.allocate([25088], "int8", "global") - T.evaluate(T.call_extern("tvmgen_default_fused_cast_subtract", input, T.lookup_param("p0", dtype="handle"), sid_9, dtype="int32")) - T.evaluate(T.call_extern("tvmgen_default_fused_nn_conv2d_add_fixed_point_multiply_clip_cast", sid_9, T.lookup_param("p1", dtype="handle"), T.lookup_param("p2", dtype="handle"), sid_8, dtype="int32")) - T.evaluate(T.call_extern("tvmgen_default_fused_nn_max_pool2d_cast", sid_8, sid_7, dtype="int32")) - T.evaluate(T.call_extern("tvmgen_default_fused_nn_conv2d_add_fixed_point_multiply_clip_cast_cast", sid_7, T.lookup_param("p3", dtype="handle"), T.lookup_param("p4", dtype="handle"), sid_6, dtype="int32")) - T.evaluate(T.call_extern("tvmgen_default_fused_nn_conv2d_add_fixed_point_multiply_clip_cast_1", sid_6, T.lookup_param("p5", dtype="handle"), T.lookup_param("p6", dtype="handle"), sid_5, dtype="int32")) - T.evaluate(T.call_extern("tvmgen_default_fused_nn_max_pool2d", sid_5, sid_4, dtype="int32")) - T.evaluate(T.call_extern("tvmgen_default_fused_cast", sid_4, sid_3, dtype="int32")) - T.evaluate(T.call_extern("tvmgen_default_fused_nn_conv2d_add_fixed_point_multiply_clip_cast_2", sid_3, T.lookup_param("p7", dtype="handle"), T.lookup_param("p8", dtype="handle"), sid_2, dtype="int32")) - T.evaluate(T.call_extern("tvmgen_default_fused_nn_conv2d_add_fixed_point_multiply_clip_cast_cast_1", sid_3, T.lookup_param("p9", dtype="handle"), T.lookup_param("p10", dtype="handle"), sid_20, dtype="int32")) - T.evaluate(T.call_extern("tvmgen_default_fused_nn_conv2d_add_fixed_point_multiply_clip_cast_cast_fixed_point_multiply_cli_4464294615199028320_", sid_20, T.lookup_param("p11", dtype="handle"), T.lookup_param("p12", dtype="handle"), sid_19, dtype="int32")) - T.evaluate(T.call_extern("tvmgen_default_fused_nn_conv2d_add_fixed_point_multiply_clip_cast_cast_2", sid_3, T.lookup_param("p13", dtype="handle"), T.lookup_param("p14", dtype="handle"), sid_26, dtype="int32")) - T.evaluate(T.call_extern("tvmgen_default_fused_nn_conv2d_add_fixed_point_multiply_clip_cast_cast_fixed_point_multiply_cli_4464294615199028320__1", sid_26, T.lookup_param("p15", dtype="handle"), T.lookup_param("p16", dtype="handle"), sid_25, dtype="int32")) - T.evaluate(T.call_extern("tvmgen_default_fused_nn_max_pool2d_cast_1", sid_4, sid_32, dtype="int32")) - T.evaluate(T.call_extern("tvmgen_default_fused_nn_conv2d_add_fixed_point_multiply_clip_cast_cast_fixed_point_multiply_cli_4464294615199028320__2", sid_32, T.lookup_param("p17", dtype="handle"), T.lookup_param("p18", dtype="handle"), sid_31, dtype="int32")) - T.evaluate(T.call_extern("tvmgen_default_fused_concatenate", sid_2, sid_19, sid_25, sid_31, output, dtype="int32")) -# fmt: on - - -def test_inception_structure(): - target = Target("c") - global_ws_pool = WorkspacePoolInfo( - pool_name="global_workspace", - targets=[target], - ) - tir_mod = InceptionStructure - tir_mod = _assign_targets_to_primfuncs_irmodule(tir_mod, target) - tir_mod = _assign_poolinfos_to_allocates_in_irmodule(tir_mod, [global_ws_pool]) - main_func = tir_mod["run_model"] - buffer_info_analysis = tvm.tir.usmp.analysis.extract_buffer_info(main_func, tir_mod) - assert buffer_info_analysis.memory_pressure == 1117718 - buffer_info_map = _replace_stmt_with_buf_var_names(buffer_info_analysis.buffer_info_stmts) - - # check conflicts - _verify_conflicts( - "sid_3", - [ - "sid_4", - "PaddedInput_2", - "sid_2", - "Conv2dOutput_2", - "PaddedInput_1", - "Conv2dOutput_1", - "sid_20", - "PaddedInput_6", - "Conv2dOutput_6", - "sid_19", - "PaddedInput_4", - ], - buffer_info_map, - ) - _verify_conflicts( - "Conv2dOutput", - [ - "sid_6", - "PaddedInput", - ], - buffer_info_map, - ) - _verify_conflicts( - "Conv2dOutput_7", - [ - "PaddedInput_7", - "sid_8", - ], - buffer_info_map, - ) - _verify_conflicts( - "sid_4", - [ - "sid_5", - "sid_3", - "PaddedInput_2", - "sid_2", - "Conv2dOutput_2", - "PaddedInput_1", - "Conv2dOutput_1", - "sid_20", - "PaddedInput_6", - "Conv2dOutput_6", - "sid_19", - "PaddedInput_4", - "Conv2dOutput_4", - "sid_26", - "PaddedInput_5", - "Conv2dOutput_5", - "sid_25", - "tensor_3", - ], - buffer_info_map, - ) - _verify_conflicts( - "sid_2", - [ - "PaddedInput_2", - "sid_3", - "sid_4", - "Conv2dOutput_2", - "PaddedInput_1", - "Conv2dOutput_1", - "sid_20", - "PaddedInput_6", - "Conv2dOutput_6", - "sid_19", - "PaddedInput_4", - "Conv2dOutput_4", - "sid_26", - "PaddedInput_5", - "Conv2dOutput_5", - "sid_25", - "tensor_3", - "sid_32", - "PaddedInput_3", - "Conv2dOutput_3", - "sid_31", - ], - buffer_info_map, - ) - _verify_conflicts( - "sid_19", - [ - "Conv2dOutput_6", - "sid_2", - "PaddedInput_6", - "sid_3", - "sid_4", - "PaddedInput_4", - "Conv2dOutput_4", - "sid_26", - "PaddedInput_5", - "Conv2dOutput_5", - "sid_25", - "tensor_3", - "sid_32", - "PaddedInput_3", - "Conv2dOutput_3", - "sid_31", - ], - buffer_info_map, - ) - _verify_conflicts( - "PaddedInput_2", - [ - "sid_3", - "sid_4", - "sid_2", - "Conv2dOutput_2", - ], - buffer_info_map, - ) - _verify_conflicts( - "Conv2dOutput_6", - [ - "sid_2", - "PaddedInput_6", - "sid_3", - "sid_4", - "sid_19", - ], - buffer_info_map, - ) - _verify_conflicts( - "sid_9", - [ - "PaddedInput_7", - ], - buffer_info_map, - ) - _verify_conflicts( - "sid_7", - [ - "tensor_2", - "PaddedInput", - ], - buffer_info_map, - ) - _verify_conflicts( - "PaddedInput_4", - [ - "sid_2", - "sid_19", - "sid_3", - "sid_4", - "Conv2dOutput_4", - "sid_26", - ], - buffer_info_map, - ) - _verify_conflicts( - "PaddedInput_3", - [ - "sid_2", - "sid_32", - "sid_25", - "sid_19", - "Conv2dOutput_3", - "sid_31", - ], - buffer_info_map, - ) - _verify_conflicts( - "sid_5", - [ - "PaddedInput_8", - "Conv2dOutput_8", - "sid_4", - ], - buffer_info_map, - ) - _verify_conflicts( - "sid_31", - [ - "Conv2dOutput_3", - "PaddedInput_3", - "sid_2", - "sid_25", - "sid_19", - ], - buffer_info_map, - ) - _verify_conflicts( - "PaddedInput", - [ - "sid_7", - "sid_6", - "Conv2dOutput", - ], - buffer_info_map, - ) - _verify_conflicts( - "Conv2dOutput_2", - [ - "sid_2", - "PaddedInput_2", - "sid_3", - "sid_4", - ], - buffer_info_map, - ) - _verify_conflicts( - "sid_32", - [ - "tensor_3", - "sid_2", - "sid_25", - "sid_19", - "PaddedInput_3", - ], - buffer_info_map, - ) - _verify_conflicts( - "tensor_2", - [ - "sid_8", - "sid_7", - ], - buffer_info_map, - ) - _verify_conflicts( - "sid_26", - [ - "Conv2dOutput_4", - "PaddedInput_4", - "sid_2", - "sid_19", - "sid_4", - "PaddedInput_5", - ], - buffer_info_map, - ) - _verify_conflicts( - "Conv2dOutput_3", - [ - "PaddedInput_3", - "sid_2", - "sid_25", - "sid_19", - "sid_31", - ], - buffer_info_map, - ) - _verify_conflicts( - "PaddedInput_6", - [ - "sid_2", - "sid_3", - "sid_20", - "sid_4", - "Conv2dOutput_6", - "sid_19", - ], - buffer_info_map, - ) - _verify_conflicts( - "sid_6", - [ - "PaddedInput", - "Conv2dOutput", - "PaddedInput_8", - ], - buffer_info_map, - ) - _verify_conflicts( - "PaddedInput_8", - [ - "sid_6", - "sid_5", - "Conv2dOutput_8", - ], - buffer_info_map, - ) - _verify_conflicts( - "Conv2dOutput_5", - [ - "PaddedInput_5", - "sid_2", - "sid_19", - "sid_4", - "sid_25", - ], - buffer_info_map, - ) - _verify_conflicts( - "Conv2dOutput_1", - [ - "PaddedInput_1", - "sid_2", - "sid_3", - "sid_4", - "sid_20", - ], - buffer_info_map, - ) - _verify_conflicts( - "tensor_3", - [ - "sid_2", - "sid_25", - "sid_19", - "sid_4", - "sid_32", - ], - buffer_info_map, - ) - _verify_conflicts( - "sid_8", - [ - "Conv2dOutput_7", - "PaddedInput_7", - "tensor_2", - ], - buffer_info_map, - ) - _verify_conflicts( - "sid_20", - [ - "Conv2dOutput_1", - "PaddedInput_1", - "sid_2", - "sid_3", - "sid_4", - "PaddedInput_6", - ], - buffer_info_map, - ) - _verify_conflicts( - "Conv2dOutput_8", - [ - "sid_5", - "PaddedInput_8", - ], - buffer_info_map, - ) - _verify_conflicts( - "PaddedInput_1", - [ - "sid_2", - "sid_3", - "sid_4", - "Conv2dOutput_1", - "sid_20", - ], - buffer_info_map, - ) - _verify_conflicts( - "Conv2dOutput_4", - [ - "PaddedInput_4", - "sid_2", - "sid_19", - "sid_4", - "sid_26", - ], - buffer_info_map, - ) - _verify_conflicts( - "sid_25", - [ - "PaddedInput_5", - "Conv2dOutput_5", - "sid_2", - "sid_19", - "sid_4", - "tensor_3", - "sid_32", - "PaddedInput_3", - "Conv2dOutput_3", - "sid_31", - ], - buffer_info_map, - ) - _verify_conflicts( - "PaddedInput_7", - [ - "sid_9", - "Conv2dOutput_7", - "sid_8", - ], - buffer_info_map, - ) - _verify_conflicts( - "PaddedInput_5", - [ - "sid_2", - "sid_19", - "sid_26", - "sid_4", - "Conv2dOutput_5", - "sid_25", - ], - buffer_info_map, - ) - - # check sizes - assert buffer_info_map["sid_20"].size_bytes == 150528 - assert buffer_info_map["tensor_2"].size_bytes == 200704 - assert buffer_info_map["sid_5"].size_bytes == 602112 - assert buffer_info_map["sid_9"].size_bytes == 301056 - assert buffer_info_map["Conv2dOutput_3"].size_bytes == 4 - assert buffer_info_map["sid_26"].size_bytes == 25088 - assert buffer_info_map["Conv2dOutput_2"].size_bytes == 256 - assert buffer_info_map["PaddedInput_5"].size_bytes == 28800 - assert buffer_info_map["sid_8"].size_bytes == 802816 - assert buffer_info_map["Conv2dOutput_5"].size_bytes == 4 - assert buffer_info_map["sid_3"].size_bytes == 301056 - assert buffer_info_map["Conv2dOutput"].size_bytes == 256 - assert buffer_info_map["PaddedInput_3"].size_bytes == 301056 - assert buffer_info_map["sid_32"].size_bytes == 301056 - assert buffer_info_map["PaddedInput_8"].size_bytes == 430592 - assert buffer_info_map["sid_4"].size_bytes == 150528 - assert buffer_info_map["PaddedInput_7"].size_bytes == 314646 - assert buffer_info_map["sid_6"].size_bytes == 401408 - assert buffer_info_map["Conv2dOutput_8"].size_bytes == 256 - assert buffer_info_map["sid_25"].size_bytes == 25088 - assert buffer_info_map["PaddedInput"].size_bytes == 401408 - assert buffer_info_map["sid_7"].size_bytes == 401408 - assert buffer_info_map["Conv2dOutput_1"].size_bytes == 4 - assert buffer_info_map["Conv2dOutput_4"].size_bytes == 4 - assert buffer_info_map["PaddedInput_2"].size_bytes == 301056 - assert buffer_info_map["sid_31"].size_bytes == 25088 - assert buffer_info_map["PaddedInput_1"].size_bytes == 301056 - assert buffer_info_map["Conv2dOutput_6"].size_bytes == 256 - assert buffer_info_map["PaddedInput_4"].size_bytes == 301056 - assert buffer_info_map["sid_2"].size_bytes == 50176 - assert buffer_info_map["tensor_3"].size_bytes == 150528 - assert buffer_info_map["Conv2dOutput_7"].size_bytes == 256 - assert buffer_info_map["sid_19"].size_bytes == 100352 - assert buffer_info_map["PaddedInput_6"].size_bytes == 172800 - - -# fmt: off -@tvm.script.ir_module -class MultipleCallsToSamePrimFuncModule: - @T.prim_func - def tvmgen_default_fused_layout_transform_1(placeholder: T.handle, T_layout_trans: T.handle) -> None: - # function attr dict - T.func_attr({"from_legacy_te_schedule": True, "global_symbol": "tvmgen_default_fused_layout_transform_1", "tir.noalias": True}) - placeholder_1 = T.match_buffer(placeholder, [864], dtype="float32") - T_layout_trans_1 = T.match_buffer(T_layout_trans, [41], dtype="float32") - # body - for ax0_ax1_fused_ax2_fused, ax3, ax4_inner in T.grid(24, 12, 3): - T_layout_trans_1[ax0_ax1_fused_ax2_fused * 36 + ax3 * 3 + ax4_inner] = placeholder_1[ax4_inner * 288 + ax0_ax1_fused_ax2_fused * 12 + ax3] - - @T.prim_func - def tvmgen_default_fused_nn_contrib_conv2d_NCHWc(placeholder_2: T.handle, placeholder_3: T.handle, conv2d_NCHWc: T.handle) -> None: - # function attr dict - T.func_attr({"from_legacy_te_schedule": True, "global_symbol": "tvmgen_default_fused_nn_contrib_conv2d_NCHWc", "tir.noalias": True}) - placeholder_4 = T.match_buffer(placeholder_2, [864], dtype="float32") - placeholder_5 = T.match_buffer(placeholder_3, [81], dtype="float32") - conv2d_NCHWc_1 = T.match_buffer(conv2d_NCHWc, [41], dtype="float32") - # body - data_pad = T.decl_buffer([1092], "float32") - for i0_i1_fused_i2_fused, i3, i4 in T.grid(26, 14, 3): - data_pad[i0_i1_fused_i2_fused * 42 + i3 * 3 + i4] = T.if_then_else(1 <= i0_i1_fused_i2_fused and i0_i1_fused_i2_fused < 25 and 1 <= i3 and i3 < 13, placeholder_4[i0_i1_fused_i2_fused * 36 + i3 * 3 + i4 - 39], T.float32(0), dtype="float32") - for n_oc_chunk_fused_oh_fused in T.serial(0, 24): - conv2d_NCHWc_global = T.decl_buffer([36], "float32") - for oc_block_c_init in T.serial(0, 3): - conv2d_NCHWc_global[oc_block_c_init] = T.float32(0) - for oc_block_c_init in T.serial(0, 3): - conv2d_NCHWc_global[oc_block_c_init + 3] = T.float32(0) - for oc_block_c_init in T.serial(0, 3): - conv2d_NCHWc_global[oc_block_c_init + 6] = T.float32(0) - for oc_block_c_init in T.serial(0, 3): - conv2d_NCHWc_global[oc_block_c_init + 9] = T.float32(0) - for oc_block_c_init in T.serial(0, 3): - conv2d_NCHWc_global[oc_block_c_init + 12] = T.float32(0) - for oc_block_c_init in T.serial(0, 3): - conv2d_NCHWc_global[oc_block_c_init + 15] = T.float32(0) - for oc_block_c_init in T.serial(0, 3): - conv2d_NCHWc_global[oc_block_c_init + 18] = T.float32(0) - for oc_block_c_init in T.serial(0, 3): - conv2d_NCHWc_global[oc_block_c_init + 21] = T.float32(0) - for oc_block_c_init in T.serial(0, 3): - conv2d_NCHWc_global[oc_block_c_init + 24] = T.float32(0) - for oc_block_c_init in T.serial(0, 3): - conv2d_NCHWc_global[oc_block_c_init + 27] = T.float32(0) - for oc_block_c_init in T.serial(0, 3): - conv2d_NCHWc_global[oc_block_c_init + 30] = T.float32(0) - for oc_block_c_init in T.serial(0, 3): - conv2d_NCHWc_global[oc_block_c_init + 33] = T.float32(0) - for kh, kw, ic_inner in T.grid(3, 3, 3): - for oc_block_c in T.serial(0, 3): - conv2d_NCHWc_global[oc_block_c] = conv2d_NCHWc_global[oc_block_c] + data_pad[kh * 42 + n_oc_chunk_fused_oh_fused * 42 + kw * 3 + ic_inner] * placeholder_5[kh * 27 + kw * 9 + ic_inner * 3 + oc_block_c] - for oc_block_c in T.serial(0, 3): - conv2d_NCHWc_global[oc_block_c + 3] = conv2d_NCHWc_global[oc_block_c + 3] + data_pad[kh * 42 + n_oc_chunk_fused_oh_fused * 42 + kw * 3 + ic_inner + 3] * placeholder_5[kh * 27 + kw * 9 + ic_inner * 3 + oc_block_c] - for oc_block_c in T.serial(0, 3): - conv2d_NCHWc_global[oc_block_c + 6] = conv2d_NCHWc_global[oc_block_c + 6] + data_pad[kh * 42 + n_oc_chunk_fused_oh_fused * 42 + kw * 3 + ic_inner + 6] * placeholder_5[kh * 27 + kw * 9 + ic_inner * 3 + oc_block_c] - for oc_block_c in T.serial(0, 3): - conv2d_NCHWc_global[oc_block_c + 9] = conv2d_NCHWc_global[oc_block_c + 9] + data_pad[kh * 42 + n_oc_chunk_fused_oh_fused * 42 + kw * 3 + ic_inner + 9] * placeholder_5[kh * 27 + kw * 9 + ic_inner * 3 + oc_block_c] - for oc_block_c in T.serial(0, 3): - conv2d_NCHWc_global[oc_block_c + 12] = conv2d_NCHWc_global[oc_block_c + 12] + data_pad[kh * 42 + n_oc_chunk_fused_oh_fused * 42 + kw * 3 + ic_inner + 12] * placeholder_5[kh * 27 + kw * 9 + ic_inner * 3 + oc_block_c] - for oc_block_c in T.serial(0, 3): - conv2d_NCHWc_global[oc_block_c + 15] = conv2d_NCHWc_global[oc_block_c + 15] + data_pad[kh * 42 + n_oc_chunk_fused_oh_fused * 42 + kw * 3 + ic_inner + 15] * placeholder_5[kh * 27 + kw * 9 + ic_inner * 3 + oc_block_c] - for oc_block_c in T.serial(0, 3): - conv2d_NCHWc_global[oc_block_c + 18] = conv2d_NCHWc_global[oc_block_c + 18] + data_pad[kh * 42 + n_oc_chunk_fused_oh_fused * 42 + kw * 3 + ic_inner + 18] * placeholder_5[kh * 27 + kw * 9 + ic_inner * 3 + oc_block_c] - for oc_block_c in T.serial(0, 3): - conv2d_NCHWc_global[oc_block_c + 21] = conv2d_NCHWc_global[oc_block_c + 21] + data_pad[kh * 42 + n_oc_chunk_fused_oh_fused * 42 + kw * 3 + ic_inner + 21] * placeholder_5[kh * 27 + kw * 9 + ic_inner * 3 + oc_block_c] - for oc_block_c in T.serial(0, 3): - conv2d_NCHWc_global[oc_block_c + 24] = conv2d_NCHWc_global[oc_block_c + 24] + data_pad[kh * 42 + n_oc_chunk_fused_oh_fused * 42 + kw * 3 + ic_inner + 24] * placeholder_5[kh * 27 + kw * 9 + ic_inner * 3 + oc_block_c] - for oc_block_c in T.serial(0, 3): - conv2d_NCHWc_global[oc_block_c + 27] = conv2d_NCHWc_global[oc_block_c + 27] + data_pad[kh * 42 + n_oc_chunk_fused_oh_fused * 42 + kw * 3 + ic_inner + 27] * placeholder_5[kh * 27 + kw * 9 + ic_inner * 3 + oc_block_c] - for oc_block_c in T.serial(0, 3): - conv2d_NCHWc_global[oc_block_c + 30] = conv2d_NCHWc_global[oc_block_c + 30] + data_pad[kh * 42 + n_oc_chunk_fused_oh_fused * 42 + kw * 3 + ic_inner + 30] * placeholder_5[kh * 27 + kw * 9 + ic_inner * 3 + oc_block_c] - for oc_block_c in T.serial(0, 3): - conv2d_NCHWc_global[oc_block_c + 33] = conv2d_NCHWc_global[oc_block_c + 33] + data_pad[kh * 42 + n_oc_chunk_fused_oh_fused * 42 + kw * 3 + ic_inner + 33] * placeholder_5[kh * 27 + kw * 9 + ic_inner * 3 + oc_block_c] - for ow_inner, oc_block in T.grid(12, 3): - conv2d_NCHWc_1[n_oc_chunk_fused_oh_fused * 36 + ow_inner * 3 + oc_block] = conv2d_NCHWc_global[ow_inner * 3 + oc_block] - - @T.prim_func - def tvmgen_default_fused_nn_softmax_add_add_multiply_add(placeholder_6: T.handle, placeholder_7: T.handle, placeholder_8: T.handle, placeholder_9: T.handle, placeholder_10: T.handle, T_add: T.handle) -> None: - # function attr dict - T.func_attr({"from_legacy_te_schedule": True, "global_symbol": "tvmgen_default_fused_nn_softmax_add_add_multiply_add", "tir.noalias": True}) - placeholder_11 = T.match_buffer(placeholder_6, [864], dtype="float32") - placeholder_12 = T.match_buffer(placeholder_7, [864], dtype="float32") - placeholder_13 = T.match_buffer(placeholder_8, [3], dtype="float32") - placeholder_14 = T.match_buffer(placeholder_9, [3], dtype="float32") - placeholder_15 = T.match_buffer(placeholder_10, [3], dtype="float32") - T_add_1 = T.match_buffer(T_add, [864], dtype="float32") - # body - for ax0_ax1_fused_ax2_fused in T.serial(0, 72): - T_softmax_norm = T.decl_buffer([12], "float32") - with T.decl_buffer([1], "float32") as T_softmax_maxelem: - T_softmax_maxelem[0] = T.float32(-3.4028234663852886e+38) - for k in T.serial(0, 12): - T_softmax_maxelem[0] = T.max(T_softmax_maxelem[0], placeholder_11[ax0_ax1_fused_ax2_fused * 12 + k]) - T_softmax_exp = T.decl_buffer([12], "float32") - for i3 in T.serial(0, 12): - T_softmax_exp[i3] = T.exp(placeholder_11[ax0_ax1_fused_ax2_fused * 12 + i3] - T_softmax_maxelem[0], dtype="float32") - T_softmax_expsum = T.decl_buffer([1], "float32") - T_softmax_expsum[0] = T.float32(0) - for k in T.serial(0, 12): - T_softmax_expsum[0] = T_softmax_expsum[0] + T_softmax_exp[k] - for i3 in T.serial(0, 12): - T_softmax_norm[i3] = T_softmax_exp[i3] / T_softmax_expsum[0] - for ax3 in T.serial(0, 12): - T_add_1[ax0_ax1_fused_ax2_fused * 12 + ax3] = (placeholder_12[ax0_ax1_fused_ax2_fused * 12 + ax3] + T_softmax_norm[ax3] + placeholder_13[T.floordiv(ax0_ax1_fused_ax2_fused, 24)]) * placeholder_14[T.floordiv(ax0_ax1_fused_ax2_fused, 24)] + placeholder_15[T.floordiv(ax0_ax1_fused_ax2_fused, 24)] - - @T.prim_func - def tvmgen_default_fused_nn_contrib_dense_pack_nn_relu(placeholder_16: T.handle, placeholder_17: T.handle, T_relu: T.handle) -> None: - # function attr dict - T.func_attr({"from_legacy_te_schedule": True, "global_symbol": "tvmgen_default_fused_nn_contrib_dense_pack_nn_relu", "tir.noalias": True}) - placeholder_18 = T.match_buffer(placeholder_16, [864], dtype="float32") - placeholder_19 = T.match_buffer(placeholder_17, [144], dtype="float32") - T_relu_1 = T.match_buffer(T_relu, [864], dtype="float32") - # body - for ax1_outer_ax0_outer_fused in T.serial(0, 18): - compute = T.decl_buffer([48], "float32") - with T.decl_buffer([48], "float32") as compute_global: - for x_c_init in T.serial(0, 6): - compute_global[x_c_init] = T.float32(0) - for x_c_init in T.serial(0, 6): - compute_global[x_c_init + 6] = T.float32(0) - for x_c_init in T.serial(0, 6): - compute_global[x_c_init + 12] = T.float32(0) - for x_c_init in T.serial(0, 6): - compute_global[x_c_init + 18] = T.float32(0) - for x_c_init in T.serial(0, 6): - compute_global[x_c_init + 24] = T.float32(0) - for x_c_init in T.serial(0, 6): - compute_global[x_c_init + 30] = T.float32(0) - for x_c_init in T.serial(0, 6): - compute_global[x_c_init + 36] = T.float32(0) - for x_c_init in T.serial(0, 6): - compute_global[x_c_init + 42] = T.float32(0) - for k_outer in T.serial(0, 12): - for x_c in T.serial(0, 6): - compute_global[x_c] = compute_global[x_c] + placeholder_18[T.floormod(ax1_outer_ax0_outer_fused, 9) * 96 + k_outer] * placeholder_19[T.floordiv(ax1_outer_ax0_outer_fused, 9) * 72 + k_outer * 6 + x_c] - for x_c in T.serial(0, 6): - compute_global[x_c + 6] = compute_global[x_c + 6] + placeholder_18[T.floormod(ax1_outer_ax0_outer_fused, 9) * 96 + k_outer + 12] * placeholder_19[T.floordiv(ax1_outer_ax0_outer_fused, 9) * 72 + k_outer * 6 + x_c] - for x_c in T.serial(0, 6): - compute_global[x_c + 12] = compute_global[x_c + 12] + placeholder_18[T.floormod(ax1_outer_ax0_outer_fused, 9) * 96 + k_outer + 24] * placeholder_19[T.floordiv(ax1_outer_ax0_outer_fused, 9) * 72 + k_outer * 6 + x_c] - for x_c in T.serial(0, 6): - compute_global[x_c + 18] = compute_global[x_c + 18] + placeholder_18[T.floormod(ax1_outer_ax0_outer_fused, 9) * 96 + k_outer + 36] * placeholder_19[T.floordiv(ax1_outer_ax0_outer_fused, 9) * 72 + k_outer * 6 + x_c] - for x_c in T.serial(0, 6): - compute_global[x_c + 24] = compute_global[x_c + 24] + placeholder_18[T.floormod(ax1_outer_ax0_outer_fused, 9) * 96 + k_outer + 48] * placeholder_19[T.floordiv(ax1_outer_ax0_outer_fused, 9) * 72 + k_outer * 6 + x_c] - for x_c in T.serial(0, 6): - compute_global[x_c + 30] = compute_global[x_c + 30] + placeholder_18[T.floormod(ax1_outer_ax0_outer_fused, 9) * 96 + k_outer + 60] * placeholder_19[T.floordiv(ax1_outer_ax0_outer_fused, 9) * 72 + k_outer * 6 + x_c] - for x_c in T.serial(0, 6): - compute_global[x_c + 36] = compute_global[x_c + 36] + placeholder_18[T.floormod(ax1_outer_ax0_outer_fused, 9) * 96 + k_outer + 72] * placeholder_19[T.floordiv(ax1_outer_ax0_outer_fused, 9) * 72 + k_outer * 6 + x_c] - for x_c in T.serial(0, 6): - compute_global[x_c + 42] = compute_global[x_c + 42] + placeholder_18[T.floormod(ax1_outer_ax0_outer_fused, 9) * 96 + k_outer + 84] * placeholder_19[T.floordiv(ax1_outer_ax0_outer_fused, 9) * 72 + k_outer * 6 + x_c] - for x_inner_inner in T.serial(0, 6): - compute[x_inner_inner] = compute_global[x_inner_inner] - for x_inner_inner in T.serial(0, 6): - compute[x_inner_inner + 6] = compute_global[x_inner_inner + 6] - for x_inner_inner in T.serial(0, 6): - compute[x_inner_inner + 12] = compute_global[x_inner_inner + 12] - for x_inner_inner in T.serial(0, 6): - compute[x_inner_inner + 18] = compute_global[x_inner_inner + 18] - for x_inner_inner in T.serial(0, 6): - compute[x_inner_inner + 24] = compute_global[x_inner_inner + 24] - for x_inner_inner in T.serial(0, 6): - compute[x_inner_inner + 30] = compute_global[x_inner_inner + 30] - for x_inner_inner in T.serial(0, 6): - compute[x_inner_inner + 36] = compute_global[x_inner_inner + 36] - for x_inner_inner in T.serial(0, 6): - compute[x_inner_inner + 42] = compute_global[x_inner_inner + 42] - for ax0_inner_inner, ax1_inner_inner in T.grid(8, 6): - T_relu_1[T.floormod(ax1_outer_ax0_outer_fused, 9) * 96 + ax0_inner_inner * 12 + T.floordiv(ax1_outer_ax0_outer_fused, 9) * 6 + ax1_inner_inner] = T.max(compute[ax0_inner_inner * 6 + ax1_inner_inner], T.float32(0)) - - @T.prim_func - def tvmgen_default_fused_reshape_1(placeholder_20: T.handle, T_reshape: T.handle) -> None: - # function attr dict - T.func_attr({"from_legacy_te_schedule": True, "global_symbol": "tvmgen_default_fused_reshape_1", "tir.noalias": True}) - placeholder_21 = T.match_buffer(placeholder_20, [864], dtype="float32") - T_reshape_1 = T.match_buffer(T_reshape, [864], dtype="float32") - # body - for ax0, ax1_inner in T.grid(72, 12): - T_reshape_1[ax0 * 12 + ax1_inner] = placeholder_21[ax0 * 12 + ax1_inner] - - @T.prim_func - def tvmgen_default_fused_layout_transform(placeholder_22: T.handle, T_layout_trans_2: T.handle) -> None: - # function attr dict - T.func_attr({"from_legacy_te_schedule": True, "global_symbol": "tvmgen_default_fused_layout_transform", "tir.noalias": True}) - placeholder_23 = T.match_buffer(placeholder_22, [864], dtype="float32") - T_layout_trans_3 = T.match_buffer(T_layout_trans_2, [864], dtype="float32") - # body - for ax0_ax1_fused, ax2, ax3_inner in T.grid(3, 24, 12): - T_layout_trans_3[ax0_ax1_fused * 288 + ax2 * 12 + ax3_inner] = placeholder_23[ax2 * 36 + ax3_inner * 3 + ax0_ax1_fused] - - @T.prim_func - def tvmgen_default_fused_reshape(placeholder_24: T.handle, T_reshape_2: T.handle) -> None: - # function attr dict - T.func_attr({"from_legacy_te_schedule": True, "global_symbol": "tvmgen_default_fused_reshape", "tir.noalias": True}) - placeholder_25 = T.match_buffer(placeholder_24, [864], dtype="float32") - T_reshape_3 = T.match_buffer(T_reshape_2, [864], dtype="float32") - # body - for ax0_ax1_fused, ax2, ax3_inner in T.grid(3, 24, 12): - T_reshape_3[ax0_ax1_fused * 288 + ax2 * 12 + ax3_inner] = placeholder_25[ax0_ax1_fused * 288 + ax2 * 12 + ax3_inner] - - @T.prim_func - def tvmgen_default_fused_nn_softmax_add(placeholder_26: T.handle, placeholder_27: T.handle, T_add_2: T.handle) -> None: - # function attr dict - T.func_attr({"from_legacy_te_schedule": True, "global_symbol": "tvmgen_default_fused_nn_softmax_add", "tir.noalias": True}) - placeholder_28 = T.match_buffer(placeholder_26, [864], dtype="float32") - placeholder_29 = T.match_buffer(placeholder_27, [864], dtype="float32") - T_add_3 = T.match_buffer(T_add_2, [864], dtype="float32") - # body - for ax0_ax1_fused_ax2_fused in T.serial(0, 72): - T_softmax_norm = T.decl_buffer([12], "float32") - with T.decl_buffer([1], "float32") as T_softmax_maxelem: - T_softmax_maxelem[0] = T.float32(-3.4028234663852886e+38) - for k in T.serial(0, 12): - T_softmax_maxelem[0] = T.max(T_softmax_maxelem[0], placeholder_28[ax0_ax1_fused_ax2_fused * 12 + k]) - T_softmax_exp= T.decl_buffer([12], "float32") - for i3 in T.serial(0, 12): - T_softmax_exp[i3] = T.exp(placeholder_28[ax0_ax1_fused_ax2_fused * 12 + i3] - T_softmax_maxelem[0], dtype="float32") - T_softmax_expsum = T.decl_buffer([1], "float32") - T_softmax_expsum[0] = T.float32(0) - for k in T.serial(0, 12): - T_softmax_expsum[0] = T_softmax_expsum[0] + T_softmax_exp[k] - for i3 in T.serial(0, 12): - T_softmax_norm[i3] = T_softmax_exp[i3] / T_softmax_expsum[0] - for ax3 in T.serial(0, 12): - T_add_3[ax0_ax1_fused_ax2_fused * 12 + ax3] = placeholder_29[ax0_ax1_fused_ax2_fused * 12 + ax3] + T_softmax_norm[ax3] - - @T.prim_func - def run_model(data: T.handle, output: T.handle) -> None: - # function attr dict - T.func_attr({"global_symbol": "tvmgen_default_run_model", "runner_function": True}) - data_buffer = T.match_buffer(data, [864], dtype="float32", align=16) - output_buffer = T.match_buffer(output, [864], dtype="float32", align=16) - # body - sid_11 = T.allocate([3456], "int8", "global.workspace") - sid_5 = T.allocate([3456], "int8", "global.workspace") - sid_10 = T.allocate([3456], "int8", "global.workspace") - sid_6 = T.allocate([3456], "int8", "global.workspace") - sid_8 = T.allocate([3456], "int8", "global.workspace") - sid_2 = T.allocate([3456], "int8", "global.workspace") - sid_7 = T.allocate([3456], "int8", "global.workspace") - sid_3 = T.allocate([3456], "int8", "global.workspace") - sid_12 = T.allocate([3456], "int8", "global.workspace") - sid_4 = T.allocate([3456], "int8", "global.workspace") - sid_18 = T.allocate([3456], "int8", "global.workspace") - sid_19 = T.allocate([3456], "int8", "global.workspace") - sid_20 = T.allocate([3456], "int8", "global.workspace") - - sid_21 = T.allocate_const([0,1,2,3,4,5,6,7,8,9], "int8", [10]) - sid_22 = T.allocate_const([1], "int8", [1]) - sid_23 = T.allocate_const([2,1], "int8", [3456]) - - T.evaluate(T.tvm_call_cpacked("tvmgen_default_fused_layout_transform_1", data_buffer.data, sid_23, dtype="int32")) - T.evaluate(T.tvm_call_cpacked("tvmgen_default_fused_nn_contrib_conv2d_NCHWc", sid_8, T.cast(T.lookup_param("p0", dtype="handle"), "handle"), sid_7, dtype="int32")) - T.evaluate(T.tvm_call_cpacked("tvmgen_default_fused_layout_transform", sid_7, sid_6, dtype="int32")) - T.evaluate(T.tvm_call_cpacked("tvmgen_default_fused_reshape_1", data_buffer.data, sid_12, dtype="int32")) - T.evaluate(T.tvm_call_cpacked("tvmgen_default_fused_nn_contrib_dense_pack_nn_relu", sid_12, T.cast(T.lookup_param("p1", dtype="handle"), "handle"), sid_11, dtype="int32")) - T.evaluate(T.tvm_call_cpacked("tvmgen_default_fused_reshape", sid_11, sid_10, dtype="int32")) - T.evaluate(T.tvm_call_cpacked("tvmgen_default_fused_nn_softmax_add_add_multiply_add", sid_6, sid_10, T.cast(T.lookup_param("p2", dtype="handle"), "handle"), T.cast(T.lookup_param("p3", dtype="handle"), "handle"), T.cast(T.lookup_param("p4", dtype="handle"), "handle"), sid_5, dtype="int32")) - T.evaluate(T.tvm_call_cpacked("tvmgen_default_fused_layout_transform_1", sid_5, sid_4, dtype="int32")) - T.evaluate(T.tvm_call_cpacked("tvmgen_default_fused_nn_contrib_conv2d_NCHWc", sid_4, T.cast(T.lookup_param("p5", dtype="handle"), "handle"), sid_3, dtype="int32")) - T.evaluate(T.tvm_call_cpacked("tvmgen_default_fused_layout_transform", sid_3, sid_2, dtype="int32")) - T.evaluate(T.tvm_call_cpacked("tvmgen_default_fused_reshape_1", sid_5, sid_20, dtype="int32")) - T.evaluate(T.tvm_call_cpacked("tvmgen_default_fused_nn_contrib_dense_pack_nn_relu", sid_20, T.cast(T.lookup_param("p6", dtype="handle"), "handle"), sid_19, dtype="int32")) - T.evaluate(T.tvm_call_cpacked("tvmgen_default_fused_reshape", sid_19, sid_18, dtype="int32")) - T.evaluate(T.tvm_call_cpacked("tvmgen_default_fused_nn_softmax_add", sid_2, sid_18, output_buffer.data, dtype="int32")) -# fmt: on - - -def test_multiple_calls_to_same_primfunc(): - target = Target("c") - global_ws_pool = WorkspacePoolInfo( - pool_name="global_workspace", - targets=[target], - ) - global_const_pool = ConstantPoolInfo( - pool_name="global_constants", - targets=[target], - ) - - tir_mod = MultipleCallsToSamePrimFuncModule - tir_mod = _assign_targets_to_primfuncs_irmodule(tir_mod, target) - tir_mod = _assign_poolinfos_to_allocates_in_irmodule( - tir_mod, [global_ws_pool], [global_const_pool] - ) - main_func = tir_mod["run_model"] - buffer_info_analysis = tvm.tir.usmp.analysis.extract_buffer_info(main_func, tir_mod) - assert buffer_info_analysis.memory_pressure == 11424 - buffer_info_map = _replace_stmt_with_buf_var_names(buffer_info_analysis.buffer_info_stmts) - - # check conflicts - _verify_conflicts("sid_23", ["sid_22", "sid_21"], buffer_info_map) - _verify_conflicts( - "sid_6", - [ - "sid_7", - "sid_12", - "compute", - "compute_global", - "sid_11", - "sid_10", - "T_softmax_exp", - "T_softmax_maxelem", - "sid_5", - "T_softmax_norm", - "T_softmax_expsum", - ], - buffer_info_map, - ) - _verify_conflicts( - "T_softmax_exp", - [ - "sid_10", - "sid_6", - "T_softmax_maxelem", - "sid_5", - "T_softmax_norm", - "T_softmax_expsum", - ], - buffer_info_map, - ) - _verify_conflicts( - "T_softmax_expsum2", - [ - "T_softmax_exp2", - "T_softmax_norm2", - "sid_18", - "T_softmax_maxelem2", - "sid_2", - ], - buffer_info_map, - ) - _verify_conflicts( - "compute", - [ - "sid_12", - "sid_6", - "compute_global", - "sid_11", - "sid_19", - "sid_20", - "sid_2", - "compute_global", - ], - buffer_info_map, - ) - _verify_conflicts( - "compute_global", - [ - "compute", - "sid_12", - "sid_6", - "sid_11", - "compute", - "sid_19", - "sid_20", - "sid_2", - ], - buffer_info_map, - ) - _verify_conflicts( - "sid_10", - [ - "sid_11", - "sid_6", - "T_softmax_exp", - "T_softmax_maxelem", - "sid_5", - "T_softmax_norm", - "T_softmax_expsum", - ], - buffer_info_map, - ) - _verify_conflicts( - "sid_2", - [ - "sid_3", - "sid_5", - "sid_20", - "sid_19", - "compute", - "compute_global", - "sid_18", - "T_softmax_norm2", - "T_softmax_exp2", - "T_softmax_maxelem2", - "T_softmax_expsum2", - ], - buffer_info_map, - ) - _verify_conflicts( - "sid_5", - [ - "T_softmax_maxelem", - "sid_10", - "T_softmax_exp", - "sid_6", - "T_softmax_norm", - "T_softmax_expsum", - "sid_4", - "data_pad", - "sid_3", - "conv2d_NCHWc_global", - "sid_2", - "sid_20", - ], - buffer_info_map, - ) - _verify_conflicts( - "T_softmax_norm2", - [ - "sid_18", - "sid_2", - "T_softmax_exp2", - "T_softmax_maxelem2", - "T_softmax_expsum2", - ], - buffer_info_map, - ) - _verify_conflicts( - "sid_20", - [ - "sid_2", - "sid_5", - "sid_19", - "compute", - "compute_global", - ], - buffer_info_map, - ) - _verify_conflicts( - "T_softmax_expsum", - [ - "sid_5", - "T_softmax_norm", - "T_softmax_maxelem", - "sid_10", - "T_softmax_exp", - "sid_6", - ], - buffer_info_map, - ) - _verify_conflicts( - "data_pad", - [ - "sid_8", - "conv2d_NCHWc_global", - "sid_7", - "sid_4", - "sid_5", - "sid_3", - "conv2d_NCHWc_global", - ], - buffer_info_map, - ) - _verify_conflicts( - "sid_19", - [ - "sid_20", - "sid_2", - "compute", - "compute_global", - "sid_18", - ], - buffer_info_map, - ) - _verify_conflicts( - "conv2d_NCHWc_global", - [ - "data_pad", - "sid_7", - "sid_3", - "data_pad", - "sid_5", - ], - buffer_info_map, - ) - _verify_conflicts( - "sid_18", - [ - "sid_19", - "sid_2", - "T_softmax_norm2", - "T_softmax_exp2", - "T_softmax_maxelem2", - "T_softmax_expsum2", - ], - buffer_info_map, - ) - _verify_conflicts( - "sid_7", - [ - "conv2d_NCHWc_global", - "data_pad", - "sid_6", - ], - buffer_info_map, - ) - _verify_conflicts( - "T_softmax_exp2", - [ - "T_softmax_norm2", - "sid_18", - "sid_2", - "T_softmax_maxelem2", - "T_softmax_expsum2", - ], - buffer_info_map, - ) - _verify_conflicts( - "sid_4", - [ - "sid_5", - "data_pad", - ], - buffer_info_map, - ) - _verify_conflicts( - "T_softmax_maxelem", - [ - "sid_10", - "T_softmax_exp", - "sid_6", - "sid_5", - "T_softmax_norm", - "T_softmax_expsum", - ], - buffer_info_map, - ) - _verify_conflicts( - "T_softmax_maxelem2", - [ - "T_softmax_exp2", - "T_softmax_norm2", - "sid_18", - "sid_2", - "T_softmax_expsum2", - ], - buffer_info_map, - ) - _verify_conflicts( - "sid_11", - [ - "compute", - "sid_12", - "compute_global", - "sid_6", - "sid_10", - ], - buffer_info_map, - ) - _verify_conflicts( - "sid_12", - [ - "sid_6", - "compute", - "compute_global", - "sid_11", - ], - buffer_info_map, - ) - _verify_conflicts( - "T_softmax_norm", - [ - "sid_5", - "T_softmax_maxelem", - "sid_10", - "T_softmax_exp", - "sid_6", - "T_softmax_expsum", - ], - buffer_info_map, - ) - _verify_conflicts( - "sid_8", - [ - "data_pad", - ], - buffer_info_map, - ) - - -if __name__ == "__main__": - tvm.testing.main() diff --git a/tests/python/tir-usmp/test_tir_usmp_transform_convert_pool_allocations_to_offsets.py b/tests/python/tir-usmp/test_tir_usmp_transform_convert_pool_allocations_to_offsets.py deleted file mode 100644 index 9e9fea7c8152..000000000000 --- a/tests/python/tir-usmp/test_tir_usmp_transform_convert_pool_allocations_to_offsets.py +++ /dev/null @@ -1,695 +0,0 @@ -# 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 sys - -import pytest -import tvm -from tvm import PoolInfoProperties, WorkspacePoolInfo -from tvm.script import tir as T, ir as I -from tvm.target import Target -from tvm.tir import stmt_functor -from tvm.tir.usmp import utils as usmp_utils - - -def _get_primfuncs_from_module(module): - primfuncs = list() - for gv, primfunc in module.functions.items(): - primfuncs.append(primfunc) - return primfuncs - - -def assign_poolinfos_to_allocates_in_primfunc(primfunc, pool_infos): - """Helper to assign poolinfos to allocate nodes in a tir.PrimFunc""" - - def set_poolinfos(stmt): - if isinstance(stmt, tvm.tir.Allocate): - return tvm.tir.Allocate( - buffer_var=stmt.buffer_var, - dtype=stmt.dtype, - extents=stmt.extents, - condition=stmt.condition, - body=stmt.body, - annotations={tvm.tir.usmp.utils.CANDIDATE_MEMORY_POOL_ATTR: pool_infos}, - ) - - return primfunc.with_body(stmt_functor.ir_transform(primfunc.body, None, set_poolinfos)) - - -def assign_poolinfos_to_allocates_in_irmodule(mod, pool_infos): - """Helper to assign poolinfos to allocate nodes in a IRModule""" - ret = tvm.IRModule() - for global_var, basefunc in mod.functions.items(): - if isinstance(basefunc, tvm.tir.PrimFunc): - ret[global_var] = assign_poolinfos_to_allocates_in_primfunc(basefunc, pool_infos) - return ret - - -def _assign_targets_to_primfuncs_irmodule(mod, target): - """Helper to assign target for PrimFunc in a IRModule""" - ret = tvm.IRModule() - for global_var, basefunc in mod.functions.items(): - if isinstance(basefunc, tvm.tir.PrimFunc): - ret[global_var] = basefunc.with_attr("target", target) - return ret - - -def _plan_and_convert(tir_mod, pools=None): - target = Target("c") - - if pools is None: - pools = [ - WorkspacePoolInfo( - "global_workspace", - [target], - ) - ] - - tir_mod = _assign_targets_to_primfuncs_irmodule(tir_mod, target) - tir_mod = assign_poolinfos_to_allocates_in_irmodule(tir_mod, pools) - main_func = tir_mod["__tvm_main__"] - buffer_analysis = tvm.tir.usmp.analysis.extract_buffer_info(main_func, tir_mod) - buffer_info_map = buffer_analysis.buffer_info_stmts - - fcreate_array_bi = tvm.get_global_func("tir.usmp.CreateArrayBufferInfo") - buffer_info_arr = fcreate_array_bi(buffer_info_map) - fusmp_algo_greedy_by_size = tvm.get_global_func("tir.usmp.algo.greedy_by_size") - buffer_pool_allocations = fusmp_algo_greedy_by_size( - buffer_info_arr, buffer_analysis.memory_pressure - ) - fassign_stmt_pool_allocations = tvm.get_global_func("tir.usmp.AssignStmtPoolAllocations") - pool_allocations = fassign_stmt_pool_allocations(buffer_info_map, buffer_pool_allocations) - tir_mod_with_offsets = tvm.tir.usmp.transform.convert_pool_allocations_to_offsets( - pool_allocations, emit_tvmscript_printable=True - )(tir_mod) - - return tir_mod_with_offsets - - -# fmt: off -@tvm.script.ir_module -class LinearStructure: - @T.prim_func - def tvmgen_default_fused_cast_subtract(placeholder_2: T.handle, placeholder_3: T.handle, T_subtract: T.handle) -> None: - # function attr dict - T.func_attr({"global_symbol": "tvmgen_default_fused_cast_subtract", "tir.noalias": True}) - placeholder_4 = T.match_buffer(placeholder_2, [150528], dtype="uint8", elem_offset=0, align=64, offset_factor=1) - placeholder_5 = T.match_buffer(placeholder_3, [1], dtype="int16", elem_offset=0, align=64, offset_factor=1) - T_subtract_1 = T.match_buffer(T_subtract, [452], dtype="int16", elem_offset=0, align=64, offset_factor=1) - # body - for ax0_ax1_fused_1 in T.serial(0, 224): - for ax2_1, ax3_inner_1 in T.grid(224, 3): - T_subtract_1[(((ax0_ax1_fused_1*672) + (ax2_1*3)) + ax3_inner_1)] = (T.cast(placeholder_4[(((ax0_ax1_fused_1*672) + (ax2_1*3)) + ax3_inner_1)], "int16") - placeholder_5[0]) - - @T.prim_func - def tvmgen_default_fused_nn_conv2d_add_fixed_point_multiply_clip_cast(placeholder_62: T.handle, placeholder_63: T.handle, placeholder_64: T.handle, T_cast_20: T.handle) -> None: - # function attr dict - T.func_attr({"global_symbol": "tvmgen_default_fused_nn_conv2d_add_fixed_point_multiply_clip_cast", "tir.noalias": True}) - placeholder_65 = T.match_buffer(placeholder_62, [150528], dtype="int16", elem_offset=0, align=64, offset_factor=1) - placeholder_66 = T.match_buffer(placeholder_63, [9408], dtype="int16", elem_offset=0, align=64, offset_factor=1) - placeholder_67 = T.match_buffer(placeholder_64, [64], dtype="int32", elem_offset=0, align=64, offset_factor=1) - T_cast_21 = T.match_buffer(T_cast_20, [289], dtype="uint8", elem_offset=0, align=64, offset_factor=1) - # body - PaddedInput_7_data = T.allocate([157323], "int16", "global") - PaddedInput_7 = T.Buffer(shape=[157323], dtype="int16", data=PaddedInput_7_data) - for i0_i1_fused_7 in T.serial(0, 229): - for i2_7, i3_7 in T.grid(229, 3): - PaddedInput_7[(((i0_i1_fused_7*687) + (i2_7*3)) + i3_7)] = T.if_then_else(((((2 <= i0_i1_fused_7) and (i0_i1_fused_7 < 226)) and (2 <= i2_7)) and (i2_7 < 226)), placeholder_65[((((i0_i1_fused_7*672) + (i2_7*3)) + i3_7) - 1350)], T.int16(0), dtype="int16") - for ax0_ax1_fused_ax2_fused_7 in T.serial(0, 12544): - Conv2dOutput_7_data = T.allocate([64], "int32", "global") - Conv2dOutput_7 = T.Buffer(shape=[64], dtype="int32", data=Conv2dOutput_7_data) - for ff_3 in T.serial(0, 64): - Conv2dOutput_7[ff_3] = 0 - for ry_2, rx_2, rc_7 in T.grid(7, 7, 3): - Conv2dOutput_7[ff_3] = (Conv2dOutput_7[ff_3] + (T.cast(PaddedInput_7[(((((T.floordiv(ax0_ax1_fused_ax2_fused_7, 112)*1374) + (ry_2*687)) + (T.floormod(ax0_ax1_fused_ax2_fused_7, 112)*6)) + (rx_2*3)) + rc_7)], "int32")*T.cast(placeholder_66[((((ry_2*1344) + (rx_2*192)) + (rc_7*64)) + ff_3)], "int32"))) - for ax3_inner_7 in T.serial(0, 64): - T_cast_21[((ax0_ax1_fused_ax2_fused_7*64) + ax3_inner_7)] = T.cast(T.max(T.min(T.q_multiply_shift((Conv2dOutput_7[ax3_inner_7] + placeholder_67[ax3_inner_7]), 1939887962, 31, -9, dtype="int32"), 255), 0), "uint8") - - @T.prim_func - def tvmgen_default_fused_nn_max_pool2d_cast(placeholder_28: T.handle, T_cast_6: T.handle) -> None: - # function attr dict - T.func_attr({"global_symbol": "tvmgen_default_fused_nn_max_pool2d_cast", "tir.noalias": True}) - placeholder_29 = T.match_buffer(placeholder_28, [802816], dtype="uint8", elem_offset=0, align=64, offset_factor=1) - T_cast_7 = T.match_buffer(T_cast_6, [177], dtype="int16", elem_offset=0, align=64, offset_factor=1) - # body - tensor_2_data = T.allocate([200704], "uint8", "global") - tensor_2 = T.Buffer(shape=[200704], dtype="uint8", data=tensor_2_data) - for ax0_ax1_fused_4 in T.serial(0, 56): - for ax2_4 in T.serial(0, 56): - for ax3_init in T.serial(0, 64): - tensor_2[(((ax0_ax1_fused_4*3584) + (ax2_4*64)) + ax3_init)] = T.uint8(0) - for rv0_rv1_fused_1, ax3_2 in T.grid(9, 64): - tensor_2[(((ax0_ax1_fused_4*3584) + (ax2_4*64)) + ax3_2)] = T.max(tensor_2[(((ax0_ax1_fused_4*3584) + (ax2_4*64)) + ax3_2)], T.if_then_else(((((ax0_ax1_fused_4*2) + T.floordiv(rv0_rv1_fused_1, 3)) < 112) and (((ax2_4*2) + T.floormod(rv0_rv1_fused_1, 3)) < 112)), placeholder_29[(((((ax0_ax1_fused_4*14336) + (T.floordiv(rv0_rv1_fused_1, 3)*7168)) + (ax2_4*128)) + (T.floormod(rv0_rv1_fused_1, 3)*64)) + ax3_2)], T.uint8(0), dtype="uint8")) - for ax0_ax1_fused_5 in T.serial(0, 56): - for ax2_5, ax3_3 in T.grid(56, 64): - T_cast_7[(((ax0_ax1_fused_5*3584) + (ax2_5*64)) + ax3_3)] = T.cast(tensor_2[(((ax0_ax1_fused_5*3584) + (ax2_5*64)) + ax3_3)], "int16") - - @T.prim_func - def __tvm_main__(input: T.handle, output: T.handle) -> None: - # function attr dict - T.func_attr({"global_symbol": "__tvm_main__", "runner_function": True}) - # body - T.attr("default", "device_id", 0) - T.attr("default", "device_type", 1) - sid_9 = T.allocate([301056], "int8", "global") - sid_8 = T.allocate([802816], "int8", "global") - T.evaluate(T.call_extern("tvmgen_default_fused_cast_subtract", input, T.lookup_param("p0", dtype="handle"), sid_9, dtype="int32")) - T.evaluate(T.call_extern("tvmgen_default_fused_nn_conv2d_add_fixed_point_multiply_clip_cast", sid_9, T.lookup_param("p1", dtype="handle"), T.lookup_param("p2", dtype="handle"), sid_8, dtype="int32")) - T.evaluate(T.call_extern("tvmgen_default_fused_nn_max_pool2d_cast", sid_8, output, dtype="int32")) -# fmt: on - - -# fmt: off -@tvm.script.ir_module -class LinearStructurePlanned: - @T.prim_func - def __tvm_main__(input: T.handle, fast_memory_0_var: T.handle("uint8"), slow_memory_1_var: T.handle("uint8"), output: T.handle) -> None: - fast_memory_0_buffer_var = T.match_buffer(fast_memory_0_var, [200704], dtype="uint8", strides=[1], elem_offset=0, align=16) - slow_memory_1_buffer_var = T.match_buffer(slow_memory_1_var, [1418528], dtype="uint8", strides=[1], elem_offset=0, align=16) - # body - T.attr("default", "device_id", 0) - T.attr("default", "device_type", 1) - sid_9_let: T.handle("int8") = T.address_of(slow_memory_1_buffer_var[1117472], dtype="handle") - sid_8_let: T.handle("int8") = T.address_of(slow_memory_1_buffer_var[0], dtype="handle") - T.evaluate(T.call_extern("tvmgen_default_fused_cast_subtract", input, T.lookup_param("p0", dtype="handle"), sid_9_let, fast_memory_0_buffer_var.data, slow_memory_1_buffer_var.data, dtype="int32")) - T.evaluate(T.call_extern("tvmgen_default_fused_nn_conv2d_add_fixed_point_multiply_clip_cast", sid_9_let, T.lookup_param("p1", dtype="handle"), T.lookup_param("p2", dtype="handle"), sid_8_let, fast_memory_0_buffer_var.data, slow_memory_1_buffer_var.data, dtype="int32")) - T.evaluate(T.call_extern("tvmgen_default_fused_nn_max_pool2d_cast", sid_8_let, output, fast_memory_0_buffer_var.data, slow_memory_1_buffer_var.data, dtype="int32")) - - @T.prim_func - def tvmgen_default_fused_nn_max_pool2d_cast(placeholder_28: T.handle, T_cast_6: T.handle, fast_memory_6_var: T.handle("uint8"), slow_memory_7_var: T.handle("uint8")) -> None: - placeholder_29 = T.match_buffer(placeholder_28, [802816], dtype="uint8") - T_cast_7 = T.match_buffer(T_cast_6, [177], dtype="int16") - fast_memory_6_buffer_var = T.match_buffer(fast_memory_6_var, [200704], dtype="uint8", strides=[1], elem_offset=0, align=16) - slow_memory_7_buffer_var = T.match_buffer(slow_memory_7_var, [1418528], dtype="uint8", strides=[1], elem_offset=0, align=16) - # body - tensor_2_let = T.Buffer([200704], dtype="uint8") - with T.LetStmt(T.address_of(fast_memory_6_buffer_var[0], dtype="handle"), var=tensor_2_let.data): - for ax0_ax1_fused_4, ax2_4 in T.grid(56, 56): - for ax3_init in T.serial(0, 64): - tensor_2_let[ax0_ax1_fused_4 * 3584 + ax2_4 * 64 + ax3_init] = T.uint8(0) - for rv0_rv1_fused_1, ax3_2 in T.grid(9, 64): - tensor_2_let[ax0_ax1_fused_4 * 3584 + ax2_4 * 64 + ax3_2] = T.max(tensor_2_let[ax0_ax1_fused_4 * 3584 + ax2_4 * 64 + ax3_2], T.if_then_else(ax0_ax1_fused_4 * 2 + rv0_rv1_fused_1 // 3 < 112 and ax2_4 * 2 + rv0_rv1_fused_1 % 3 < 112, placeholder_29[ax0_ax1_fused_4 * 14336 + rv0_rv1_fused_1 // 3 * 7168 + ax2_4 * 128 + rv0_rv1_fused_1 % 3 * 64 + ax3_2], T.uint8(0), dtype="uint8")) - for ax0_ax1_fused_5, ax2_5, ax3_3 in T.grid(56, 56, 64): - T_cast_7[ax0_ax1_fused_5 * 3584 + ax2_5 * 64 + ax3_3] = T.cast(tensor_2_let[ax0_ax1_fused_5 * 3584 + ax2_5 * 64 + ax3_3], "int16") - - @T.prim_func - def tvmgen_default_fused_cast_subtract(placeholder_2: T.handle, placeholder_3: T.handle, T_subtract: T.handle, fast_memory_2_var: T.handle("uint8"), slow_memory_3_var: T.handle("uint8")) -> None: - placeholder_4 = T.match_buffer(placeholder_2, [150528], dtype="uint8") - placeholder_5 = T.match_buffer(placeholder_3, [1], dtype="int16") - T_subtract_1 = T.match_buffer(T_subtract, [452], dtype="int16") - fast_memory_2_buffer_var = T.match_buffer(fast_memory_2_var, [200704], dtype="uint8", strides=[1], elem_offset=0, align=16) - slow_memory_3_buffer_var = T.match_buffer(slow_memory_3_var, [1418528], dtype="uint8", strides=[1], elem_offset=0, align=16) - # body - for ax0_ax1_fused_1, ax2_1, ax3_inner_1 in T.grid(224, 224, 3): - T_subtract_1[ax0_ax1_fused_1 * 672 + ax2_1 * 3 + ax3_inner_1] = T.cast(placeholder_4[ax0_ax1_fused_1 * 672 + ax2_1 * 3 + ax3_inner_1], "int16") - placeholder_5[0] - - @T.prim_func - def tvmgen_default_fused_nn_conv2d_add_fixed_point_multiply_clip_cast(placeholder_62: T.handle, placeholder_63: T.handle, placeholder_64: T.handle, T_cast_20: T.handle, fast_memory_4_var: T.handle("uint8"), slow_memory_5_var: T.handle("uint8")) -> None: - placeholder_65 = T.match_buffer(placeholder_62, [150528], dtype="int16") - placeholder_66 = T.match_buffer(placeholder_63, [9408], dtype="int16") - placeholder_67 = T.match_buffer(placeholder_64, [64], dtype="int32") - T_cast_21 = T.match_buffer(T_cast_20, [289], dtype="uint8") - fast_memory_4_buffer_var = T.match_buffer(fast_memory_4_var, [200704], dtype="uint8", strides=[1], elem_offset=0, align=16) - slow_memory_5_buffer_var = T.match_buffer(slow_memory_5_var, [1418528], dtype="uint8", strides=[1], elem_offset=0, align=16) - # body - PaddedInput_7_let = T.Buffer([157323], "int16") - with T.LetStmt(T.address_of(slow_memory_5_buffer_var[802816], dtype="handle"), var=PaddedInput_7_let.data): - for i0_i1_fused_7, i2_7, i3_7 in T.grid(229, 229, 3): - PaddedInput_7_let[i0_i1_fused_7 * 687 + i2_7 * 3 + i3_7] = T.if_then_else(2 <= i0_i1_fused_7 and i0_i1_fused_7 < 226 and 2 <= i2_7 and i2_7 < 226, placeholder_65[i0_i1_fused_7 * 672 + i2_7 * 3 + i3_7 - 1350], T.int16(0), dtype="int16") - for ax0_ax1_fused_ax2_fused_7 in T.serial(0, 12544): - Conv2dOutput_7_let = T.Buffer([64], "int32") - with T.LetStmt(T.address_of(fast_memory_4_buffer_var[0], dtype="handle"), var=Conv2dOutput_7_let.data): - for ff_3 in T.serial(0, 64): - Conv2dOutput_7_let[ff_3] = 0 - for ry_2, rx_2, rc_7 in T.grid(7, 7, 3): - Conv2dOutput_7_let[ff_3] = Conv2dOutput_7_let[ff_3] + T.cast(PaddedInput_7_let[ax0_ax1_fused_ax2_fused_7 // 112 * 1374 + ry_2 * 687 + ax0_ax1_fused_ax2_fused_7 % 112 * 6 + rx_2 * 3 + rc_7], "int32") * T.cast(placeholder_66[ry_2 * 1344 + rx_2 * 192 + rc_7 * 64 + ff_3], "int32") - for ax3_inner_7 in T.serial(0, 64): - T_cast_21[ax0_ax1_fused_ax2_fused_7 * 64 + ax3_inner_7] = T.cast(T.max(T.min(T.q_multiply_shift(Conv2dOutput_7_let[ax3_inner_7] + placeholder_67[ax3_inner_7], 1939887962, 31, -9, dtype="int32"), 255), 0), "uint8") -# fmt: on - - -def test_mobilenet_subgraph(): - before = LinearStructure - - expected = LinearStructurePlanned - - target = Target("c") - pools = [ - WorkspacePoolInfo( - "fast_memory", - [target], - PoolInfoProperties(size_hint_bytes=200704), - ), - WorkspacePoolInfo( - "slow_memory", - [target], - ), - ] - after = _plan_and_convert(before, pools=pools) - tvm.ir.assert_structural_equal(after, expected) - - -# fmt: off -@tvm.script.ir_module -class ResnetStructure: - @T.prim_func - def tvmgen_default_fused_cast_subtract_fixed_point_multiply_add_clip_cast_cast(placeholder: T.handle, placeholder_1: T.handle, T_cast: T.handle) -> None: - # function attr dict - T.func_attr({"global_symbol": "tvmgen_default_fused_cast_subtract_fixed_point_multiply_add_clip_cast_cast", "tir.noalias": True}) - placeholder_2 = T.match_buffer(placeholder, [360000], dtype="uint8") - placeholder_3 = T.match_buffer(placeholder_1, [64], dtype="int32") - T_cast_1 = T.match_buffer(T_cast, [215], dtype="int16") - # body - for ax0_ax1_fused, ax2, ax3_outer, ax3_inner in T.grid(75, 75, 4, 16): - T_cast_1[ax0_ax1_fused * 4800 + ax2 * 64 + ax3_outer * 16 + ax3_inner] = T.cast(T.cast(T.max(T.min(T.q_multiply_shift(T.cast(placeholder_2[ax0_ax1_fused * 4800 + ax2 * 64 + ax3_outer * 16 + ax3_inner], "int32") - 94, 1843157232, 31, 1, dtype="int32") + placeholder_3[ax3_outer * 16 + ax3_inner], 255), 0), "uint8"), "int16") - - @T.prim_func - def tvmgen_default_fused_nn_conv2d_add_fixed_point_multiply_clip_cast_cast_1(placeholder_10: T.handle, placeholder_11: T.handle, placeholder_12: T.handle, T_cast_4: T.handle) -> None: - # function attr dict - T.func_attr({"global_symbol": "tvmgen_default_fused_nn_conv2d_add_fixed_point_multiply_clip_cast_cast_1", "tir.noalias": True}) - placeholder_13 = T.match_buffer(placeholder_10, [360000], dtype="int16") - placeholder_14 = T.match_buffer(placeholder_11, [36864], dtype="int16") - placeholder_15 = T.match_buffer(placeholder_12, [64], dtype="int32") - T_cast_5 = T.match_buffer(T_cast_4, [215], dtype="int16") - # body - PaddedInput_1_data = T.allocate([379456], "int16", "global") - PaddedInput_1 = T.Buffer(shape=[379456], dtype="int16", data=PaddedInput_1_data) - for i0_i1_fused_1, i2_1, i3_1 in T.grid(77, 77, 64): - PaddedInput_1[i0_i1_fused_1 * 4928 + i2_1 * 64 + i3_1] = T.if_then_else(1 <= i0_i1_fused_1 and i0_i1_fused_1 < 76 and 1 <= i2_1 and i2_1 < 76, placeholder_13[i0_i1_fused_1 * 4800 + i2_1 * 64 + i3_1 - 4864], T.int16(0), dtype="int16") - for ax0_ax1_fused_ax2_fused_1 in T.serial(0, 5625): - Conv2dOutput_1_data = T.allocate([64], "int32", "global") - Conv2dOutput_1 = T.Buffer(shape=[64], dtype="int32", data=Conv2dOutput_1_data) - for ff_1 in T.serial(0, 64): - Conv2dOutput_1[ff_1] = 0 - for ry, rx, rc_1 in T.grid(3, 3, 64): - Conv2dOutput_1[ff_1] = Conv2dOutput_1[ff_1] + T.cast(PaddedInput_1[T.floordiv(ax0_ax1_fused_ax2_fused_1, 75) * 4928 + ry * 4928 + rx * 64 + T.floormod(ax0_ax1_fused_ax2_fused_1, 75) * 64 + rc_1], "int32") * T.cast(placeholder_14[ry * 12288 + rx * 4096 + rc_1 * 64 + ff_1], "int32") - for ax3_inner_2 in T.serial(0, 64): - T_cast_5[ax0_ax1_fused_ax2_fused_1 * 64 + ax3_inner_2] = T.cast(T.cast(T.max(T.min(T.q_multiply_shift(Conv2dOutput_1[ax3_inner_2] + placeholder_15[ax3_inner_2], 1608879842, 31, -7, dtype="int32"), 255), 0), "uint8"), "int16") - - @T.prim_func - def tvmgen_default_fused_nn_conv2d_add_fixed_point_multiply_add_clip_cast_cast_subtract_fixed_point_15934180698220515269_(placeholder_16: T.handle, placeholder_17: T.handle, placeholder_18: T.handle, T_add: T.handle) -> None: - # function attr dict - T.func_attr({"global_symbol": "tvmgen_default_fused_nn_conv2d_add_fixed_point_multiply_add_clip_cast_cast_subtract_fixed_point_15934180698220515269_", "tir.noalias": True}) - placeholder_19 = T.match_buffer(placeholder_16, [360000], dtype="int16") - placeholder_20 = T.match_buffer(placeholder_17, [16384], dtype="int16") - placeholder_21 = T.match_buffer(placeholder_18, [256], dtype="int32") - T_add_1 = T.match_buffer(T_add, [407], dtype="int32") - # body - PaddedInput_2_data = T.allocate([360000], "int16", "global") - PaddedInput_2 = T.Buffer(shape=[360000], dtype="int16", data=PaddedInput_2_data) - for i0_i1_fused_2, i2_2, i3_2 in T.grid(75, 75, 64): - PaddedInput_2[i0_i1_fused_2 * 4800 + i2_2 * 64 + i3_2] = placeholder_19[i0_i1_fused_2 * 4800 + i2_2 * 64 + i3_2] - for ax0_ax1_fused_ax2_fused_2 in T.serial(0, 5625): - Conv2dOutput_2_data = T.allocate([64], "int32", "global") - Conv2dOutput_2 = T.Buffer(shape=[64], dtype="int32", data=Conv2dOutput_2_data) - for ax3_outer_1 in T.serial(0, 4): - for ff_2 in T.serial(0, 64): - Conv2dOutput_2[ff_2] = 0 - for rc_2 in T.serial(0, 64): - Conv2dOutput_2[ff_2] = Conv2dOutput_2[ff_2] + T.cast(PaddedInput_2[ax0_ax1_fused_ax2_fused_2 * 64 + rc_2], "int32") * T.cast(placeholder_20[rc_2 * 256 + ax3_outer_1 * 64 + ff_2], "int32") - for ax3_inner_3 in T.serial(0, 64): - T_add_1[ax0_ax1_fused_ax2_fused_2 * 256 + ax3_outer_1 * 64 + ax3_inner_3] = T.q_multiply_shift(T.cast(T.cast(T.max(T.min(T.q_multiply_shift(Conv2dOutput_2[ax3_inner_3] + placeholder_21[ax3_outer_1 * 64 + ax3_inner_3], 1711626602, 31, -8, dtype="int32") + 132, 255), 0), "uint8"), "int32") - 132, 2094289803, 31, -2, dtype="int32") + 136 - - @T.prim_func - def tvmgen_default_fused_nn_conv2d_add_fixed_point_multiply_add_clip_cast_cast_subtract_fixed_point_4200876283395191415_(placeholder_22: T.handle, placeholder_23: T.handle, placeholder_24: T.handle, placeholder_25: T.handle, T_cast_6: T.handle) -> None: - # function attr dict - T.func_attr({"global_symbol": "tvmgen_default_fused_nn_conv2d_add_fixed_point_multiply_add_clip_cast_cast_subtract_fixed_point_4200876283395191415_", "tir.noalias": True}) - placeholder_29 = T.match_buffer(placeholder_22, [360000], dtype="int16") - placeholder_27 = T.match_buffer(placeholder_23, [16384], dtype="int16") - placeholder_26 = T.match_buffer(placeholder_24, [256], dtype="int32") - placeholder_28 = T.match_buffer(placeholder_25, [1440000], dtype="int32") - T_cast_7 = T.match_buffer(T_cast_6, [407], dtype="uint8") - # body - PaddedInput_3_data = T.allocate([360000], "int16", "global") - PaddedInput_3 = T.Buffer(shape=[360000], dtype="int16", data=PaddedInput_3_data) - for i0_i1_fused_3, i2_3, i3_3 in T.grid(75, 75, 64): - PaddedInput_3[i0_i1_fused_3 * 4800 + i2_3 * 64 + i3_3] = placeholder_29[i0_i1_fused_3 * 4800 + i2_3 * 64 + i3_3] - for ax0_ax1_fused_ax2_fused_3 in T.serial(0, 5625): - Conv2dOutput_3_data = T.allocate([64], "int32", "global") - Conv2dOutput_3 = T.Buffer(shape=[64], dtype="int32", data=Conv2dOutput_3_data) - for ax3_outer_2 in T.serial(0, 4): - for ff_3 in T.serial(0, 64): - Conv2dOutput_3[ff_3] = 0 - for rc_3 in T.serial(0, 64): - Conv2dOutput_3[ff_3] = Conv2dOutput_3[ff_3] + T.cast(PaddedInput_3[ax0_ax1_fused_ax2_fused_3 * 64 + rc_3], "int32") * T.cast(placeholder_27[rc_3 * 256 + ax3_outer_2 * 64 + ff_3], "int32") - for ax3_inner_4 in T.serial(0, 64): - T_cast_7[ax0_ax1_fused_ax2_fused_3 * 256 + ax3_outer_2 * 64 + ax3_inner_4] = T.cast(T.max(T.min(T.q_multiply_shift(T.cast(T.cast(T.max(T.min(T.q_multiply_shift(Conv2dOutput_3[ax3_inner_4] + placeholder_26[ax3_outer_2 * 64 + ax3_inner_4], 1343014664, 31, -8, dtype="int32") + 136, 255), 0), "uint8"), "int32") - 136, 1073903788, 31, 1, dtype="int32") + placeholder_28[ax0_ax1_fused_ax2_fused_3 * 256 + ax3_outer_2 * 64 + ax3_inner_4], 255), 0), "uint8") - - @T.prim_func - def __tvm_main__(input: T.handle, output: T.handle) -> None: - # function attr dict - T.func_attr({"global_symbol": "__tvm_main__", "runner_function": True}) - # body - T.attr("default", "device_id", 0) - T.attr("default", "device_type", 1) - sid_2 = T.allocate([720000], "int8", "global") - sid_6 = T.allocate([5760000], "int8", "global") - sid_7 = T.allocate([720000], "int8", "global") - sid_8 = T.allocate([720000], "int8", "global") - T.evaluate(T.call_extern("tvmgen_default_fused_cast_subtract_fixed_point_multiply_add_clip_cast_cast", input, T.lookup_param("p0", dtype="handle"), sid_2, dtype="int32")) - T.evaluate(T.call_extern("tvmgen_default_fused_nn_conv2d_add_fixed_point_multiply_clip_cast_cast", sid_2, T.lookup_param("p3", dtype="handle"), T.lookup_param("p4", dtype="handle"), sid_8, dtype="int32")) - T.evaluate(T.call_extern("tvmgen_default_fused_nn_conv2d_add_fixed_point_multiply_clip_cast_cast_1", sid_8, T.lookup_param("p5", dtype="handle"), T.lookup_param("p6", dtype="handle"), sid_7, dtype="int32")) - T.evaluate(T.call_extern("tvmgen_default_fused_nn_conv2d_add_fixed_point_multiply_add_clip_cast_cast_subtract_fixed_point_15934180698220515269_", sid_7, T.lookup_param("p7", dtype="handle"), T.lookup_param("p8", dtype="handle"), sid_6, dtype="int32")) - T.evaluate(T.call_extern("tvmgen_default_fused_nn_conv2d_add_fixed_point_multiply_add_clip_cast_cast_subtract_fixed_point_4200876283395191415_", sid_2, T.lookup_param("p1", dtype="handle"), T.lookup_param("p2", dtype="handle"), sid_6, output, dtype="int32")) - - @T.prim_func - def tvmgen_default_fused_nn_conv2d_add_fixed_point_multiply_clip_cast_cast(placeholder_4: T.handle, placeholder_5: T.handle, placeholder_6: T.handle, T_cast_2: T.handle) -> None: - # function attr dict - T.func_attr({"global_symbol": "tvmgen_default_fused_nn_conv2d_add_fixed_point_multiply_clip_cast_cast", "tir.noalias": True}) - placeholder_7 = T.match_buffer(placeholder_4, [360000], dtype="int16") - placeholder_8 = T.match_buffer(placeholder_5, [4096], dtype="int16") - placeholder_9 = T.match_buffer(placeholder_6, [64], dtype="int32") - T_cast_3 = T.match_buffer(T_cast_2, [215], dtype="int16") - # body - PaddedInput_data = T.allocate([360000], "int16", "global") - PaddedInput = T.Buffer([360000], "int16", data=PaddedInput_data) - for i0_i1_fused, i2, i3 in T.grid(75, 75, 64): - PaddedInput[i0_i1_fused * 4800 + i2 * 64 + i3] = placeholder_7[i0_i1_fused * 4800 + i2 * 64 + i3] - for ax0_ax1_fused_ax2_fused in T.serial(0, 5625): - Conv2dOutput_data = T.allocate([64], "int32", "global") - Conv2dOutput = T.Buffer([64], "int32", data=Conv2dOutput_data) - for ff in T.serial(0, 64): - Conv2dOutput[ff] = 0 - for rc in T.serial(0, 64): - Conv2dOutput[ff] = Conv2dOutput[ff] + T.cast(PaddedInput[ax0_ax1_fused_ax2_fused * 64 + rc], "int32") * T.cast(placeholder_8[rc * 64 + ff], "int32") - for ax3_inner_1 in T.serial(0, 64): - T_cast_3[ax0_ax1_fused_ax2_fused * 64 + ax3_inner_1] = T.cast(T.cast(T.max(T.min(T.q_multiply_shift(Conv2dOutput[ax3_inner_1] + placeholder_9[ax3_inner_1], 1843106743, 31, -6, dtype="int32"), 255), 0), "uint8"), "int16") -# fmt: on - - -# fmt: off -@tvm.script.ir_module -class ResnetStructurePlanned: - @T.prim_func - def tvmgen_default_fused_cast_subtract_fixed_point_multiply_add_clip_cast_cast(placeholder: T.handle, placeholder_1: T.handle, T_cast: T.handle, global_workspace_1_var: T.handle("uint8")) -> None: - placeholder_2 = T.match_buffer(placeholder, [360000], dtype="uint8") - placeholder_3 = T.match_buffer(placeholder_1, [64], dtype="int32") - T_cast_1 = T.match_buffer(T_cast, [215], dtype="int16") - global_workspace_1_buffer_var = T.match_buffer(global_workspace_1_var, [7920256], dtype="uint8", strides=[1], elem_offset=0, align=16) - # body - for ax0_ax1_fused, ax2, ax3_outer, ax3_inner in T.grid(75, 75, 4, 16): - T_cast_1[ax0_ax1_fused * 4800 + ax2 * 64 + ax3_outer * 16 + ax3_inner] = T.cast(T.cast(T.max(T.min(T.q_multiply_shift(T.cast(placeholder_2[ax0_ax1_fused * 4800 + ax2 * 64 + ax3_outer * 16 + ax3_inner], "int32") - 94, 1843157232, 31, 1, dtype="int32") + placeholder_3[ax3_outer * 16 + ax3_inner], 255), 0), "uint8"), "int16") - - @T.prim_func - def tvmgen_default_fused_nn_conv2d_add_fixed_point_multiply_add_clip_cast_cast_subtract_fixed_point_4200876283395191415_(placeholder_22: T.handle, placeholder_23: T.handle, placeholder_24: T.handle, placeholder_25: T.handle, T_cast_6: T.handle, global_workspace_5_var: T.handle("uint8")) -> None: - placeholder_29 = T.match_buffer(placeholder_22, [360000], dtype="int16") - placeholder_27 = T.match_buffer(placeholder_23, [16384], dtype="int16") - placeholder_26 = T.match_buffer(placeholder_24, [256], dtype="int32") - placeholder_28 = T.match_buffer(placeholder_25, [1440000], dtype="int32") - T_cast_7 = T.match_buffer(T_cast_6, [407], dtype="uint8") - global_workspace_5_buffer_var = T.match_buffer(global_workspace_5_var, [7920256], dtype="uint8", strides=[1], elem_offset=0, align=16) - # body - PaddedInput_3_let = T.Buffer([360000], 'int16') - with T.LetStmt(T.address_of(global_workspace_5_buffer_var[6480000], dtype="handle"), var=PaddedInput_3_let.data): - for i0_i1_fused_3, i2_3, i3_3 in T.grid(75, 75, 64): - PaddedInput_3_let[i0_i1_fused_3 * 4800 + i2_3 * 64 + i3_3] = placeholder_29[i0_i1_fused_3 * 4800 + i2_3 * 64 + i3_3] - for ax0_ax1_fused_ax2_fused_3 in T.serial(0, 5625): - Conv2dOutput_3_let = T.Buffer([64], 'int32') - with T.LetStmt(T.address_of(global_workspace_5_buffer_var[7200000], dtype="handle"), var=Conv2dOutput_3_let.data): - for ax3_outer_2 in T.serial(0, 4): - for ff_3 in T.serial(0, 64): - Conv2dOutput_3_let[ff_3] = 0 - for rc_3 in T.serial(0, 64): - Conv2dOutput_3_let[ff_3] = Conv2dOutput_3_let[ff_3] + T.cast(PaddedInput_3_let[ax0_ax1_fused_ax2_fused_3 * 64 + rc_3], "int32") * T.cast(placeholder_27[rc_3 * 256 + ax3_outer_2 * 64 + ff_3], "int32") - for ax3_inner_4 in T.serial(0, 64): - T_cast_7[ax0_ax1_fused_ax2_fused_3 * 256 + ax3_outer_2 * 64 + ax3_inner_4] = T.cast(T.max(T.min(T.q_multiply_shift(T.cast(T.cast(T.max(T.min(T.q_multiply_shift(Conv2dOutput_3_let[ax3_inner_4] + placeholder_26[ax3_outer_2 * 64 + ax3_inner_4], 1343014664, 31, -8, dtype="int32") + 136, 255), 0), "uint8"), "int32") - 136, 1073903788, 31, 1, dtype="int32") + placeholder_28[ax0_ax1_fused_ax2_fused_3 * 256 + ax3_outer_2 * 64 + ax3_inner_4], 255), 0), "uint8") - - @T.prim_func - def tvmgen_default_fused_nn_conv2d_add_fixed_point_multiply_add_clip_cast_cast_subtract_fixed_point_15934180698220515269_(placeholder_16: T.handle, placeholder_17: T.handle, placeholder_18: T.handle, T_add: T.handle, global_workspace_4_var: T.handle("uint8")) -> None: - placeholder_19 = T.match_buffer(placeholder_16, [360000], dtype="int16") - placeholder_20 = T.match_buffer(placeholder_17, [16384], dtype="int16") - placeholder_21 = T.match_buffer(placeholder_18, [256], dtype="int32") - T_add_1 = T.match_buffer(T_add, [407], dtype="int32") - global_workspace_4_buffer_var = T.match_buffer(global_workspace_4_var, [7920256], dtype="uint8", strides=[1], elem_offset=0, align=16) - # body - PaddedInput_2_let = T.Buffer([360000], "int16") - with T.LetStmt(T.address_of(global_workspace_4_buffer_var[7200000], dtype="handle"), var=PaddedInput_2_let.data): - for i0_i1_fused_2, i2_2, i3_2 in T.grid(75, 75, 64): - PaddedInput_2_let[i0_i1_fused_2 * 4800 + i2_2 * 64 + i3_2] = placeholder_19[i0_i1_fused_2 * 4800 + i2_2 * 64 + i3_2] - for ax0_ax1_fused_ax2_fused_2 in T.serial(0, 5625): - Conv2dOutput_2_let = T.Buffer([64], 'int32') - with T.LetStmt(T.address_of(global_workspace_4_buffer_var[7920000], dtype="handle"), var=Conv2dOutput_2_let.data): - for ax3_outer_1 in T.serial(0, 4): - for ff_2 in T.serial(0, 64): - Conv2dOutput_2_let[ff_2] = 0 - for rc_2 in T.serial(0, 64): - Conv2dOutput_2_let[ff_2] = Conv2dOutput_2_let[ff_2] + T.cast(PaddedInput_2_let[ax0_ax1_fused_ax2_fused_2 * 64 + rc_2], "int32") * T.cast(placeholder_20[rc_2 * 256 + ax3_outer_1 * 64 + ff_2], "int32") - for ax3_inner_3 in T.serial(0, 64): - T_add_1[ax0_ax1_fused_ax2_fused_2 * 256 + ax3_outer_1 * 64 + ax3_inner_3] = T.q_multiply_shift(T.cast(T.cast(T.max(T.min(T.q_multiply_shift(Conv2dOutput_2_let[ax3_inner_3] + placeholder_21[ax3_outer_1 * 64 + ax3_inner_3], 1711626602, 31, -8, dtype="int32") + 132, 255), 0), "uint8"), "int32") - 132, 2094289803, 31, -2, dtype="int32") + 136 - - @T.prim_func - def tvmgen_default_fused_nn_conv2d_add_fixed_point_multiply_clip_cast_cast(placeholder_4: T.handle, placeholder_5: T.handle, placeholder_6: T.handle, T_cast_2: T.handle, global_workspace_2_var: T.handle("uint8")) -> None: - placeholder_7 = T.match_buffer(placeholder_4, [360000], dtype="int16") - placeholder_8 = T.match_buffer(placeholder_5, [4096], dtype="int16") - placeholder_9 = T.match_buffer(placeholder_6, [64], dtype="int32") - T_cast_3 = T.match_buffer(T_cast_2, [215], dtype="int16") - global_workspace_2_buffer_var = T.match_buffer(global_workspace_2_var, [7920256], dtype="uint8", strides=[1], elem_offset=0, align=16) - # body - PaddedInput_let = T.Buffer([360000], "int16") - with T.LetStmt(T.address_of(global_workspace_2_buffer_var[7200000], dtype="handle"), var=PaddedInput_let.data): - for i0_i1_fused, i2, i3 in T.grid(75, 75, 64): - PaddedInput_let[i0_i1_fused * 4800 + i2 * 64 + i3] = placeholder_7[i0_i1_fused * 4800 + i2 * 64 + i3] - for ax0_ax1_fused_ax2_fused in T.serial(0, 5625): - Conv2dOutput_let = T.Buffer([64], "int32") - with T.LetStmt(T.address_of(global_workspace_2_buffer_var[7920000], dtype="handle"), var=Conv2dOutput_let.data): - for ff in T.serial(0, 64): - Conv2dOutput_let[ff] = 0 - for rc in T.serial(0, 64): - Conv2dOutput_let[ff] = Conv2dOutput_let[ff] + T.cast(PaddedInput_let[ax0_ax1_fused_ax2_fused * 64 + rc], "int32") * T.cast(placeholder_8[rc * 64 + ff], "int32") - for ax3_inner_1 in T.serial(0, 64): - T_cast_3[ax0_ax1_fused_ax2_fused * 64 + ax3_inner_1] = T.cast(T.cast(T.max(T.min(T.q_multiply_shift(Conv2dOutput_let[ax3_inner_1] + placeholder_9[ax3_inner_1], 1843106743, 31, -6, dtype="int32"), 255), 0), "uint8"), "int16") - - @T.prim_func - def tvmgen_default_fused_nn_conv2d_add_fixed_point_multiply_clip_cast_cast_1(placeholder_10: T.handle, placeholder_11: T.handle, placeholder_12: T.handle, T_cast_4: T.handle, global_workspace_3_var: T.handle("uint8")) -> None: - placeholder_13 = T.match_buffer(placeholder_10, [360000], dtype="int16") - placeholder_14 = T.match_buffer(placeholder_11, [36864], dtype="int16") - placeholder_15 = T.match_buffer(placeholder_12, [64], dtype="int32") - T_cast_5 = T.match_buffer(T_cast_4, [215], dtype="int16") - global_workspace_3_buffer_var = T.match_buffer(global_workspace_3_var, [7920256], dtype="uint8", strides=[1], elem_offset=0, align=16) - # body - PaddedInput_1_let = T.Buffer([379456], "int16") - with T.LetStmt(T.address_of(global_workspace_3_buffer_var[0], dtype="handle"), var=PaddedInput_1_let.data): - for i0_i1_fused_1, i2_1, i3_1 in T.grid(77, 77, 64): - PaddedInput_1_let[i0_i1_fused_1 * 4928 + i2_1 * 64 + i3_1] = T.if_then_else(1 <= i0_i1_fused_1 and i0_i1_fused_1 < 76 and 1 <= i2_1 and i2_1 < 76, placeholder_13[i0_i1_fused_1 * 4800 + i2_1 * 64 + i3_1 - 4864], T.int16(0), dtype="int16") - for ax0_ax1_fused_ax2_fused_1 in T.serial(0, 5625): - Conv2dOutput_1_let = T.Buffer([64], "int32") - with T.LetStmt(T.address_of(global_workspace_3_buffer_var[7200000], dtype="handle"), var=Conv2dOutput_1_let.data): - for ff_1 in T.serial(0, 64): - Conv2dOutput_1_let[ff_1] = 0 - for ry, rx, rc_1 in T.grid(3, 3, 64): - Conv2dOutput_1_let[ff_1] = Conv2dOutput_1_let[ff_1] + T.cast(PaddedInput_1_let[ax0_ax1_fused_ax2_fused_1 // 75 * 4928 + ry * 4928 + rx * 64 + ax0_ax1_fused_ax2_fused_1 % 75 * 64 + rc_1], "int32") * T.cast(placeholder_14[ry * 12288 + rx * 4096 + rc_1 * 64 + ff_1], "int32") - for ax3_inner_2 in T.serial(0, 64): - T_cast_5[ax0_ax1_fused_ax2_fused_1 * 64 + ax3_inner_2] = T.cast(T.cast(T.max(T.min(T.q_multiply_shift(Conv2dOutput_1_let[ax3_inner_2] + placeholder_15[ax3_inner_2], 1608879842, 31, -7, dtype="int32"), 255), 0), "uint8"), "int16") - - @T.prim_func - def __tvm_main__(input: T.handle, global_workspace_0_var: T.handle("uint8"), output: T.handle) -> None: - global_workspace_0_buffer_var = T.match_buffer(global_workspace_0_var, [7920256], dtype="uint8", strides=[1], elem_offset=0, align=16) - # body - T.attr("default", "device_id", 0) - T.attr("default", "device_type", 1) - sid_2_let: T.handle("int8") = T.address_of(global_workspace_0_buffer_var[5760000], dtype="handle") - sid_6_let: T.handle("int8") = T.address_of(global_workspace_0_buffer_var[0], dtype="handle") - sid_7_let: T.handle("int8") = T.address_of(global_workspace_0_buffer_var[6480000], dtype="handle") - sid_8_let: T.handle("int8") = T.address_of(global_workspace_0_buffer_var[6480000], dtype="handle") - T.evaluate(T.call_extern("tvmgen_default_fused_cast_subtract_fixed_point_multiply_add_clip_cast_cast", input, T.lookup_param("p0", dtype="handle"), sid_2_let, global_workspace_0_buffer_var.data, dtype="int32")) - T.evaluate(T.call_extern("tvmgen_default_fused_nn_conv2d_add_fixed_point_multiply_clip_cast_cast", sid_2_let, T.lookup_param("p3", dtype="handle"), T.lookup_param("p4", dtype="handle"), sid_8_let, global_workspace_0_buffer_var.data, dtype="int32")) - T.evaluate(T.call_extern("tvmgen_default_fused_nn_conv2d_add_fixed_point_multiply_clip_cast_cast_1", sid_8_let, T.lookup_param("p5", dtype="handle"), T.lookup_param("p6", dtype="handle"), sid_7_let, global_workspace_0_buffer_var.data, dtype="int32")) - T.evaluate(T.call_extern("tvmgen_default_fused_nn_conv2d_add_fixed_point_multiply_add_clip_cast_cast_subtract_fixed_point_15934180698220515269_", sid_7_let, T.lookup_param("p7", dtype="handle"), T.lookup_param("p8", dtype="handle"), sid_6_let, global_workspace_0_buffer_var.data, dtype="int32")) - T.evaluate(T.call_extern("tvmgen_default_fused_nn_conv2d_add_fixed_point_multiply_add_clip_cast_cast_subtract_fixed_point_4200876283395191415_", sid_2_let, T.lookup_param("p1", dtype="handle"), T.lookup_param("p2", dtype="handle"), sid_6_let, output, global_workspace_0_buffer_var.data, dtype="int32")) -# fmt: on - - -def test_resnet_subgraph(): - before = ResnetStructure - expected = ResnetStructurePlanned - after = _plan_and_convert(before) - tvm.ir.assert_structural_equal(after, expected) - - -@tvm.script.ir_module -class TensorIntrinStructure: - @T.prim_func - def tensor_intrin_primfunc() -> None: - dense_data = T.allocate([10], "int32", "global") - T.evaluate( - T.call_extern( - "intrin_function", - T.tvm_access_ptr( - T.type_annotation(dtype="int32"), dense_data, 0, 1, 2, dtype="handle" - ), - dtype="int32", - ) - ) - - dense = T.Buffer([10], "int32", data=dense_data) - dense[0] = T.q_multiply_shift(dense[0], 1608879842, 31, -7, dtype="int32") - - @T.prim_func - def __tvm_main__(input: T.handle, output: T.handle) -> None: - T.evaluate(T.call_extern("tensor_intrin_primfunc", dtype="int32")) - - -@tvm.script.ir_module -class TensorIntrinStructurePlanned: - @T.prim_func - def tensor_intrin_primfunc(global_workspace_1_var: T.handle("uint8")) -> None: - global_workspace_1_buffer_var = T.match_buffer( - global_workspace_1_var, [40], dtype="uint8", strides=[1], elem_offset=0, align=16 - ) - dense_let = T.Buffer([10], "int32") - with T.LetStmt( - T.address_of(global_workspace_1_buffer_var[0], dtype="handle"), var=dense_let.data - ): - T.evaluate( - T.call_extern( - "intrin_function", - T.tvm_access_ptr( - T.type_annotation(dtype="int32"), dense_let.data, 0, 1, 2, dtype="handle" - ), - dtype="int32", - ) - ) - dense_let[0] = T.q_multiply_shift(dense_let[0], 1608879842, 31, -7, dtype="int32") - - @T.prim_func - def __tvm_main__( - input: T.handle, global_workspace_1_var: T.handle("uint8"), output: T.handle - ) -> None: - global_workspace_1_buffer_var = T.match_buffer( - global_workspace_1_var, [40], dtype="uint8", strides=[1], elem_offset=0, align=16 - ) - T.evaluate( - T.call_extern( - "tensor_intrin_primfunc", global_workspace_1_buffer_var.data, dtype="int32" - ) - ) - - -def test_tensor_intrin(): - before = TensorIntrinStructure - after = _plan_and_convert(before) - expected = TensorIntrinStructurePlanned - tvm.ir.assert_structural_equal(after, expected) - - -class TestMergeAllocations(tvm.testing.CompareBeforeAfter): - def transform(self): - return _plan_and_convert - - def before(self): - @I.ir_module - class mod: - @T.prim_func - def __tvm_main__(A: T.Buffer(256, "int8"), D: T.Buffer(256, "int8")): - B = T.allocate([256], "int8") - T.call_extern("subroutine", A.data, B, dtype="int32") - C = T.allocate([256], "int8") - T.call_extern("subroutine", B, C, dtype="int32") - T.call_extern("subroutine", C, D.data, dtype="int32") - - @T.prim_func - def subroutine(A: T.Buffer(256, "int8"), B: T.Buffer(256, "int8")): - for i in range(256): - B[i] = A[i] - - return mod - - def expected(self): - @I.ir_module - class mod: - @T.prim_func - def __tvm_main__( - A: T.Buffer(256, "int8"), - D: T.Buffer(256, "int8"), - workspace_var: T.handle("uint8"), - ): - workspace = T.match_buffer(workspace_var, 512, "uint8", strides=[1], align=16) - B: T.handle("int8") = T.address_of(workspace[256]) - T.call_extern("subroutine", A.data, B, workspace.data, dtype="int32") - C: T.handle("int8") = T.address_of(workspace[0]) - T.call_extern("subroutine", B, C, workspace.data, dtype="int32") - T.call_extern("subroutine", C, D.data, workspace.data, dtype="int32") - - @T.prim_func - def subroutine( - A: T.Buffer(256, "int8"), - B: T.Buffer(256, "int8"), - workspace_var: T.handle("uint8"), - ): - workspace = T.match_buffer(workspace_var, 512, "uint8", strides=[1], align=16) - for i in range(256): - B[i] = A[i] - - return mod - - -class TestMergeAllocationsWithDeclBuffer(tvm.testing.CompareBeforeAfter): - """Like TestMergeAllocations, but using T.decl_buffer""" - - def transform(self): - return _plan_and_convert - - def before(self): - @I.ir_module - class mod: - @T.prim_func - def __tvm_main__(A: T.Buffer(256, "int8"), D: T.Buffer(256, "int8")): - B = T.decl_buffer([256], "int8") - T.call_extern("subroutine", A.data, B.data, dtype="int32") - C = T.decl_buffer([256], "int8") - T.call_extern("subroutine", B.data, C.data, dtype="int32") - T.call_extern("subroutine", C.data, D.data, dtype="int32") - - @T.prim_func - def subroutine(A: T.Buffer(256, "int8"), B: T.Buffer(256, "int8")): - for i in range(256): - B[i] = A[i] - - return mod - - def expected(self): - @I.ir_module - class mod: - @T.prim_func - def __tvm_main__( - A: T.Buffer(256, "int8"), - D: T.Buffer(256, "int8"), - workspace_var: T.handle("uint8"), - ): - workspace = T.match_buffer(workspace_var, 512, "uint8", strides=[1], align=16) - B_data: T.handle("int8") = T.address_of(workspace[256]) - B = T.decl_buffer(256, "int8", data=B_data) - T.call_extern("subroutine", A.data, B.data, workspace.data, dtype="int32") - C_data: T.handle("int8") = T.address_of(workspace[0]) - C = T.decl_buffer(256, "int8", data=C_data) - T.call_extern("subroutine", B.data, C.data, workspace.data, dtype="int32") - T.call_extern("subroutine", C.data, D.data, workspace.data, dtype="int32") - - @T.prim_func - def subroutine( - A: T.Buffer(256, "int8"), - B: T.Buffer(256, "int8"), - workspace_var: T.handle("uint8"), - ): - workspace = T.match_buffer(workspace_var, 512, "uint8", strides=[1], align=16) - for i in range(256): - B[i] = A[i] - - return mod - - -if __name__ == "__main__": - tvm.testing.main() diff --git a/tests/python/tir-usmp/test_tir_usmp_transform_create_io_allocates.py b/tests/python/tir-usmp/test_tir_usmp_transform_create_io_allocates.py deleted file mode 100644 index 53a381c82b14..000000000000 --- a/tests/python/tir-usmp/test_tir_usmp_transform_create_io_allocates.py +++ /dev/null @@ -1,206 +0,0 @@ -# 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 pytest -from typing import NamedTuple, List - -import tvm -from tvm.script import tir as T - - -# fmt: off -@tvm.script.ir_module -class SingleInputSingleOutput: - @T.prim_func - def tvmgen_default_fused_cast_subtract(placeholder_2: T.handle, placeholder_3: T.handle, T_subtract: T.handle) -> None: - # function attr dict - T.func_attr({"global_symbol": "tvmgen_default_fused_cast_subtract", "tir.noalias": True}) - placeholder_4 = T.match_buffer(placeholder_2, [150528], dtype="uint8", elem_offset=0, align=64, offset_factor=1) - placeholder_5 = T.match_buffer(placeholder_3, [1], dtype="int16", elem_offset=0, align=64, offset_factor=1) - T_subtract_1 = T.match_buffer(T_subtract, [452], dtype="int16", elem_offset=0, align=64, offset_factor=1) - # body - for ax0_ax1_fused_1 in T.serial(0, 224): - for ax2_1, ax3_inner_1 in T.grid(224, 3): - T_subtract_1[(((ax0_ax1_fused_1*672) + (ax2_1*3)) + ax3_inner_1)] = (T.cast(placeholder_4[(((ax0_ax1_fused_1*672) + (ax2_1*3)) + ax3_inner_1)], "int16") - placeholder_5[0]) - - @T.prim_func - def __tvm_main__(input: T.handle, output: T.handle) -> None: - # function attr dict - T.func_attr({"global_symbol": "__tvm_main__", "runner_function": True}) - input_buffer_var = T.match_buffer(input, [150528], dtype="uint8", elem_offset=0, align=64, offset_factor=1) - output_buffer_var = T.match_buffer(output, [452], dtype="int16", elem_offset=0, align=64, offset_factor=1) - # body - T.evaluate(T.call_extern("tvmgen_default_fused_cast_subtract", input_buffer_var.data, T.lookup_param("p0", dtype="handle"), output_buffer_var.data, dtype="int32")) -# fmt: on - - -# fmt: off -@tvm.script.ir_module -class TwoInputSingleOutput: - @T.prim_func - def tvmgen_default_fused_cast_subtract(placeholder_2: T.handle, placeholder_3: T.handle, T_subtract: T.handle) -> None: - # function attr dict - T.func_attr({"global_symbol": "tvmgen_default_fused_cast_subtract", "tir.noalias": True}) - placeholder_4 = T.match_buffer(placeholder_2, [150528], dtype="uint8", elem_offset=0, align=64, offset_factor=1) - placeholder_5 = T.match_buffer(placeholder_3, [1], dtype="int16", elem_offset=0, align=64, offset_factor=1) - T_subtract_1 = T.match_buffer(T_subtract, [452], dtype="int16", elem_offset=0, align=64, offset_factor=1) - # body - for ax0_ax1_fused_1 in T.serial(0, 224): - for ax2_1, ax3_inner_1 in T.grid(224, 3): - T_subtract_1[(((ax0_ax1_fused_1*672) + (ax2_1*3)) + ax3_inner_1)] = (T.cast(placeholder_4[(((ax0_ax1_fused_1*672) + (ax2_1*3)) + ax3_inner_1)], "int16") - placeholder_5[0]) - - @T.prim_func - def __tvm_main__(input1: T.handle, input2: T.handle, output: T.handle) -> None: - # function attr dict - T.func_attr({"global_symbol": "__tvm_main__", "runner_function": True}) - input1_buffer_var = T.match_buffer(input1, [150528], dtype="uint8", elem_offset=0, align=64, offset_factor=1) - input2_buffer_var = T.match_buffer(input2, [1], dtype="int16", elem_offset=0, align=64, offset_factor=1) - output_buffer_var = T.match_buffer(output, [452], dtype="int16", elem_offset=0, align=64, offset_factor=1) - # body - T.evaluate(T.call_extern("tvmgen_default_fused_cast_subtract", input1_buffer_var.data, input2_buffer_var.data, output_buffer_var.data, dtype="int32")) -# fmt: on - - -# fmt: off -@tvm.script.ir_module -class TwoInputTwoOutput: - @T.prim_func - def tvmgen_default_fused_cast_subtract(placeholder_2: T.handle, placeholder_3: T.handle, T_subtract: T.handle) -> None: - # function attr dict - T.func_attr({"global_symbol": "tvmgen_default_fused_cast_subtract", "tir.noalias": True}) - placeholder_4 = T.match_buffer(placeholder_2, [150528], dtype="uint8", elem_offset=0, align=64, offset_factor=1) - placeholder_5 = T.match_buffer(placeholder_3, [1], dtype="int16", elem_offset=0, align=64, offset_factor=1) - T_subtract_1 = T.match_buffer(T_subtract, [452], dtype="int16", elem_offset=0, align=64, offset_factor=1) - # body - for ax0_ax1_fused_1 in T.serial(0, 224): - for ax2_1, ax3_inner_1 in T.grid(224, 3): - T_subtract_1[(((ax0_ax1_fused_1*672) + (ax2_1*3)) + ax3_inner_1)] = (T.cast(placeholder_4[(((ax0_ax1_fused_1*672) + (ax2_1*3)) + ax3_inner_1)], "int16") - placeholder_5[0]) - - @T.prim_func - def __tvm_main__(input1: T.handle, input2: T.handle, output1: T.handle, output2: T.handle) -> None: - # function attr dict - T.func_attr({"global_symbol": "__tvm_main__", "runner_function": True}) - input1_buffer_var = T.match_buffer(input1, [150528], dtype="uint8", elem_offset=0, align=64, offset_factor=1) - input2_buffer_var = T.match_buffer(input2, [150528], dtype="uint8", elem_offset=0, align=64, offset_factor=1) - output1_buffer_var = T.match_buffer(output1, [452], dtype="int16", elem_offset=0, align=64, offset_factor=1) - output2_buffer_var = T.match_buffer(output2, [452], dtype="int16", elem_offset=0, align=64, offset_factor=1) - # body - T.evaluate(T.call_extern("tvmgen_default_fused_cast_subtract", input1_buffer_var.data, T.lookup_param("p0", dtype="handle"), output1_buffer_var.data, dtype="int32")) - T.evaluate(T.call_extern("tvmgen_default_fused_cast_subtract", input2_buffer_var.data, T.lookup_param("p1", dtype="handle"), output2_buffer_var.data, dtype="int32")) -# fmt: on - - -# fmt: off -@tvm.script.ir_module -class SingleInputTwoOutput: - @T.prim_func - def tvmgen_default_fused_cast_subtract(placeholder_2: T.handle, placeholder_3: T.handle, T_subtract: T.handle) -> None: - # function attr dict - T.func_attr({"global_symbol": "tvmgen_default_fused_cast_subtract", "tir.noalias": True}) - placeholder_4 = T.match_buffer(placeholder_2, [150528], dtype="uint8", elem_offset=0, align=64, offset_factor=1) - placeholder_5 = T.match_buffer(placeholder_3, [1], dtype="int16", elem_offset=0, align=64, offset_factor=1) - T_subtract_1 = T.match_buffer(T_subtract, [452], dtype="int16", elem_offset=0, align=64, offset_factor=1) - # body - for ax0_ax1_fused_1 in T.serial(0, 224): - for ax2_1, ax3_inner_1 in T.grid(224, 3): - T_subtract_1[(((ax0_ax1_fused_1*672) + (ax2_1*3)) + ax3_inner_1)] = (T.cast(placeholder_4[(((ax0_ax1_fused_1*672) + (ax2_1*3)) + ax3_inner_1)], "int16") - placeholder_5[0]) - - @T.prim_func - def __tvm_main__(input: T.handle, output1: T.handle, output2: T.handle) -> None: - # function attr dict - T.func_attr({"global_symbol": "__tvm_main__", "runner_function": True}) - input_buffer_var = T.match_buffer(input, [150528], dtype="uint8", elem_offset=0, align=64, offset_factor=1) - output1_buffer_var = T.match_buffer(output1, [452], dtype="int16", elem_offset=0, align=64, offset_factor=1) - output2_buffer_var = T.match_buffer(output2, [452], dtype="int16", elem_offset=0, align=64, offset_factor=1) - # body - T.evaluate(T.call_extern("tvmgen_default_fused_cast_subtract", input_buffer_var.data, T.lookup_param("p0", dtype="handle"), output1_buffer_var.data, dtype="int32")) - T.evaluate(T.call_extern("tvmgen_default_fused_cast_subtract", input_buffer_var.data, T.lookup_param("p1", dtype="handle"), output2_buffer_var.data, dtype="int32")) -# fmt: on - - -class IOInfo(NamedTuple): - """A data structure to hold test outputs per I/O tensor""" - - name: str - shape: list - dtype: str - - -def check_io_allocations(mod: tvm.IRModule, inputs: List[IOInfo], outputs: List[IOInfo]): - """This function checks whether outer most allocates correspond to I/O tensors""" - found_non_io_allocate_node = False - - input_name_to_info = {} - for input in inputs: - input_name_to_info[input.name] = input - output_name_to_info = {} - for output in outputs: - output_name_to_info[output.name] = output - - def _visit(stmt): - nonlocal found_non_io_allocate_node - if isinstance(stmt, tvm.tir.Allocate) and not found_non_io_allocate_node: - allocate = stmt - if dict(allocate.annotations).get("input_tensor"): - input_tensor_name = str(dict(allocate.annotations).get("input_tensor")) - assert input_tensor_name in input_name_to_info.keys() - assert input_name_to_info[input_tensor_name].shape == list(allocate.extents) - assert input_name_to_info[input_tensor_name].dtype == str(allocate.dtype) - del input_name_to_info[input_tensor_name] - if dict(allocate.annotations).get("output_tensor"): - output_tensor_name = str(dict(allocate.annotations).get("output_tensor")) - assert output_tensor_name in output_name_to_info.keys() - assert output_name_to_info[output_tensor_name].shape == list(allocate.extents) - assert output_name_to_info[output_tensor_name].dtype == str(allocate.dtype) - del output_name_to_info[output_tensor_name] - else: - found_non_io_allocate_node = True - - main = mod["__tvm_main__"] - tvm.tir.stmt_functor.ir_transform(main.body, _visit, None, ["tir.Allocate", "tir.Call"]) - assert len(input_name_to_info) == 0 - assert len(output_name_to_info) == 0 - - -@pytest.mark.parametrize( - "test_mod, input_names, output_names", - [ - ( - SingleInputSingleOutput, - [IOInfo("input", [150528], "uint8")], - [IOInfo("output", [452], "int16")], - ), - ( - SingleInputTwoOutput, - [IOInfo("input", [150528], "uint8")], - [IOInfo("output1", [452], "int16"), IOInfo("output2", [452], "int16")], - ), - ( - TwoInputSingleOutput, - [IOInfo("input1", [150528], "uint8"), IOInfo("input2", [1], "int16")], - [IOInfo("output", [452], "int16")], - ), - ( - TwoInputTwoOutput, - [IOInfo("input1", [150528], "uint8"), IOInfo("input2", [150528], "uint8")], - [IOInfo("output1", [452], "int16"), IOInfo("output2", [452], "int16")], - ), - ], -) -def test_mobilenet_subgraph(test_mod, input_names, output_names): - CreateAllocatesForIO = tvm.get_global_func("tir.usmp.transform.CreateAllocatesForIO") - test_mod = CreateAllocatesForIO()(test_mod) - check_io_allocations(test_mod, input_names, output_names) diff --git a/tests/python/tir-usmp/test_tir_usmp_utils.py b/tests/python/tir-usmp/test_tir_usmp_utils.py deleted file mode 100644 index 635c9a760f87..000000000000 --- a/tests/python/tir-usmp/test_tir_usmp_utils.py +++ /dev/null @@ -1,200 +0,0 @@ -# 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 pytest -import sys - -import tvm -from tvm.script import tir as T -from tvm.tir import stmt_functor -from tvm.tir.usmp import utils as usmp_utils -from tvm.target import Target -from tvm import WorkspacePoolInfo, PoolInfoProperties - -# fmt: off -@tvm.script.ir_module -class LinearStructure: - @T.prim_func - def tvmgen_default_fused_cast_subtract(placeholder_2: T.handle, placeholder_3: T.handle, T_subtract: T.handle) -> None: - # function attr dict - T.func_attr({"global_symbol": "tvmgen_default_fused_cast_subtract", "tir.noalias": True}) - placeholder_4 = T.match_buffer(placeholder_2, [150528], dtype="uint8", elem_offset=0, align=64, offset_factor=1) - placeholder_5 = T.match_buffer(placeholder_3, [1], dtype="int16", elem_offset=0, align=64, offset_factor=1) - T_subtract_1 = T.match_buffer(T_subtract, [150528], dtype="int16", elem_offset=0, align=64, offset_factor=1) - # body - for ax0_ax1_fused_1 in T.serial(0, 224): - for ax2_1, ax3_inner_1 in T.grid(224, 3): - T_subtract_1[(((ax0_ax1_fused_1*672) + (ax2_1*3)) + ax3_inner_1)] = (T.cast(placeholder_4[(((ax0_ax1_fused_1*672) + (ax2_1*3)) + ax3_inner_1)], "int16") - placeholder_5[0]) - - @T.prim_func - def tvmgen_default_fused_nn_conv2d_add_fixed_point_multiply_clip_cast(placeholder_62: T.handle, placeholder_63: T.handle, placeholder_64: T.handle, T_cast_20: T.handle) -> None: - # function attr dict - T.func_attr({"global_symbol": "tvmgen_default_fused_nn_conv2d_add_fixed_point_multiply_clip_cast", "tir.noalias": True}) - placeholder_65 = T.match_buffer(placeholder_62, [150528], dtype="int16", elem_offset=0, align=64, offset_factor=1) - placeholder_66 = T.match_buffer(placeholder_63, [9408], dtype="int16", elem_offset=0, align=64, offset_factor=1) - placeholder_67 = T.match_buffer(placeholder_64, [64], dtype="int32", elem_offset=0, align=64, offset_factor=1) - T_cast_21 = T.match_buffer(T_cast_20, [289], dtype="uint8", elem_offset=0, align=64, offset_factor=1) - # body - PaddedInput_7 = T.decl_buffer([157323], "int16") - for i0_i1_fused_7 in T.serial(0, 229): - for i2_7, i3_7 in T.grid(229, 3): - PaddedInput_7[(((i0_i1_fused_7*687) + (i2_7*3)) + i3_7)] = T.if_then_else(((((2 <= i0_i1_fused_7) and (i0_i1_fused_7 < 226)) and (2 <= i2_7)) and (i2_7 < 226)), placeholder_65[((((i0_i1_fused_7*672) + (i2_7*3)) + i3_7) - 1350)], T.int16(0), dtype="int16") - for ax0_ax1_fused_ax2_fused_7 in T.serial(0, 12544): - Conv2dOutput_7 = T.decl_buffer([64], "int32") - for ff_3 in T.serial(0, 64): - Conv2dOutput_7[ff_3] = 0 - for ry_2, rx_2, rc_7 in T.grid(7, 7, 3): - Conv2dOutput_7[ff_3] = (Conv2dOutput_7[ff_3] + (T.cast(PaddedInput_7[(((((T.floordiv(ax0_ax1_fused_ax2_fused_7, 112)*1374) + (ry_2*687)) + (T.floormod(ax0_ax1_fused_ax2_fused_7, 112)*6)) + (rx_2*3)) + rc_7)], "int32")*T.cast(placeholder_66[((((ry_2*1344) + (rx_2*192)) + (rc_7*64)) + ff_3)], "int32"))) - for ax3_inner_7 in T.serial(0, 64): - T_cast_21[((ax0_ax1_fused_ax2_fused_7*64) + ax3_inner_7)] = T.cast(T.max(T.min(T.q_multiply_shift((Conv2dOutput_7[ax3_inner_7] + placeholder_67[ax3_inner_7]), 1939887962, 31, -9, dtype="int32"), 255), 0), "uint8") - - @T.prim_func - def tvmgen_default_fused_nn_max_pool2d_cast(placeholder_28: T.handle, T_cast_6: T.handle) -> None: - # function attr dict - T.func_attr({"global_symbol": "tvmgen_default_fused_nn_max_pool2d_cast", "tir.noalias": True}) - placeholder_29 = T.match_buffer(placeholder_28, [802816], dtype="uint8", elem_offset=0, align=64, offset_factor=1) - T_cast_7 = T.match_buffer(T_cast_6, [177], dtype="int16", elem_offset=0, align=64, offset_factor=1) - # body - tensor_2 = T.decl_buffer([200704], "uint8") - for ax0_ax1_fused_4 in T.serial(0, 56): - for ax2_4 in T.serial(0, 56): - for ax3_init in T.serial(0, 64): - tensor_2[(((ax0_ax1_fused_4*3584) + (ax2_4*64)) + ax3_init)] = T.uint8(0) - for rv0_rv1_fused_1, ax3_2 in T.grid(9, 64): - tensor_2[(((ax0_ax1_fused_4*3584) + (ax2_4*64)) + ax3_2)] = T.max(tensor_2[(((ax0_ax1_fused_4*3584) + (ax2_4*64)) + ax3_2)], T.if_then_else(((((ax0_ax1_fused_4*2) + T.floordiv(rv0_rv1_fused_1, 3)) < 112) and (((ax2_4*2) + T.floormod(rv0_rv1_fused_1, 3)) < 112)), placeholder_29[(((((ax0_ax1_fused_4*14336) + (T.floordiv(rv0_rv1_fused_1, 3)*7168)) + (ax2_4*128)) + (T.floormod(rv0_rv1_fused_1, 3)*64)) + ax3_2)], T.uint8(0), dtype="uint8")) - for ax0_ax1_fused_5 in T.serial(0, 56): - for ax2_5, ax3_3 in T.grid(56, 64): - T_cast_7[(((ax0_ax1_fused_5*3584) + (ax2_5*64)) + ax3_3)] = T.cast(tensor_2[(((ax0_ax1_fused_5*3584) + (ax2_5*64)) + ax3_3)], "int16") - - @T.prim_func - def tvmgen_default_run_model(input: T.handle, output: T.handle) -> None: - # function attr dict - T.func_attr({"global_symbol": "tvmgen_default_run_model", "runner_function": True}) - # body - T.attr("default", "device_id", 0) - T.attr("default", "device_type", 1) - sid_9 = T.allocate([301056], "int8", "global") - sid_8 = T.allocate([802816], "int8", "global") - T.evaluate(T.call_extern("tvmgen_default_fused_cast_subtract", input, T.lookup_param("p0", dtype="handle"), sid_9, dtype="int32")) - T.evaluate(T.call_extern("tvmgen_default_fused_nn_conv2d_add_fixed_point_multiply_clip_cast", sid_9, T.lookup_param("p1", dtype="handle"), T.lookup_param("p2", dtype="handle"), sid_8, dtype="int32")) - T.evaluate(T.call_extern("tvmgen_default_fused_nn_max_pool2d_cast", sid_8, output, dtype="int32")) -# fmt: on - - -def test_create_pool_info(): - target = Target("c") - pool_info = WorkspacePoolInfo( - "foo_workspace", - [target], - ) - assert pool_info.pool_name == "foo_workspace" - # default pool size constraint - assert pool_info.size_hint_bytes == -1 - - pool_info = WorkspacePoolInfo( - "bar_workspace", - [target], - PoolInfoProperties(size_hint_bytes=1425), - ) - assert pool_info.pool_name == "bar_workspace" - assert pool_info.size_hint_bytes == 1425 - - -def test_create_buffer_info(): - global_ws_pool = WorkspacePoolInfo( - "global_workspace", - [Target("c")], - ) - buffer_info_obj = tvm.tir.usmp.BufferInfo( - name_hint="buf1", size_bytes=256, pool_candidates=[global_ws_pool] - ) - assert buffer_info_obj.name_hint == "buf1" - assert buffer_info_obj.size_bytes == 256 - assert list(buffer_info_obj.pool_candidates) == [global_ws_pool] - # default workspace alignment - assert buffer_info_obj.alignment == 1 - - buffer_info_obj = tvm.tir.usmp.BufferInfo("buf2", 512, [global_ws_pool], 8) - assert buffer_info_obj.name_hint == "buf2" - assert buffer_info_obj.size_bytes == 512 - assert list(buffer_info_obj.pool_candidates) == [global_ws_pool] - assert buffer_info_obj.alignment == 8 - - -def test_create_pool_allocation(): - pool_info = WorkspacePoolInfo( - "foo_workspace", - [Target("c")], - ) - pool_allocation = usmp_utils.PoolAllocation(pool_info=pool_info, byte_offset=64) - assert pool_allocation.pool_info == pool_info - assert pool_allocation.byte_offset == 64 - - -def _assign_poolinfos_to_allocates_in_primfunc(primfunc, pool_infos): - """helper to assing poolinfos to allocate nodes in a tir.PrimFunc""" - - def set_poolinfos(stmt): - if isinstance(stmt, tvm.tir.Allocate): - return tvm.tir.Allocate( - buffer_var=stmt.buffer_var, - dtype=stmt.dtype, - extents=stmt.extents, - condition=stmt.condition, - body=stmt.body, - annotations={tvm.tir.usmp.utils.CANDIDATE_MEMORY_POOL_ATTR: pool_infos}, - ) - - return primfunc.with_body(stmt_functor.ir_transform(primfunc.body, None, set_poolinfos)) - - -def _assign_poolinfos_to_allocates_in_irmodule(mod, pool_infos): - """helper to assing poolinfos to allocate nodes in a IRModule""" - ret = tvm.IRModule() - for global_var, basefunc in mod.functions.items(): - if isinstance(basefunc, tvm.tir.PrimFunc): - ret[global_var] = _assign_poolinfos_to_allocates_in_primfunc(basefunc, pool_infos) - return ret - - -def _assign_targets_to_primfuncs_irmodule(mod, target): - """helper to assign target for PrimFunc in a IRModule""" - ret = tvm.IRModule() - for global_var, basefunc in mod.functions.items(): - if isinstance(basefunc, tvm.tir.PrimFunc): - ret[global_var] = basefunc.with_attr("target", target) - return ret - - -def test_create_array_buffer_info(): - target = Target("c") - global_ws_pool = WorkspacePoolInfo( - "global_workspace", - [target], - ) - fcreate_array_bi = tvm.get_global_func("tir.usmp.CreateArrayBufferInfo") - tir_mod = LinearStructure - tir_mod = _assign_targets_to_primfuncs_irmodule(tir_mod, target) - tir_mod = _assign_poolinfos_to_allocates_in_irmodule(tir_mod, [global_ws_pool]) - main_func = tir_mod["tvmgen_default_run_model"] - buffer_info_analysis = tvm.tir.usmp.analysis.extract_buffer_info(main_func, tir_mod) - buffer_info_array = fcreate_array_bi(buffer_info_analysis.buffer_info_stmts) - for buffer_info in buffer_info_array: - assert buffer_info in buffer_info_analysis.buffer_info_stmts.keys() - - -if __name__ == "__main__": - tvm.testing.main() diff --git a/tests/scripts/task_build_adreno_bins.sh b/tests/scripts/task_build_adreno_bins.sh index 412af4928123..e5775c10ec34 100755 --- a/tests/scripts/task_build_adreno_bins.sh +++ b/tests/scripts/task_build_adreno_bins.sh @@ -40,7 +40,6 @@ fi echo set\(USE_RPC ON\) >> config.cmake echo set\(USE_CPP_RPC ON\) >> config.cmake echo set\(USE_CPP_RTVM ON\) >> config.cmake -echo set\(USE_GRAPH_EXECUTOR ON\) >> config.cmake echo set\(USE_LIBBACKTRACE AUTO\) >> config.cmake echo set\(USE_KALLOC_ALIGNMENT 32\) >> config.cmake diff --git a/tests/scripts/task_config_build_adreno.sh b/tests/scripts/task_config_build_adreno.sh index cf8917c9a546..10fefefbe800 100755 --- a/tests/scripts/task_config_build_adreno.sh +++ b/tests/scripts/task_config_build_adreno.sh @@ -29,6 +29,5 @@ echo set\(USE_CLML ${ADRENO_OPENCL}\) >> config.cmake fi echo set\(USE_OPENCL ON\) >> config.cmake echo set\(USE_RPC ON\) >> config.cmake -echo set\(USE_GRAPH_EXECUTOR ON\) >> config.cmake echo set\(USE_LIBBACKTRACE AUTO\) >> config.cmake echo set\(USE_LLVM ON\) >> config.cmake diff --git a/tests/scripts/task_config_build_arm.sh b/tests/scripts/task_config_build_arm.sh index 48ce67f9f790..3f0505310ba3 100755 --- a/tests/scripts/task_config_build_arm.sh +++ b/tests/scripts/task_config_build_arm.sh @@ -25,7 +25,6 @@ cp ../cmake/config.cmake . echo set\(USE_SORT ON\) >> config.cmake echo set\(USE_RPC ON\) >> config.cmake -echo set\(USE_PROFILER ON\) >> config.cmake echo set\(USE_LLVM llvm-config-17\) >> config.cmake echo set\(CMAKE_CXX_FLAGS -Werror\) >> config.cmake echo set\(USE_ARM_COMPUTE_LIB ON\) >> config.cmake diff --git a/tests/scripts/task_config_build_cpu.sh b/tests/scripts/task_config_build_cpu.sh index 9e195de9bc17..f9065ece6e5f 100755 --- a/tests/scripts/task_config_build_cpu.sh +++ b/tests/scripts/task_config_build_cpu.sh @@ -24,7 +24,6 @@ cd "$BUILD_DIR" cp ../cmake/config.cmake . echo set\(USE_SORT ON\) >> config.cmake -echo set\(USE_PROFILER ON\) >> config.cmake echo set\(USE_DNNL ON\) >> config.cmake echo set\(USE_ARM_COMPUTE_LIB ON\) >> config.cmake echo set\(USE_LLVM \"/usr/bin/llvm-config-17 --link-static\"\) >> config.cmake diff --git a/tests/scripts/task_config_build_gpu.sh b/tests/scripts/task_config_build_gpu.sh index e411ee2c5e87..2d5600c51369 100755 --- a/tests/scripts/task_config_build_gpu.sh +++ b/tests/scripts/task_config_build_gpu.sh @@ -33,9 +33,7 @@ echo set\(USE_OPENCL_GTEST \"/googletest\"\) >> config.cmake echo set\(USE_LLVM \"/usr/bin/llvm-config-15 --link-static\"\) >> config.cmake echo set\(USE_RPC ON\) >> config.cmake echo set\(USE_SORT ON\) >> config.cmake -echo set\(USE_GRAPH_EXECUTOR ON\) >> config.cmake echo set\(USE_STACKVM_RUNTIME ON\) >> config.cmake -echo set\(USE_PROFILER ON\) >> config.cmake echo set\(USE_ANTLR ON\) >> config.cmake echo set\(USE_BLAS openblas\) >> config.cmake echo set\(CMAKE_CXX_FLAGS -Werror\) >> config.cmake diff --git a/tests/scripts/task_config_build_gpu_other.sh b/tests/scripts/task_config_build_gpu_other.sh index 747e1006e507..115089207d91 100755 --- a/tests/scripts/task_config_build_gpu_other.sh +++ b/tests/scripts/task_config_build_gpu_other.sh @@ -27,7 +27,6 @@ cp ../cmake/config.cmake . echo set\(USE_OPENCL ON\) >> config.cmake echo set\(USE_ROCM ON\) >> config.cmake -echo set\(USE_PROFILER ON\) >> config.cmake echo set\(USE_LIBBACKTRACE OFF\) >> config.cmake echo set\(CMAKE_CXX_FLAGS -Werror\) >> config.cmake echo set\(BACKTRACE_ON_SEGFAULT ON\) >> config.cmake diff --git a/tests/scripts/task_config_build_i386.sh b/tests/scripts/task_config_build_i386.sh index f5cbad42bbf2..675b604f35a7 100755 --- a/tests/scripts/task_config_build_i386.sh +++ b/tests/scripts/task_config_build_i386.sh @@ -25,7 +25,6 @@ cp ../cmake/config.cmake . echo set\(USE_SORT ON\) >> config.cmake echo set\(USE_RPC ON\) >> config.cmake -echo set\(USE_PROFILER ON\) >> config.cmake echo set\(USE_LLVM llvm-config-10\) >> config.cmake echo set\(CMAKE_CXX_FLAGS -Werror\) >> config.cmake echo set\(USE_CCACHE OFF\) >> config.cmake diff --git a/tests/scripts/task_config_build_jvm.sh b/tests/scripts/task_config_build_jvm.sh index cf23c848127d..593f226f6ca2 100755 --- a/tests/scripts/task_config_build_jvm.sh +++ b/tests/scripts/task_config_build_jvm.sh @@ -27,7 +27,6 @@ cp ../cmake/config.cmake . echo set\(USE_SORT ON\) >> config.cmake echo set\(USE_RPC ON\) >> config.cmake -echo set\(USE_PROFILER ON\) >> config.cmake echo set\(CMAKE_CXX_FLAGS -Werror\) >> config.cmake echo set\(USE_CCACHE OFF\) >> config.cmake echo set\(SUMMARIZE ON\) >> config.cmake diff --git a/tests/scripts/task_config_build_static.sh b/tests/scripts/task_config_build_static.sh index 7d5c1f59592d..ed4ce73a283b 100755 --- a/tests/scripts/task_config_build_static.sh +++ b/tests/scripts/task_config_build_static.sh @@ -30,6 +30,3 @@ echo set\(BUILD_STATIC_RUNTIME ON\) >> config.cmake echo set\(USE_FALLBACK_STL_MAP ON\) >> config.cmake echo set\(USE_MSVC_MT ON\) >> config.cmake echo set\(USE_RPC OFF\) >> config.cmake -echo set\(USE_GRAPH_EXECUTOR OFF\) >> config.cmake -echo set\(USE_PROFILER OFF\) >> config.cmake -echo set\(USE_AOT_EXECUTOR OFF\) >> config.cmake diff --git a/tests/scripts/task_config_build_wasm.sh b/tests/scripts/task_config_build_wasm.sh index d92bb83deba4..daac6c0a0c34 100755 --- a/tests/scripts/task_config_build_wasm.sh +++ b/tests/scripts/task_config_build_wasm.sh @@ -24,7 +24,6 @@ cd "$BUILD_DIR" cp ../cmake/config.cmake . echo set\(USE_SORT ON\) >> config.cmake -echo set\(USE_PROFILER ON\) >> config.cmake echo set\(USE_LLVM llvm-config-15\) >> config.cmake echo set\(USE_ANTLR ON\) >> config.cmake echo set\(CMAKE_CXX_FLAGS -Werror\) >> config.cmake