在 Day 4 InstructionTranslator 的結尾有一句話說:Graph Break 不是查表失敗,而是 handler 做到一半主動舉手。前面幾天我們把「乖乖翻完」的路整條走通了,今天回頭走那條一直沒走的岔路。什麼樣的程式碼會踩到斷點、一個函式怎麼被切成三段、resume function 怎麼做到「從函式中間開始跑」,以及為什麼深層函式裡的一個 print,會劈開最外層的圖。
正文開始!
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 | print、input |
效果發生在真實世界,圖裡建不了模 |
| 沒有符號模型的 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 名單,名單上要跳過的函式,呼叫它就是潛在的斷點。
handler 丟錯的方式不是回傳一個錯誤碼,而是直接丟一個 Unsupported exception(exc.py),而且丟的時候固定要附上分類、現場、解釋和修法建議。這就是為什麼你看到的每一條 graph break 訊息都長同一個樣子,等一下的實驗就會看到實例。
那丟出來之後誰來接呢?CALL 這類最容易出事的 handler 外面包著一層接手的 decorator(symbolic_convert.py 的 break_graph_if_unsupported)。接住之後分兩條路,fullgraph=True 的話不收拾,原樣往上丟變成使用者看到的錯誤,否則就進入 graph break 的收拾流程。
不過收拾流程有個麻煩,exception 飛出來的時候,翻譯常常做到一半,stack 疊了一半、SideEffects 的帳本記了一半,直接在這裡切圖就會切出錯的東西。在這裡 Dynamo 的解法很乾脆,毫不留情全部砍掉重練。記下「走到第 N 條 instruction 會失敗」之後整個重翻一遍,第二遍走到 N 就不往裡鑽,先把圖收乾淨、把斷點擺在乾淨的 instruction 邊界上。這就是為什麼等一下的實驗明明只編了一次,dynamo.explain 裡卻會出現兩筆 compile attempt。
收拾流程本身做的就是三件事。翻譯進行到某條 instruction、Unsupported 被接住之後,Dynamo 並不會放棄整個函式,而是照下面三步處理。
compile_subgraph 收成一張圖,該結的 SideEffects 結清,再交給後端編譯。而關鍵就在第三步的後續。resume function 也是函式,一被呼叫,eval hook 照樣攔截它,於是斷點之後的程式碼就編成了第二張圖。一次 break 的結果,就是兩張圖,中間夾一小段 eager。
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 才是續集的起點。
接著 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:斷點前的四條 instruction 坍縮成 __compiled_fn_2、print 那三條原樣掉回 eager、offset 32 之後包成 __resume_at_32_3 再被 eval hook 攔截編成 __compiled_fn_5,最後執行順序把三段縫回一條路。
那如果是被 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) 這整條 CALL 在 big 裡變成斷點。前半段收成圖一(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>
三張圖(mul、add、mul 各一張)、兩次 break,User Stack 指著兩層,斷點被回報在 big 的第 39 行,但兇手在 util 的第 34 行。所以抓 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,依賴資料的 if 用 torch.cond 改寫(Day 3 的 Hint 就這麼說)。第三方 C 函式則把它挪出被編譯的函式,讓斷點斷在便宜的地方。
最後來收個帳,一次 break 至少要付四筆。圖被切小,跨不過斷點的 fusion 機會全沒了。中間那段 eager 本身就慢,SideEffects 的帳本被迫提前結算,resume function 還要多編一次(上面 util 的例子,一個 print 換來 big、util、兩個 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。那我們明天見!