完整指南:從 TorchDynamo 到 AOTAutograd 的接入實戰)
PyTorch 自定義編譯后端Custom Backends完整指南從 TorchDynamo 到 AOTAutograd 的接入實戰【免費下載鏈接】pytorchTensors and Dynamic neural networks in Python with strong GPU acceleration項目地址: https://gitcode.com/GitHub_Trending/py/pytorchtorch.compile為用戶提供了一條直接定義自定義編譯后端backend的通道一個后端本質上是一個以torch.fx.GraphModule和示例輸入為參數、返回等價可調用對象的函數。本文以 torch.compiler_custom_backends.md 為骨架結合當前倉庫中的注冊表實現registry.py、AOTAutograd 封裝common.py與后端初始化鉤子eval_frame.py源碼系統講解后端的契約、注冊方式、AOTAutograd 訓練后端接入、eager 初始化鉤子以及 Debug/Speedy/Composable 三類典型后端示例。讀完本文你將能夠獨立編寫、注冊并組合自己的torch.compile后端。后端契約一個函數連接 Dynamo 與編譯產物torch.compile的圖追蹤組件 TorchDynamo 在完成字節碼分析并抽取出一張 FX 圖之后會調用用戶提供的后端函數。后端函數必須滿足如下契約(gm: torch.fx.GraphModule, example_inputs: List[torch.Tensor]) - Callable其中gm是 Dynamo 從用戶代碼中抽取出的 FXGraphModuleexample_inputs是用于推導形狀等信息的示例輸入張量列表在 Dynamo 內部通常為 FakeTensor返回值是一個“已編譯函數”其行為必須與傳入的 FX 圖等價。返回的可調用對象契約與原torch.fx.GraphModule的forward一致(*args: torch.Tensor) - List[torch.Tensor]在 registry.py 中這兩類簽名被形式化定義為CompiledFn與CompilerFn兩個類型class CompiledFn(Protocol): def __call__(self, *args: torch.Tensor) - tuple[torch.Tensor, ...]: ... CompilerFn Callable[[fx.GraphModule, list[torch.Tensor]], CompiledFn]要讓 TorchDynamo 調用你的后端只需把后端函數作為backend關鍵字參數傳給torch.compile既支持函數式調用也支持裝飾器形式import torch def my_custom_backend(gm, example_inputs): return gm.forward def f(...): ... f_opt torch.compile(f, backendmy_custom_backend) torch.compile(backendmy_custom_backend) def g(...): ...最樸素的后端就是直接返回gm.forward即“不優化、僅捕獲”的 eager 語義后端這也是后續所有示例的起點。注冊自定義后端裝飾器與 Entry Points 雙通道使用register_backend裝飾器你可以用torch._dynamo.register_backend裝飾器將后端注冊進全局注冊表from torch._dynamo import register_backend register_backend def my_compiler(gm, example_inputs): ...從源碼看register_backend 的核心邏輯如下若不傳name則默認使用compiler_fn.__name__作為后端名tags參數可對后端打標簽分類倉庫預置了register_debug_backendtags(debug,)與register_experimental_backendtags(experimental,)兩個便捷變體見 registry.py注冊后后端函數會同時被寫入_BACKENDS、_COMPILER_FNS、_BACKEND_TAGS三個全局字典若名字重復會拋出AssertionError(duplicate name: ...)。注冊帶來的直接收益是你可以用字符串代替函數本身來調用后端例如torch.compile(model, backendmy_compiler)。lookup_backendregistry.py會把字符串展開為真正的函數先在_BACKENDS中查找未命中則觸發_lazy_import()與_discover_entrypoint_backends()仍找不到時會借助difflib.get_close_matches給出相近名字建議并拋出InvalidBackend。通過 Python 包 Entry Points 注冊如果你的后端位于獨立的 Python 包中可以通過 setuptools 的 entry points 機制注冊插件這也是“一個包為另一個包注冊插件”的標準做法。在包的setup.py中加入torch_dynamo_backends分組即可setup( ... entry_points{ torch_dynamo_backends: [ my_compiler your_module.submodule:my_compiler, ], }, ... )規則前是后端名稱后是后端函數所在的模塊路徑與函數名。包安裝后 entry point 即注入當前 Python 環境調用torch.compile(model, backendmy_compiler)時PyTorch 會先搜索register_backend注冊的同名后端未命中再掃描 entry points 注冊的后端。該機制的底層實現在 registry.py 的_discover_entrypoint_backends通過importlib.metadata.entry_points(grouptorch_dynamo_backends)枚舉該分組的全部 entry point存入_BACKENDS字典lookup_backend在首次按名訪問時通過entry_point.load()惰性加載并完成注冊。注冊有兩個明確的用途允許以字符串形式傳入torch.compile供 minifier最小化復現工具使用——minifier 生成的任何代碼都必須調用注冊你后端函數的代碼通常是通過一條import語句實現。可用后端一覽倉庫內置大量后端可用torch._dynamo.list_backends()列出。其實現registry.py默認排除debug與experimental兩類標簽的后端返回按字母序排序的名字列表def list_backends(exclude_tags(debug, experimental)) - list[str]: ... return sorted(backends)面向訓練的 AOTAutograd 后端TorchDynamo 之外你還可以定義由 AOTAutograd 調用的自定義后端。這有兩個核心價值支持訓練AOTAutograd 能生成用于編譯的反向圖因此這類后端天然支持模型訓練更小的算子集合AOTAutograd 產出的 FX 圖僅由 core Aten 算子 構成后端只需支持遠小于完整 torch/Aten 算子集的 core Aten opset實現成本顯著降低。接入方式是用torch._dynamo.backends.common.aot_autograd包裝你的后端再照常通過backend關鍵字傳給torch.compile。被包裝的后端函數契約與之前完全一致。后端通過fw_compiler前向編譯器與bw_compiler反向編譯器兩個關鍵字參數傳入aot_autograd若未指定bw_compiler反向編譯函數默認復用前向編譯函數。一個必須注意的細節AOTAutograd 要求后端返回的編譯函數是“boxed”的即用functorch.compile.make_boxed_func包裝from torch._dynamo.backends.common import aot_autograd from functorch.compile import make_boxed_func def my_compiler(gm, example_inputs): return make_boxed_func(gm.forward) my_backend aot_autograd(fw_compilermy_compiler) # bw_compilermy_compiler model_opt torch.compile(model, backendmy_backend)從源碼看AotAutograd 的實現要點包括構造時把所有關鍵字參數存入self.kwargs運行期轉發給torch._functorch.aot_autograd.aot_module_simplified見 common.py若示例輸入含 list/tuple/dict 等結構會先經flatten_graph_inputs展平common.py反向編譯器會被雙重disable包裝既阻止 Dynamo 追蹤 bw_compiler 函數本身也阻止追蹤其生成的 backward passcommon.pybw_compiler缺省時回退到fw_compilerinference_compiler同理common.py可通過decompositions關鍵字傳入分解表也支持返回表的零參 thunk以規避循環導入問題并在運行期解析為具體表common.py。其余可用關鍵字參數定義于 AotAutogradKwargs還包括partition_fn前反向圖切分策略如min_cut_rematerialization_partition、keep_inference_input_mutations、ignore_shape_env、disable_functionalization、pre_grad_passes、compile_region_name等完整列表可查看該 TypedDict 定義。Eager 后端初始化_dynamo_backend_init鉤子有些后端需要在torch.compile()時刻執行 eager 初始化例如加載原生庫或初始化設備上下文。此時可以給后端定義一個_dynamo_backend_init屬性——一個無參可調用對象在后端被解析時任何一次調用發生之前觸發def my_backend(gm, example_inputs): return gm.forward def my_backend_init(): load_native_libs() # 在 compile() 時刻運行先于任何調用 my_backend._dynamo_backend_init my_backend_init torch.compile(backendmy_backend) def fn(x): return x 1該鉤子的觸發點在 eval_frame.py 的_maybe_fire_backend_initget_compiler_fn在lookup_backend解析后端之后、wrap_backend_debug包裝之前調用它實現為“每次解析都觸發”。由于屬性通過getattr從后端對象讀取因此無論是實例屬性還是可經 MRO 解析的類方法都能命中。鉤子行為要點生效范圍廣無論后端是直接傳入、經register_backend按名注冊還是通過torch.compiler.set_stance(force_backend...)強制指定鉤子都會觸發。torch._TorchCompileWrapper與AotAutograd都通過property把該屬性轉發給它們所包裝的后端見 torch/init.py 與 common.py。AOTAutograd 場景使用aot_autograd(fw_compiler...)時把鉤子設在內部的fw_compiler上即可——AotAutograd在觸發時刻而非構造時刻讀取它因此鉤子可以在aot_autograd()構造之前或之后設置。只有fw_compiler上的鉤子會被查詢設在bw_compiler或inference_compiler上的鉤子會被忽略。每次解析都會觸發正常路徑與fullgraphTrue路徑都會觸發且發生在任何調用之前因此環境損壞時torch.compile()會快速失敗fail fast。后端被重復解析如set_stance(force_backend...)或compiled_autograd重建路徑時每次解析觸發一次需要一次性初始化的后端應在鉤子內部自行去重import functools functools.cache # 每個進程僅執行一次 def my_backend_init(): load_native_libs() my_backend._dynamo_backend_init my_backend_init異常傳播若鉤子拋出異常異常會從torch.compile()中直接向外傳播在torch.compiler.set_stance(force_backend...)場景下解析發生在首次調用時因此鉤子觸發以及失敗從該次調用處浮出。實戰示例三類典型自定義后端Debugging Backend打印 Dynamo 抽取的 FX 圖想了解編譯過程中發生了什么可以寫一個打印 FX 圖并返回forward()的后端from typing import List import torch def my_compiler(gm: torch.fx.GraphModule, example_inputs: List[torch.Tensor]): print(my_compiler() called with FX graph:) gm.graph.print_tabular() return gm.forward # 返回一個 python callable torch.compile(backendmy_compiler) def fn(x, y): a torch.cos(x) b torch.sin(y) return a b fn(torch.randn(10), torch.randn(10))運行輸出如下表格化的 FX IRmy_compiler() called with FX graph: opcode name target args kwargs ------------- ------ ------------------------------------------------------ ---------- -------- placeholder x x () {} placeholder y y () {} call_function cos built-in method cos of type object at 0x7f1a894649a8 (x,) {} call_function sin built-in method sin of type object at 0x7f1a894649a8 (y,) {} call_function add built-in function add (cos, sin) {} output output output ((add,),) {}同樣的后端同樣適用于torch.nn.Modulefrom typing import List import torch def my_compiler(gm: torch.fx.GraphModule, example_inputs: List[torch.Tensor]): print(my_compiler() called with FX graph:) gm.graph.print_tabular() return gm.forward # 返回一個 python callable class MockModule(torch.nn.Module): def __init__(self): super().__init__() self.relu torch.nn.ReLU() def forward(self, x): return self.relu(torch.cos(x)) mod MockModule() optimized_mod torch.compile(mod, backendmy_compiler) optimized_mod(torch.randn(10))再看一個含控制流的例子它直觀展示了 Dynamo 對條件分支的圖切分能力from typing import List import torch def my_compiler(gm: torch.fx.GraphModule, example_inputs: List[torch.Tensor]): print(my_compiler() called with FX graph:) gm.graph.print_tabular() return gm.forward # 返回一個 python callable torch.compile(backendmy_compiler) def toy_example(a, b): x a / (torch.abs(a) 1) if b.sum() 0: b b * -1 return x * b for _ in range(100): toy_example(torch.randn(10), torch.randn(10))運行會依次輸出三張子圖其中if b.sum() 0的分支被拆成獨立圖最后兩張圖的輸出順序取決于 JIT 編譯器先遇到哪一個是非確定性的my_compiler() called with FX graph: opcode name target args kwargs ------------- ------- ------------------------------------------------------ ---------------- -------- placeholder a a () {} placeholder b b () {} call_function abs_1 built-in method abs of type object at 0x7f8d259298a0 (a,) {} call_function add built-in function add (abs_1, 1) {} call_function truediv built-in function truediv (a, add) {} call_method sum_1 sum (b,) {} call_function lt built-in function lt (sum_1, 0) {} output output output ((truediv, lt),) {} my_compiler() called with FX graph: opcode name target args kwargs ------------- ------ ----------------------- ----------- -------- placeholder b b () {} placeholder x x () {} call_function mul built-in function mul (b, -1) {} call_function mul_1 built-in function mul (x, mul) {} output output output ((mul_1,),) {} my_compiler() called with FX graph: opcode name target args kwargs ------------- ------ ----------------------- --------- -------- placeholder b b () {} placeholder x x () {} call_function mul built-in function mul (x, b) {} output output output ((mul,),) {}Speedy Backend接入真實推理優化器接入一個性能更優的后端同樣簡單下面把torch.jit.optimize_for_inference集成進自定義后端def optimize_for_inference_compiler(gm: torch.fx.GraphModule, example_inputs: List[torch.Tensor]): scripted torch.jit.script(gm) return torch.jit.optimize_for_inference(scripted)隨后即可用它加速任意既有代碼torch.compile(backendoptimize_for_inference_compiler) def code_to_accelerate(): ...注意這里把 FX 圖torch.jit.script成 TorchScript 模塊再走 JIT 推理優化展示的是一條把 Dynamo 圖轉交給既有優化器管線的通用集成模式。Composable Backends后端組合與優雅降級TorchDynamo 內置了大量后端可用torch._dynamo.list_backends()列出。你可以組合多個后端實現“優先使用高性能后端、失敗則逐級降級”的策略from torch._dynamo import lookup_backend def my_compiler(gm: torch.fx.GraphModule, example_inputs: List[torch.Tensor]): try: trt_compiled lookup_backend(tensorrt)(gm, example_inputs) if trt_compiled is not None: return trt_compiled except Exception: pass # 第一個后端失敗嘗試其他后端... try: inductor_compiled lookup_backend(inductor)(gm, example_inputs) if inductor_compiled is not None: return inductor_compiled except Exception: pass return gm.forward該示例先用lookup_backend(tensorrt)嘗試 TensorRT 后端失敗拋異常或返回None則回退到inductor最終兜底返回gm.forward保證可用性。結合 registry.py 的實現可知lookup_backend對字符串會完成注冊表查找、entry point 惰性加載與名字建議等全套解析邏輯因此既可按名取用內置后端也可取用你自己注冊的后端使“組合后端”天然可擴展。小結與進一步閱讀自定義后端是torch.compile開放生態的關鍵接口核心是(gm, example_inputs) - callable這一簡潔契約register_backend與torch_dynamo_backendsentry points 提供了函數級與包級兩種注冊通道aot_autograd包裝把后端擴展到了前向/反向圖同時編譯的訓練場景并以 core Aten opset 大幅降低實現成本_dynamo_backend_init鉤子則讓后端能在編譯時刻完成 eager 環境準備并快速失敗。三者配合即可構建從調試、推理優化到多后端降級組合的完整后端體系。倉庫中可供繼續深挖的相關實現與測試后端注冊表全量實現torch/_dynamo/backends/registry.pyAOTAutograd 后端封裝torch/_dynamo/backends/common.py后端解析與初始化鉤子觸發點torch/_dynamo/eval_frame.py后端初始化鉤子轉發屬性torch/init.py后端注冊與 entry points 的測試用例test/dynamo/test_backends.py相關 IR 概念core Aten IRtorch.compiler_ir.md【免費下載鏈接】pytorchTensors and Dynamic neural networks in Python with strong GPU acceleration項目地址: https://gitcode.com/GitHub_Trending/py/pytorch創作聲明:本文部分內容由AI輔助生成(AIGC),僅供參考