import math, struct, sys
from tinygrad.renderer import tc
from tinygrad.renderer import Renderer
from tinygrad.renderer.cstyle import HIPRenderer, create_non_native_float_pats, pm_manual_bf16_cast, fp8_index, amd_fp8s
from tinygrad.codegen.decomp.transcendental import xexp2, xlog2
from tinygrad.uop.ops import UOp, PatternMatcher, UPat, Ops, GroupOp, range_str
from tinygrad.dtype import dtypes, float_to_fp8, DType, truncate, AddrSpace
from tinygrad.helpers import prod, Target, OSX

def is_volatile(u:UOp) -> bool: return (buf:=u.buf_uop).op is Ops.PARAM and buf.arg.volatile

def ldt(dt:DType, count=1, ptr=False):
  if ptr: return ldt(dt, count) + "*"
  if count > 1: return f"<{count} x {ldt(dt, 1, ptr)}>"
  return {dtypes.void: "void", dtypes.bool: "i1", dtypes.int8: "i8", dtypes.int16: "i16", dtypes.int32: "i32", dtypes.int64: "i64",
          dtypes.uint8: "i8", dtypes.uint16: "i16", dtypes.uint32: "i32", dtypes.uint64: "i64", **{d: "i8" for d in dtypes.fp8s},
          dtypes.float16: "half", dtypes.bfloat16: "bfloat", dtypes.float32: "float", dtypes.float64: "double"}[dt]

def lconst(x, dtype:DType):
  if dtype in dtypes.floats:
    if dtype in dtypes.fp8s: return float_to_fp8(x, dtype)
    if math.isinf(x) or math.isnan(x): return "0x%02X%02X%02X%02X%02X%02X%02X%02X" % tuple(struct.pack("d",x)[::-1])
    return truncate[dtype](x)
  return int(x)

def lcast(input_type:DType, output_type:DType):
  if dtypes.is_float(input_type):
    if dtypes.is_float(output_type): return 'fpext' if output_type.itemsize > input_type.itemsize else 'fptrunc'
    if dtypes.is_int(output_type): return 'fptoui' if dtypes.is_unsigned(output_type) else 'fptosi'
  if dtypes.is_unsigned(input_type) or dtypes.is_bool(input_type):
    if dtypes.is_float(output_type): return 'uitofp'
    if dtypes.is_int(output_type): return 'trunc' if output_type.itemsize < input_type.itemsize else 'zext'
  if dtypes.is_int(input_type):
    if dtypes.is_float(output_type): return 'sitofp'
    if dtypes.is_int(output_type): return 'trunc' if output_type.itemsize < input_type.itemsize else 'sext'
  raise NotImplementedError(f"cast from {input_type} -> {output_type} not implemented")

def render_wmma_amd(ctx, wmma: UOp, cdna=False, rdna4=False) -> str:
  dt_map = {dtypes.half: "f16", dtypes.float: "f32", dtypes.ushort: "bf16.1k" if cdna else "bf16", dtypes.bfloat16: "bf16.1k" if cdna else "bf16",
            **{d: ("." if cdna else "") + ("fp8.fp8", "bf8.bf8")[fp8_index(d)] for d in dtypes.fp8s},
            dtypes.int8: "iu8", dtypes.int32: "i32", dtypes.uint32: "i32"}
  # https://github.com/llvm/llvm-project/blob/main/clang/test/CodeGenOpenCL/builtins-amdgcn-mfma.cl
  N,M,K = wmma.arg[0]
  if cdna:
    if K == 32: dt_map.update({dtypes.half: ".f16", dtypes.bfloat16: ".bf16"})
    scaled = K == 128
    args = [f"{ldt(w.dtype, w.max_numel())} {ctx[w]}" for w in wmma.src]
    # scaled mfma call require E8M0 scale args, byte = 0x7F = 127, scale = 2^(127 - 127) = 1.0
    if scaled:
      _fmt = { dtypes.fp8e5m2:1, dtypes.fp8e4m3:0 }
      # (a_fp8_fmt, b_fp8_fmt, opsel, scale_a, opsel, scale_b)
      args.extend([f"i32 {_fmt[wmma.arg[1]]}", f"i32 {_fmt[wmma.arg[1]]}", "i32 0", "i32 127", "i32 0", "i32 127"])
    else: args.extend(["i32 0", "i32 0", "i32 0"]) # (cbsz, blgp, ?)

    scale = "scale." if scaled else ""
    dt_in = dt_map[wmma.arg[1]] if not scaled else ".f8f6f4"
    return f"  {ctx[wmma]} = call {ldt(wmma.dtype, wmma.max_numel())} @llvm.amdgcn.mfma.{scale}{dt_map[wmma.src[-1].dtype]}" + \
           f".{N}x{M}x{K}{dt_in}(" + ", ".join(args) + ")"
  # https://github.com/llvm/llvm-project/blob/main/llvm/test/CodeGen/AMDGPU/GlobalISel/llvm.amdgcn.wmma_32.ll
  # example: %wmma0 = call <8 x float> @llvm.amdgcn.wmma.f32.16x16x16.f16(<16 x half> %v99,<16 x half> %v100,<8 x float> %v101)
  args = [f"{ldt(w.dtype, w.max_numel())} {ctx[w]}" for w in wmma.src]
  if wmma.arg[1] == dtypes.int8: args = ["i1 true", args[0], "i1 true", args[1], args[2]]  # iu8 flags A/B signed
  if wmma.dtype != dtypes.float: args.append("i1 false") # opsel
  suffix = f".v{wmma.max_numel()}{dt_map[wmma.dtype]}.v{wmma.src[0].max_numel()}{dt_map[wmma.src[0].dtype]}" if rdna4 else ""
  return f"  {ctx[wmma]} = call {ldt(wmma.dtype, wmma.max_numel())} @llvm.amdgcn.wmma.{dt_map[wmma.src[-1].dtype]}.16x16x16." + \
    f"{dt_map[wmma.arg[1]]}{suffix}(" + ", ".join(args) + ")"

# llvm ops, lop[<dtype>][<op>]
unsigned_lop = { Ops.ADD: "add", Ops.MUL: "mul", Ops.CDIV: "udiv", Ops.CMOD: "urem",
                 Ops.CMPLT: "icmp ult", Ops.CMPNE: "icmp ne", Ops.CMPEQ: "icmp eq", Ops.OR: "or", Ops.AND: "and", Ops.XOR: "xor",
                 Ops.SHL: "shl", Ops.SHR: "lshr",}
signed_lop = {**unsigned_lop, Ops.ADD: "add nsw", Ops.CMPLT: "icmp slt", Ops.CDIV: "sdiv", Ops.CMOD: "srem", Ops.SHR: "ashr"}
flags = " nsz arcp contract afn"
float_lop = {Ops.ADD: "fadd"+flags, Ops.MUL: "fmul"+flags, Ops.CMPLT: f"fcmp{flags} olt",
    Ops.CMPNE: f"fcmp{flags} une", Ops.CMPEQ: f"fcmp{flags} oeq", Ops.FDIV: "fdiv"+flags}
lop = {**{x:unsigned_lop for x in (dtypes.bool,)+dtypes.uints}, **{x:signed_lop for x in dtypes.sints}, **{x:float_lop for x in dtypes.floats}}

base_rewrite = PatternMatcher([
  # memory load/store
  (UPat((Ops.INDEX, Ops.SHRINK), src=(UPat((Ops.BUFFER, Ops.PARAM, Ops.AFTER)),), allow_any_len=True, name="x"), lambda ctx,x:
   f"  {ctx[x]} = getelementptr inbounds {ldt(x.dtype)}, {ldt(x.dtype, ptr=True)} {ctx[x.src[0]]}, {ldt(x.src[1].dtype)} {ctx[x.src[1]]}"),
  # register index
  (UPat(Ops.INDEX, src=(UPat.var("buf"), UPat.cvar("c").cast()), name="x"), lambda ctx,buf,c,x:
   f"  {ctx[x]} = extractelement {ldt(buf.dtype, buf.max_numel())} {ctx[buf]}, i32 {c.val}" if buf.addrspace == AddrSpace.ALU else None),

  # load/store
  (UPat(Ops.LOAD, src=(UPat.var("idx"), UPat.var("alt"), UPat.var("mask")), name="x"),
   lambda ctx,x,idx,alt,mask:
   f"  br label {ctx[x]}_entry\n{ctx[x][1:]}_entry:\n"
   f"  br i1 {ctx[mask]}, label {ctx[x]}_load, label {ctx[x]}_exit\n{ctx[x][1:]}_load:\n"
   f"  {ctx[x]}_yes = load {'volatile ' if is_volatile(idx) else ''}{ldt(idx.dtype, idx.max_numel())}, "
   f"{ldt(idx.dtype, idx.max_numel(), True)} {ctx[idx]}\n"
   f"  br label {ctx[x]}_exit\n{ctx[x][1:]}_exit:\n"
   f"  {ctx[x]} = phi {ldt(x.dtype, x.max_numel())} [{ctx[x]}_yes, {ctx[x]}_load], [{ctx[alt]}, {ctx[x]}_entry]"),
  (UPat.var('idx').load(name="x"), lambda ctx,x,idx:
   f"  {ctx[x]} = load {'volatile ' if is_volatile(idx) else ''}{ldt(idx.dtype, idx.max_numel())}, "
   f"{ldt(idx.dtype, idx.max_numel(), True)} {ctx[idx]}"),
  (UPat.var('idx').store(UPat.var("var")), lambda ctx,idx,var:
   f"  store {'volatile ' if is_volatile(idx) else ''}{ldt(var.dtype, idx.max_numel())} {ctx[var]}, "
   f"{ldt(idx.dtype, idx.max_numel(), True)} {ctx[idx]}"),

  # GEP/VECTORIZE/CAST for float4 support
  (UPat(Ops.STACK, name="x"), lambda ctx,x:
   "\n".join([(f"  {ctx[x]}_{i}" if i+1 != len(x.src) else f"  {ctx[x]}")+
               f" = insertelement {ldt(x.dtype, x.max_numel())} "+(f"{ctx[x]}_{i-1}" if i != 0 else "poison")+
               f", {ldt(u.dtype)} {ctx[u]}, i32 {i}" for i,u in enumerate(x.src)])),
  # unary/binary/ternary ops
  (UPat(Ops.BITCAST, name="x"), lambda ctx,x:
   f"  {ctx[x]} = bitcast {ldt(x.src[0].dtype, x.src[0].max_numel())} {ctx[x.src[0]]} to {ldt(x.dtype, x.max_numel())}"),
  (UPat(Ops.CAST, name="x"), lambda ctx,x: f"  {ctx[x]} = {lcast(x.src[0].dtype, x.dtype)} {ldt(x.src[0].dtype)} {ctx[x.src[0]]} to {ldt(x.dtype)}"),
  (UPat(Ops.TRUNC, name="x"),
   lambda ctx,x: f"  {ctx[x]} = call {ldt(x.dtype)} @llvm.trunc.{ldt(x.dtype)}({ldt(x.src[0].dtype)} {ctx[x.src[0]]})"),
  (UPat(GroupOp.Binary, name="x"), lambda ctx,x:
   f"  {ctx[x]} = {lop[x.src[0].dtype][x.op]} {ldt(x.src[0].dtype)} {ctx[x.src[0]]}, {ctx[x.src[1]]}"),
  (UPat(Ops.WHERE, name="x"), lambda ctx,x:
   f"  {ctx[x]} = select {ldt(x.src[0].dtype)} {ctx[x.src[0]]}, {ldt(x.src[1].dtype)} {ctx[x.src[1]]}, {ldt(x.src[2].dtype)} {ctx[x.src[2]]}"),

  # loop (a void RANGE is an unbounded loop header)
  (UPat(Ops.RANGE, dtypes.void, name="l"), lambda ctx,l: f"  br label %loop_{ctx[l][1:]}\nloop_{ctx[l][1:]}:"),
  (UPat(Ops.BACKEDGE, src=(UPat(), UPat(Ops.RANGE, dtypes.void, name="l"), UPat(name="c"))), lambda ctx,l,c:
    f"  br i1 {ctx[c]}, label %loop_{ctx[l][1:]}, label %loop_exit_{ctx[l][1:]}\nloop_exit_{ctx[l][1:]}:"),

  # range
  (UPat(Ops.RANGE, name="r"), lambda ctx,r:
   f"  br label %loop_entry_{range_str(r)}\n"
   f"loop_entry_{range_str(r)}:\n"
   f"  br label %loop_latch_{range_str(r)}\n"
   f"loop_latch_{range_str(r)}:\n"
   f"  {ctx[r]} = phi {ldt(r.dtype)} [ 0, %loop_entry_{range_str(r)} ], [ {ctx[r]}phi, %loop_footer_{range_str(r)} ]\n"
   f"  {ctx[r]}phi = add {ldt(r.dtype)} {ctx[r]}, 1\n"
   f"  {ctx[r]}cmp = icmp ult {ldt(r.dtype)} {ctx[r]}, {ctx[r.src[0]]}\n"
   f"  br i1 {ctx[r]}cmp, label %loop_body_{range_str(r)}, label %loop_exit_{range_str(r)}\n"
   f"loop_body_{range_str(r)}:"),
  (UPat(Ops.END, src=(UPat(), UPat(Ops.RANGE, name="r"))), lambda r:
   f"  br label %loop_footer_{range_str(r)}\n"
   f"loop_footer_{range_str(r)}:\n"
   f"  br label %loop_latch_{range_str(r)}\n"
   f"loop_exit_{range_str(r)}:"),

  # if
  (UPat(Ops.IF, name="x"), lambda ctx,x: f"  br i1 {ctx[x.src[0]]}, label %ifbody_{ctx[x][1:]}, label %ifskip_{ctx[x][1:]}\nifbody_{ctx[x][1:]}:"),
  (UPat(Ops.ENDIF, name="x"), lambda ctx,x: f"  br label %ifskip_{ctx[x.src[0]][1:]}\nifskip_{ctx[x.src[0]][1:]}:"),

  (UPat(Ops.BARRIER), lambda ctx: "  fence seq_cst"),

  # call a function by the name of its body, the args in param order
  (UPat(Ops.CALL, src=(UPat(Ops.LINEAR, name="body"),), allow_any_len=True, name="x"), lambda ctx,x,body: f"  call void @{ctx[body]}(" +
   ", ".join(f"{ldt(p.dtype, ptr=p.addrspace != AddrSpace.ALU)} {ctx[x.src[p.arg.slot+1]]}" for p in body.src if p.op is Ops.PARAM) + ")"),
])

class LLVMRenderer(Renderer):
  abi: str | None
  string_rewrite: PatternMatcher
  code_for_op = {k:lambda:None for v in lop.values() for k in v.keys()}

  extra_matcher = create_non_native_float_pats((dtypes.bfloat16,)) + pm_manual_bf16_cast + PatternMatcher([
    # GEP sign-extends indices to pointer width, so unsigned indices must be widened first.
    (UPat((Ops.INDEX, Ops.SHRINK), src=(UPat(), UPat(dtype=(dtypes.uint8, dtypes.uint16, dtypes.uint32), name="i")),
          allow_any_len=True, name="x"), lambda x,i: x.replace(src=(x.src[0], i.cast(dtypes.int64), *x.src[2:]))),
  ])
  def _render_fn(self, name:str, args:list[tuple[str,UOp]], kernel:list[str], prefix:list[str]|None=None) -> str:
    # Buffer views may start at an offset from the aligned allocation.
    sargs = ", ".join([f"{ldt(u.dtype, ptr=u.addrspace != AddrSpace.ALU)}{' noalias' if u.addrspace == AddrSpace.GLOBAL else ''} " + \
      name for name,u in args])
    return "\n".join((prefix or []) + [f"define{' ' + self.abi if self.abi else ''} void @{name}({sargs}) #0", "{"] + kernel + ["  ret void\n}"])
  def _render_kernel(self, uops: list[UOp], prefix:list[str]|None=None, name="test", fns:dict[UOp, str]|None=None) -> tuple[tuple[str, ...], str]:
    r: dict[UOp, str] = dict(fns or {}) # the functions the kernel calls, by name
    args: list[tuple[str, UOp]] = []
    kernel: list[str] = []
    vc = -1

    local_args: list[str] = []
    for u in uops:
      if u.op in {Ops.NOOP, Ops.CONST} or (u.op is Ops.STACK and not u.src): continue
      if u.op is Ops.AFTER:
        if u.dtype is dtypes.void: continue
        r[u] = r[u.src[0]]
        continue
      if u.op is Ops.SINK:
        if u.arg is not None: name = u.arg.function_name
        continue
      if u.op is Ops.PARAM:
        r[u] = f"%data{u.arg.slot}"
        args.append((r[u], u))
      elif u.op is Ops.BUFFER:
        r[u] = f"%{'local' if u.addrspace == AddrSpace.LOCAL else 'reg'}_{str(u.arg.slot)}"
        size = u.max_numel()
        if u.addrspace == AddrSpace.REG:
          kernel.append(f"  {r[u]} = alloca [{size} x {ldt(u.dtype)}]")
        elif self.has_local:
          local_args.append(f"@{r[u][1:]} = internal unnamed_addr addrspace(3) global [{size} x {ldt(u.dtype)}] undef, align 16")
          kernel.append(f"  {r[u]} = addrspacecast [{size} x {ldt(u.dtype)}] addrspace(3)* @{r[u][1:]} to [{size} x {ldt(u.dtype)}]*")
        else:
          kernel.append(f"  {r[u]} = alloca [{size} x {ldt(u.dtype)}], align 16")
      elif u.op is Ops.CAST and u.src[0].op is Ops.CONST: r[u] = lconst(u.src[0].val, u.dtype)
      elif u.op is Ops.CAST and ldt(u.dtype) == ldt(u.src[0].dtype):
        r[u] = r[u.src[0]] # cast from signed to unsigned of the same size is a noop, or pointer cast
      else:
        vc += 1
        r[u] = f"%v{vc}"

        # do the rendering of the llvm ir code
        l: str|None = self.string_rewrite.rewrite(u, ctx=r)
        if l is None:
          raise RuntimeError(f"failed to render {u.op} with {u.dtype} srcs {[x.dtype for x in u.src]}")
        kernel.append(l)
    return tuple(local_args), self._render_fn(name, args, kernel, prefix)

class CPULLVMRenderer(LLVMRenderer):
  has_local = False
  global_max = (1, 0, 0)
  abi = 'win64cc' if sys.platform == 'win32' else None
  string_rewrite = base_rewrite
  def render(self, uops: list[UOp]) -> str: # the kernel is first, it is the entry. its functions follow, a name traced with other args gets a suffix
    fns = {b: f"{b.arg}_{i}" for i, b in enumerate(b for b in UOp.sink(*uops).toposort() if b.op is Ops.LINEAR)}
    defs = [self._render_kernel(b.src, name=n, fns=fns)[1] for b, n in fns.items()]
    return "\n".join((k:=self._render_kernel(uops, fns=fns))[0] + (k[1], *defs, self._render_footer(uops)))
  def _render_footer(self, uops: list[UOp]) -> str: return 'attributes #0 = { alwaysinline nounwind "no-builtins" "no-trapping-math"="true" }'
  def __init__(self, target:Target):
    super().__init__(target)
    from tinygrad.runtime.support.compiler_llvm import CPULLVMCompiler
    self.compiler = CPULLVMCompiler(target.arch.split(","))

  # FIXME: fp16 works on non-osx, but only if the cpu supports it
  def supported_dtypes(self):
    return {d for d in super().supported_dtypes() if
            (d != dtypes.bfloat16 or self.target.arch.startswith(("x86", "arm"))) and (d != dtypes.half or OSX) and d not in dtypes.fp8s}

barrier = 'fence syncscope("workgroup") release\ntail call void @llvm.amdgcn.s.barrier()\nfence syncscope("workgroup") acquire\n'
code_for_workitem = {"g": lambda x: f"tail call i32 @llvm.amdgcn.workgroup.id.{chr(120+int(x))}()",
                     "l": lambda x: f"tail call i32 @llvm.amdgcn.workitem.id.{chr(120+int(x))}()"}
# https://rocm.docs.amd.com/projects/llvm-project/en/latest/LLVM/llvm/html/AMDGPUUsage.html#llvm-ir-intrinsics
llvm_intrinsics = {Ops.SQRT: "sqrt", Ops.LOG2: "log2", Ops.EXP2: "exp2"}
class AMDLLVMRenderer(LLVMRenderer):
  shared_max = HIPRenderer.shared_max
  global_max = HIPRenderer.global_max
  global_prod_max = HIPRenderer.global_prod_max
  abi = "amdgpu_kernel"
  code_for_op = {**LLVMRenderer.code_for_op, **{op: lambda: None for op in llvm_intrinsics}}
  string_rewrite = PatternMatcher([
    (UPat(Ops.SPECIAL, name="x"), lambda ctx, x: f"  {ctx[x]} = " + f"{ code_for_workitem[x.arg[0]](x.arg[-1])}; "),
    (UPat(tuple(llvm_intrinsics), name="x"),
    lambda ctx, x: f"  {ctx[x]} = call {ldt(x.dtype)} @llvm.{llvm_intrinsics[x.op]}.{ldt(x.dtype)}({ldt(x.src[0].dtype)} {ctx[x.src[0]]})"),
    (UPat(Ops.BARRIER), lambda ctx: barrier),
    (UPat(Ops.CAST, dtypes.fp8s, (UPat(dtype=dtypes.float),), name="x",), lambda ctx,x:
      f"  {ctx[x]} = call i8 @f32_to_fp8({ldt(x.src[0].dtype)}  {ctx[x.src[0]]}, i1 {fp8_index(x.dtype)})"),
    (UPat(Ops.CAST, dtypes.float, (UPat.var("y", dtypes.fp8s),), name="x",), lambda ctx,x,y:
      f"  {ctx[x]}_i32 = zext i8 {ctx[x.src[0]]} to i32\n"
      f"  {ctx[x]} = call float @llvm.amdgcn.cvt.f32.{('fp8', 'bf8')[fp8_index(y.dtype)]}(i32 {ctx[x]}_i32, i32 0)"),
  ]) + base_rewrite
  extra_matcher = LLVMRenderer.extra_matcher + create_non_native_float_pats(dtypes.fp8s) + PatternMatcher([
    # amd llvm intrinsics llvm.log2/llvm.exp2 don't support double
    (UPat(Ops.LOG2, dtype=dtypes.double, src=(UPat.var("d"),)), xlog2),
    (UPat(Ops.EXP2, dtype=dtypes.double, src=(UPat.var("d"),)), xexp2),
  ])
  def asm(self, prg: UOp, lin: UOp) -> bytes:
    from tinygrad.renderer.amd.elf import assemble_linear
    return assemble_linear(prg, lin, self.target.arch)
  def render(self, uops: list[UOp]) -> str:
    fp8_max = amd_fp8s(self.target.arch)[0].max if (has_fp8:=any(u.dtype in dtypes.fp8s for u in uops)) else None
    prefix = [f"""define i8 @f32_to_fp8(float %val, i1 %is_bf8) {{
entry: %ival = bitcast float %val to i32\n  %exp = and i32 %ival, 2139095040\n  %is_special = icmp eq i32 %exp, 2139095040
br i1 %is_special, label %select_clip, label %clip
clip: br i1 %is_bf8, label %bf8_clip, label %fp8_clip
bf8_clip: %clamped_bf8 = call float @llvm.amdgcn.fmed3.f32(float %val, float 57344.0, float -57344.0)\n  br label %select_clip
fp8_clip: %clamped_fp8 = call float @llvm.amdgcn.fmed3.f32(float %val, float {fp8_max}, float -{fp8_max})    \n  br label %select_clip
select_clip: %phi_val = phi float [%val, %entry], [%clamped_bf8, %bf8_clip], [%clamped_fp8, %fp8_clip]\n  br i1 %is_bf8, label %do_bf8, label %do_fp8
do_bf8: %packed_bf8 = call i32 @llvm.amdgcn.cvt.pk.bf8.f32(float %phi_val, float %phi_val, i32 0, i1 false)\n  br label %exit
do_fp8: %packed_fp8 = call i32 @llvm.amdgcn.cvt.pk.fp8.f32(float %phi_val, float %phi_val, i32 0, i1 false)\n  br label %exit
exit: %packed = phi i32 [%packed_bf8, %do_bf8], [%packed_fp8, %do_fp8]\n  %trunc = trunc i32 %packed to i8\n  ret i8 %trunc
}}""".replace(": ", ":\n  ")] if has_fp8 else []
    return "\n".join((k:=self._render_kernel(uops, prefix))[0] + (k[1], self._render_footer(uops)))
  def _render_footer(self, uops: list[UOp]) -> str:
    # TODO: this is copied from cstyle
    local_dims = [u.src[0] for u in uops if u.op is Ops.SPECIAL and u.arg[0] == "l"]
    requiredMaxThreadsPerBlock = prod([d.vmax for d in local_dims])
    attributes = ["alwaysinline", "nounwind", '"no-builtins"',
                  f'"amdgpu-flat-work-group-size"="1,{requiredMaxThreadsPerBlock}"', '"no-trapping-math"="true"']
    return 'attributes #0 = { ' + ' '.join(attributes) + ' }'
  @staticmethod
  def is_rdna4(arch): return arch.split(':')[0] in {'gfx1200', 'gfx1201'}
  def __init__(self, target:Target):
    super().__init__(target)
    from tinygrad.runtime.support.compiler_llvm import AMDLLVMCompiler
    self.compiler, self.tensor_cores, self.is_cdna = AMDLLVMCompiler(target.arch), tc.get_amd(target.arch), HIPRenderer.is_cdna(target.arch)
    self.string_rewrite += PatternMatcher([
      (UPat(Ops.WMMA, name="wmma"), lambda ctx, wmma, rdna4=AMDLLVMRenderer.is_rdna4(target.arch), cdna=self.is_cdna:
        render_wmma_amd(ctx, wmma, cdna, rdna4))
    ])
    if self.is_cdna: self.extra_matcher += tc.pm_validate_wmma_cdna
    if target.arch in {"gfx1100", "gfx1151"}: self.extra_matcher += tc.pm_validate_wmma_rdna3
    if target.arch in {"gfx1200", "gfx1201"}: self.extra_matcher += tc.pm_validate_wmma_rdna4

  def supported_dtypes(self): return {d for d in super().supported_dtypes() if d not in dtypes.fp8s or d in amd_fp8s(self.target.arch)}
