Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
138 changes: 115 additions & 23 deletions benchmarks/python/gated_delta_bench.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,23 +3,16 @@
import itertools
import os
import time
from datetime import datetime
from typing import Optional, Tuple

import mlx.core as mx
import numpy as np

RED_BOLD = "\033[1;31m"
GREEN = "\033[0;32m"
RESET = "\033[0m"


N_warmup = 8
N_iter_bench = 80
N_iter_func = 5


# similar to ./blas/bench_gemm.py
def bench(f, *args):
for _ in range(N_warmup):
f(*args)
Expand All @@ -30,7 +23,7 @@ def bench(f, *args):
f(*args)
mx.synchronize()
e = time.perf_counter_ns()
return (e - s) * 1e-9 # total seconds for N_iter_bench * N_iter_func calls
return (e - s) * 1e-9


def do_kernel_bench(f, *args):
Expand All @@ -43,20 +36,42 @@ def do_kernel_bench(f, *args):
return ys


def benchmark_shape(B, T, Hk, Hv, Dk, Dv, chunk_sizes):
def do_grad_bench(f, *args):
ys = []
for _ in range(N_iter_func):
ys.extend(f(*args))
mx.eval(ys)
return ys


def make_grad_fn():
def f(q, k, v, g, b, h0):
out, state = mx.fast.gated_delta_update(q, k, v, g, b, h0)
return out.sum() + state.sum()

return mx.grad(f, argnums=(0, 1, 2, 3, 4, 5))


def make_inputs(B, T, Hk, Hv, Dk, Dv):
mx.random.seed(42)
q = mx.random.normal(shape=(B, T, Hk, Dk))
k = mx.random.normal(shape=(B, T, Hk, Dk))
k = k / (mx.linalg.norm(k, axis=-1, keepdims=True) + 1e-6)
v = mx.random.normal(shape=(B, T, Hv, Dv))
g = mx.random.normal(shape=(B, T, Hv)) * 0.1 - 1.0
b = mx.sigmoid(mx.random.normal(shape=(B, T, Hv)))
h0 = mx.zeros((B, Hv, Dv, Dk), dtype=mx.float32)
mx.eval(q, k, v, g, b, h0)
return q, k, v, g, b, h0


def benchmark_shape(B, T, Hk, Hv, Dk, Dv, chunk_sizes):
q, k, v, g, b, h0 = make_inputs(B, T, Hk, Hv, Dk, Dv)

shape_str = f"B={B} T={T} Hk={Hk} Hv={Hv} Dk={Dk} Dv={Dv}"
denom = N_iter_bench * N_iter_func

os.environ["GATED_DELTA_CHUNK"] = "0"
h0 = mx.zeros((B, Hv, Dv, Dk), dtype=mx.float32)
mx.eval(*mx.fast.gated_delta_update(q, k, v, g, b, initial_state=h0))
ms_seq = (
bench(do_kernel_bench, mx.fast.gated_delta_update, q, k, v, g, b, h0)
Expand All @@ -68,7 +83,6 @@ def benchmark_shape(B, T, Hk, Hv, Dk, Dv, chunk_sizes):
for C in (c for c in chunk_sizes if c != 0):
try:
os.environ["GATED_DELTA_CHUNK"] = str(C)
h0 = mx.zeros((B, Hv, Dv, Dk), dtype=mx.float32)
mx.eval(*mx.fast.gated_delta_update(q, k, v, g, b, initial_state=h0))
ms_c = (
bench(do_kernel_bench, mx.fast.gated_delta_update, q, k, v, g, b, h0)
Expand All @@ -86,18 +100,11 @@ def benchmark_shape(B, T, Hk, Hv, Dk, Dv, chunk_sizes):
def run_benchmark(run_full, to_csv=False, csv_path="benchmark_results.csv"):
if run_full:
Bs = [1, 4, 8, 16]
Ts = [8, 64, 256, 512, 1024, 2048, 4096]
Hks = [16]
Hvs = [32]
Dks = [128]
Dvs = [128]
Ts = [8, 64, 256, 512, 1024, 2048]
else:
Bs = [1, 8, 16]
Ts = [8, 512, 1024, 2048]
Hks = [16]
Hvs = [32]
Dks = [128]
Dvs = [128]
Hks, Hvs, Dks, Dvs = [16], [32], [128], [128]

chunk_sizes = [0, 8, 16]
non_zero_Cs = [C for C in chunk_sizes if C != 0]
Expand All @@ -106,13 +113,13 @@ def run_benchmark(run_full, to_csv=False, csv_path="benchmark_results.csv"):
f"C={C} (speedup)" for C in non_zero_Cs
]

col_widths = [6, 6, 6, 6, 6, 6, 15] + [25] * (len(non_zero_Cs))
col_widths = [6, 6, 6, 6, 6, 6, 15] + [25] * len(non_zero_Cs)
fmt = "".join(f"{{:<{w}}}" for w in col_widths)

rows = []

print(fmt.format(*headers))
print("-" * (sum(col_widths)))
print("-" * sum(col_widths))

for B, T, Hk, Hv, Dk, Dv in itertools.product(Bs, Ts, Hks, Hvs, Dks, Dvs):
shapes_s, base_time_s, speedups, base_time = benchmark_shape(
Expand All @@ -135,11 +142,96 @@ def run_benchmark(run_full, to_csv=False, csv_path="benchmark_results.csv"):
print(f"\nResults also written to {csv_path}")


def benchmark_variants_shape(B, T, Hk, Hv, Dk, Dv, do_backward, variants):
q, k, v, g, b, h0 = make_inputs(B, T, Hk, Hv, Dk, Dv)
denom = N_iter_bench * N_iter_func

if do_backward:
fn = make_grad_fn()
runner = do_grad_bench
else:
fn = mx.fast.gated_delta_update
runner = do_kernel_bench

def time_one(variant):
os.environ["GATED_DELTA_VJP_FALLBACK"] = "1" if variant == "fallback" else "0"
C = "16" if variant == "nax" else "0"
if do_backward:
os.environ["GATED_DELTA_CHUNK"] = "16"
os.environ["GATED_DELTA_CHUNK_VJP"] = C
else:
os.environ["GATED_DELTA_CHUNK"] = C
mx.eval(*fn(q, k, v, g, b, h0))
return bench(runner, fn, q, k, v, g, b, h0) / denom * 1e3

times = []
for variant in variants:
try:
times.append(time_one(variant))
except Exception as ex:
print(f" {variant} failed: {ex}")
times.append(float("nan"))
mx.clear_cache()
return times


def run_variants_benchmark(run_full, do_backward=False, do_fallback=False):
if run_full:
Bs = [1, 4, 8, 16]
Ts = [8, 32, 64, 128, 256, 512, 1024, 2048, 4096]
else:
Bs = [1, 8]
Ts = [8, 512, 1024]
if do_fallback:
Ts = [8, 32, 64, 128, 256, 512]
Hks, Hvs, Dks, Dvs = [16], [32], [128], [128]

variants = ["seq", "nax"]
if do_fallback:
variants = ["fallback"] + variants

headers = ["B", "T", "Hk", "Hv", "Dk", "Dv"]
headers += [f"{v} (ms)" for v in variants]
headers += [f"{v} (speedup)" for v in variants[1:]]

col_widths = [6, 6, 6, 6, 6, 6] + [16] * len(variants) + [16] * (len(variants) - 1)
fmt = "".join(f"{{:<{w}}}" for w in col_widths)

mode = "BACKWARD" if do_backward else "FORWARD"
print(f"\n=== {mode}: {' vs '.join(variants)} ===")
print(fmt.format(*headers))
print("-" * sum(col_widths))

for B, T, Hk, Hv, Dk, Dv in itertools.product(Bs, Ts, Hks, Hvs, Dks, Dvs):
try:
times = benchmark_variants_shape(
B, T, Hk, Hv, Dk, Dv, do_backward, variants
)
except Exception as ex:
print(f" B={B} T={T} failed: {ex}")
mx.clear_cache()
continue

base = times[0]
row = [f"{B}", f"{T}", f"{Hk}", f"{Hv}", f"{Dk}", f"{Dv}"]
row += [f"{t:.3f}" for t in times]
row += [f"{base / t:.2f}x" if t > 0 else "nan" for t in times[1:]]
print(fmt.format(*row))
print(RESET, end="")


if __name__ == "__main__":
parser = argparse.ArgumentParser(description="Gated delta benchmark")
parser.add_argument("--full", "-f", action="store_true")
parser.add_argument("--csv", "-c", action="store_true")
parser.add_argument("--csv_out", "-co", default="benchmark_results.csv")
parser.add_argument("--fallback", "-fb", action="store_true")
parser.add_argument("--backward", "-bw", action="store_true")
args = parser.parse_args()

run_benchmark(args.full, to_csv=args.csv, csv_path=args.csv_out)
if args.backward or args.fallback:
run_variants_benchmark(
args.full, do_backward=args.backward, do_fallback=args.fallback
)
else:
run_benchmark(args.full, to_csv=args.csv, csv_path=args.csv_out)
12 changes: 11 additions & 1 deletion mlx/backend/cuda/primitives.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -34,6 +34,15 @@ bool fast::GatedDeltaUpdate::use_fallback(
return true;
}

bool fast::GatedDeltaUpdateVJP::use_fallback(
const int Hk,
const int Dk,
const int Hv,
const int Dv,
Stream s) {
return true;
}

NO_GPU_MULTI(LUF)
NO_GPU_MULTI(QRF)
NO_GPU_MULTI(SVD)
Expand All @@ -43,7 +52,8 @@ NO_GPU_MULTI(Eigh)

namespace fast {
NO_GPU_MULTI(GatedDeltaUpdate)
}
NO_GPU_MULTI(GatedDeltaUpdateVJP)
} // namespace fast

namespace distributed {
NO_GPU_MULTI(Send)
Expand Down
3 changes: 2 additions & 1 deletion mlx/backend/metal/CMakeLists.txt
Original file line number Diff line number Diff line change
Expand Up @@ -103,7 +103,8 @@ if(MLX_METAL_JIT)
kernels/fp4.h)

make_jit_source(steel/attn/kernels/steel_attention_nax)
make_jit_source(gated_delta_update_nax)
make_jit_source(gated_delta_update_nax kernels/gated_delta_nax_ops.h)
make_jit_source(gated_delta_update_nax_vjp kernels/gated_delta_nax_ops.h)

else()
message(
Expand Down
Loading
Loading