forked from tinygrad/tinygrad
* move shape into arg for param/buffer * no param_from_shape * drop gratuitous syntax changes * image is a in-graph view, folded into the param arg at render; drop dead multi param sharding * view_as helper, simpler resolve_function, spec update * spec: param/buffer are flat storage, no shape input * image dims live in the param arg from transform_to_image; tighten kernel graph spec * kernel graph spec: only RESHAPE/SHRINK over storage values, not all movement * kernel graph: call args are storage, not views (pm_no_view_args); assert in spec * strip views at the kernel graph level (pm_no_views), move into rangeify * touchups
46 lines
1.8 KiB
Python
46 lines
1.8 KiB
Python
from __future__ import annotations
|
|
import functools, pathlib
|
|
from tinygrad import Tensor
|
|
from tinygrad.runtime.support.compiler_amd import HIPCCCompiler
|
|
|
|
FP8_MAX = 448.0
|
|
NUM_WG, THREADS_PER_WG = 1024, 256
|
|
|
|
# per-device abs max without allreduce
|
|
@functools.cache
|
|
def _local_abs_max_fxn(x_p, device):
|
|
x = Tensor(x_p, device=device)
|
|
inner = Tensor(x.uop.src[0]) if x.uop.axis is not None else x # the per-shard view of the flat param
|
|
return (inner.abs().max(),)
|
|
|
|
def local_abs_max(x:Tensor) -> Tensor:
|
|
param = x.as_param(0)
|
|
fxn = _local_abs_max_fxn(param.uop, x.device)
|
|
return Tensor(fxn[0].uop.call(x.uop).gettuple(0))
|
|
|
|
def shard_shape(shape:tuple, axis:int, ndev:int) -> list:
|
|
s = list(shape)
|
|
s[axis] //= ndev
|
|
return s
|
|
|
|
def dname_of(device) -> str:
|
|
if isinstance(device, tuple): return device[0].split(":")[0]
|
|
return device.split(":")[0] if isinstance(device, str) else device
|
|
|
|
def alloc_like(shape, dtype, device, axis=None) -> Tensor:
|
|
if isinstance(device, tuple) and axis is not None:
|
|
return Tensor(Tensor.invalids(*shard_shape(shape, axis, len(device)), dtype=dtype, device=device).uop.unshard(axis), device=device)
|
|
return Tensor.invalids(*shape, dtype=dtype, device=device)
|
|
|
|
def alloc_local(shape, dtype, device, axis=None) -> Tensor:
|
|
if isinstance(device, tuple) and axis is not None:
|
|
return Tensor(Tensor.invalids(*shape, dtype=dtype, device=device).uop.unshard(0), device=device)
|
|
return Tensor.invalids(*shape, dtype=dtype, device=device)
|
|
|
|
def compile_hip(src:str, defines:list[str]):
|
|
return HIPCCCompiler("gfx950", ["-std=c++20", "-ffast-math", *defines]).compile_cached(src)
|
|
|
|
def compile_cpp(cpp_dir:pathlib.Path, cpp_name:str, n_elems:int, hidden:int):
|
|
src = (cpp_dir/cpp_name).read_text()
|
|
return src, compile_hip(src, [f"-DN_ELEMS={n_elems}", f"-DHIDDEN={hidden}", f"-DNUM_WG={NUM_WG}", f"-DTHREADS_PER_WG={THREADS_PER_WG}"])
|