from __future__ import annotations
from typing import Final, ClassVar, Callable, Literal
import math, struct, ctypes, functools
from dataclasses import dataclass, fields
from tinygrad.helpers import getenv, DEFAULT_FLOAT, DEFAULT_INT
from enum import IntEnum, auto

class ConstFloat(float):
  """Float subclass that distinguishes -0.0 from 0.0 and where nan == nan."""
  __slots__ = ('bits',)
  bits: int
  def __new__(cls, v:float):
    obj = super().__new__(cls, v)
    obj.bits = struct.unpack('<Q', struct.pack('<d', v))[0]
    return obj
  def __eq__(self, other):
    if self is other: return True
    if isinstance(other, float) and math.isnan(self) and math.isnan(other): return True
    return float.__eq__(self, other)
  def __ne__(self, other): return res if (res:=self.__eq__(other)) is NotImplemented else not res  # float.__ne__ disagrees with __eq__ on nan
  def __hash__(self): return hash(self.bits)
  def __repr__(self): return f"ConstFloat({float.__repr__(self)})"
  def __str__(self): return float.__repr__(self)

class InvalidType:
  _instance: ClassVar[InvalidType|None] = None
  def __new__(cls):
    if cls._instance is None: cls._instance = object.__new__(cls)
    return cls._instance
  def __eq__(self, other): return self is other if isinstance(other, InvalidType) else NotImplemented  # foreign types get the reflected eq
  def __hash__(self): return id(self)
  def __repr__(self): return "Invalid"
  def __reduce__(self): return (InvalidType, ())  # unpickle returns the singleton
  def __format__(self, spec): return "Invalid"

Invalid = InvalidType()

PyConst = float|int|bool
ConstType = PyConst|InvalidType

FmtStr = Literal['?', 'b', 'B', 'h', 'H', 'i', 'I', 'q', 'Q', 'e', 'f', 'd']

# all DTypes should only be created once
class DTypeMetaClass(type):
  dcache: dict[tuple, DType] = {}
  def __call__(cls, *args, **kwargs):
    if (ret:=DTypeMetaClass.dcache.get(args, None)) is not None: return ret
    DTypeMetaClass.dcache[args] = ret = super().__call__(*args)
    return ret

class AddrSpace(IntEnum):
  def __str__(self): return repr(self)
  def __repr__(self): return f"{self.__class__.__name__}.{self.name}"
  GLOBAL = auto(); LOCAL = auto(); REG = auto(); ALU = auto()  # noqa: E702

@dataclass(frozen=True, eq=False)
class DType(metaclass=DTypeMetaClass):
  priority: int  # this determines when things get upcasted
  bitsize: int
  name: str
  fmt: FmtStr|None
  @property
  def itemsize(self) -> int: return (self.bitsize + 7) // 8
  @staticmethod
  def new(priority:int, bitsize:int, name:str, fmt:FmtStr|None): return DType(priority, bitsize, name, fmt)
  def __reduce__(self): return type(self), tuple(getattr(self, f.name) for f in fields(self))
  def __repr__(self): return f"dtypes.{self.name}"
  def __lt__(self, o:DType): return (self.priority, self.bitsize, self.name, self.fmt) < (o.priority, o.bitsize, o.name, o.fmt)
  @functools.cached_property
  def min(self):
    if dtypes.is_int(self): return 0 if dtypes.is_unsigned(self) else -2**(self.bitsize-1)
    return -self.max if dtypes.is_float(self) else False
  @functools.cached_property
  def max(self):
    if dtypes.is_int(self): return 2**(self.bitsize)-1+self.min
    # e4m3 and the fnuz fp8s have no inf: their largest value is the largest normal
    if self in dtypes.fp8s and self is not dtypes.fp8e5m2: return fp8_to_float(_fp8_cfg[self][5], self)
    return float("inf") if dtypes.is_float(self) else True
  def const(self, val: ConstType):
    if isinstance(val, InvalidType): return val
    # NOTE: float('nan') != float('nan'), so we canonicalize here
    if isinstance(val, float) and math.isnan(val): val = math.nan
    # int is the default. wrap floats in ConstFloat to distinguish -0.0 from 0.0 in cache
    return ConstFloat(truncate.get(self, float)(float(val))) if dtypes.is_float(self) else bool(val) if dtypes.is_bool(self) else int(val)


class DTypes:
  @staticmethod
  @functools.cache
  def is_float(x: DType) -> bool: return x in (dtypes.floats + (dtypes.weakfloat,))
  @staticmethod # static methods on top, or bool in the type info will refer to dtypes.bool
  @functools.cache
  def is_int(x: DType) -> bool: return x in (dtypes.ints + (dtypes.weakint,))
  @staticmethod
  @functools.cache
  def is_unsigned(x: DType) -> bool: return x in dtypes.uints
  @staticmethod
  def is_bool(x: DType) -> bool: return x == dtypes.bool
  @staticmethod
  def from_py(x) -> DType:
    # NOTE: isinstance(True, int) is True, so bool must be checked before int
    if isinstance(x, (bool, InvalidType)): return dtypes.bool
    if isinstance(x, float): return dtypes.weakfloat
    if isinstance(x, int): return dtypes.weakint
    # put this in the last is faster because there are more items than lists/tuples to check
    if isinstance(x, (list, tuple)):
      dt = max(dtypes.from_py(xi) for xi in x) if x else dtypes.weakfloat
      if dt is not dtypes.weakint: return strong_dtype(dt)
      ints = [xi for xi in x if isinstance(xi, int)]  # a vconst also holds Invalid
      return commit_int(min(ints), max(ints))
    raise RuntimeError(f"Could not infer dtype of {x} with type {type(x)}")
  @staticmethod
  def finfo(dtype:DType) -> tuple[int, int]:
    """(exponent, mantissa)"""
    if not dtypes.is_float(dtype): raise ValueError(f"{dtype} is not a floating point type")
    return {dtypes.float16: (5, 10), dtypes.bfloat16: (8, 7), dtypes.float32: (8, 23), dtypes.float64: (11, 52),
            dtypes.fp8e4m3: (4, 3), dtypes.fp8e5m2: (5, 2), dtypes.fp8e4m3fnuz: (4, 3), dtypes.fp8e5m2fnuz: (5, 2)}[dtype]
  void: Final[DType] = DType.new(-1, 0, "void", None)
  weakint: Final[DType] = DType.new(0, 800, "weakint", None)  # the weak int position in the promo lattice
  bool: Final[DType] = DType.new(0, 1, "bool", '?')
  i8: Final[DType] = DType.new(1, 8, "i8", 'b')
  u8: Final[DType] = DType.new(2, 8, "u8", 'B')
  i16: Final[DType] = DType.new(3, 16, "i16", 'h')
  u16: Final[DType] = DType.new(4, 16, "u16", 'H')
  i32: Final[DType] = DType.new(5, 32, "i32", 'i')
  u32: Final[DType] = DType.new(6, 32, "u32", 'I')
  i64: Final[DType] = DType.new(7, 64, "i64", 'q')
  u64: Final[DType] = DType.new(8, 64, "u64", 'Q')
  weakfloat: Final[DType] = DType.new(9, 800, "weakfloat", None)
  fp8e4m3: Final[DType] = DType.new(10, 8, "fp8e4m3", None)
  fp8e5m2: Final[DType] = DType.new(11, 8, "fp8e5m2", None)
  fp8e4m3fnuz: Final[DType] = DType.new(10, 8, "fp8e4m3fnuz", None)
  fp8e5m2fnuz: Final[DType] = DType.new(11, 8, "fp8e5m2fnuz", None)
  f16: Final[DType] = DType.new(12, 16, "f16", 'e')
  bf16: Final[DType] = DType.new(13, 16, "bf16", None)
  f32: Final[DType] = DType.new(14, 32, "f32", 'f')
  f64: Final[DType] = DType.new(15, 64, "f64", 'd')

  # legacy dtype aliases
  float16 = half = f16; bfloat16 = bf16; float32 = float = f32; float64 = double = f64 # noqa: E702
  uint8 = uchar = u8; uint16 = ushort = u16; uint32 = uint = u32; uint64 = ulong = u64 # noqa: E702
  int8 = char = i8; int16 = short = i16; int32 = int = i32; int64 = long = i64 # noqa: E702

  @property
  def default_float(self) -> DType: return to_dtype(DEFAULT_FLOAT.value)
  @property
  def default_int(self) -> DType: return to_dtype(DEFAULT_INT.value)

  fp8_ocp = (fp8e4m3, fp8e5m2)
  fp8_fnuz = (fp8e4m3fnuz, fp8e5m2fnuz)
  fp8s = fp8_ocp + fp8_fnuz
  floats = fp8s + (float16, bfloat16, float32, float64)
  int8s = (uint8, int8)
  int16s = (uint16, int16)
  int32s = (uint32, int32)
  int64s = (uint64, int64)
  uints = (uint8, uint16, uint32, uint64)
  sints = (int8, int16, int32, int64)
  ints = uints + sints
  weaks = (weakint, weakfloat)
  all = floats + ints + (bool,) # noqa: A003

dtypes = DTypes()

DTypeLike = str|DType
def to_dtype(dtype:DTypeLike) -> DType: return dtype if isinstance(dtype, DType) else getattr(dtypes, dtype.lower())
assert dtypes.is_float(dtypes.default_float), f"{DEFAULT_FLOAT.value} is not a float dtype"
assert dtypes.is_int(dtypes.default_int), f"{DEFAULT_INT.value} is not an int dtype"
def strong_dtype(dtype:DType) -> DType:
  return {dtypes.weakint: dtypes.default_int, dtypes.weakfloat: dtypes.default_float}.get(dtype, dtype)
def commit_int(lo:int|float, hi:int|float, default_int:DType|None=None) -> DType:
  if lo == hi and not dtypes.long.min <= lo <= dtypes.ulong.max: raise OverflowError(f"{lo} does not fit any int")
  ladder = (dtypes.default_int if default_int is None else default_int, dtypes.int, dtypes.long, dtypes.ulong)
  return next((dt for dt in ladder if dt.min <= lo and hi <= dt.max), dtypes.long)
def weak_dtype(dtype:DType) -> DType:
  return dtypes.weakfloat if dtypes.is_float(dtype) else dtypes.weakint if dtypes.is_int(dtype) else dtype

# https://jax.readthedocs.io/en/latest/jep/9407-type-promotion.html
# we don't support complex type
promo_lattice = { dtypes.bool: [dtypes.weakint], dtypes.weakint: [dtypes.int8, dtypes.uint8],
  dtypes.int8: [dtypes.int16], dtypes.int16: [dtypes.int32], dtypes.int32: [dtypes.int64],
  dtypes.int64: [dtypes.weakfloat], dtypes.uint8: [dtypes.int16, dtypes.uint16], dtypes.uint16: [dtypes.int32, dtypes.uint32],
  dtypes.uint32: [dtypes.int64, dtypes.uint64], dtypes.uint64: [dtypes.weakfloat],
  dtypes.weakfloat: [dtypes.fp8e4m3, dtypes.fp8e5m2, dtypes.fp8e4m3fnuz, dtypes.fp8e5m2fnuz],
  dtypes.fp8e4m3: [dtypes.float16, dtypes.bfloat16], dtypes.fp8e5m2: [dtypes.float16, dtypes.bfloat16],
  dtypes.fp8e4m3fnuz: [dtypes.float16, dtypes.bfloat16], dtypes.fp8e5m2fnuz: [dtypes.float16, dtypes.bfloat16],
  dtypes.float16: [dtypes.float32], dtypes.bfloat16: [dtypes.float32], dtypes.float32: [dtypes.float64], }

@functools.cache
def _get_recursive_parents(dtype:DType) -> set[DType]:
  return set.union(*[_get_recursive_parents(d) for d in promo_lattice[dtype]], {dtype}) if dtype != dtypes.float64 else {dtypes.float64}
@functools.cache
def least_upper_dtype(*ds:DType) -> DType:
  return min(set.intersection(*[_get_recursive_parents(d) for d in ds]))
def least_upper_float(dt:DType) -> DType:
  return dtypes.weakfloat if dt is dtypes.weakint else dt if dtypes.is_float(dt) else least_upper_dtype(dt, dtypes.default_float)

DTYPES_DICT = {k: v for k, v in DTypes.__dict__.items() if isinstance(v, DType) and not k.startswith(("void", "weak", "_"))}

@functools.cache
def can_lossless_cast(dt0:DType, dt1:DType) -> bool:
  # return if dt1 preserves value of dt0
  # similar to https://numpy.org/doc/stable/reference/generated/numpy.can_cast.html
  if dt0 == dt1 or dt0 == dtypes.bool: return True
  match dt1:
    case dtypes.weakint: return dt0 in dtypes.ints
    case dtypes.double: return dt0 in (dtypes.float, dtypes.half, dtypes.bfloat16, *dtypes.fp8s,
      dtypes.uint32, dtypes.uint16, dtypes.uint8, dtypes.int32, dtypes.int16, dtypes.int8)
    case dtypes.float: return dt0 in (dtypes.half, dtypes.bfloat16, *dtypes.fp8s, dtypes.uint16, dtypes.uint8, dtypes.int16, dtypes.int8)
    case dtypes.half: return dt0 in (*dtypes.fp8s, dtypes.uint8, dtypes.int8)
    case dtypes.uint64: return dt0 in (dtypes.uint32, dtypes.uint16, dtypes.uint8)
    case dtypes.uint32: return dt0 in (dtypes.uint16, dtypes.uint8)
    case dtypes.uint16: return dt0 in (dtypes.uint8,)
    case dtypes.int64: return dt0 in (dtypes.uint32, dtypes.uint16, dtypes.uint8, dtypes.int32, dtypes.int16, dtypes.int8)
    case dtypes.int32: return dt0 in (dtypes.uint16, dtypes.uint8, dtypes.int16, dtypes.int8)
    case dtypes.int16: return dt0 in (dtypes.uint8, dtypes.int8)
    case _: return False

def sum_acc_dtype(dt:DType):
  # default acc dtype for sum
  if dtypes.is_unsigned(dt): return least_upper_dtype(dt, dtypes.uint)
  if dtypes.is_int(dt) or dt == dtypes.bool: return least_upper_dtype(dt, dtypes.int)
  return least_upper_dtype(dt, to_dtype(getenv("SUM_DTYPE", "float32")))

def float_to_fp16(x):
  try: return struct.unpack('e', struct.pack('e', float(x)))[0]
  except OverflowError: return math.copysign(math.inf, x)

def float_to_bf16(x):
  if not math.isfinite(x): return x
  u = struct.unpack('I', struct.pack('f', truncate[dtypes.float](x)))[0]
  u = (u + 0x7FFF + ((u >> 16) & 1)) & 0xFFFF0000
  return struct.unpack('f', struct.pack('I', u))[0]

# fp8-float conversions based on https://gitlab.com/nvidia/headers/cuda-individual/cudart/-/blob/main/cuda_fp8.hpp
# (bias, sig_bits, mant_mask, min_denorm_half, ovf_threshold, max_norm, min_norm)
_fp8_cfg = {
  dtypes.fp8e4m3: (7, 4, 0x7, 0x3F50000000000000, 0x407D000000000000, 0x7E, 0x3F90000000000000),
  dtypes.fp8e5m2: (15, 3, 0x3, 0x3EE0000000000000, 0x40EE000000000000-1, 0x7B, 0x3F10000000000000),
  dtypes.fp8e4m3fnuz: (8, 4, 0x7, 0x3F40000000000000, 0x406F000000000000-1, 0x7F, 0x3F80000000000000),
  dtypes.fp8e5m2fnuz: (16, 3, 0x3, 0x3ED0000000000000, 0x40EE000000000000-1, 0x7F, 0x3F00000000000000),
}

def float_to_fp8(x: float, dtype: DType) -> int:
  assert dtype in dtypes.fp8s, "Only for fp8s"
  if dtype in dtypes.fp8_fnuz and not math.isfinite(x): return 0x80
  # e4m3 don't support inf, return 0x7f(+NaN) and 0xff(-NaN) to match jax
  # NaN is unordered, can't compare with zero, use math.copysign to get sign
  if dtype == dtypes.fp8e4m3 and not math.isfinite(x): return 0x7f if math.copysign(1, x) > 0 else 0xff
  if dtype == dtypes.fp8e5m2 and not math.isfinite(x): return (0 if math.copysign(1, x) > 0 else 0x80) | (0x7c if math.isinf(x) else 0x7f)
  bias, sig_bits, mant_mask, min_denorm_half, ovf_threshold, max_norm, min_norm = _fp8_cfg[dtype]
  xbits, = struct.unpack('Q', struct.pack('d', x))
  half_ulp = 1 << (52 - sig_bits)
  sign, exp, mantissa, absx = ((xbits>>63)&1)<<7, ((xbits>>52)&0x7FF)-1023+bias, (xbits>>(53-sig_bits))&mant_mask, xbits&0x7FFFFFFFFFFFFFFF
  if absx <= min_denorm_half: res = 0
  elif absx > ovf_threshold: res = max_norm
  elif absx >= min_norm:
    res, round_bits = (exp << (sig_bits - 1)) | mantissa, xbits & ((half_ulp << 1) - 1)
    if round_bits > half_ulp or (round_bits == half_ulp and mantissa & 1): res += 1
  else:
    shift = 1 - exp
    mantissa |= 1 << (sig_bits - 1)
    res, half = mantissa >> shift, half_ulp << shift
    round_bits = (xbits | (1 << 52)) & ((half << 1) - 1)
    if round_bits > half or (round_bits == half and res & 1): res += 1
  return 0 if dtype in dtypes.fp8_fnuz and res == 0 else res | sign  # fnuz has no negative zero

def fp8_to_float(x: int, dtype: DType) -> float:
  assert dtype in dtypes.fp8s, "Only for fp8s"
  if dtype in dtypes.fp8_fnuz and x == 0x80: return math.nan
  if (x & 0x7F) == 0: return -0.0 if x & 0x80 else 0.0
  bias, sig_bits, *_ = _fp8_cfg[dtype]
  mant_bits, exp_bits = sig_bits - 1, 8 - sig_bits
  exp_max, mant_max = (1 << exp_bits) - 1, (1 << mant_bits) - 1
  sign, exp, mantissa = (x >> 7) & 1, (x >> mant_bits) & exp_max, x & mant_max
  if dtype not in dtypes.fp8_fnuz and exp == exp_max:
    if dtype == dtypes.fp8e5m2: return math.copysign(math.nan if mantissa else math.inf, -1 if sign else 1)
    if mantissa == mant_max: return math.nan
  val = (mantissa / (mant_max + 1)) * 2 ** (1 - bias) if exp == 0 else (1 + mantissa / (mant_max + 1)) * 2 ** (exp - bias)
  return -val if sign else val

def storage_fmt_for_dtype(dtype:DType): return 'H' if dtype == dtypes.bfloat16 else 'B' if dtype in dtypes.fp8s else dtype.fmt

def to_storage_scalar(x, dtype:DType):
  if dtype == dtypes.half: return float_to_fp16(x)
  if dtype == dtypes.bfloat16: return (struct.unpack('I', struct.pack('f', float_to_bf16(x)))[0] >> 16) & 0xFFFF
  if dtype in dtypes.fp8s: return float_to_fp8(float(x), dtype)
  return x

def from_storage_scalar(x, dtype:DType):
  if dtype == dtypes.bfloat16: return struct.unpack('f', struct.pack('I', (x & 0xFFFF) << 16))[0]
  if dtype in dtypes.fp8s: return fp8_to_float(int(x), dtype)
  return x

truncate: dict[DType, Callable] = {dtypes.bool: bool,
  dtypes.float16: float_to_fp16, dtypes.bfloat16: lambda x: float_to_bf16(float(x)),
  **{fp8: (lambda x, dtype=fp8: fp8_to_float(float_to_fp8(x, dtype), dtype)) for fp8 in dtypes.fp8s},
  **{getattr(dtypes, n): (lambda x, c=getattr(ctypes, f'c_{n}'): c(x).value)
     for n in ('float', 'double', 'int8', 'int16', 'int32', 'int64', 'uint8', 'uint16', 'uint32', 'uint64')}}

def bitcast(x, in_dtype:DType, out_dtype:DType):
  assert in_dtype.itemsize == out_dtype.itemsize, "bitcast itemsize mismatch"
  packed = struct.pack(storage_fmt_for_dtype(in_dtype), to_storage_scalar(x, in_dtype))
  out_val = struct.unpack(storage_fmt_for_dtype(out_dtype), packed)[0]
  return from_storage_scalar(out_val, out_dtype)

# numpy and torch dtype interop

def _to_np_dtype(dtype:DType) -> type|None:
  import numpy as np
  if dtype in { dtypes.bfloat16, *dtypes.fp8s }: return np.float32
  return np.dtype(dtype.fmt).type if dtype.fmt is not None else None
def _from_np_dtype(npdtype:'np.dtype') -> DType: # type: ignore [name-defined] # noqa: F821
  import numpy as np
  return DTYPES_DICT[np.dtype(npdtype).name]

@functools.cache
def _to_torch_dtype(dtype:DType) -> 'torch.dtype'|None:  # type: ignore [name-defined] # noqa: F821
  import numpy as np, torch
  dtype = strong_dtype(dtype)
  if dtype == dtypes.uint64: return torch.uint64
  if dtype == dtypes.bfloat16: return torch.bfloat16
  if dtype in dtypes.fp8s: return torch.uint8
  # NOTE: torch doesn't expose this mapping with a stable API
  try: return torch.from_numpy(np.array([], dtype=_to_np_dtype(dtype))).dtype
  except TypeError: return None
@functools.cache
def _from_torch_dtype(torchdtype:'torch.dtype') -> DType: # type: ignore [name-defined] # noqa: F821
  return {v:k for k in DTYPES_DICT.values() if (v:=_to_torch_dtype(k)) is not None}[torchdtype]
