Sakhanda Wire
NVDA $230.86 +1.09% MSFT $512.80 -0.02% GOOGL $338.24 -1.70% META $725.93 +0.10% AMZN $248.23 -0.37%
← К новостям

Внутри Graph API NVIDIA cuDNN: фьюжн, автотюнинг и повторное использование планов с cuDNN Frontend

В этом руководстве мы пошагово разбираем API графов cuDNN Frontend на уровне ниже фреймворка: описываем вычисления как граф операций, позволяем cuDNN выбрать движок для их выполнения, а затем сами берём управление этим выбором. Каждое создаваемое здесь ядро описывается одинаково: мы объявляем тензоры по их размерностям и шагам, объединяем операции в цепочку, запускаем пятиэтапный конвейер сборки — проверку, построение графа операций, создание планов выполнения, проверку поддержки и сборку планов — а затем выполняем его для набора указателей variant pack. Всё это запускается на одном GPU Colab, причём каждый результат проверяется по эталону PyTorch, чтобы увидеть и корректность fusion, и её стоимость. Темы последовательно усложняются: от одной свёртки с fusion до автотюнинга конфигураций движка, эпилогов в стиле FP8, attention, сериализации планов, динамических размеров и захвата CUDA-графов.

Копировать кодСкопированоИспользовать другой браузер
import os
import sys
import glob
import math
import time
import ctypes
import traceback
import subprocess
RESULTS = {}
def banner(title):
   print("\n" + "=" * 78)
   print(title)
   print("=" * 78)
def section(name):
   def wrap(fn):
       def run(*a, **kw):
           banner(name)
           try:
               out = fn(*a, **kw)
               RESULTS[name] = out if isinstance(out, str) else "ok"
               return out
           except Exception as e:
               RESULTS[name] = f"SKIPPED / FAILED -> {type(e).__name__}: {e}"
               print(f"\n[!] {name} did not complete: {type(e).__name__}: {e}")
               traceback.print_exc(limit=3)
               return None
       return run
   return wrap
banner("0. Install nvidia-cudnn-frontend and locate libcudnn")
subprocess.run(
   [sys.executable, "-m", "pip", "install", "-q", "nvidia-cudnn-frontend"],
   check=True,
)
import torch
assert torch.cuda.is_available(), "No GPU. Runtime -> Change runtime type -> GPU."
torch.backends.cudnn.enabled = True
_ = torch.nn.functional.conv2d(
   torch.randn(1, 1, 8, 8, device="cuda"), torch.randn(1, 1, 3, 3, device="cuda")
)
torch.cuda.synchronize()
try:
   import nvidia.cudnn
   _libdir = os.path.join(os.path.dirname(nvidia.cudnn.__file__), "lib")
   os.environ["CUDNN_PATH"] = os.path.dirname(nvidia.cudnn.__file__)
   os.environ["LD_LIBRARY_PATH"] = _libdir + ":" + os.environ.get("LD_LIBRARY_PATH", "")
   for _so in sorted(glob.glob(os.path.join(_libdir, "libcudnn*.so*"))):
       try:
           ctypes.CDLL(_so, mode=ctypes.RTLD_GLOBAL)
       except OSError:
           pass
except Exception as _e:
   print(f"  (no pip cuDNN package found, relying on system cuDNN: {_e})")
import cudnn
print("  cuDNN frontend imported successfully.")
banner("1. Environment")
DEV = torch.device("cuda")
MAJOR, MINOR = torch.cuda.get_device_capability()
SM = MAJOR * 10 + MINOR
CUDNN_VER = cudnn.backend_version()
print(f"  GPU                 : {torch.cuda.get_device_name(0)}")
print(f"  Compute capability  : sm_{SM}")
print(f"  Torch / CUDA        : {torch.__version__} / {torch.version.cuda}")
print(f"  cuDNN backend       : {CUDNN_VER}")
try:
   print(f"  cuDNN version str   : {cudnn.backend_version_string()}")
except Exception:
   pass
DTYPE = torch.bfloat16 if SM >= 80 else torch.float16
HAS_SDPA = SM >= 80
print(f"  Working dtype       : {DTYPE}")
print(f"  Fused SDPA usable   : {HAS_SDPA}")
HANDLE = cudnn.create_handle()
TORCH2CUDNN = {
   torch.float16: cudnn.data_type.HALF,
   torch.bfloat16: cudnn.data_type.BFLOAT16,
   torch.float32: cudnn.data_type.FLOAT,
   torch.int32: cudnn.data_type.INT32,
   torch.int64: cudnn.data_type.INT64,
   torch.int8: cudnn.data_type.INT8,
   torch.uint8: cudnn.data_type.UINT8,
}
def tensor_of(graph, t, name):
   return graph.tensor(
       name=name,
       dim=list(t.size()),
       stride=list(t.stride()),
       data_type=TORCH2CUDNN[t.dtype],
   )
def scalar_of(graph, name):
   return graph.tensor(
       name=name,
       dim=[1, 1, 1],
       stride=[1, 1, 1],
       data_type=cudnn.data_type.FLOAT,
       is_pass_by_value=True,
   )
def build(graph, heur=None, policy=None):
   heur = heur or [cudnn.heur_mode.A, cudnn.heur_mode.FALLBACK]
   graph.validate()
   graph.build_operation_graph()
   graph.create_execution_plans(heur)
   graph.check_support()
   if policy is None:
       graph.build_plans()
   else:
       graph.build_plans(policy)
   return graph
def workspace_for(graph):
   n = graph.get_workspace_size()
   return torch.empty(max(n, 1), device=DEV, dtype=torch.uint8)
def bench(fn, warmup=10, iters=50):
   for _ in range(warmup):
       fn()
   torch.cuda.synchronize()
   s, e = torch.cuda.Event(True), torch.cuda.Event(True)
   s.record()
   for _ in range(iters):
       fn()
   e.record()
   torch.cuda.synchronize()
   return s.elapsed_time(e) / iters
def tflops(flops, ms):
   return flops / (ms * 1e-3) / 1e12
def report(tag, ms, flops=None):
   extra = f"   ({tflops(flops, ms):7.2f} TFLOP/s)" if flops else ""
   print(f"    {tag:<34s} {ms:8.3f} ms{extra}")

Мы начинаем с установки nvidia-cudnn-frontend и решаем проблему, которая чаще всего возникает при первых запусках: делаем libcudnn.so видимой для динамического загрузчика frontend. Сначала принудительно загружаем встроенную в PyTorch версию cuDNN, а затем явно предварительно загружаем общие библиотеки, чтобы собственный вызов dlopen во frontend находил библиотеку, уже присутствующую в процессе. Затем выводим вычислительную способность, выбираем bfloat16 или float16 в зависимости от неё, создаём дескриптор cuDNN и определяем вспомогательные функции для описания тензоров, построения графов, выделения рабочей области и бенчмаркинга на основе событий, которые повторно используются в остальной части ноутбука.

Копировать кодСкопированоИспользовать другой браузер
N, C, H, W = 32, 128, 56, 56
K, R, S = 256, 3, 3
PAD, STR, DIL = 1, 1, 1
P = (H + 2 * PAD - DIL * (R - 1) - 1) // STR + 1
Q = (W + 2 * PAD - DIL * (S - 1) - 1) // STR + 1
CONV_FLOPS = 2 * N * K * P * Q * C * R * S
CONV_STATE = {}
@section("2. Fused Conv -> Bias -> ReLU")
def conv_fusion():
   x = torch.randn(N, C, H, W, device=DEV, dtype=DTYPE).to(memory_format=torch.channels_last)
   w = torch.randn(K, C, R, S, device=DEV, dtype=DTYPE).to(memory_format=torch.channels_last)
   b = torch.randn(1, K, 1, 1, device=DEV, dtype=DTYPE)
   y = torch.empty(N, K, P, Q, device=DEV, dtype=DTYPE).to(memory_format=torch.channels_last)
   g = cudnn.pygraph(
       handle=HANDLE,
       name="conv_bias_relu",
       io_data_type=TORCH2CUDNN[DTYPE],
       intermediate_data_type=cudnn.data_type.FLOAT,
       compute_data_type=cudnn.data_type.FLOAT,
   )
   X = tensor_of(g, x, "X")
   Wt = tensor_of(g, w, "W")
   Bt = tensor_of(g, b, "bias")
   conv = g.conv_fprop(
       image=X, weight=Wt,
       padding=[PAD, PAD], stride=[STR, STR], dilation=[DIL, DIL],
       compute_data_type=cudnn.data_type.FLOAT,
   )
   biased = g.bias(input=conv, bias=Bt)
   Y = g.relu(input=biased)
   Y.set_output(True).set_data_type(TORCH2CUDNN[DTYPE])
   Y.set_dim(list(y.size())).set_stride(list(y.stride()))
   t0 = time.perf_counter()
   build(g)
   build_ms = (time.perf_counter() - t0) * 1e3
   ws = workspace_for(g)
   pack = {X: x, Wt: w, Bt: b, Y: y}
   g.execute(pack, ws)
   torch.cuda.synchronize()
   ref = torch.relu(torch.nn.functional.conv2d(x, w, bias=b.flatten(), padding=PAD))
   err = (y.float() - ref.float()).abs().max().item()
   scale = ref.float().abs().max().item()
   print(f"    problem  : N{N} C{C} {H}x{W} -> K{K} {R}x{S}  ({DTYPE})")
   print(f"    build    : {build_ms:.1f} ms   workspace: {ws.numel()/1024:.1f} KiB")
   print(f"    max |err|: {err:.4f}  (ref max {scale:.2f}, rel {err/max(scale,1e-9):.2e})")
   assert err / max(scale, 1e-9) < 5e-2, "numerical mismatch vs PyTorch"
   ms_cudnn = bench(lambda: g.execute(pack, ws))
   ms_torch = bench(lambda: torch.relu(
       torch.nn.functional.conv2d(x, w, bias=b.flatten(), padding=PAD)))
   print()
   report("cuDNN FE (single fused kernel)", ms_cudnn, CONV_FLOPS)
   report("PyTorch (conv+bias, then relu)", ms_torch, CONV_FLOPS)
   print(f"    speedup: {ms_torch/ms_cudnn:.2f}x")
   CONV_STATE.update(graph=g, pack=pack, ws=ws, x=x, w=w, b=b, y=y)
   return f"{ms_cudnn:.3f} ms, {tflops(CONV_FLOPS, ms_cudnn):.1f} TFLOP/s"
conv_fusion()

Мы строим наш первый граф — свёртку, за которой следуют сложение смещения и ReLU, объединённые в одно ядро. Все тензоры хранятся в формате channels_last, поскольку именно он предоставляет cuDNN шаги NHWC, необходимые его тензорным движкам. Размерности и шаги выходного тензора явно фиксируются, поэтому результат записывается обратно в том же формате. Затем мы проверяем результат с помощью torch.nn.functional.conv2d и сравниваем производительность объединённого графа с PyTorch, который выполняет свёртку и активацию как отдельные ядра.

Копировать кодСкопированоИспользовать другой браузер
@section("3. Autotuning: build ALL plans, time each engine config")
def autotune():
   x, w, b, y = CONV_STATE["x"], CONV_STATE["w"], CONV_STATE["b"], CONV_STATE["y"]
   g = cudnn.pygraph(
       handle=HANDLE, name="conv_autotune",
       io_data_type=TORCH2CUDNN[DTYPE],
       intermediate_data_type=cudnn.data_type.FLOAT,
       compute_data_type=cudnn.data_type.FLOAT,
   )
   X = tensor_of(g, x, "X")
   Wt = tensor_of(g, w, "W")
   Bt = tensor_of(g, b, "bias")
   Y = g.relu(input=g.bias(
       input=g.conv_fprop(image=X, weight=Wt, padding=[PAD, PAD],
                          stride=[STR, STR], dilation=[DIL, DIL],
                          compute_data_type=cudnn.data_type.FLOAT),
       bias=Bt))
   Y.set_output(True).set_data_type(TORCH2CUDNN[DTYPE])
   Y.set_dim(list(y.size())).set_stride(list(y.stride()))
   g.validate()
   g.build_operation_graph()
   g.create_execution_plans([cudnn.heur_mode.A, cudnn.heur_mode.B, cudnn.heur_mode.FALLBACK])
   g.check_support()
   g.build_plans(cudnn.build_plan_policy.ALL)
   n_plans = g.get_execution_plan_count()
   print(f"    {n_plans} candidate engine configs survived support checks\n")
   pack = {X: x, Wt: w, Bt: b, Y: y}
   timings = []
   for i in range(n_plans):
       try:
           g.build_plan_at_index(i)
           ws_sz = max(g.get_workspace_size_plan_at_index(i), 1)
           ws = torch.empty(ws_sz, device=DEV, dtype=torch.uint8)
           ms = bench(lambda: g.execute_plan_at_index(pack, ws, i), warmup=3, iters=15)
           timings.append((ms, i, ws_sz))
           print(f"      plan {i:>3d}: {ms:8.3f} ms  "
                 f"{tflops(CONV_FLOPS, ms):7.2f} TFLOP/s  ws={ws_sz/1024:8.1f} KiB")
       except Exception as e:
           print(f"      plan {i:>3d}: unusable ({type(e).__name__})")
   assert timings, "no plan executed"
   timings.sort()
   best_ms, best_i, best_ws = timings[0]
   worst_ms = timings[-1][0]
   print(f"\n    fastest = plan {best_i} @ {best_ms:.3f} ms")
   print(f"    slowest = {worst_ms:.3f} ms  -> {worst_ms/best_ms:.1f}x spread across engines")
   print("    Takeaway: heuristics are good, but for a hot shape you ship the")
   print("    autotuned index (or the serialized plan from section 6).")
   return f"best plan {best_i} @ {best_ms:.3f} ms ({worst_ms/best_ms:.1f}x spread)"
autotune()

Мы заново строим ту же свёртку, но перестаём полагаться на эвристику: запрашиваем планы в режимах A, B и FALLBACK и компилируем их все с помощью build_plan_policy.ALL. Затем перебираем список планов, собираем каждую конфигурацию, выделяем для неё рабочую область и измеряем время через execute_plan_at_index, выводя пропускную способность и размер рабочей области для каждого кандидата. Разница между самым быстрым и самым медленным движком — ключевой результат эксперимента: она показывает, насколько выгоднее поставлять индекс после автотюнинга, а не принимать выбор по умолчанию.

Копировать кодСкопированоИспользовать другой браузер
@section("4. Matmul -> scale -> bias -> activation -> AMAX")
def matmul_epilogue():
   Bsz, M, Kd, Nd = 16, 512, 1024, 512
   MM_FLOPS = 2 * Bsz * M * Nd * Kd
   a = torch.randn(Bsz, M, Kd, device=DEV, dtype=DTYPE)
   bm = torch.randn(Bsz, Kd, Nd, device=DEV, dtype=DTYPE)
   bias = torch.randn(1, 1, Nd, device=DEV, dtype=DTYPE)
   out = torch.empty(Bsz, M, Nd, device=DEV, dtype=DTYPE)
   amax = torch.empty(1, 1, 1, device=DEV, dtype=torch.float32)
   alpha_val = 0.125
   alpha = torch.full((1, 1, 1), alpha_val, dtype=torch.float32)
   g = cudnn.pygraph(
       handle=HANDLE, name="matmul_epilogue",
       io_data_type=TORCH2CUDNN[DTYPE],
       intermediate_data_type=cudnn.data_type.FLOAT,
       compute_data_type=cudnn.data_type.FLOAT,
   )
   A = tensor_of(g, a, "A")
   Bt = tensor_of(g, bm, "B")
   BIAS = tensor_of(g, bias, "bias")
   ALPHA = scalar_of(g, "alpha")
   acc = g.matmul(A=A, B=Bt, compute_data_type=cudnn.data_type.FLOAT)
   scaled = g.mul(a=acc, b=ALPHA)
   biased = g.bias(input=scaled, bias=BIAS)
   act_name = "relu"
   if hasattr(g, "gelu"):
       try:
           act = g.gelu(input=biased)
           act_name = "gelu"
       except Exception:
           act = g.relu(input=biased)
   else:
       act = g.relu(input=biased)
   print(f"    activation used: {act_name}")
   OUT = act
   OUT.set_output(True).set_data_type(TORCH2CUDNN[DTYPE])
   have_amax = True
   try:
       AMAX = g.reduction(input=act, mode=cudnn.reduction_mode.AMAX,
                          compute_data_type=cudnn.data_type.FLOAT)
       AMAX.set_output(True).set_data_type(cudnn.data_type.FLOAT)
       AMAX.set_dim([1, 1, 1]).set_stride([1, 1, 1])
   except Exception as e:
       have_amax = False
       print(f"    (AMAX reduction unavailable here: {e})")
   build(g)
   ws = workspace_for(g)
   pack = {A: a, Bt: bm, BIAS: bias, ALPHA: alpha, OUT: out}
   if have_amax:
       pack[AMAX] = amax
   g.execute(pack, ws)
   torch.cuda.synchronize()
   ref = torch.matmul(a.float(), bm.float()) * alpha_val + bias.float()
   ref = torch.nn.functional.gelu(ref) if act_name == "gelu" else torch.relu(ref)
   rel = ((out.float() - ref).abs().max() / ref.abs().max()).item()
   print(f"    shape    : ({Bsz},{M},{Kd}) x ({Bsz},{Kd},{Nd})")
   print(f"    rel err  : {rel:.2e}")
   if have_amax:
       print(f"    fused AMAX {amax.item():.4f} vs torch {ref.abs().max().item():.4f}")
   ms = bench(lambda: g.execute(pack, ws))
   def torch_ref():
       r = torch.baddbmm(bias.expand(Bsz, M, Nd), a, bm, beta=1.0, alpha=alpha_val)
       r = torch.nn.functional.gelu(r) if act_name == "gelu" else torch.relu(r)
       return r.abs().amax()
   ms_t = bench(torch_ref)
   print()
   report("cuDNN FE (one fused kernel)", ms, MM_FLOPS)
   report("PyTorch (bmm + act + amax)", ms_t, MM_FLOPS)
   print(f"    speedup: {ms_t/ms:.2f}x  -- the win is the epilogue traffic, not the GEMM")
   return f"{ms:.3f} ms, {tflops(MM_FLOPS, ms):.1f} TFLOP/s, {ms_t/ms:.2f}x vs torch"
matmul_epilogue()

Далее мы переходим к пакетному матричному умножению и добавляем к нему полноценный эпилог: масштабирование alpha, передаваемое как скаляр хоста по значению, сложение смещения, активацию и редукцию AMAX по результату. AMAX в том же ядре — это схема, на которой основано обучение в формате FP8: она собирает коэффициент масштабирования для следующего шага квантования без второго прохода по выходному тензору. Мы сравниваем этот вариант с цепочкой PyTorch из baddbmm, активации и amax, что наглядно показывает: ускорение достигается за счёт устранения обращений к памяти для эпилога, а не за счёт более быстрого GEMM.

Копировать кодСкопированоИспользовать другой браузер
@section("5. SDPA (Flash Attention) with causal masking")
def sdpa_demo():
   if not HAS_SDPA:
       raise RuntimeError(f"fused SDPA needs SM80+ (Ampere), this GPU is sm_{SM}")
   b, h, s, d = 4, 16, 1024, 64
   scale = 1.0 / math.sqrt(d)
   SDPA_FLOPS = 4 * b * h * s * s * d * 0.5
   q = torch.randn(b, h, s, d, device=DEV, dtype=DTYPE)
   k = torch.randn(b, h, s, d, device=DEV, dtype=DTYPE)
   v = torch.randn(b, h, s, d, device=DEV, dtype=DTYPE)
   o = torch.empty(b, h, s, d, device=DEV, dtype=DTYPE)
   g = cudnn.pygraph(
       handle=HANDLE, name="sdpa",
       io_data_type=TORCH2CUDNN[DTYPE],
       intermediate_data_type=cudnn.data_type.FLOAT,
       compute_data_type=cudnn.data_type.FLOAT,
   )
   Q, Kt, V = tensor_of(g, q, "Q"), tensor_of(g, k, "K"), tensor_of(g, v, "V")
   causal = True
   try:
       O, _stats = g.sdpa(name="sdpa", q=Q, k=Kt, v=V,
                          is_inference=True, attn_scale=scale, use_causal_mask=True)
   except TypeError:
       try:
           O, _stats = g.sdpa(name="sdpa", q=Q, k=Kt, v=V,
                              is_inference=True, attn_scale=scale,
                              diagonal_alignment=cudnn.diagonal_alignment.TOP_LEFT,
                              right_bound=0)
       except Exception:
           causal = False
           O, _stats = g.sdpa(name="sdpa", q=Q, k=Kt, v=V,
                              is_inference=True, attn_scale=scale)
   print(f"    causal masking: {causal}")
   O.set_output(True).set_data_type(TORCH2CUDNN[DTYPE])
   O.set_dim(list(o.size())).set_stride(list(o.stride()))
   build(g)
   ws = workspace_for(g)
   pack = {Q: q, Kt: k, V: v, O: o}
   g.execute(pack, ws)
   torch.cuda.synchronize()
   ref = torch.nn.functional.scaled_dot_product_attention(q, k, v, is_causal=causal, scale=scale)
   rel = ((o.float() - ref.float()).abs().max() / ref.float().abs().max()).item()
   print(f"    shape   : b{b} h{h} s{s} d{d}   workspace {ws.numel()/1024:.1f} KiB")
   print(f"    rel err : {rel:.2e}")
   ms = bench(lambda: g.execute(pack, ws))
   ms_t = bench(lambda: torch.nn.functional.scaled_dot_product_attention(
       q, k, v, is_causal=causal, scale=scale))
   print()
   report("cuDNN FE SDPA", ms, SDPA_FLOPS)
   report("torch SDPA (backend's choice)", ms_t, SDPA_FLOPS)
   print("    Note: torch may already be dispatching to cuDNN or FlashAttention,")
   print("    so parity here is the expected, healthy outcome.")
   return f"{ms:.3f} ms, {tflops(SDPA_FLOPS, ms):.1f} TFLOP/s"
sdpa_demo()
@section("6. Serialize a built graph, reload it, execute by UID")
def serialization():
   Bsz, M, Kd, Nd = 8, 256, 512, 256
   a = torch.randn(Bsz, M, Kd, device=DEV, dtype=DTYPE)
   bm = torch.randn(Bsz, Kd, Nd, device=DEV, dtype=DTYPE)
   out = torch.empty(Bsz, M, Nd, device=DEV, dtype=DTYPE)
   UID_A, UID_B, UID_C = 1, 2, 3
   g = cudnn.pygraph(
       handle=HANDLE, name="serializable_mm",
       io_data_type=TORCH2CUDNN[DTYPE],
       intermediate_data_type=cudnn.data_type.FLOAT,
       compute_data_type=cudnn.data_type.FLOAT,
   )
   A = tensor_of(g, a, "A").set_uid(UID_A)
   Bt = tensor_of(g, bm, "B").set_uid(UID_B)
   C = g.matmul(A=A, B=Bt, compute_data_type=cudnn.data_type.FLOAT)
   C.set_output(True).set_data_type(TORCH2CUDNN[DTYPE]).set_uid(UID_C)
   t0 = time.perf_counter()
   build(g)
   cold_ms = (time.perf_counter() - t0) * 1e3
   blob = g.serialize()
   print(f"    cold build      : {cold_ms:.1f} ms")
   print(f"    serialized plan : {len(blob)} bytes (cache this to disk / ship it)")
   t0 = time.perf_counter()
   g2 = cudnn.pygraph()
   try:
       g2.deserialize(HANDLE, blob)
   except TypeError:
       g2.deserialize(blob)
   warm_ms = (time.perf_counter() - t0) * 1e3
   print(f"    deserialize     : {warm_ms:.1f} ms  -> {cold_ms/max(warm_ms,1e-6):.1f}x faster startup")
   ws = torch.empty(max(g2.get_workspace_size(), 1), device=DEV, dtype=torch.uint8)
   g2.execute({UID_A: a, UID_B: bm, UID_C: out}, ws, handle=HANDLE)
   torch.cuda.synchronize()
   ref = torch.bmm(a.float(), bm.float())
   rel = ((out.float() - ref).abs().max() / ref.abs().max()).item()
   print(f"    rel err after reload: {rel:.2e}")
   return f"{len(blob)} B blob, reload {cold_ms/max(warm_ms,1e-6):.1f}x faster than rebuild"
serialization()

Мы строим граф fused scaled dot-product attention с каузальной маской и проверяем его по torch.nn.functional.scaled_dot_product_attention, защищая весь раздел проверкой SM80, поскольку объединённым ядрам нужен Ampere или более новый GPU. Аргумент каузальности задаётся с резервными вариантами, так как в разных версиях frontend 1.x интерфейс менялся от use_causal_mask к аргументам diagonal_alignment и bound. Затем мы сериализуем построенный граф матричного умножения в байты, загружаем его в новый объект графа и выполняем через целочисленные UID, что позволяет полностью избежать затрат на компиляцию при запуске процесса.

Копировать кодСкопированоИспользовать другой браузер
@section("7. Dynamic shapes with a shared kernel cache")
def dynamic_shapes():
   kc = cudnn.create_kernel_cache()
   def make(n):
       x = torch.randn(n, 64, 32, 32, device=DEV, dtype=DTYPE).to(memory_format=torch.channels_last)
       w = torch.randn(64, 64, 3, 3, device=DEV, dtype=DTYPE).to(memory_format=torch.channels_last)
       y = torch.empty(n, 64, 32, 32, device=DEV, dtype=DTYPE).to(memory_format=torch.channels_last)
       g = cudnn.pygraph(
           handle=HANDLE, name=f"dyn_{n}",
           io_data_type=TORCH2CUDNN[DTYPE],
           intermediate_data_type=cudnn.data_type.FLOAT,
           compute_data_type=cudnn.data_type.FLOAT,
           kernel_cache=kc,
           is_dynamic_shape_enabled=True,
       )
       X, Wt = tensor_of(g, x, "X"), tensor_of(g, w, "W")
       Y = g.conv_fprop(image=X, weight=Wt, padding=[1, 1], stride=[1, 1],
                        dilation=[1, 1], compute_data_type=cudnn.data_type.FLOAT)
       Y.set_output(True).set_data_type(TORCH2CUDNN[DTYPE])
       Y.set_dim(list(y.size())).set_stride(list(y.stride()))
       t0 = time.perf_counter()
       build(g)
       ms = (time.perf_counter() - t0) * 1e3
       ws = workspace_for(g)
       g.execute({X: x, Wt: w, Y: y}, ws)
       torch.cuda.synchronize()
       return ms
   times = [(n, make(n)) for n in (8, 16, 24, 32)]
   for n, ms in times:
       print(f"      batch {n:>3d}: build {ms:7.1f} ms")
   first, rest = times[0][1], [m for _, m in times[1:]]
   print(f"\n    first shape {first:.1f} ms, later shapes avg {sum(rest)/len(rest):.1f} ms")
   print("    The cache lets shape-variant graphs reuse an already-JIT'd kernel,")
   print("    which is what keeps variable batch/seqlen serving out of rebuild hell.")
   return f"first {first:.0f} ms vs subsequent {sum(rest)/len(rest):.0f} ms"
dynamic_shapes()
@section("8. CUDA Graph capture around a cuDNN execution plan")
def cuda_graph_capture():
   if not CONV_STATE:
       raise RuntimeError("section 2 did not run, nothing to capture")
   g, pack, ws = CONV_STATE["graph"], CONV_STATE["pack"], CONV_STATE["ws"]
   eager_ms = bench(lambda: g.execute(pack, ws))
   side = torch.cuda.Stream()
   side.wait_stream(torch.cuda.current_stream())
   with torch.cuda.stream(side):
       cudnn.set_stream(handle=HANDLE, stream=side.cuda_stream)
       for _ in range(3):
           g.execute(pack, ws, handle=HANDLE)
   torch.cuda.current_stream().wait_stream(side)
   torch.cuda.synchronize()
   cg = torch.cuda.CUDAGraph()
   with torch.cuda.graph(cg):
       cudnn.set_stream(handle=HANDLE, stream=torch.cuda.current_stream().cuda_stream)
       g.execute(pack, ws, handle=HANDLE)
   cudnn.set_stream(handle=HANDLE, stream=torch.cuda.current_stream().cuda_stream)
   replay_ms = bench(lambda: cg.replay())
   report("plain execute()", eager_ms)
   report("cuda graph replay()", replay_ms)
   print(f"    launch overhead removed: {(eager_ms-replay_ms)*1e3:.1f} us/iter")
   print("    Pointers are frozen at capture time -- reuse the same buffers and")
   print("    copy new data into them, or re-capture.")
   return f"{eager_ms:.3f} -> {replay_ms:.3f} ms via replay"
cuda_graph_capture()
banner("SUMMARY")
for name, res in RESULTS.items():
   print(f"  {name:<58s} {res}")
print("""
Where to go next
 - samples/python in the repo: FP8/MXFP8 attention, paged KV cache, MoE grouped GEMM
 - python/cudnn/: the open-sourced CuTe DSL kernels (SDPA, grouped GEMM + SwiGLU,
   block-sparse and native sparse attention) you can read and modify
 - debugging: CUDNN_FRONTEND_LOG_INFO=1 and CUDNN_FRONTEND_LOG_FILE=stdout
   (use level 10 during CUDA graph capture -- level 1 dumps tensors and is not
   capture-safe)
""")

В завершение мы рассматриваем два производственных аспекта. Сначала используем общий кэш ядер для четырёх графов, отличающихся только размером пакета, и измеряем время каждой сборки, чтобы увидеть, как последующие формы повторно используют уже скомпилированное ядро вместо повторной оплаты JIT-компиляции. Затем захватываем план свёртки в CUDA-граф, устанавливая поток дескриптора cuDNN в поток захвата. В результате операции попадают в граф, и мы измеряем, насколько уменьшаются накладные расходы на запуск при каждом повторном воспроизведении.

В итоге созданный нами код невелик, но охватывает широкий спектр задач: ядра свёртки, матричного умножения и attention, каждое из которых выражено как граф, а не как вызов библиотеки. Работа на этом уровне изменила то, чем мы могли управлять. Мы сами выбирали, какие операции объединять в одно ядро, поэтому сложения смещения, активации и редукции AMAX, включённые в эпилоги, не записывали промежуточные результаты в память. Мы самостоятельно выбирали движок вместо того, чтобы принимать эвристическое решение, а измерение времени каждой конфигурации кандидата показывало ценность этого выбора. Мы также выбирали момент оплаты компиляции, вынося её из горячего пути с помощью сериализованных планов, общего для разных форм кэша ядер и захвата CUDA-графа. Проверки по PyTorch были не менее важны, чем измерения времени: там, где мы лишь достигали сопоставимых результатов, PyTorch обычно уже вызывал cuDNN внутри. Это показало, где данный API действительно полезен: для fusion, не имеющих эквивалента на уровне фреймворка, для достаточно часто используемых форм, оправдывающих автотюнинг, и для небольших ядер, где доминируют затраты на запуск и инициализацию.


Ознакомьтесь с ПОЛНЫМ КОДОМ здесь. Все заслуги принадлежат исследователю этого проекта. Также подписывайтесь на нас в Twitter и не забудьте присоединиться к нашему ML SubReddit с более чем 150 тыс. участников и подписаться на нашу рассылку. Стоп! Вы есть в Telegram? теперь вы также можете присоединиться к нам в Telegram.

Хотите сотрудничать с нами в продвижении вашего репозитория GitHub, страницы Hugging Face, релиза продукта, вебинара и т. д.? Свяжитесь с нами

Переведено автоматически с английского. Оригинал статьи — по ссылке ниже.

Впервые опубликовано изданием MarkTechPost

Читать оригинал на MarkTechPost ↗

Текст и изображения принадлежат MarkTechPost и приводятся здесь с указанием авторства и ссылкой на оригинальную публикацию.

← К новостям

Ещё новости

Все последние новости