import math, time, traceback, signal
from dataclasses import replace
from tinygrad.uop.ops import sym_infer, AxisType, UOp, Ops
from tinygrad.device import Device, Buffer
from tinygrad.dtype import AddrSpace
from tinygrad.helpers import prod, flatten, DEBUG, CACHELEVEL, diskcache_get, diskcache_put, getenv, colored, time_to_str
from tinygrad.helpers import IGNORE_BEAM_CACHE
from tinygrad.codegen.opt import Opt, OptOps, KernelOptError
from tinygrad.engine.realize import time_call
from tinygrad.engine.worker import get_worker_pool, terminate_worker_pool
from tinygrad.codegen import to_program
from tinygrad.codegen.opt.postrange import Scheduler

actions = [Opt(op=OptOps.SPLIT, axis=axis, arg=(amt, AxisType.UPCAST)) for amt in [0,2,3,4,5,7] for axis in range(10)]
actions += [Opt(op=OptOps.SPLIT, axis=axis, arg=(amt, AxisType.LOCAL)) for amt in [0,2,3,4,8,13,16,29] for axis in range(8)]
actions += [Opt(op=OptOps.SPLIT, axis=axis, arg=(amt, AxisType.LOCAL, True)) for amt in [13,16,28,29,32,49,64,256] for axis in range(8)]
if getenv("BEAM_PADTO", 0): actions += [Opt(op=OptOps.PADTO, axis=axis, arg=amt) for amt in [32] for axis in range(7)]
actions += [Opt(op=OptOps.SPLIT, axis=0, arg=(32, AxisType.LOCAL))]
actions += [Opt(op=OptOps.TC, axis=0, arg=(-1, 0, getenv("TC", 1)))]
# covers resnet kernels (3 global * 3 reduce)
actions += [Opt(op=OptOps.TC, axis=axis, arg=(-1, getenv("TC_OPT", 2), getenv("TC", 1))) for axis in range(9)]
actions += [Opt(op=OptOps.SWAP, axis=axis_0, arg=axis_1) for axis_0 in range(5) for axis_1 in range(axis_0+1, 5)]

def get_test_global_size(global_size, max_global_size, var_vals):
  test_global_size = [sym_infer(sz, var_vals) for sz in global_size]
  input_size = prod(test_global_size)
  while prod(test_global_size) > max_global_size:
    for j in range(len(global_size)-1,-1,-1):
      if test_global_size[j] > 16:
        test_global_size[j] //= 2
        break
  return test_global_size, input_size / prod(test_global_size)

def _time_program(prg:UOp, var_vals:dict[str, int], rawbufs:list[Buffer], early_stop:float|None=None,
                  allow_test_size:int=True, max_global_size:int|None=65536, clear_l2=False, cnt=3, name="test", dev_timeout=False) -> list[float]:
  timeout = int(early_stop * 1e3) if dev_timeout and early_stop is not None and early_stop < math.inf else None
  factor = 1
  if allow_test_size and max_global_size is not None:
    global_size, factor = get_test_global_size(prg.arg.global_size, max_global_size, var_vals)
    prg = prg.replace(arg=replace(prg.arg, global_size=tuple(global_size)))
  call = prg.call(*[UOp.from_buffer(b) for b in rawbufs])
  tms, timer = [], time_call(call, var_vals, timeout=timeout, clear_l2=clear_l2)
  for _ in range(cnt):
    try: tms.append(next(timer) * factor)
    except AssertionError: return [math.inf] * cnt
    if early_stop is not None and early_stop < min(tms): break
  return tms

class TimeoutException(Exception): pass
def timeout_handler(signum, frame):
  if DEBUG >= 2: print("*** BEAM COMPILE TIMEOUT")
  raise TimeoutException()

def _try_compile(x:tuple[int,Scheduler]) -> tuple[int, tuple[UOp, float]|None]:
  if hasattr(signal, "alarm"):
    signal.signal(getattr(signal, 'SIGALRM'), timeout_handler)
    # set timeout
    signal.alarm(getenv("BEAM_TIMEOUT_SEC", 10))
  ret = None
  try:
    st = time.perf_counter()
    ast, dev = x[1].copy().get_optimized_ast(name_override="test"), x[1].ren.target.device
    prg = to_program(ast.substitute({p: p.replace(arg=replace(p.arg, device=dev)) for p in ast.toposort()
                                     if p.op is Ops.PARAM and p.addrspace is not AddrSpace.ALU}), x[1].ren)
    et = time.perf_counter() - st
    uops = prg.src[1].src
    if len(uops) >= (uops_max:=getenv("BEAM_UOPS_MAX", 3000)) > 0:
      if getenv("BEAM_LOG_SURPASS_MAX"): print(f"too many uops. {len(uops)=}, {uops_max=}")
      raise RuntimeError("too many uops")
    ret = (prg, et)
  except RuntimeError:
    if DEBUG >= 4: traceback.print_exc()
  except Exception as e:
    if getenv("BEAM_STRICT_MODE"): raise e
  finally:
    if hasattr(signal, "alarm"): signal.alarm(0)
  return x[0], ret

def _ensure_buffer_alloc(bufs:list[Buffer]) -> list[Buffer]: return [buf.ensure_allocated() if buf is not None else buf for buf in bufs]

# *** external API ***

# get dictionary of all possible actions
def get_kernel_actions(s:Scheduler, include_0=True, max_up:int|None=None) -> dict[int, Scheduler]:
  acted, max_up, max_lcl = {0:s} if include_0 else {}, getenv("BEAM_UPCAST_MAX", 256) if max_up is None else max_up, getenv("BEAM_LOCAL_MAX", 1024)
  for i,a in enumerate(actions):
    if a.axis is not None and a.op is not OptOps.TC:
      if (a.axis >= s.shape_len) or (a.op is OptOps.SPLIT and isinstance(arg:=a.arg, tuple) and s.full_shape[a.axis] == arg[0]
                                     and replace(a, arg=(0,)+arg[1:]) in actions): continue
    s2 = s.copy()
    try:
      s2.apply_opt(a)
      up, lcl, tc_up = 1, 1, next((prod(u.arg[0])//u.arg[2] for u in s2.ast.backward_slice if u.op is Ops.WMMA), 1)
      for x,t in zip(s2.full_shape, s2.axis_types):
        if t is AxisType.UPCAST: up *= x
        elif t in (AxisType.WARP, AxisType.LOCAL): lcl *= x
      if up//tc_up > max_up or lcl > max_lcl:
        if getenv("BEAM_LOG_SURPASS_MAX"): print(f"too many upcast/local. {up//tc_up=}, {max_up=}, {lcl=}, {max_lcl=}")
        continue
      acted[i+1] = s2
    except KernelOptError: pass
  return acted

BEAM_DEBUG = getenv("BEAM_DEBUG")
def beam_search(s:Scheduler, rawbufs:list[Buffer], var_vals:dict[str,int], amt:int, allow_test_size=True):
  key = {"ast": s.ast.key, "amt": amt, "allow_test_size": allow_test_size, "device": s.ren.target.device, "suffix": s.ren.suffix}
  if not IGNORE_BEAM_CACHE and CACHELEVEL >= 1 and (val:=diskcache_get("beam_search", key)) is not None:
    ret = s.copy()
    for o in val[len(s.applied_opts):]: ret.apply_opt(o)
    return ret

  beam: list[tuple[Scheduler, float]] = [(s, float("inf"))]
  seen_libs = set()

  pool = get_worker_pool()

  min_progress = getenv("BEAM_MIN_PROGRESS", 0.01)/1e6
  if BEAM_DEBUG:
    print("BEAM_SEARCH:")
    print(s.ast.replace(arg=None).render_uir())
  if DEBUG >= 2: print(f"   0.00s:                from   1 ->   1 actions {s.colored_shape()}")

  try:
    rawbufs = _ensure_buffer_alloc(rawbufs)
    exiting, st = False, time.perf_counter()
    dev = Device[s.ren.target.device]
    while not exiting:
      candidates: list[Scheduler] = flatten([get_kernel_actions(si, include_0=False).values() for si,_ in beam])
      timed: list[tuple[Scheduler, float]] = []
      least_compute_ops = math.inf
      for i, proc in ((map if pool is None else pool.imap_unordered)(_try_compile, enumerate(candidates))):
        if proc is None: continue
        prg, compile_et = proc
        if (lib:=prg.src[3].arg) in seen_libs: continue
        # filter out kernels that use 1000x more compute than the smallest
        estimates = prg.src[0].arg.estimates
        least_compute_ops = min(this_compute_ops:=sym_infer(estimates.ops if estimates is not None else 0, var_vals), least_compute_ops)
        if least_compute_ops*1000 < this_compute_ops:
          if getenv("BEAM_LOG_SURPASS_MAX"): print(f"too much compute. {this_compute_ops} when least is {least_compute_ops}")
          continue
        seen_libs.add(lib)
        try: tms = _time_program(prg, var_vals, rawbufs, early_stop=beam[0][1]*3 if len(beam) else 1.0,
                                 allow_test_size=allow_test_size, clear_l2=hasattr(dev, 'invalidate_caches'),
                                 dev_timeout=getenv("BEAM_DEV_TIMEOUT", 1))
        except Exception as e:
          if BEAM_DEBUG: print(f"BEAM failed for opts: {candidates[i].applied_opts}\n{e}")
          if isinstance(e, RuntimeError): continue
          raise
        timed.append((candidates[i], min(tms)))
        if BEAM_DEBUG > 1:
          print(f"{time.perf_counter() - st:7.2f}s: {i:5d} {len(prg.src[1].src):5d} uops",
                f"{time_to_str(compile_et, w=12)} compile/{time_to_str(timed[-1][1], w=12)} run",
                f"      {len(timed):4d}/{len(candidates):4d}         {timed[-1][0].colored_shape()}")
        elif DEBUG >= 2:
          print(f"\r{time.perf_counter() - st:7.2f}s: {time_to_str(timed[-1][1], w=12)}",
                f"      {len(timed):4d}/{len(candidates):4d}         {timed[-1][0].colored_shape()}\033[K", end="")

      # done
      opts = sorted(timed, key=lambda x: x[1])
      exiting = len(opts) == 0 or (opts[0][1] < min_progress) or (len(beam) > 0 and ((beam[0][1]-opts[0][1]) < min_progress))
      if not exiting: beam = opts[:amt]
      elif len(opts) > 0 and opts[0][1] < beam[0][1]: beam = opts[:1]
      if DEBUG >= 2:
        print(f"\r{time.perf_counter() - st:7.2f}s:", colored(time_to_str(beam[0][1], w=12), "green" if exiting else None),
              f"from {len(candidates):3d} -> {len(opts):3d} actions\033[K", beam[0][0].colored_shape())
  except KeyboardInterrupt as e:
    terminate_worker_pool()
    raise e

  if CACHELEVEL >= 1: diskcache_put("beam_search", key, beam[0][0].applied_opts)
  if BEAM_DEBUG: print(f"BEAM_SEARCH: final tm={time_to_str(beam[0][1], w=0)}, applied_opts={beam[0][0].applied_opts}")
  return beam[0][0]
