import functools, itertools
from tinygrad.helpers import all_int, prod, DEBUG, RING, ALL2ALL, ALLREDUCE_NODE_NDEVS, getenv
from tinygrad.uop.ops import UOp, Ops, ParamArg

# *** allreduce implementation ***
def handle_allreduce(buf:UOp, red:UOp) -> UOp|None:
  if not isinstance(buf.device, tuple): return None
  ndev, shape, numel = len(buf.device), buf.shape, prod(buf.shape)
  op, device = red.arg

  # ring allreduce doesn't provide a benefit with only 2 nodes or where number of elements is less than 256k (empirically)
  # fallback to naive allreduce to save on kernel dispatch, chunking and reassembling chunks.
  concrete = all_int(shape)
  use_all2all = concrete and (ALL2ALL >= 2 or (ndev > 2 and numel > getenv("RING_ALLREDUCE_THRESHOLD", 256_000) and ALL2ALL >= 1))
  use_ring = concrete and not use_all2all and (RING >= 2 or (ndev > 2 and numel > getenv("RING_ALLREDUCE_THRESHOLD", 256_000) and RING >= 1))
  if DEBUG >= 2: print(f"{'ALL2ALL' if use_all2all else 'RING' if use_ring else 'NAIVE'} ALLREDUCE {ndev}x{numel} | {buf.dtype}")

  buf = buf.pad_to(buf.max_shape)
  # contiguous before we copy it
  buf = buf.contiguous()

  if concrete and (hdev:=ALLREDUCE_NODE_NDEVS.value) > 0 and ndev % hdev == 0:
    d, flat, fold = buf.device, buf.reshape((numel,)), functools.partial(functools.reduce, lambda x, y: x.alu(op, y))
    boxes, cs = [range(b, b + hdev) for b in range(0, ndev, hdev)], [(numel * k // hdev, numel * (k + 1) // hdev) for k in range(hdev)]
    owned = {i: fold([flat.mselect(j).shrink((cs[k],)).copy_to_device(d[i]) for j in box]) for box in boxes for k, i in enumerate(box)}
    summed = {i: fold([owned[i], *(owned[j].copy_to_device(d[i]) for j in rank if j != i)]) for rank in zip(*boxes) for i in rank}
    gathered = [UOp.mstack(*(summed[box[k]].copy_to_device(d[j]) for box in boxes for j in box)) for k in range(hdev)]
    return UOp.usum(*[c.pad(((s, numel - e),)) for (s, e), c in zip(cs, gathered)]).reshape(shape)

  # naive: copy to all devices. if you shrink later, that'll be handled
  if not use_ring and not use_all2all:
    return functools.reduce(lambda x,y: x.alu(op, y), [buf.mselect(i).copy_to_device(device) for i in range(ndev)]).shrink_to(shape)

  # chunk data into ndev pieces
  assert isinstance(numel, int)
  factor = next((f for f in [32, 16, 8, 4, 2] if numel % f == 0), 1)
  base, left = divmod(numel // factor,  ndev)
  chunks = list(itertools.pairwise(itertools.accumulate([(base + 1) * factor] * left + [base * factor] * (ndev - left), initial=0)))

  # reduce-scatter
  reduced_chunks:list[UOp] = []
  for i,(s,e) in enumerate(chunks):
    if use_all2all:
      chunks_on_i = [buf.mselect(j).reshape((numel,)).shrink(((s,e),)).copy_to_device(buf.device[i]) for j in range(ndev)]
      reduced_chunks.append(functools.reduce(lambda x,y: x.alu(op, y), chunks_on_i))
    else:
      chunk, reduced = buf.reshape((numel,)).shrink(((s,e),)), buf.reshape((numel,)).shrink(((s,e),))
      for step in range(ndev-1):
        src, dest = (i+step)%ndev, (i+step+1)%ndev
        cp = reduced.copy_to_device(buf.device[dest], src if isinstance(reduced.device, tuple) else None)
        reduced = cp.alu(op, chunk.copy_to_device(buf.device[dest], dest))
      reduced_chunks.append(reduced)

  # allgather
  copied_chunks:list[UOp] = []
  for i,rc in enumerate(reduced_chunks):
    if isinstance(device, str): copied_chunks.append(rc.copy_to_device(device))
    elif use_all2all: copied_chunks.append(UOp.mstack(*(rc.copy_to_device(buf.device[j]) for j in range(ndev))))
    else:
      chain:list[UOp] = [rc]
      for step in range(ndev-1):
        chain.append(rc := rc.copy_to_device(buf.device[(i+step)%ndev]))
      copied_chunks.append(UOp.mstack(*(chain[(j-i+1)%ndev] for j in range(ndev))))

  # reassemble
  return UOp.usum(*[c.pad(((s,numel-e),)) for (s,e),c in zip(chunks, copied_chunks)]).reshape(shape)

def create_allreduce_function(buf:UOp, red:UOp, output:UOp|None=None) -> UOp|None:
  if output is None:
    output = UOp(Ops.ALLOC, src=(UOp.const(red.max_numel()),)+UOp.device_range_src(red.device),
                 arg=ParamArg(next(UOp.unique_num), red.dtype, device=red.device))
    output = output.reshape(red.max_shape).shrink_to(red.shape)
  to = red.param_like(0)
  src = buf.param_like(1)
  red = src.allreduce(*red.arg)
  return output.after(to.after(to.store(handle_allreduce(src, red))).sink().call(output.base, buf.contiguous(), name="allreduce", precompile=True))
