PyTorchの@torch.compileでエラーになる

PyTorchで学習させようとしたら、これで死んでしまった。@torch.compileを外すと動く

治し方

本当はnightlyとかで治ってるはずなんですが、修正されたのが最近っぽいのでまだ上手くいかなかった。

このissueの通り、llvmのコンパイラの変換が間違った命令を出力しているのを踏んでいたので、llvm->tritonを差し替えるとうまくいくです。

tritonの使っているllvmを差し替えるにはこの辺を参照して

インストール方法

sudo apt install ninja-build
git clone https://github.com/llvm/llvm-project
cd llvm-project
mkdir build
cd build
cmake -G Ninja -DCMAKE_BUILD_TYPE=Release -DLLVM_ENABLE_ASSERTIONS=ON ../llvm -DLLVM_ENABLE_PROJECTS="mlir;llvm" -DLLVM_TARGETS_TO_BUILD="host;NVPTX;AMDGPU"
ninja
export LLVM_BUILD_DIR=$HOME/llvm-project/build
cd
git clone https://github.com/triton-lang/triton
cd triton
LLVM_INCLUDE_DIRS=$LLVM_BUILD_DIR/include LLVM_LIBRARY_DIR=$LLVM_BUILD_DIR/lib LLVM_SYSPATH=$LLVM_BUILD_DIR pip install -e python

こんな感じで、llvm-projectのmasterを持ってきて(修正が反映されてるやつ)ビルドしたものを、tritonで参照させてインストールします。

いいなと思ったら応援しよう!