mirror of
https://github.com/tinygrad/tinygrad.git
synced 2026-08-29 12:36:07 +00:00
135 lines
7.8 KiB
Python
135 lines
7.8 KiB
Python
import functools, pathlib
|
|
from tinygrad import Tensor, dtypes
|
|
from tinygrad.uop.ops import UOp, Ops, KernelInfo, AxisType
|
|
from tinygrad.helpers import getenv
|
|
from tinygrad.renderer import Estimates
|
|
from tinygrad.runtime.support.compiler_amd import HIPCCCompiler
|
|
from extra.gemm.cdna_asm_gemm import quantize_mxfp8, _mx_block_scale, _mx_block_scale_3d
|
|
|
|
ZERO_OPTIM = getenv("ZERO_OPTIM", 0)
|
|
|
|
def reduce_scatter_devaxis(out:Tensor, shard_axis:int=0) -> Tensor:
|
|
# out: sharded on the device axis, shape (ndev, *rest); return the device-axis sum left sharded on shard_axis.
|
|
u = out.uop
|
|
devs, rest = u.device, u.shape[1:]
|
|
assert rest[shard_axis] % len(devs) == 0, f"reduce_scatter needs even shards: {rest[shard_axis]} % {len(devs)}"
|
|
# reach the raw per-device buffer below the UNSHARD, keeping the AFTERs so reads stay ordered after the kernel writes
|
|
node, barriers = u, []
|
|
while node.op is not Ops.UNSHARD:
|
|
if node.op is Ops.AFTER: barriers += node.src[1:]
|
|
node = node.src[0]
|
|
mbuf = node.src[0].after(*barriers) if barriers else node.src[0]
|
|
sz = rest[shard_axis] // len(devs)
|
|
shards = []
|
|
for i in range(len(devs)):
|
|
bounds = tuple((0,s) if a != shard_axis else (i*sz,(i+1)*sz) for a,s in enumerate(rest))
|
|
contribs = [mbuf.mselect(j).reshape(rest).shrink(bounds).copy_to_device(devs[i]) for j in range(len(devs))]
|
|
shards.append(functools.reduce(lambda a,b: a.alu(Ops.ADD, b), contribs))
|
|
return Tensor(UOp.mstack(*shards).unshard(shard_axis, UOp.range(len(devs), -1, AxisType.DEVICE)), device=devs)
|
|
|
|
@functools.cache
|
|
def custom_hk_grouped_mxfp8_gemm(C:UOp, A:UOp, B:UOp, scale_A:UOp, scale_B:UOp, *extra:UOp, dname:str, n_experts:int) -> UOp:
|
|
M, K = A.shape
|
|
E, N, K2 = B.shape
|
|
assert K == K2, f"{A.shape} {B.shape}"
|
|
assert E == n_experts, f"{E} != {n_experts}"
|
|
threads = UOp.special(64 * 8, "lidx0")
|
|
workgroups = UOp.special((M // 256) * (N // 256), "gidx0")
|
|
sink_inputs = (C.base, A.base, B.base, scale_A.base, scale_B.base, extra[0].base, extra[1].base, extra[2].base, threads, workgroups)
|
|
sink = UOp.sink(*sink_inputs,
|
|
arg=KernelInfo(f"hk_grouped_mxfp8_gemm_{E}_{M}_{N}_{K}",
|
|
estimates=Estimates(ops=2*M*N*K, mem=(M*K+E*N*K)*A.dtype.itemsize+M*N*C.dtype.itemsize)))
|
|
kittens_path = pathlib.Path(__file__).parent.parent/"thunder"/"amd"
|
|
src = (kittens_path/"grouped_mxfp8_gemm.cpp").read_text()
|
|
lib = HIPCCCompiler("gfx950", [f"-I{(kittens_path/'include').as_posix()}", "-std=c++20", "-DKITTENS_CDNA4", "-ffast-math",
|
|
"-DHIP_ENABLE_WARP_SYNC_BUILTINS", f"-DGEMM_M={M}", f"-DGEMM_N={N}", f"-DGEMM_K={K}",
|
|
f"-DGEMM_E={E}"]).compile_cached(src)
|
|
return UOp(Ops.PROGRAM, src=(sink, UOp(Ops.LINEAR, src=(*sink.src, sink)), UOp(Ops.SOURCE, arg=src),
|
|
UOp(Ops.BINARY, arg=lib)))
|
|
|
|
@functools.cache
|
|
def custom_hk_grouped_mxfp8_wgrad(C:UOp, A:UOp, B:UOp, scale_A:UOp, scale_B:UOp, expert_off:UOp, *, dname:str, n_experts:int) -> UOp:
|
|
N, M = A.shape
|
|
K, M2 = B.shape
|
|
assert M == M2, f"{A.shape} {B.shape}"
|
|
E = n_experts
|
|
threads = UOp.special(64 * 8, "lidx0")
|
|
workgroups = UOp.special(E * (N // 256) * (K // 256), "gidx0")
|
|
sink = UOp.sink(C.base, A.base, B.base, scale_A.base, scale_B.base, expert_off.base, threads, workgroups,
|
|
arg=KernelInfo(f"hk_grouped_mxfp8_wgrad_{E}_{M}_{N}_{K}",
|
|
estimates=Estimates(ops=2*M*N*K, mem=(N*M+K*M)*A.dtype.itemsize+E*N*K*C.dtype.itemsize)))
|
|
kittens_path = pathlib.Path(__file__).parent.parent/"thunder"/"amd"
|
|
src = (kittens_path/"grouped_mxfp8_wgrad.cpp").read_text()
|
|
lib = HIPCCCompiler("gfx950", [f"-I{(kittens_path/'include').as_posix()}", "-std=c++20", "-DKITTENS_CDNA4", "-ffast-math",
|
|
"-DHIP_ENABLE_WARP_SYNC_BUILTINS", f"-DWGRAD_M={M}", f"-DWGRAD_N={N}", f"-DWGRAD_K={K}",
|
|
f"-DWGRAD_E={E}"]).compile_cached(src)
|
|
return UOp(Ops.PROGRAM, src=(sink, UOp(Ops.LINEAR, src=(*sink.src, sink)), UOp(Ops.SOURCE, arg=src),
|
|
UOp(Ops.BINARY, arg=lib)))
|
|
|
|
def grouped_mx_wgrad(g:Tensor, xg:Tensor, expert_off:Tensor, n_experts:int) -> Tensor:
|
|
from extra.llama_kernels.transpose_quantize_mxfp8 import transpose_quantize_mxfp8
|
|
M, N = g.shape
|
|
M2, K = xg.shape
|
|
assert M == M2, f"{g.shape} {xg.shape}"
|
|
assert M % 128 == 0 and N % 256 == 0 and K % 256 == 0, f"wgrad needs M%128,N%256,K%256, got {g.shape} {xg.shape}"
|
|
gT, _, g_si = transpose_quantize_mxfp8(g.contiguous())
|
|
xT, _, x_si = transpose_quantize_mxfp8(xg.contiguous())
|
|
dname = (g.device[0] if isinstance(g.device, tuple) else g.device).split(":")[0]
|
|
is_multi = isinstance(g.device, tuple)
|
|
inv = Tensor.invalids(1, n_experts * N, K, dtype=dtypes.bfloat16, device=g.device)
|
|
out = Tensor(inv.uop.unshard(0), device=g.device) if is_multi else inv
|
|
out = Tensor.custom_kernel(out, gT, xT, g_si, x_si, expert_off,
|
|
fxn=functools.partial(custom_hk_grouped_mxfp8_wgrad, dname=dname, n_experts=n_experts))[0]
|
|
if is_multi and ZERO_OPTIM: out = reduce_scatter_devaxis(out, 0)
|
|
else: out = out.sum(0) if is_multi else out.squeeze(0)
|
|
return out.reshape(n_experts, N, K)
|
|
|
|
def mx_pack_3d(e8:Tensor) -> Tensor:
|
|
E, rows, scale_K = e8.shape
|
|
return e8.reshape(E, rows, scale_K // 4, 4).bitcast(dtypes.uint32).reshape(E, rows, scale_K // 4).permute(0, 2, 1).contiguous()
|
|
|
|
@functools.cache
|
|
def custom_grouped_mx_gemm_bw(gradient:UOp, kernel:UOp, w_stored:bool=False) -> tuple:
|
|
inputs = kernel.src[1:]
|
|
aq = Tensor(inputs[1], device=inputs[1].device)
|
|
bq = Tensor(inputs[2], device=inputs[2].device)
|
|
ae8 = Tensor(inputs[5], device=inputs[5].device)
|
|
be8 = Tensor(inputs[6], device=inputs[6].device)
|
|
E, N = bq.shape[0], bq.shape[1]
|
|
M, K = aq.shape
|
|
g = Tensor(gradient, device=aq.device).reshape(M, N).cast(dtypes.bfloat16)
|
|
x_phys = (aq.cast(dtypes.bfloat16) * _mx_block_scale(ae8).cast(dtypes.bfloat16))
|
|
w_phys = (bq.cast(dtypes.bfloat16) * _mx_block_scale_3d(be8).cast(dtypes.bfloat16))
|
|
expert_off = Tensor(inputs[7], device=inputs[7].device)
|
|
grad_x = grouped_mx_gemm(g, w_phys.transpose(1, 2), expert_off)
|
|
grad_w = grouped_mx_wgrad(g, x_phys, expert_off, E)
|
|
grad_xq = grad_x * _mx_block_scale(ae8).cast(dtypes.bfloat16)
|
|
grad_wq = grad_w.contiguous() if w_stored else (grad_w * _mx_block_scale_3d(be8).cast(dtypes.bfloat16)).contiguous()
|
|
return (None, grad_xq.uop, grad_wq.uop) + tuple(None for _ in inputs[3:])
|
|
|
|
_grouped_bw_stored = functools.partial(custom_grouped_mx_gemm_bw, w_stored=True)
|
|
|
|
def grouped_mx_gemm(x:Tensor, w:Tensor|tuple[Tensor, Tensor], expert_off:Tensor) -> Tensor:
|
|
if (pre_quantized := isinstance(w, tuple)):
|
|
w_q, w_e8 = w
|
|
E, N, K2 = w_q.shape
|
|
else:
|
|
E, N, K2 = w.shape
|
|
M, K = x.shape
|
|
assert K == K2, f"shape mismatch {x.shape} {w.shape}"
|
|
assert M % 256 == 0 and N % 256 == 0 and K % 128 == 0, f"grouped mxfp8 needs M%256,N%256,K%128, got {x.shape} {w.shape}"
|
|
dname = (x.device[0] if isinstance(x.device, tuple) else x.device).split(":")[0]
|
|
x_q, x_e8, x_si = quantize_mxfp8(x)
|
|
if not pre_quantized: w_q, w_e8, _ = quantize_mxfp8(w)
|
|
w_si = mx_pack_3d(w_e8)
|
|
xe_in, out_shape = x_e8.reshape(M, K // 32), (M, N)
|
|
if isinstance(x.device, tuple) and (row_axis := x.uop.axis) is not None:
|
|
ndev = len(x.device)
|
|
out = Tensor(Tensor.invalids(*(s // ndev if i == row_axis else s for i, s in enumerate(out_shape)),
|
|
dtype=dtypes.bfloat16, device=x.device).uop.unshard(row_axis), device=x.device)
|
|
else:
|
|
out = Tensor.invalids(*out_shape, dtype=dtypes.bfloat16, device=x.device)
|
|
return Tensor.custom_kernel(out, x_q, w_q, x_si, w_si, xe_in, w_e8, expert_off,
|
|
fxn=functools.partial(custom_hk_grouped_mxfp8_gemm, dname=dname, n_experts=E),
|
|
grad_fxn=(_grouped_bw_stored if pre_quantized else custom_grouped_mx_gemm_bw))[0]
|