from dataclasses import replace
import itertools, functools
from tinygrad.helpers import DISABLE_FAST_IDIV, TRANSCENDENTAL, SPEC, DEBUG, VIZ, IMAGE, NOOPT, EMULATED_DTYPES, USE_TC
from tinygrad.helpers import ALLOW_TF32, DEFAULT_FLOAT, DEFAULT_INT, TC_SELECT, TC_OPT, TC_MIN_GLOBALS, TracingKey, Context, panic
from tinygrad.uop.ops import PatternMatcher, graph_rewrite, UOp, Ops, UPat, rewrite_group, KernelInfo, ProgramInfo, GroupOp, AxisType
from tinygrad.uop.weak import pm_lower_weak, pm_commit_weak, pm_cast_const
from tinygrad.uop.render import render_uir
from tinygrad.uop.spec import type_verify, spec_tensor, spec_program
from tinygrad.renderer import Renderer, Estimates
from tinygrad.renderer.isa import ISARenderer, IselContext
from tinygrad.dtype import dtypes, AddrSpace

# import all pattern matchers here
from tinygrad.codegen.gpudims import pm_group_gpudims, pm_range_to_special
from tinygrad.uop.symbolic import sym, symbolic_simple, symbolic, pm_move_where_on_load, pm_clean_up_group_sink, pm_remove_invalid, invalid_gate
from tinygrad.uop.movement import mop_cleanup
from tinygrad.codegen.decomp.dtype import pm_dtype_decomps
from tinygrad.codegen.decomp.op import get_late_rewrite_patterns, get_simplifying_rewrite_patterns
from tinygrad.codegen.decomp.transcendental import get_transcendental_patterns
from tinygrad.codegen.late.coalesce import indexing_simplify
from tinygrad.codegen.opt.postrange import apply_opts
from tinygrad.codegen.late.gater import pm_move_gates_from_index
from tinygrad.codegen.simplify import pm_simplify_ranges, pm_flatten_range, pm_split_ranges, pm_load_collapse, pm_reduce_unparented
from tinygrad.schedule.multi import multi_pm
from tinygrad.schedule.prepare import pm_mops
from tinygrad.codegen.late.linearizer import CFGContext, pm_split_ends, pm_add_control_flow, linearize
from tinygrad.codegen.late.regalloc import LinearScanRegallocContext, pm_regalloc_rewrite
from tinygrad.codegen.late.coalesce import memory_coalescing, pm_simplify_add_image
from tinygrad.helpers import all_same, all_int, argsort, partition, to_function_name
from tinygrad.uop.ops import _broadcast_shape, identity_element
from tinygrad.schedule.rangeify import BufferizeOpts

def do_number_param(ctx:tuple[int, dict[str, int]], x:UOp): # after the params, one slot per name
  if x.is_variable: return x.replace(arg=replace(x.arg, slot=ctx[0] + ctx[1].setdefault(x.arg.name, len(ctx[1]))))

pm_number_params = PatternMatcher([
  (UPat(Ops.PARAM, name="x"), do_number_param),
])

def build_range_map(sink:UOp) -> dict[tuple, int]:
  ctx: dict[tuple, int] = {}
  for x in sink.toposort():
    if x.op is Ops.RANGE and x.axis_type is AxisType.UPCAST:
      ctx[x.arg] = len(ctx)
  return ctx

def expand_reduce(r:UOp):
  range_srcs = []
  new_axes = []
  for u in r.src[1:]:
    if u.op == Ops.RANGE:
      range_srcs.append(u)
    else:
      for i,s in enumerate(u.shape):
        if s > 1: new_axes.append(i)
  if len(new_axes) == 0: return None
  assert r.arg[1] == 0
  # permute so new_axes come to front, then reduce
  perm = tuple(new_axes) + tuple(i for i in range(len(r.src[0].shape)) if i not in new_axes)
  out_shape = tuple([1 if i in new_axes else s for i,s in enumerate(r.src[0].shape)])
  return r.src[0].permute(perm).reduce(*range_srcs, arg=(r.arg[0], len(new_axes))).reshape(out_shape)

def contract_axis(u:UOp, dims:list[int]) -> UOp:
  return u.permute([i for i in range(u.ndim) if i not in dims]+dims).flatten(-len(dims))

def unroll_axis(u:UOp, dims:list[int], sizes:list[int]) -> UOp:
  out = u.unflatten(-1, tuple(sizes))
  return out.permute(argsort([i for i in range(out.ndim) if i not in dims]+dims))

def expand_wmma(ctx:dict[tuple, int], u:UOp):
  if u.arg[3] is None: return None
  in0, in1, out0 = [[ctx[rn] for rn,_ in upcast_axes] for upcast_axes in u.arg[3]]
  wmma = u.replace(src=(contract_axis(u.src[0], in0), contract_axis(u.src[1], in1), u.src[2]), arg=(*u.arg[:3], None))
  return unroll_axis(wmma, out0, [sz for _,sz in u.arg[3][2]])

expander = PatternMatcher([
  (UPat(Ops.REDUCE, name="r"), expand_reduce),
  (UPat(Ops.RANGE, name="r"),
   lambda ctx, r: UOp.const(tuple(range(r.vmax+1)), r.dtype) \
    .reshape(tuple([r.vmax+1 if i == ctx[r.arg] else 1 for i in range(len(ctx))])) if r.arg in ctx else None),
  (UPat(Ops.WMMA, name="u"), expand_wmma),
])+pm_flatten_range+mop_cleanup

def expand_broadcast(x:UOp):
  shapes = [u._shape for u in x.src]
  if any(s is None for s in shapes) or all_same(shapes): return None
  shape = _broadcast_shape(*shapes)
  return x.replace(src=tuple([u.expand(shape) for u in x.src]))

def broadcast_and_devec_wmma(b:UOp):
  shapes = [u.shape[:-1] for u in b.src]
  if not any(shapes): return None
  shape = _broadcast_shape(*shapes)
  src_expanded = tuple([u.expand(shape+(u.shape[-1],)) for u in b.src])
  src = []
  for idx in itertools.product(*[range(i) for i in b.shape[:-1]]):
    src.append(b.replace(src=tuple([x.index(*idx) for x in src_expanded])))
  return UOp.stack(*src).reshape(b.shape)

pm_wmma_add = PatternMatcher([
  (UPat(Ops.WMMA, name="wmma") + UPat.var("add"),
   lambda add, wmma: UOp(wmma.op, src=(wmma.src[0], wmma.src[1], wmma.src[2]+add), arg=wmma.arg)),
  # push permute/reshape to the other side of the add
  (UPat(Ops.PERMUTE, src=(UPat(Ops.WMMA, name="wmma"),), name="permute") + UPat.var("add"),
    lambda wmma,permute,add: (wmma + add.permute(argsort(permute.arg))).permute(permute.arg)),
  (UPat(Ops.PERMUTE, src=(UPat(Ops.RESHAPE, src=(UPat(Ops.WMMA, name="wmma"), UPat()), name="reshape"),), name="permute") + UPat.var("add"),
    lambda wmma,reshape,permute,add: (wmma + add.permute(argsort(permute.arg)).reshape(wmma.shape)).reshape(reshape.shape).permute(permute.arg)),
])

pm_expand_broadcast = pm_wmma_add+PatternMatcher([
  (UPat(GroupOp.Binary|GroupOp.Ternary|{Ops.STORE}, name="x"), expand_broadcast),
  (UPat(Ops.WMMA, name="b"), broadcast_and_devec_wmma),
])

def do_devectorize(b:UOp):
  if b.shape == (): return None
  # broadcasting needs to be already unpacked, Invalid matches any dtype and shape
  if not all(x.shape == b.shape or x.base.is_invalid for x in b.src): return None
  src = []
  for idx_c in itertools.product(*[[UOp.const(i) for i in range(x)] for x in b.shape]):
    src.append(b.replace(src=tuple(x.base if x.base.is_invalid else x.index(*idx_c) for x in b.src)))
  return UOp.stack(*src).reshape(b.shape) if b.op is not Ops.STORE else UOp.group(*src)

def do_stack_wmma(u:UOp):
  if all(x.op in (Ops.STACK, Ops.WMMA) for x in u.src): return None
  assert len(u.shape) == 1
  src = []
  for b in u.src:
    if b.op != Ops.STACK:
      src.append(UOp.stack(*[b.index(i) for i in range(b.max_numel())]))
    else:
      src.append(b)
  return u.replace(src=tuple(src))

devectorizer2 = pm_mops+PatternMatcher([
  # unpack broadcasting
  (UPat(GroupOp.Elementwise|{Ops.LOAD,Ops.STORE}, name="b"), do_devectorize),
  # INDEX without src is nothing (TODO: this should be in mop_cleanup)
  (UPat(Ops.INDEX, src=(UPat.var('x'),)), lambda x: x),
  # unpack WMMA
  (UPat(Ops.WMMA, name="u"), do_stack_wmma),
  # stacked INDEX is many INDEX
  (UPat(Ops.INDEX, src=(UPat(GroupOp.Defines, name="b"), UPat(Ops.STACK, name="s")), name="x"),
   lambda b,s,x: UOp.stack(*[x.replace(src=(b,u)) for u in s.src])),
  # INDEX into RESHAPE moves the RESHAPE
  (UPat(Ops.INDEX, src=(UPat(GroupOp.Defines, name="b"), UPat(Ops.RESHAPE, name="s"))),
   lambda b,s: b.index(s.src[0]).reshape(s.shape)),
  # RESHAPE a void is removed (hack for AFTER)
  (UPat(Ops.RESHAPE, dtype=dtypes.void, name="x"), lambda x: x.src[0]),
  # reshape of a single element shaped value to scalar is an index
  (UPat(Ops.RESHAPE, name="x"), lambda x: x.src[0].index(0) if x.marg == () and x.src[0].shape == (1,) else None),
  # EXPAND on scalar -> nested STACKs with the same shape
  (UPat(Ops.EXPAND, src=(UPat.var("x"), UPat()), name="out"),
   lambda x,out: functools.reduce(lambda x,s: UOp.stack(*([x]*s)), reversed(out.shape), x)
   if x.shape == () and all_int(out.shape) and 0 not in out.shape else None),
])

def fix_group_for_reduce(x:UOp):
  threads = (AxisType.WARP, AxisType.LOCAL)
  reduce_gfr, reduce_r = partition(x.src[1:], lambda u: u.op is Ops.RANGE and u.axis_type in threads)
  if len(reduce_gfr) == 0: return None

  # do only the non grouped reduces early
  ret = x.replace(src=(x.src[0],)+tuple(reduce_r))
  reduce_loop = [x.replace(arg=(AxisType.WEAK, x.axis_id[0]+100, *x.axis_id[1:])) for x in reduce_gfr]
  buf = ret.bufferize(*reduce_gfr, arg=BufferizeOpts(None, AddrSpace.LOCAL)).index(*reduce_loop)

  # do the final reduce (if/barrier are added in gpudims step)
  # NOTE: we remove all horizontal reduces here, they remain in the first reduce
  return buf.reduce(*reduce_loop, arg=(x.arg[0], 0))

def merge_reduce_ends(sink:UOp):
  # merge ENDs that share the same range and nesting context (only those created by reduce_to_acc)
  # ENDs at different nesting depths get cloned RANGEs so each RANGE maps to one END
  range_to_ends: dict[tuple[UOp, ...], list[UOp]] = {}
  for u in sink.backward_slice:
    if u.op is Ops.END and u.tag == "mergeable": range_to_ends.setdefault(u.src[1:], []).append(u)
  subs: dict[UOp, UOp] = {}
  next_axis = max((u.axis_id[0] for u in sink.backward_slice if u.op is Ops.RANGE), default=-1) + 1
  for r, ends in range_to_ends.items():
    if len(ends) <= 1: continue
    by_ctx: dict[frozenset[UOp], list[UOp]] = {}
    for e in ends: by_ctx.setdefault(frozenset(e.ranges), []).append(e)
    for i, group in enumerate(by_ctx.values()):
      tr = r if i == 0 else tuple(rr.replace(arg=(rr.axis_type, next_axis + j, *rr.axis_id[1:])) for j, rr in enumerate(r))
      if i > 0: next_axis += len(r)
      mapped = [e.substitute(dict(zip(r, tr))) if i > 0 else e for e in group]
      merged = mapped[0] if len(mapped) == 1 else UOp.group(*(e.src[0] for e in mapped)).end(*tr)
      for e in group: subs[e] = merged
  return sink.substitute(subs) if subs else None

def reduce_ranges_to_acc(ctx:itertools.count, r:UOp):
  acc = UOp.alloc_like(r, next(ctx), AddrSpace.REG)
  input_ranges = tuple(x for x in r.src[0].ranges if x not in r.src[1:])
  acc_init = acc.after(*input_ranges).store(UOp.const(identity_element(r.arg[0], r.dtype)))
  acc_initted = acc.after(acc_init, *r.src[1:])
  inp = r.src[0].reduce(arg=r.arg) if r.arg[1] else r.src[0]
  acc_out = acc_initted.store(acc_initted.alu(r.arg[0], inp)).end(*r.src[1:]).rtag("mergeable")
  return acc.after(acc_out)

def expand_horizontal_reduce(r:UOp):
  inp = r.src[0]
  vals = [inp.index(*idx) for idx in itertools.product(*[range(inp.max_shape[a]) for a in range(r.arg[1])])]
  return functools.reduce(lambda x,y: x.alu(r.arg[0], y), vals)

# an Invalid in a REDUCE source is that reduce's identity
pm_reduce_identity = PatternMatcher([
  (invalid_gate.reduce(allow_any_len=True, name="red"), lambda red,cond,x,i:
   red.replace(src=(cond.where(x, x.const_like(identity_element(red.arg[0], red.dtype))),)+red.src[1:])),
])

pm_reduce_local = pm_wmma_add+PatternMatcher([
  # fix group for reduce
  (UPat(Ops.REDUCE, name="x"), fix_group_for_reduce),
  # remove reduces
  (UPat(Ops.REDUCE, src=(UPat(), UPat()), allow_any_len=True, name="r"), reduce_ranges_to_acc),
  (UPat(Ops.REDUCE, src=(UPat(),), name="r"), expand_horizontal_reduce),
  (UPat(Ops.SINK, name="sink"), merge_reduce_ends),
])+pm_clean_up_group_sink

def is_shape_changing_bitcast(u:UOp): return u.op is Ops.BITCAST and u.shape != u.src[0].shape
def maybe_load(u:UOp): return u.load() if u.addrspace in (AddrSpace.GLOBAL, AddrSpace.LOCAL, AddrSpace.REG) else u
pm_add_loads = PatternMatcher([
  (UPat(GroupOp.Elementwise|{Ops.REDUCE,Ops.WMMA,Ops.STACK}, name="x"),
   lambda x: None if is_shape_changing_bitcast(x) else x.replace(src=tuple(map(maybe_load, x.src)))),
  (UPat(Ops.STORE, name="x"), lambda x: x.replace(src=(x.src[0], maybe_load(x.src[1]))+x.src[2:])),
])

def add_local_buffer(ctx, x:UOp):
  upstream_locals = [u for u in x.ranges if u.axis_type in (AxisType.WARP, AxisType.LOCAL)] if x.arg.addrspace is AddrSpace.LOCAL else []
  buf = UOp.alloc(tuple(int(u.vmax+1) for u in upstream_locals)+x.max_shape, x.dtype, slot=next(ctx), addrspace=x.arg.addrspace)
  return buf.after(buf.index(*upstream_locals, *x.src[1:]).store(x.src[0]).end(*x.src[1:])).index(*upstream_locals)

pm_add_local_buffers = PatternMatcher([
  (UPat(Ops.STAGE, name="x"), add_local_buffer),
])+pm_mops

# float ALUs need a float operand
# make that cast explicit before the decomps, which expand SIN/LOG2/EXP2 into float polynomials and assert a float operand
pm_cast_float_alu = PatternMatcher([
  (UPat((Ops.SIN, Ops.LOG2, Ops.EXP2, Ops.SQRT, Ops.RECIPROCAL), src=(UPat(name="x"),), name="u"),
   lambda u,x: u.replace(src=(x.cast(u.dtype),)) if x.dtype != u.dtype else None),
])

def _is_local_store(x:UOp): return x.op is Ops.STORE and x.addrspace is AddrSpace.LOCAL

def add_raw_barrier(after:UOp):
  # loads from a LOCAL buffer that depend (via AFTER) on stores to LOCAL memory need a workgroup barrier
  if after.addrspace is not AddrSpace.LOCAL: return None
  # one toposort over all the deps
  deps = UOp.sink(*after.src[1:]).toposort(gate=lambda x: x.op is not Ops.BARRIER)
  if not any(_is_local_store(x) for x in deps): return None
  return after.src[0].after(UOp(Ops.BARRIER, src=after.src[1:]))

def add_war_barrier(end:UOp):
  # a LOCAL buffer stored and loaded in the same loop needs a barrier at the end of the loop body
  rngs = [r for r in end.ended_ranges if r.axis_type in (AxisType.WEAK, AxisType.LOOP) and r.vmax > 0]
  if not rngs or end.src[0].op is Ops.BARRIER: return None
  sl = end.src[0].backward_slice_with_self
  # only stores that are inside this loop body (not in the backward slice through AFTER chains from other loops)
  store_bufs = {x.buf_uop for x in sl if _is_local_store(x) and any(r in x.ranges for r in rngs)}
  # a load whose buffer matches a local store's buffer is necessarily a local load
  if not any(x.op is Ops.LOAD and x.src[0].buf_uop in store_bufs for x in sl): return None
  return end.replace(src=(UOp(Ops.BARRIER, src=(end.src[0],)),)+end.src[1:])

pm_implicit_barriers = PatternMatcher([
  (UPat(Ops.AFTER, name="after"), add_raw_barrier),
  (UPat((Ops.END, Ops.BACKEDGE), name="end"), add_war_barrier),
])

def full_rewrite_to_sink(ast:UOp, ren:Renderer, optimize:bool=True) -> UOp:
  if DEBUG >= 5: print(render_uir(list(ast.toposort())))
  if SPEC: type_verify(ast, spec_tensor)

  # resolve UNSHARDs (multi-device UNSHARDs are already resolved by the scheduler; this handles in-kernel shards, e.g. fragments)
  sink = graph_rewrite(ast, multi_pm, name="multi_pm")

  # preprocess
  sink = graph_rewrite(sink, pm_mops, name="early movement ops", bottom_up=True)

  # first we optimize
  if optimize:
    # collapse loads reduce (indexing by a tensor)
    sink = graph_rewrite(sink, pm_load_collapse, name="load collapse")

    # split ranges
    sink = graph_rewrite(sink, pm_split_ranges+pm_flatten_range, ctx={}, name="split ranges")

    # symbolic (NOTE: this is a requirement for pm_simplify_ranges to be correct)
    sink = graph_rewrite(sink, sym+pm_flatten_range, name="initial symbolic")

    # optimize (schedule) the AST
    sink = graph_rewrite(sink, pm_flatten_range+pm_simplify_ranges, ctx={}, name="simplify ranges")

    # do postrange optimization, BEAM or hand_coded_optimizations
    sink = apply_opts(sink, ren, beam=ast.arg.beam)

  # ** expander (expand_rewrite) **
  # reduce_unparented: a REDUCE whose src folded to a CONST (e.g. x*0) has no parented ranges, collapse it before the expander
  sink = graph_rewrite(sink, sym+pm_move_where_on_load+pm_flatten_range+pm_reduce_unparented+pm_reduce_identity, name="postopt symbolic")

  # expand
  sink = graph_rewrite(sink, expander, ctx=build_range_map(sink), name="expander")

  slots = itertools.count(max([u.arg.slot+1 for u in sink.toposort() if u.op in {Ops.BUFFER, Ops.ALLOC}], default=0))

  # remove reduce
  sink = graph_rewrite(sink, mop_cleanup+pm_reduce_local, ctx=slots, name="remove reduces")

  # add locals
  sink = graph_rewrite(sink, pm_add_local_buffers, ctx=slots, name="add local buffers")

  # group GPU dimensions early so their index arithmetic goes through normal lowering
  sink = graph_rewrite(sink, pm_group_gpudims, ctx=ren, name="group gpudims", walk=True)

  # **** optimizations are done, now we lower to actual code ****

  sink = graph_rewrite(sink, symbolic_simple+pm_expand_broadcast+pm_add_loads, name="*** expand broadcast / add loads")

  # devectorize
  sink = graph_rewrite(sink, symbolic_simple+devectorizer2+indexing_simplify, name="devectorize2")

  # some coalescing misses without this
  sink = graph_rewrite(sink, sym, name="early symbolic")

  # do memory coalescing (late)
  sink = memory_coalescing(sink, ren)
  sink = graph_rewrite(sink, symbolic_simple+pm_simplify_add_image, name="add images", ctx=({}, ren), bottom_up=True)

  # extra symbolic before decomp. crashes without this?
  # NOTE: also run indexing_simplify here, while the index is still weakint and (x+y)*c -> x*c+y*c applies
  # commit widths minted in this fixpoint before lowering inspects INDEX shapes
  sink = graph_rewrite(sink, sym+indexing_simplify+pm_commit_weak, name="extra symbolic")

  # the boundary: required compute dtypes settle here; derivable const edges may stay bare
  # NOTE: we need indexing_simplify to remove the cast to long using the Invalid
  # NOTE: symbolic must NOT be composed here -- pm_data_invalid pushes the weak result CAST into a gated WHERE, remaking the weak node, and it cycles
  sink = graph_rewrite(sink, pm_lower_weak+indexing_simplify, name="lower all index dtypes")

  # final symbolic before decomp
  sink = graph_rewrite(sink, symbolic, name="final symbolic")

  sink = graph_rewrite(sink, pm_cast_float_alu, name="cast float alu operands")

  # **** decomps ****

  # floordiv+mod / dtype decomp (early)
  supported_ops = tuple(ren.code_for_op.keys())
  pm_decomp = symbolic_simple+get_simplifying_rewrite_patterns(supported_ops)
  sink = graph_rewrite(sink, pm_decomp, name="early decompositions")

  # late decomps + move gates from unrenderable INVALID where
  sink = graph_rewrite(sink, pm_dtype_decomps+pm_commit_weak, ctx=(set(), ren), name="decomp dtypes")
  pm_decomp = pm_decomp+\
    get_late_rewrite_patterns(supported_ops, bool(DISABLE_FAST_IDIV))+\
    get_transcendental_patterns(supported_ops, TRANSCENDENTAL>=2)
  sink = graph_rewrite(sink, pm_decomp, ctx=ren, name="late decompositions")
  sink = graph_rewrite(sink, pm_move_gates_from_index, name="move gates from index")

  # final rules for the renderer (without sym)
  extra_matcher = ren.extra_matcher if ren.extra_matcher is not None else PatternMatcher([])
  pm_final_rewrite = pm_commit_weak+pm_decomp+extra_matcher+pm_split_ends
  sink = graph_rewrite(sink, pm_final_rewrite+pm_remove_invalid, ctx=ren, name="final rewrite")

  # commit every const still bare so no renderer reads one
  sink = graph_rewrite(sink, pm_cast_const, name="cast consts")

  # add implicit barriers (stores/loads through LOCAL memory ordered by AFTER or across loop iterations need workgroup barriers)
  sink = graph_rewrite(sink, pm_implicit_barriers, name="add implicit barriers")

  # hardware ranges are no longer loops; preserve their already lowered bounds
  sink = graph_rewrite(sink, pm_range_to_special, name="range to special")

  # this was the linearizer
  sink = graph_rewrite(sink, pm_add_control_flow, ctx=CFGContext(sink), name="add control flow", bottom_up=True)

  # put the variables in slots
  num_params = max([x.arg.slot + 1 for x in sink.toposort() if x.op is Ops.PARAM and not x.is_variable], default=0)
  sink = graph_rewrite(sink, pm_number_params, ctx=(num_params, {}), name="number variables", walk=True)

  if VIZ: graph_rewrite(sink, PatternMatcher([]), name="View Output AST")
  if SPEC:
    import os
    if os.environ.get("DBGTV"):
      try: type_verify(sink, spec_program)
      except RuntimeError:
        print(render_uir(list(sink.toposort())))
        raise
    else: type_verify(sink, spec_program)

  # return the rewritten sink
  return sink

# inject IF/ENDIF. only needed if device doesn't support gated stores
pm_linearize_cleanups = PatternMatcher([
  # if statements are not allowed in the graph
  (UPat((Ops.IF, Ops.ENDIF)), lambda: panic(RuntimeError, "if not allowed in graph")),
  # gated STORE becomes IF-STORE-ENDIF. this is the only use of IF-ENDIF
  (UPat(Ops.STORE, name="u", src=(UPat((Ops.INDEX, Ops.SHRINK)).or_casted(), UPat(), UPat(name="gate", dtype=dtypes.bool))),
   lambda u, gate: ((st:=u.replace(src=u.src[0:2])), [mif:=UOp(Ops.IF, src=(gate, u.src[0])), st, UOp(Ops.ENDIF, src=(mif,))]))
])

pm_alloc_to_buf = PatternMatcher([(UPat(Ops.ALLOC, name="x"), lambda x: ((buf:=x.replace(op=Ops.BUFFER)), [buf])),])

# requires lst be toposorted. like graph rewrite, but for lines
def line_rewrite(lst:list[UOp], pm:PatternMatcher, ctx=None) -> list[UOp]:
  newlst = []
  replaced: dict[UOp, UOp] = {}
  for u in lst:
    nu = u.replace(src=tuple([replaced.get(x, x) for x in u.src]))
    ret: tuple[UOp, list[UOp]] = pm.rewrite(nu, ctx) or (nu, [nu])
    replaced[u] = ret[0]
    newlst.extend(ret[1])
  return newlst

pm_lower_calls = PatternMatcher([
  (UPat(Ops.CALL, src=(UPat(Ops.SINK),), allow_any_len=True, name="call"),
   lambda ctx,call: call.replace(src=(full_rewrite_to_sink(call.body, ctx, optimize=False),)+call.src[1:])),
])

pm_call_fixup = PatternMatcher([
  (UPat(Ops.CALL, src=(UPat(Ops.SINK, name="sink"),), allow_any_len=True, name="call"),
   lambda call,sink: call.replace(src=(UOp(Ops.LINEAR, src=tuple(line_rewrite(linearize(sink), pm_linearize_cleanups+pm_alloc_to_buf)),
                                           arg=to_function_name(call.arg.name)),)+call.src[1:])),
])

def do_linearize(ctx:Renderer, prg:UOp, sink:UOp) -> UOp:
  if DEBUG >= 3 and sink.arg.applied_opts: print(f"{sink.arg.function_name:<25} opts: {sink.arg.applied_opts}")
  sink = graph_rewrite(sink, pm_call_fixup, name="call fixup", enter_calls=True)
  lst = line_rewrite(linearize(sink), pm_linearize_cleanups+pm_alloc_to_buf)
  prg = prg.replace(src=(lst[-1],))
  # isa renderers need to allocate registers
  if isinstance(ctx, ISARenderer):
    lin_ctx = ctx.linear_ctx_type(ctx)
    lst = line_rewrite(lst, ctx.pre_regalloc_matcher, lin_ctx)
    # register definitions (INS without srcs) move to the top so regalloc sees their live ranges span the whole program (callee saved regs)
    lst = sorted(lst, key=lambda u: u.op is not Ops.INS or bool(u.src))
    regalloc_ctx = LinearScanRegallocContext(lin_ctx, lst, ctx)
    lst = line_rewrite(lst, pm_regalloc_rewrite, regalloc_ctx)
    lst = line_rewrite(lst, ctx.post_regalloc_matcher, lin_ctx)
    if DEBUG >= 4: print(ctx.asm_str(lst, sink.arg.function_name))
  return prg.replace(src=prg.src + (UOp(Ops.LINEAR, src=tuple(lst)),))

def do_estimates(prg:UOp, sink:UOp, lin:UOp) -> UOp|None:
  if sink.arg.estimates is not None: return None
  return prg.replace(src=(sink.replace(arg=replace(sink.arg, estimates=Estimates.from_uops(lin.src, ignore_indexing=True))),)+prg.src[1:])

def do_assemble(ctx:Renderer, prg:UOp, lin:UOp) -> UOp:
  src = "\n".join(str(u.arg[0]) for u in lin.src)
  if DEBUG >= 4: print(src)
  binary = ctx.asm(prg, lin)
  return prg.replace(src=prg.src[:2]+(UOp(Ops.SOURCE, arg=src), UOp(Ops.BINARY, arg=binary)))

def do_render(ctx:Renderer, prg:UOp, lin:UOp) -> UOp:
  src = ctx.render(list(lin.src))
  return prg.replace(src=prg.src + (UOp(Ops.SOURCE, arg=src),))

def do_compile(ctx:Renderer, prg:UOp, source:UOp) -> UOp|None:
  if DEBUG >= 4: print(source.arg)
  lib = ctx.compiler.compile_cached(source.arg)
  if DEBUG >= 7: ctx.compiler.disassemble(lib)
  return prg.replace(src=prg.src + (UOp(Ops.BINARY, arg=lib),))

pm_to_program = PatternMatcher([
  (UPat(Ops.PROGRAM, src=(UPat(Ops.SINK, name="sink"),), name="prg"), do_linearize),
  (UPat(Ops.PROGRAM, src=(UPat(Ops.SINK, name="sink"), UPat(Ops.LINEAR, name="lin")), name="prg"), do_estimates),
  (UPat(Ops.PROGRAM, src=(UPat(), UPat(Ops.LINEAR, src=UPat(Ops.INS), name="lin")), name="prg"), do_assemble),
  (UPat(Ops.PROGRAM, src=(UPat(), UPat(Ops.LINEAR, name="lin")), name="prg"), do_render),
  (UPat(Ops.PROGRAM, src=(UPat(), UPat(Ops.LINEAR), UPat(Ops.SOURCE, name="source")), name="prg"), do_compile),
])

@rewrite_group(name=lambda ast,renderer,ret,**_: TracingKey((k:=ret.src[0].arg).name,(k.function_name, ast, ret.key),ret=renderer), replay=True)
@Context(ALLOW_DEVICE_USAGE=0)
def do_to_program(ast:UOp, renderer:Renderer) -> UOp:
  """
  Transform an AST into a compiled PROGRAM. May trigger BEAM search.

  Args:
    ast: The Ops.SINK/Ops.PROGRAM rooted AST
    renderer: The renderer used to generate the code

  Returns:
    The Ops.PROGRAM with SINK/LINEAR/SOURCE/BINARY.
  """
  if ast.op is Ops.PROGRAM: prg = ast
  elif ast.op is Ops.SINK:
    assert isinstance(ast.arg, KernelInfo), "requires KernelInfo on arg to to_program"
    if VIZ: graph_rewrite(ast, PatternMatcher([]), name="View Base AST")
    ast = graph_rewrite(ast, pm_lower_calls, ctx=renderer, name="lower calls", walk=True, enter_calls=True)
    full_sink = full_rewrite_to_sink(ast, renderer, optimize=ast.tag is None)
    prog_info = ProgramInfo.from_sink(full_sink, renderer.target)
    # instruction selection
    if isinstance(renderer, ISARenderer):
      # Instruction selection replaces LOAD/STORE/ALU with INS, so estimate while their meaning is still available.
      if full_sink.arg.estimates is None:
        full_sink = full_sink.replace(arg=replace(full_sink.arg, estimates=Estimates.from_uops(tuple(linearize(full_sink)), ignore_indexing=True)))
      full_sink = graph_rewrite(full_sink, renderer.pre_isel_matcher, ctx=itertools.count(-1, -1), name="pre instruction selection", bottom_up=True)
      full_sink = graph_rewrite(full_sink, renderer.isel_matcher, ctx=IselContext(full_sink), name="instruction selection", bottom_up=True)
    prg = UOp(Ops.PROGRAM, src=(full_sink,), arg=prog_info)
  else: raise RuntimeError(f"can't call to_program on {ast.op}")
  if not isinstance(prg.arg, ProgramInfo): prg = prg.replace(arg=ProgramInfo.from_sink(prg.src[0], renderer.target))
  prg = graph_rewrite(prg, pm_to_program, ctx=renderer, name="linearize/render")
  if VIZ: graph_rewrite(prg, PatternMatcher([]), name="View Program")
  return prg

# config affects generated programs and cache keys; context also carries compile-only behavior to workers
to_program_config = (NOOPT, EMULATED_DTYPES, USE_TC, IMAGE, DISABLE_FAST_IDIV, TRANSCENDENTAL, ALLOW_TF32,
                     DEFAULT_FLOAT, DEFAULT_INT, TC_SELECT, TC_OPT, TC_MIN_GLOBALS)
to_program_context = (*to_program_config, SPEC, DEBUG)
def to_program_key(ast:UOp, renderer:Renderer) -> tuple:
  return (ast.key, type(renderer), renderer.target, *[x.value for x in to_program_config])

to_program_cache: dict[tuple, UOp] = {}
def to_program(ast:UOp, renderer:Renderer) -> UOp:
  if (prg:=to_program_cache.get(key:=to_program_key(ast, renderer))) is None: to_program_cache[key] = prg = do_to_program(ast, renderer)
  return prg
