import atexit, functools, math, pathlib
from tinygrad import Tensor, Device, dtypes
from tinygrad.dtype import AddrSpace
from tinygrad.uop.ops import UOp, Ops, KernelInfo, AxisType
from tinygrad.renderer import Estimates
from tinygrad.helpers import getenv, all_same, DEBUG, ceildiv
from tinygrad.runtime.support.compiler_amd import HIPCCCompiler
from extra.llama_kernels.quantize_mxfp4 import quantize_mxfp4

FP8_DTYPE = dtypes.fp8e4m3

TILE_M, TILE_N, TILE_K = 256, 256, 64

# ** MXFP8 GEMM custom kernel

@functools.cache
def custom_hk_mxfp8_gemm(C:UOp, A:UOp, B:UOp, scale_A:UOp, scale_B:UOp, *extra:UOp, dname:str) -> UOp:
  # mxfp8 block-scaled gemm: A(M,K) @ B(N,K).T, e8m0 1x32 microscales packed (k_iters,dim) uint32
  M, K = A.shape[0]*A.shape[1], A.shape[2]
  N, K2 = B.shape[(1 if B.ndim == 3 else 0):]
  assert K == K2, f"{A.shape} {B.shape}"
  block_size = 256
  threads = UOp.special(64 * 8, "lidx0")
  workgroups = UOp.special((M // block_size) * (N // block_size), "gidx0")
  e_a = extra[0].base if len(extra) >= 1 else scale_A.base
  e_b = extra[1].base if len(extra) >= 2 else scale_B.base
  sink_inputs = (C.base, A.base, B.base, scale_A.base, scale_B.base, e_a, e_b, threads, workgroups)
  sink = UOp.sink(*sink_inputs,
                  arg=KernelInfo(f"hk_mxfp8_gemm_{M}_{N}_{K}", estimates=Estimates(ops=2*M*N*K, mem=(M*K+N*K)*A.dtype.itemsize+M*N*C.dtype.itemsize)))
  kittens_path = pathlib.Path(__file__).parent.parent/"thunder"/"amd"
  src = (kittens_path/"gemm_mxfp8.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}"]).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)))

# ** MXFP4 GEMM custom kernel

MXFP4_TILES = ((256, 256), (192, 256), (128, 512))
MXFP4_TILE_MAP = {(6144, 4096, 16384):(192, 256), (16384, 4096, 6144):(128, 512), (16384, 6144, 4096):(128, 512)}

def select_mxfp4_tile(a_q:Tensor, b_q:Tensor) -> tuple[int, int]:
  a_shape, b_shape = a_q.uop.shard_shape, b_q.uop.shard_shape
  M, N, K = math.prod(a_shape[:-1]), math.prod(b_shape[:-1]), a_shape[-1]*2
  if (tile:=MXFP4_TILE_MAP.get((M, N, K))) is not None: return tile
  return next((tile_m, tile_n) for tile_m, tile_n in MXFP4_TILES if M % tile_m == N % tile_n == 0)

@functools.cache
def custom_mxfp4_gemm(C:UOp, A:UOp, B:UOp, scale_a:UOp, scale_b:UOp, *extra:UOp, tile_m:int, tile_n:int) -> UOp:
  from extra.gemm.gemm_mxfp4 import build_kernel
  M, half_k = math.prod(A.shape[:-1]), A.shape[-1]
  N, half_k_b = math.prod(B.shape[:-1]), B.shape[-1]
  K = half_k * 2
  assert half_k == half_k_b and math.prod(C.shape[:-1]) == M and C.shape[-1] == N
  threads = UOp.special(256, "lidx0")
  groups_x, groups_y = UOp.special(ceildiv(N, tile_n), "gidx0"), UOp.special(ceildiv(M, tile_m), "gidx1")
  lds = UOp.placeholder((163840,), dtypes.uint8, 0, AddrSpace.LOCAL)
  sink = UOp.sink(C.base, A.base, B.base, scale_a.base, scale_b.base, *(x.base for x in extra), lds, threads, groups_x, groups_y,
                  arg=KernelInfo(f"mxfp4_gemm_{M}_{N}_{K}_{tile_m}x{tile_n}",
                                 estimates=Estimates(ops=2*M*N*K, mem=(M*half_k+N*half_k)*A.dtype.itemsize+M*N*C.dtype.itemsize)))
  insts = build_kernel(M, N, K, tile_m, tile_n)
  return UOp(Ops.PROGRAM, src=(sink, UOp(Ops.LINEAR, src=tuple(UOp(Ops.INS, arg=(x, dtypes.void)) for x in insts))))

def _mxfp4_gemm_quantized(a_q:Tensor, b_q:Tensor, scale_a:Tensor, scale_b:Tensor) -> Tensor:
  M, half_k = a_q.shape
  N, half_k_b = b_q.shape
  assert half_k == half_k_b, f"MXFP4 K mismatch: A {a_q.shape}, B {b_q.shape}"
  is_multi = isinstance(a_q.device, tuple)
  reduce_out = is_multi and (a_q.uop.axis == 1 or b_q.uop.axis == 1)
  if not is_multi: out = Tensor.invalids(1, M, N, dtype=dtypes.bfloat16, device=a_q.device)
  elif reduce_out: out = Tensor(Tensor.invalids(1, M, N, dtype=dtypes.bfloat16, device=a_q.device).uop.unshard(0), device=a_q.device)
  elif a_q.uop.axis == 0:
    out = Tensor(Tensor.invalids(1, M//len(a_q.device), N, dtype=dtypes.bfloat16, device=a_q.device).uop.unshard(1), device=a_q.device)
  elif b_q.uop.axis == 0:
    out = Tensor(Tensor.invalids(1, M, N//len(a_q.device), dtype=dtypes.bfloat16, device=a_q.device).uop.unshard(2), device=a_q.device)
  else: out = Tensor.invalids(1, M, N, dtype=dtypes.bfloat16, device=a_q.device)
  tile_m, tile_n = select_mxfp4_tile(a_q, b_q)
  out = Tensor.custom_kernel(out, a_q, b_q, scale_a, scale_b,
                             fxn=functools.partial(custom_mxfp4_gemm, tile_m=tile_m, tile_n=tile_n))[0]
  if reduce_out: out = out.sum(0)
  return out.squeeze(0)

def quantize_mxfp8(x:Tensor) -> tuple[Tensor, Tensor, Tensor]:
  # 1x32 block scaling along the last axis
  *batch, K = x.shape
  scale_K = K // 32
  amax = x.detach().float().reshape(*batch, scale_K, 32).abs().max(axis=-1)
  e8 = (amax.maximum(1e-38).log2().floor() + 127).clamp(0, 254).cast(dtypes.uint8)
  qscale = (127.0 - e8.cast(dtypes.float32)).exp2().reshape(*batch, scale_K, 1).expand(*batch, scale_K, 32).reshape(*batch, K)
  x_scaled = x.float() * qscale
  x_clamped = x_scaled + (x_scaled.detach().clamp(-448.0, 448.0) - x_scaled.detach())  # STE
  packed = mx_pack(e8) if len(batch) == 1 and scale_K % 4 == 0 else None
  return x_clamped.cast(FP8_DTYPE), e8, packed

def mx_pack(e8:Tensor) -> Tensor:
  rows, scale_K = e8.shape
  return e8.reshape(rows, scale_K // 4, 4).bitcast(dtypes.uint32).reshape(rows, scale_K // 4).permute(1, 0).contiguous()

def _mx_block_scale(e8:Tensor) -> Tensor:
  # dequant scale 2^(e8-127) broadcast back to element shape
  rows, scale_K = e8.shape
  return (e8.cast(dtypes.float32) - 127.0).exp2().reshape(rows, scale_K, 1).expand(rows, scale_K, 32).reshape(rows, scale_K*32)

def _mx_block_scale_3d(e8:Tensor) -> Tensor:
  # batched (E, rows, scale_K) dequant scale 2^(e8-127) broadcast to (E, rows, scale_K*32)
  E, rows, scale_K = e8.shape
  return (e8.cast(dtypes.float32) - 127.0).exp2().reshape(E, rows, scale_K, 1).expand(E, rows, scale_K, 32).reshape(E, rows, scale_K*32)

counters = {"used":0, "todos":[]}
def todo(msg:str) -> bool: counters["todos"].append(msg); return False
def _asm_gemm_report():
  print(f'asm_gemm: {counters["used"]} used, {len(counters["todos"])} not used')
  if DEBUG >= 2 and counters["todos"]:
    from collections import Counter
    for msg, cnt in Counter(counters["todos"]).most_common(): print(f'  {cnt:3d}x {msg}')
atexit.register(_asm_gemm_report)

def can_use_asm_gemm(a:Tensor, b:Tensor) -> bool:
  if a.dtype != b.dtype: return todo(f"dtypes must match {a.dtype} != {b.dtype}")
  if a.dtype not in {dtypes.bfloat16, dtypes.float16, FP8_DTYPE}: return todo(f"only bfloat16/float16/fp8, got {a.dtype}")
  batch, M, K = (1, *a.shape) if a.ndim == 2 else a.shape
  N = b.shape[1]
  if isinstance(a.device, tuple):
    if a.ndim == 2 and a.uop.axis == 0 and b.uop.axis is None: M //= len(a.device)
    elif a.ndim == 2 and a.uop.axis == 1 and b.uop.axis == 0: K //= len(a.device)
    elif a.ndim == 2 and a.uop.axis is None and b.uop.axis == 1: N //= len(a.device)
    elif a.ndim == 3 and a.uop.axis == 0 and b.uop.axis is None: batch //= len(a.device)
    elif a.ndim == 3 and a.uop.axis == 1 and b.uop.axis is None: M //= len(a.device)
    elif a.ndim == 3 and a.uop.axis is None and b.uop.axis == 1: N //= len(a.device)
    elif a.ndim == 3 and a.uop.axis == 2 and b.uop.axis == 0: K //= len(a.device)
    else: return todo(f"sharding mismatch a.ndim={a.ndim} a.uop.axis={a.uop.axis} b.uop.axis={b.uop.axis}")
    dname = a.device[0]
  else: dname = a.device
  arch = Device[dname].renderer.target.arch
  if batch not in {1, 2}: return todo(f"GEMM batch size {batch}")
  if (M % TILE_M != 0 or N % TILE_N != 0 or K % TILE_K != 0) and arch == "gfx950":
    return todo(f"GEMM shape ({M},{N},{K}) not a multiple of ({TILE_M},{TILE_N},{TILE_K})")
  return True

# ** UOp gemm to test Tensor.custom_kernel multi and backward correctness on non cdna4
# note: this can be removed after we have GEMM on mixins

def custom_uop_gemm(C:UOp, A:UOp, B:UOp) -> UOp:
  M, K = A.shape[0]*A.shape[1], A.shape[2]
  K2, N = B.shape[(1 if B.ndim == 3 else 0):]
  assert K == K2
  m = UOp.range(M, 1)
  n = UOp.range(N, 2)
  k = UOp.range(K, 0, AxisType.LOOP)
  mul = (A.flatten().index((m*UOp.const(K)+k))*
         B.flatten().index((k*UOp.const(N)+n))).cast(dtypes.float32)
  red = mul.reduce(k, arg=Ops.ADD).cast(C.dtype)
  store = C.flatten().index((m*UOp.const(N)+n)).store(red).end(m, n)
  return store.sink(arg=KernelInfo(name=f'uop_gemm_{M}_{N}_{K}'))

# ** bf16 A @ B.T kernel in C

@functools.cache
def custom_hk_bf16_gemm(C:UOp, A:UOp, B:UOp, *args:UOp, dname:str) -> UOp:
  M, K = A.shape[0]*A.shape[1], A.shape[2]
  N, K2 = B.shape[(1 if B.ndim == 3 else 0):]
  assert K == K2, f"{A.shape} {B.shape}"
  block_m, block_n, block_k, num_warps = 256, 256, 64, 8
  assert M % block_m == 0 and N % block_n == 0 and K % block_k == 0, f"invalid bf16 tile {(block_m, block_n, block_k)} for {(M, N, K)}"
  threads = UOp.special(64 * num_warps, "lidx0")
  workgroups = UOp.special((M // block_m) * (N // block_n), "gidx0")
  b_extra = args[0].base if len(args) >= 1 else B.base
  sink = UOp.sink(C.base, A.base, B.base, b_extra, threads, workgroups,
                  arg=KernelInfo(f"hk_bf16_gemm_{M}_{N}_{K}", estimates=Estimates(ops=2*M*N*K, mem=(M*K+N*K+M*N)*A.dtype.itemsize)))
  kittens_path = pathlib.Path(__file__).parent.parent/"thunder"/"amd"
  src = (kittens_path/"gemm_bf16.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}"]).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_bf16_atb_gemm(C:UOp, A:UOp, B:UOp, dname:str) -> UOp:
  K, M = A.shape[0]*A.shape[1], A.shape[2]
  K2, N = B.shape[0]*B.shape[1], B.shape[2]
  assert K == K2, f"{A.shape} {B.shape}"
  block_m, block_n, block_k, num_warps = 256, 256, 64, 8
  assert M % block_m == 0 and N % block_n == 0 and K % block_k == 0, f"invalid bf16 atb tile {(block_m, block_n, block_k)} for {(M, N, K)}"
  threads = UOp.special(64 * num_warps, "lidx0")
  workgroups = UOp.special((M // block_m) * (N // block_n), "gidx0")
  sink = UOp.sink(C.base, A.base, B.base, threads, workgroups,
                  arg=KernelInfo(f"hk_bf16_atb_gemm_{M}_{N}_{K}", estimates=Estimates(ops=2*M*N*K, mem=(M*K+N*K+M*N)*A.dtype.itemsize)))
  kittens_path = pathlib.Path(__file__).parent.parent/"thunder"/"amd"
  src = (kittens_path/"gemm_bf16_atb.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}"]).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 hk_bf16_atb_gemm(a:Tensor, b:Tensor) -> Tensor:
  assert a.dtype == b.dtype == dtypes.bfloat16, f"expected bf16, got {a.dtype} {b.dtype}"
  assert a.ndim == b.ndim == 3 and a.shape[:2] == b.shape[:2], f"{a.shape} {b.shape}"
  batch, rows, M = a.shape
  N = b.shape[2]
  assert M % TILE_M == 0 and N % TILE_N == 0 and (batch * rows) % TILE_K == 0, \
    f"atb shape {a.shape} {b.shape} must produce (M,N,K) multiples of ({TILE_M},{TILE_N},{TILE_K})"
  is_multi = isinstance(a.device, tuple)
  reduce_out = False
  if is_multi:
    ndev = len(a.device)
    if a.uop.axis in (0, 1) or b.uop.axis in (0, 1): inv, out_axis, reduce_out = Tensor.invalids(1, M, N, dtype=a.dtype, device=a.device), 0, True
    elif b.uop.axis == 2: inv, out_axis = Tensor.invalids(1, M, N // ndev, dtype=a.dtype, device=a.device), 2
    elif a.uop.axis == 2: inv, out_axis = Tensor.invalids(1, M // ndev, N, dtype=a.dtype, device=a.device), 1
    else: inv, out_axis, reduce_out = Tensor.invalids(1, M, N, dtype=a.dtype, device=a.device), 0, True
    out = Tensor(inv.uop.unshard(out_axis), device=a.device)
    dname = a.device[0]
  else:
    out = Tensor.invalids(1, M, N, dtype=a.dtype, device=a.device)
    dname = a.device
  dname = dname.split(":")[0]
  out = Tensor.custom_kernel(out, a, b, fxn=functools.partial(custom_hk_bf16_atb_gemm, dname=dname))[0]
  if reduce_out: out = out.sum(0)
  return out.squeeze(0) if out.ndim == 3 else out

# ** backward gemm, might use the asm gemm

def custom_gemm_bw(gradient:UOp, kernel:UOp):
  inputs = kernel.src[1:]
  hk_bf16 = len(inputs) == 4 and inputs[1].dtype == dtypes.bfloat16
  if hk_bf16:
    out, a, b_t, b = inputs
    assert all_same([gradient.device, a.device, b_t.device, b.device, out.device])
  else:
    assert len(inputs) == 3, f"regular gemm must have exactly 3 sources, got: {len(inputs)}"
    out, a, b = inputs
    assert all_same([gradient.device, a.device, b.device, out.device])
  a_t, b_t, g_t = Tensor(a, device=a.device), Tensor(b, device=a.device), Tensor(gradient, device=a.device)
  g_t = g_t[:a.shape[0]]
  if hk_bf16 and g_t.dtype != b_t.dtype: g_t = g_t.cast(b_t.dtype)
  if can_use_asm_gemm(g_t, b_t.T): grad_a = asm_gemm(g_t, b_t.T).uop
  else: grad_a = (g_t @ b_t.T).uop
  if hk_bf16:
    grad_b = hk_bf16_atb_gemm(a_t, g_t).uop
  else:
    a_t_flat, g_t_flat = a_t.permute(2, 0, 1).reshape(a_t.shape[2], -1), g_t.reshape(-1, g_t.shape[-1])
    if can_use_asm_gemm(a_t_flat, g_t_flat): grad_b = asm_gemm(a_t_flat, g_t_flat).uop
    else: grad_b = (a_t_flat @ g_t_flat).uop
  # hk_bf16 uses b.T, writes gradients only for a and b
  return (None, grad_a, None, grad_b) if hk_bf16 else (None, grad_a, grad_b)

# ** mxfp8 gemm backward

def custom_mx_gemm_bw(gradient:UOp, kernel:UOp, has_w_post:bool, w_stored:bool=False):
  inputs = kernel.src[1:]  # (out, a_q, b_q, a_si, b_si, a_e8, b_e8, [w_post])
  aq, bq = Tensor(inputs[1], device=inputs[1].device), Tensor(inputs[2], device=inputs[2].device)
  ae8, be8 = Tensor(inputs[5], device=inputs[5].device), Tensor(inputs[6], device=inputs[6].device)
  wp = Tensor(inputs[7], device=inputs[7].device) if has_w_post else None

  a_phys = (aq.reshape(-1, aq.shape[-1]).cast(dtypes.bfloat16) * _mx_block_scale(ae8)).cast(dtypes.bfloat16)
  b_phys = (bq.cast(dtypes.bfloat16) * _mx_block_scale(be8)).cast(dtypes.bfloat16)

  g = Tensor(gradient, device=aq.device)[:aq.shape[0]].reshape(aq.shape[0]*aq.shape[1], bq.shape[0]).cast(dtypes.bfloat16)
  grad_a = asm_gemm(g, b_phys, mx=True)
  grad_b = asm_gemm(g.T, a_phys, mx=True, a_pretranspose=g)

  grad_a = (grad_a * _mx_block_scale(ae8)).reshape(aq.shape)
  if not w_stored: grad_b = grad_b * _mx_block_scale(be8)
  if wp is not None: grad_b = grad_b / wp.reshape(-1, 1)
  return (None, grad_a.uop, grad_b.uop) + tuple(None for _ in inputs[3:])

# ** mxfp4 gemm backward

def _producer_mxfp4_outputs(gradient:UOp, expected_half_k:int) -> tuple[UOp, UOp, UOp, UOp]|None:
  """Recover quantized sibling outputs from a fused gradient producer without mutable mailboxes or core gradient changes."""
  for call in reversed(gradient.toposort()):
    if call.op is not Ops.CALL or call.src[0].op is not Ops.PROGRAM or not call.src[0].src: continue
    info = call.src[0].src[0].arg
    if (isinstance(info, KernelInfo) and info.name.startswith("swiglu_bwd_mxfp4_")
        and call.src[2].shape[-1] == expected_half_k):
      assert len(call.src) >= 6
      return tuple(call.src[i].after(call) for i in range(2, 6))  # type: ignore[return-value]
  return None

def custom_mxfp4_gemm_bw(gradient:UOp, kernel:UOp):
  inputs = kernel.src[1:]  # out, row operands/scales, BF16 operands, column operands/scales
  assert len(inputs) == 11
  a, w = Tensor(inputs[5], device=inputs[5].device), Tensor(inputs[6], device=inputs[6].device)
  a_col, scale_a_col = Tensor(inputs[7], device=a.device), Tensor(inputs[8], device=a.device)
  w_col, scale_w_col = Tensor(inputs[9], device=a.device), Tensor(inputs[10], device=a.device)
  g = Tensor(gradient, device=a.device)[:a.shape[0]].cast(dtypes.bfloat16)
  if (prequant:=_producer_mxfp4_outputs(gradient, w.shape[0]//2)) is None:
    g_row, scale_g_row, g_col, scale_g_col = quantize_mxfp4(g, flatten_row=True)
  else:
    g_row, scale_g_row, g_col, scale_g_col = (Tensor(x, device=a.device) for x in prequant)
  grad_a = _mxfp4_gemm_quantized(g_row, w_col, scale_g_row, scale_w_col).reshape(*a.shape[:-1], w.shape[-1])
  grad_w = _mxfp4_gemm_quantized(g_col, a_col, scale_g_col, scale_a_col).reshape(w.shape)
  return (None, None, None, None, None, grad_a.uop, grad_w.uop, None, None, None, None)

# ** main gemm function

def asm_gemm(a:Tensor, b:Tensor, w_post_scale:Tensor|None=None, mx:bool=False, mx_scales:tuple|None=None, mx_w_stored:bool=False,
             a_pretranspose:Tensor|None=None, mxfp4:bool=False, mxfp4_w:tuple[Tensor, Tensor, Tensor, Tensor]|None=None,
             mxfp4_x:tuple[Tensor|None, Tensor|None, Tensor|None, Tensor|None]|None=None) -> Tensor:
  assert can_use_asm_gemm(a, b), f"{counters['todos'][-1]}"
  assert mx or a.dtype != FP8_DTYPE, "FP8 GEMM requires MXFP8 block scaling"
  if mxfp4:
    assert not mx and mx_scales is None, "mxfp4 owns quantization; mx/mx_scales are for mxfp8"
    assert a.dtype == dtypes.bfloat16, f"cannot quantize {a.dtype} to mxfp4"
  counters["used"] += 1
  unfold_batch = a.ndim == 3 and isinstance(a.device, tuple) and a.uop.axis == 2 and b.uop.axis == 0
  if unfold_batch:
    orig_batch = a.shape[0]
    a = a.reshape(a.shape[0]*a.shape[1], a.shape[2])
  squeeze = a.ndim == 2
  if squeeze: a = a.unsqueeze(0)
  out_dtype = dtypes.bfloat16 if a.dtype == FP8_DTYPE or mxfp4 else a.dtype

  batch, M, K = a.shape
  N = b.shape[1]
  is_multi = isinstance(a.device, tuple)
  if (k_sharded:=is_multi and a.uop.axis == 2): K //= len(a.device)
  if (m_sharded:=is_multi and a.uop.axis == 1): M //= len(a.device)
  n_sharded = is_multi and b.uop.axis == 1

  if is_multi:
    if n_sharded:
      out = Tensor(Tensor.invalids(batch, M, N//len(a.device), dtype=out_dtype, device=a.device).uop.unshard(2), device=a.device)
    elif m_sharded:
      out = Tensor(Tensor.invalids(batch, M, N, dtype=out_dtype, device=a.device).uop.unshard(1), device=a.device)
    else:
      out = Tensor(Tensor.invalids(batch//len(a.device) if a.uop.axis==0 else batch, M, N, dtype=out_dtype, device=a.device).uop.unshard(0),
                   device=a.device)
  else:
    out = Tensor.invalids(batch, M, N, dtype=out_dtype, device=a.device)

  renderer = Device[dname:=(a.device[0] if is_multi else a.device)].renderer
  dname, arch = dname.split(":")[0], renderer.target.arch
  if arch.startswith("gfx950") and getenv("USE_ASM", 1):
    if mxfp4:
      w = b.T
      if mxfp4_x is None: a_q, scale_a, a_col, scale_a_col = quantize_mxfp4(a, shuffle_col=True)
      else:
        a_q, scale_a, a_col, scale_a_col = mxfp4_x
        assert a_q is not None and scale_a is not None
        if a_col is None or scale_a_col is None:
          assert a_col is scale_a_col is None
          _, _, a_col, scale_a_col = quantize_mxfp4(a, shuffle_col=True)
        assert a_col is not None and scale_a_col is not None
      b_q, scale_b, b_col, scale_b_col = quantize_mxfp4(w, shuffle_row=True, shuffle_col=True) if mxfp4_w is None else mxfp4_w
      tile_m, tile_n = select_mxfp4_tile(a_q, b_q)
      fxn = functools.partial(custom_mxfp4_gemm, tile_m=tile_m, tile_n=tile_n)
      out = Tensor.custom_kernel(out, a_q, b_q, scale_a, scale_b, a, w,
                                 a_col, scale_a_col, b_col, scale_b_col, fxn=fxn, grad_fxn=custom_mxfp4_gemm_bw)[0]
    elif mx:
      # mxfp8 1x32 block scaling
      if mx_scales is not None:
        a_si, a_e8, b_si, b_e8 = mx_scales
        a_q, b_q = a.reshape(-1, a.shape[-1]), b.T
      elif (a_pretranspose is not None and getenv("FUSED_GRAD_QUANTIZE", 0) and a_pretranspose.dtype == dtypes.bfloat16
            and a_pretranspose.shape[0] % 32 == 0 and a_pretranspose.shape[1] % 256 == 0):
        from extra.llama_kernels.transpose_quantize_mxfp8 import transpose_quantize_mxfp8
        a_q, a_e8, a_si = transpose_quantize_mxfp8(a_pretranspose)
        b_q, b_e8, b_si = quantize_mxfp8(b.T)
      else:
        a_q, a_e8, a_si = quantize_mxfp8(a.reshape(-1, a.shape[-1]))
        b_q, b_e8, b_si = quantize_mxfp8(b.T)
      has_w_post = w_post_scale is not None
      fxn = functools.partial(custom_hk_mxfp8_gemm, dname=dname)
      grad_fxn = functools.partial(custom_mx_gemm_bw, has_w_post=has_w_post, w_stored=mx_w_stored)
      extra = [w_post_scale] if w_post_scale is not None else []
      out = Tensor.custom_kernel(out, a_q.reshape(a.shape), b_q, a_si, b_si, a_e8, b_e8, *extra, fxn=fxn, grad_fxn=grad_fxn)[0]
    elif a.dtype == dtypes.bfloat16:
      out = Tensor.custom_kernel(out, a, b.T, b, fxn=functools.partial(custom_hk_bf16_gemm, dname=dname), grad_fxn=custom_gemm_bw)[0]
  else:
    out = Tensor.custom_kernel(out, a, b, fxn=custom_uop_gemm, grad_fxn=custom_gemm_bw)[0]
  if k_sharded: out = out.sum(0)
  out = out.squeeze(0) if squeeze else out
  if unfold_batch: out = out.reshape(orig_batch, -1, out.shape[-1])
  if w_post_scale is not None: out = (out * w_post_scale.reshape(*([1]*(out.ndim-1)), -1)).cast(out.dtype)
  return out
