iT邦幫忙

2026 iThome 鐵人賽

DAY 10
0
AI Engineering

一行 torch.compile 背後發生了什麼?30 天深度拆解 PyTorch 編譯器系列 第 10

Day 10 | TorchDynamo 的破「圖」重接,Graph Break and Resume!

  • 分享至 

  • xImage
  •  

前言

在 Day 4 InstructionTranslator 的結尾有一句話說:Graph Break 不是查表失敗,而是 handler 做到一半主動舉手。前面幾天我們把「乖乖翻完」的路整條走通了,今天回頭走那條一直沒走的岔路。什麼樣的程式碼會踩到斷點、一個函式怎麼被切成三段、resume function 怎麼做到「從函式中間開始跑」,以及為什麼深層函式裡的一個 print,會劈開最外層的圖。

正文開始!

會斷的不是 instruction,是 operand

Day 4 講過判斷發生的位置。dispatch_table 對幾乎每條 opcode 都有同名 handler,查表這一步不會落空。真正斷開的時機,是 handler 接下 instruction、看了 operand,發現這個操作沒辦法只靠符號值走下去。所以「會不會斷」不是 instruction 說了算,而是 operand 說了算。同一條 CALL,呼叫 torch.sin 進圖、呼叫自己寫的函式被 inline、呼叫 print 就斷。實務上最常撞到的大概是下面這幾類。

類別 例子 為什麼走不下去
依賴資料的控制流 if x.sum() > 0: 要有真值才知道往哪跳,符號 Tensor 沒有真假
把 Tensor 值抽成 Python 純量 .item()int(t).tolist() 值要離開圖回到 Python 世界,預設不追
有 side effect 的 builtin printinput 效果發生在真實世界,圖裡建不了模
沒有符號模型的 C 函式 第三方套件的 C extension 進不去 bytecode,也沒有對應的 handler

第一類我們在 Day 3 就親眼看過了,generic_jump 發現 stack 頂端是 TensorVariable,丟出 attempted to jump with TensorVariable()。第二類打開 torch._dynamo.config.capture_scalar_outputs 之後 .item() 不再斷,代價留到明天 Symbolic Shapes 那篇再算。第三類的話就是今天的主角。第四類對應 Day 5 講過的 trace_rules.py 名單,名單上要跳過的函式,呼叫它就是潛在的斷點。

Unsupported exception 是怎麼丟出來的

handler 丟錯的方式不是回傳一個錯誤碼,而是直接丟一個 Unsupported exception(exc.py),而且丟的時候固定要附上分類、現場、解釋和修法建議。這就是為什麼你看到的每一條 graph break 訊息都長同一個樣子,等一下的實驗就會看到實例。

那丟出來之後誰來接呢?CALL 這類最容易出事的 handler 外面包著一層接手的 decorator(symbolic_convert.pybreak_graph_if_unsupported)。接住之後分兩條路,fullgraph=True 的話不收拾,原樣往上丟變成使用者看到的錯誤,否則就進入 graph break 的收拾流程。

不過收拾流程有個麻煩,exception 飛出來的時候,翻譯常常做到一半,stack 疊了一半、SideEffects 的帳本記了一半,直接在這裡切圖就會切出錯的東西。在這裡 Dynamo 的解法很乾脆,毫不留情全部砍掉重練。記下「走到第 N 條 instruction 會失敗」之後整個重翻一遍,第二遍走到 N 就不往裡鑽,先把圖收乾淨、把斷點擺在乾淨的 instruction 邊界上。這就是為什麼等一下的實驗明明只編了一次,dynamo.explain 裡卻會出現兩筆 compile attempt。

斷開之後 Dynamo 做的三件事

收拾流程本身做的就是三件事。翻譯進行到某條 instruction、Unsupported 被接住之後,Dynamo 並不會放棄整個函式,而是照下面三步處理。

  1. 收前半段:到斷點為止的節點照 OutputGraph 的 compile_subgraph 收成一張圖,該結的 SideEffects 結清,再交給後端編譯。
  2. 讓那條 instruction 回 eager:生成的 bytecode 裡,斷點 instruction 原樣保留,讓 CPython 自己跑。
  3. 把剩下的包成 resume function:斷點之後的 bytecode 被包成一個新函式,用 PyCodegen 的工具箱直接生出來。

而關鍵就在第三步的後續。resume function 也是函式,一被呼叫,eval hook 照樣攔截它,於是斷點之後的程式碼就編成了第二張圖。一次 break 的結果,就是兩張圖,中間夾一小段 eager。

用 print 實際跑一次

def f(x):
    x = x * 2
    print("mid")
    return x + 1

TORCH_LOGS="graph_breaks,graph_code,bytecode" 就可以一次把完整現場跑出來。先看還沒被動過手腳的原始 bytecode,print 那條 CALL 在 offset 24、它算完的下一條 POP_TOP 在 offset 32,這兩個數字待會都還會再出現。

ORIGINAL BYTECODE f
 15           2 LOAD_FAST     0 (x)
              4 LOAD_CONST    1 (2)
              6 BINARY_OP     5 (*)
             10 STORE_FAST    0 (x)
 16          12 LOAD_GLOBAL   1 (NULL + print)
             22 LOAD_CONST    2 ('mid')
             24 CALL          1
             32 POP_TOP
 17          34 LOAD_FAST     0 (x)
             36 LOAD_CONST    3 (1)
             38 BINARY_OP     0 (+)
             42 RETURN_VALUE

翻譯走到 offset 24 的 CALL,handler 舉手,log 印出來的正好就是 unimplemented_v2 的那四件套。

Graph break in user code at graph_break.py:16
Graph Break Reason: Failed to trace builtin operator
  Explanation: Dynamo does not know how to trace builtin operator `print`
               with argument types ['str'] (has_kwargs False)
  Hint: Avoid calling builtin `print` with argument types ['str']. Consider
        using an equivalent alternative function/method to `print`.
  Hint: If you are attempting to call a logging function (e.g. `print`), you
        can try adding it to `torch._dynamo.config.reorderable_logging_functions`.

第二條 Hint 值得停一下。如果只是想留 log,把 print 加進 reorderable_logging_functions,Dynamo 會把這類呼叫挪到圖跑完再執行,圖就不用斷。這是官方給的正規修法之一。

前半段收成圖一,內容就只有 x * 2

 ===== __compiled_fn_2 =====
def forward(self, L_x_: "f32[4][1]cuda:0"):
    x: "f32[4][1]cuda:0" = l_x_ * 2
    return (x,)

接著來看改寫後的 bytecode 是怎麼把三段接起來的,下面是節錄(profiler 標記已省略)。

MODIFIED BYTECODE f
   2 LOAD_GLOBAL   5 (NULL + __compiled_fn_2_...)   <- 圖一
  56 LOAD_FAST     0 (x)
 104 CALL          1
 112 STORE_FAST    1 (graph_out_0)
 116 LOAD_GLOBAL   2 (__builtins_dict___1)          <- 斷點 instruction 回 eager
 126 LOAD_CONST    4 ('print')
 128 BINARY_SUBSCR
 132 LOAD_CONST    2 ('mid')
 134 LOAD_FAST     1 (graph_out_0)
 136 LOAD_CONST    5 (0)
 138 BINARY_SUBSCR
 142 STORE_FAST    0 (x)                            <- 圖一的輸出放回 x
 146 CALL          1
 154 LOAD_GLOBAL  13 (NULL + __resume_at_32_3_...)  <- 剩下包成 resume fn
 172 LOAD_FAST     0 (x)
 174 CALL          2
 182 RETURN_VALUE

三段的接縫全都在這裡了。先呼叫 __compiled_fn_2 把前半算完、輸出放回 x,接著讓 CPython 真的呼叫 print,最後呼叫 __resume_at_32_3,把 print 的回傳值和 x 一起傳進去。這個名字的意思是「從原函式第 32 個 byte 繼續」,正好就是上面 POP_TOP 的位置,斷點 instruction 自己留在外面跑,它之後的第一條 instruction 才是續集的起點。

resume function:從函式中間開始跑

接著 log 裡就會出現 resume function 自己的 ORIGINAL BYTECODE,code object 的名字叫 torch_dynamo_resume_in_f_at_16(f 的第 16 行)。

ORIGINAL BYTECODE torch_dynamo_resume_in_f_at_16
 16           0 RESUME         0
              2 LOAD_FAST      0 (___stack0)
              4 JUMP_FORWARD  16 (to 38)
              6 RESUME         0
              8 LOAD_FAST      1 (x)
             ...(原函式的 bytecode 原樣跟在後面)
        >>   38 POP_TOP
 17          40 LOAD_FAST      1 (x)
             42 LOAD_CONST     3 (1)
             44 BINARY_OP      0 (+)
             48 RETURN_VALUE

它的參數表就是斷點當下所有還活著的狀態,也就是 stack 上的值(___stack0,這裡是 print 的回傳值)加上之後還會被讀的 locals(x)。開場白把 ___stack0 擺回原位,然後一條 JUMP_FORWARD 直接跳到斷點之後的第一條 instruction。合法的 Python 是寫不出「從函式第 32 個 byte 開始跑」這種函式的,但 bytecode 層沒有這個限制,resume_execution.py 就直接把它拼了出來。同一個斷點的 resume function 只生成一次,之後重用。

然後劇本就重演一遍。resume function 一被呼叫,eval hook 攔截它,return x + 1 編成下面的圖二。

 ===== __compiled_fn_5 =====
def forward(self, L_x_: "f32[4][1]cuda:0"):
    add: "f32[4][1]cuda:0" = l_x_ + 1
    return (add,)

dynamo.explain(f)(x) 的摘要也印證了整件事。

Graph Count: 2
Graph Break Count: 1
Op Count: 2
Ops per Graph:
  Ops 1: <built-in function mul>
  Ops 2: <built-in function add>

兩張圖各領一個 op,中間夾著那聲 mid。斷開、掉落、縫回這一整段旅程,用動畫走一遍就是下面這張圖。

翻譯游標沿 f 的 bytecode 逐條往下走,在 print 的 CALL 丟出 Unsupported,前半坍縮成圖一、print 三條掉回 eager、剩下包成 resume function 又被 eval hook 攔截編成圖二,最後執行路徑把三段縫回一條

圖一:翻譯游標逐條走過 f 的 bytecode,在 printCALL 丟出 Unsupported:斷點前的四條 instruction 坍縮成 __compiled_fn_2print 那三條原樣掉回 eager、offset 32 之後包成 __resume_at_32_3 再被 eval hook 攔截編成 __compiled_fn_5,最後執行順序把三段縫回一條路。

inline 裡的斷點會往上傳

那如果是被 inline 的函式(Day 5)內部發生 Unsupported 呢?Dynamo 是沒辦法在子函式裡斷的,因為子函式沒有自己的 frame,斷在裡面沒有地方可以 resume。所以 inline 途中的斷點會讓整個 call 在 caller 端變成 break 點,一層層往上傳,直到真正有 frame 的邊界為止。

def util(t):
    print("log")
    return t + 1

def big(x):
    y = x * 2
    z = util(y)
    return z * 3

翻譯 big 時,Dynamo 鑽進 util 撞到 print,於是 z = util(y) 這整條 CALLbig 裡變成斷點。前半段收成圖一(mul),util(y) 改由 CPython 真的呼叫。但 util 自己是一個 frame,eval hook 又攔截它,在它自己的 frame 裡再斷一次(print 前面沒有任何 Tensor 運算,所以沒有圖)、再 resume(t + 1 收成圖二)。最後 big 的 resume function 把 z * 3 收成圖三。explain 的結果也印證了這件事。

Graph Count: 3
Graph Break Count: 2
Op Count: 3
User Stack:
  <FrameSummary graph_break.py, line 39 in big>
  <FrameSummary graph_break.py, line 34 in util>

三張圖(muladdmul 各一張)、兩次 break,User Stack 指著兩層,斷點被回報在 big 的第 39 行,但兇手在 util 的第 34 行。所以抓 break 的時候,兇手常常不在報告的第一行,而是躲在它呼叫的函式深處。

抓 break 的工具箱

Break 不會報錯,只會默默變慢,所以我們得主動去抓,工具有三件。

  • torch.compile(f, fullgraph=True):把所有 break 升級成錯誤。前面說過 break_graph_if_unsupported 在這個模式下不收拾、直接往上丟,實跑就是一個 Unsupported exception,訊息跟 log 裡的一模一樣。

    Unsupported : Failed to trace builtin operator
      Explanation: Dynamo does not know how to trace builtin operator `print`
                   with argument types ['str'] (has_kwargs False)
      Hint: Avoid calling builtin `print` with argument types ['str']. ...
    

    上線前用它掃一輪最省事,修到能跑等於保證整個函式一張圖。

  • torch._dynamo.explain(f)(*args):不改行為,列出每張圖、每個 break 的位置、原因和 User Stack,適合拿來看全貌。

  • TORCH_LOGS="graph_breaks":執行時逐個印,適合掛在長跑的 job 上。

抓到之後該怎麼修,訊息裡的 Hint 通常已經把方向給好了。print 移出熱路徑或加進 reorderable_logging_functions.item() 考慮 capture_scalar_outputs,依賴資料的 iftorch.cond 改寫(Day 3 的 Hint 就這麼說)。第三方 C 函式則把它挪出被編譯的函式,讓斷點斷在便宜的地方。

結語

最後來收個帳,一次 break 至少要付四筆。圖被切小,跨不過斷點的 fusion 機會全沒了。中間那段 eager 本身就慢,SideEffects 的帳本被迫提前結算,resume function 還要多編一次(上面 util 的例子,一個 print 換來 bigutil、兩個 resume 共四次編譯)。再加上 Day 3 說過的,compiled function 執行期間每個 frame 都要繞進 Dynamo 判斷一次,圖越碎、繞的次數就越多。所以效能調校的第一課永遠是先數 break,再談其他。

今天我們看到,一次 break 會把函式切成三段,前半編成圖一、斷點那條 instruction 回 eager、剩下的包成 resume function 再被攔截編成圖二。而斷點必須落在有 frame 的地方,所以 inline 子函式裡的 break 會傳染到 caller 的那條 CALL,連鎖出一串圖和 resume function。

到今天為止,Dynamo 的主線就完整了,攔截、翻譯、包裝、驗票、記帳、收圖、寫碼、斷了再接。不過還剩下最後一塊拼圖。那條 EQUALS_MATCH: L['n'] == 3 的 Guard 實在太窄了,值一變就得重編。明天就來講 Symbolic Shapes,看 SymInt 怎麼把具體的 4 換成符號 s0、ShapeEnv 又是怎麼管理符號之間的約束,讓一張圖吃下所有 batch size。那我們明天見!

參考資料


上一篇
Day 9 | 新的 bytecode 誰來寫?TorchDynamo PyCodegen!
下一篇
Day 11 | TorchDynamo 的伸縮量尺 Symbolic Shapes
系列文
一行 torch.compile 背後發生了什麼?30 天深度拆解 PyTorch 編譯器11
圖片
  熱門推薦
圖片
{{ item.channelVendor }} | {{ item.webinarstarted }} |
{{ formatDate(item.duration) }}
直播中

尚未有邦友留言

立即登入留言