Compare commits

..
Author SHA1 Message Date
geohot ac20b5e984 add Ops.LOOP + conditional Ops.END (kimi) 2026-07-21 10:50:38 -07:00
chenyuandGitHub f64f96ec59 broadcast_axes [PR] (#17112)
prerequisite to simplify broadcasting logic and make it implicit
2026-07-21 11:53:12 -04:00
qazalandGitHub 7b05caf5c5 viz: do not crash on sym_infer err (#17109) 2026-07-21 13:04:49 +09:00
chenyuandGitHub 34bcc5ad63 logcumsumexp mask is bool (#17108) 2026-07-20 22:42:58 -04:00
chenyuandGitHub 40f0d4af14 clean up _broadcasted [PR] (#16974)
no more ptr issue
2026-07-20 22:28:19 -04:00
chenyuandGitHub 2864036e8e correct more spelling of coalesce (#17107) 2026-07-20 21:55:37 -04:00
sirhcmandGitHub f3a5337825 correct spelling of coalesce (#17103) 2026-07-20 21:44:04 -04:00
chenyuandGitHub 95f5c85bf3 some realize and corealize for slow tests (#17106) 2026-07-20 21:43:23 -04:00
George HotzandGitHub 2a81616492 update rules for INDEX mops (#17105) 2026-07-20 18:43:00 -07:00
chenyuandGitHub 13ca9bd8a6 remove dtypes.index again (#17104)
also reverted some dtype change, the split made things needlessly complicated
2026-07-20 20:30:04 -04:00
George HotzandGitHub 636a43722d add END and GROUP to addrspace (#17102) 2026-07-20 17:26:02 -07:00
sirhcmandGitHub 980748ccfc add multiple_of to ParamArg (#17101) 2026-07-20 20:11:54 -04:00
George HotzandGitHub f7ce7f330d llm: minor fixes + tests (#17099)
* llm: minor fixes + tests

* error
2026-07-20 14:31:20 -07:00
chenyuandGitHub 4b8db13e01 rdna int8 wmma (#17098)
nice to fix _wmma_name, also more generic tests
2026-07-20 17:00:50 -04:00
sirhcmandGitHub b1cbd1a43f pytest: use timeout_method signal (#17094) 2026-07-20 15:19:24 -04:00
chenyuandGitHub dbb0f6067e clean up ALU rules in spec.py (#17095) 2026-07-20 15:18:48 -04:00
chenyuandGitHub 8481eba866 allow-unsafe-pr-checkout for szdiff.yml (#17096)
it uses sz.py on master to parse the change, should be safe
2026-07-20 15:08:54 -04:00
nimlgenandGitHub 2b96d64496 hcq2: tiny opts and fixes (#17092) 2026-07-20 18:46:52 +03:00
Pol Puigdemont PlanaandGitHub ef77963cfd derivative of logsumexp is independent of max (#17088)
same as #7009 but for logsumexp and logcumsumexp.
fwd+bwd kernel count 5 -> 3 for both. gradients unchanged
(ties, -inf masks, torch-compared at grad_atol=1e-7).
2026-07-20 06:52:16 -07:00
qazalandGitHub abba2aebda llama: correct fused qkv shape assert (#17086) 2026-07-20 15:51:21 +09:00
qazalandGitHub 1cf8f2f68c llama: inplace amax update (#17064)
* llama: inplace amax update

* remove amax_out return

* work

* fit

* work

* work

* keep

* diff cleanup
2026-07-20 15:05:41 +09:00
chenyuandGitHub ac3f56a1a2 more shift tests (#17083) 2026-07-19 16:05:13 -04:00
chenyuandGitHub 89117d8b9e use real shift in l2i decomp [pr] (#17080)
works for variable shift distace too, also fixed signed arithmetic fill
2026-07-19 13:15:06 -04:00
chenyuandGitHub 9970a0aad0 fix Tensor << Tensor for x86 (#17082)
* fix Tensor << Tensor for x86

* torch
2026-07-19 12:31:05 -04:00
chenyuandGitHub 0146a30125 improve cast to unsign min_max [pr] (#17078) 2026-07-18 21:58:41 -04:00
George HotzandGitHub b53cd35cff llm: make tokenizer fast (kimi) (#17077)
* llm: make tokenizer fast

* simpler

* re.escape + qcom mypy fix
2026-07-18 17:31:59 -07:00
Rick WierengaandGitHub 82debb4557 only allow x86_64 target arch on X86Renderer (#17076) 2026-07-18 19:51:01 -04:00
wozeparrotandGitHub ee290b3e39 optim: mxfp8 zero 1 allgathers in fp8 (#17073) 2026-07-18 07:44:50 -07:00
nimlgenandGitHub 232529ce88 hcq2: simpler sync (#17069)
* x

* y

* n
2026-07-18 16:27:44 +03:00
qazalandGitHub 24d8681be7 viz: better sidebar collapse ux (#17072) 2026-07-18 18:07:58 +09:00
chenyuandGitHub 47629f4bcf more weak dtype materialization raise (#17071) 2026-07-17 23:15:14 -04:00
chenyuandGitHub f315df29a0 no weak Tensor from and to real buffer (#17067)
* no weak Tensor from and to real buffer

creation, assign, safe_save

* is_numpy_ndarray to tensor

* one more
2026-07-17 16:09:10 -04:00
George HotzandGitHub 86a6ad8ed2 llm: split cli.py into serve.py with the HTTP server (#17065)
* llm: split cli.py into serve.py with the HTTP server

* min edit
2026-07-17 10:45:37 -07:00
George HotzandGitHub 3ee2baf71d llm: add tool calling support (kimi) (#17061)
* llm: add tool calling support

* simpler

* cls

* gpt cleanup

* more gpt cleanups

* tests for tools calling
2026-07-17 10:20:01 -07:00
qazalandGitHub 7dd3422c63 llama: replace two stage amax with atomics (#17063)
* atomic amax in c kernels

* quantize fp8 UOp kernel

* diff
2026-07-17 19:27:10 +09:00
wozeparrotandGitHub a836c3822a gptoss: 3d mx block scale (#17062) 2026-07-16 23:30:24 -07:00
sirhcmandGitHub 6f1176ea90 benchmarks: test usbgpu copy speeds on comma (#17060) 2026-07-17 02:02:21 -04:00
George HotzandGitHub 46172bb7c7 llm: add optional jinja template support (kimi) (#17058)
* add jinja template support (kimi)

* fix tests

* lil

* more crap to fallback
2026-07-16 19:02:58 -07:00
chenyuandGitHub 88826a6f35 no weak dtype for randn_like either (#17055) 2026-07-16 18:29:12 -04:00
chenyuandGitHub 3bfd62e915 fix 0 size tolist to match numpy (#17054) 2026-07-16 17:42:52 -04:00
nimlgenandGitHub 709babb97c system: remove sibling functions of PCIDevice (#17052) 2026-07-17 00:15:32 +03:00
George HotzandGitHub d8b83daac6 set tc_upcast_axes to None when done with it (#17053)
* set tc_upcast_axes to None when done with it

* no tag needed
2026-07-16 14:15:21 -07:00
stylishvoidandGitHub c74149c973 avoid repeated parsing and toposort in _valid_priority [PR] (#17049)
* avoid repeated parsing and toposort in _valid_priority

* use backward_slice_with_self instead
2026-07-16 16:24:08 -04:00
chenyuandGitHub 6fa0b2b19e materialize weak dtype casts to default (#17051)
in clone and _buffer
2026-07-16 16:12:33 -04:00
George HotzandGitHub 4d8c3d3fc9 add test_hgemm to test_tiny (#17050)
* add test_hgemm to test_tiny

* dsp skip
2026-07-16 13:12:10 -07:00
George HotzandGitHub 2b1146b3f4 further clean up wmma (#17048)
* further clean up wmma

* comment
2026-07-16 11:43:23 -07:00
chenyuandGitHub f6a92d0a16 sum_acc_dtype(weak) is weak (#17047)
also no explicit weak for rand
2026-07-16 14:32:37 -04:00
George HotzandGitHub 61e104bdfb use UOp.wmma everywhere (#17045)
* use UOp.wmma everywhere

* fix
2026-07-16 10:40:48 -07:00
chenyuandGitHub 5a4156c5d1 bitcast and element_size raise for weak dtypes (#17046) 2026-07-16 13:07:45 -04:00
nimlgenandGitHub 7eb197b1bb nv: always wait for reset (#17043)
* nv: always wait for reset

* x
2026-07-16 16:35:12 +03:00
chenyuandGitHub dba8b6b505 allow weak alu operands (#17044) 2026-07-16 09:33:20 -04:00
nimlgenandGitHub e33e96415f hcq2: tiny cleanupg (#17042) 2026-07-16 16:14:54 +03:00
810d8732f9 fix n^2 in limit_bufs by memoizing reachable loads [PR] (#17017)
* fix n^2 in limit_bufs by memoizing reachable loads [pr]

* Update test_schedule.py

---------

Co-authored-by: Jacob Kitchen <[email protected]>
2026-07-15 23:54:04 -07:00
1c74e044a4 search /usr/lib/wsl/lib first for linux (#17027)
Co-authored-by: George Hotz <[email protected]>
2026-07-15 23:28:17 -07:00
George HotzandGitHub 8b0dd870ce use wmma helper (#17038) 2026-07-15 23:25:17 -07:00
qazalandGitHub 783042d216 viz: graph stays in place when sidebars resize (#17037)
* viz: sidebars can resize independent of main graph

* both sidebars

* fix device-list

* more work

* no variables

* raw 15%

* fix custom view

* minor detail
2026-07-16 11:47:22 +09:00
chenyuandGitHub e8d3047a50 dtype_from_uop cleanup [PR] (#17036) 2026-07-15 21:52:21 -04:00
chenyuandGitHub 6b7fee7d9f minor lower_alu_dtype cleanup [PR] (#17034) 2026-07-15 17:42:23 -04:00
chenyuandGitHub be075b200a weak dtypes in dtype_from_uop [PR] (#17032)
* weak dtypes in dtype_from_uop [PR]

* no weak in spec_program

* weak const fold tests
2026-07-15 16:54:31 -04:00
chenyuandGitHub 3ffb4dc4bc unify lower index in lower_alu_dtype [PR] (#17033)
will work for weak types too
2026-07-15 16:38:42 -04:00
nimlgenandGitHub d6fddb066f usb: keep only custom (#17029)
* usb: keep only custom

* mockgpu by gpt

* gpt said sorry

* revert

* reset

* fix

* flash
2026-07-15 22:40:52 +03:00
wozeparrotandGitHub 0d30f97584 mlperf: make v6.1 dir (#17031) 2026-07-15 10:40:55 -07:00
chenyuandGitHub c23d8188e1 remove _ensure_float [pr] (#17030)
do this cast late. allow `SQRT(int)`
2026-07-15 11:19:20 -04:00
chenyuandGitHub 0d19970edc least_upper_dtype in dtype_from_uop [PR] (#17028) 2026-07-15 09:27:53 -04:00
chenyuandGitHub ebe26420a7 update where Invalid rules [pr] (#17026)
fixed TestInvalidTensor.test_tensor_index
2026-07-15 00:00:30 -04:00
wozeparrotandGitHub 06169f5013 gptoss: small fixes (#17025) 2026-07-14 20:40:23 -07:00
chenyuandGitHub 47ddf94f17 remove InvalidType lt and gt (#17023)
not really used
2026-07-14 21:59:15 -04:00
sirhcmandGitHub c9baa2ef79 use pattern matcher in contiguous_view_offset [PR] (#17022) 2026-07-14 19:37:36 -04:00
nimlgenandGitHub 4257939e50 remove copyin/copyout from Buffer (#17020)
* remove copyin/copyout from Buffer

* x

* x

* x

* x
2026-07-14 19:47:22 +03:00
qazalandGitHub 939f28d571 fused qkv rope custom kernel (#17021)
* work

* fused qkv_norm

* work

* speed

* not that yet

* test cleanup

* just clone

* remove .realize()

* cleanup tests
2026-07-15 01:08:42 +09:00
chenyuandGitHub 82fbca43c5 fix Tensor(np) dtype and support fp8 safetensor (#17019) 2026-07-14 09:31:21 -04:00
chenyuandGitHub 872225e47d update dtype tests for small dtypes (#17016) 2026-07-14 08:07:00 -04:00
qazalandGitHub edfef062ed skip viz.cli -t in null device (#17018) 2026-07-14 19:24:58 +09:00
chenyuandGitHub 55bb251130 add pm_manual_bf16_cast to Metal [pr] (#17015)
mitigate metal compiler bug for
`as_type<half>( (bfloat)(const) )`
2026-07-13 21:53:47 -04:00
sirhcmandGitHub a9fbc7db7b expect _offset support, CL and WEBGPU are outliers (#17014) 2026-07-13 18:53:32 -04:00
chenyuandGitHub 9ce96c2628 fix subnormal in test_dtype (#17013)
* fix subnormal in test_dtype

should fix flaky test/backend/test_dtype.py::TestFp8e4m3::test_casts_from

* better
2026-07-13 18:53:13 -04:00
chenyuandGitHub c898dfe150 remove UOp.contiguous override [PR] (#17012)
also cleaned up max_shard_shape
2026-07-13 16:10:58 -04:00
chenyuandGitHub 681a5e0cfd remove UOp cast and bitcast override [PR] (#17011) 2026-07-13 14:23:11 -04:00
chenyuandGitHub 0410c9325d make test/null follow the SPEC (#17010) 2026-07-13 14:01:41 -04:00
nimlgenandGitHub e4bdc529c4 hcq2 ci (#17008)
* hcq2 ci

* x
2026-07-13 19:29:08 +03:00
nimlgenandGitHub 4536a57f79 hcq rename map (#17009)
* hcq rename map

* x
2026-07-13 19:23:12 +03:00
qazalandGitHub 62ad646d1c llama: gemm/fa backward speedups (gpt 5.6) (#17007)
* fp8 atb gemm speedup

* work

* revert

* fa bw faster
2026-07-14 00:27:00 +09:00
nimlgenandGitHub 4d2becddf8 hcq2: spec=2 (#17006)
* hcq2: spec=2

* hcq: isolate HCQ spec rules

* chq

* move
2026-07-13 18:13:33 +03:00
nimlgenandGitHub ab9dde04a9 amd: do not spam with traps (#17004) 2026-07-13 16:36:01 +03:00
chenyuandGitHub 223c6d74c3 remove unused get_empty_input_data (#17002) 2026-07-12 22:13:30 -04:00
180 changed files with 3081 additions and 2612 deletions
+2
View File
@@ -521,6 +521,8 @@ jobs:
run: BENCHMARK_LOG=usbgpu_openpilot_0_10_1_vision_load_pickle PYTHONPATH="." GMMU=0 DEV=USB+AMD ASSERT_MIN_LOAD_TIME=15 python3 examples/openpilot/load_pickle.py
- name: openpilot run_pickle 0.10.1 driving_vision
run: BENCHMARK_LOG=usbgpu_openpilot_0_10_1_vision_run_pickle RUN_PICKLE=1 PYTHONPATH="." GMMU=0 DEV=USB+AMD ASSERT_MIN_STEP_TIME=50 python3 examples/openpilot/compile3.py
- name: Test copy speeds
run: SIZE=64e6 PYTHONPATH=. GMMU=0 DEV=USB+AMD python3 test/external/external_test_usb_asm24.py TestDevCopySpeeds
driverbenchmarks:
name: PCI Driver Benchmark (DEV=${{ matrix.dev }})
+7 -1
View File
@@ -14,12 +14,15 @@ jobs:
outputs:
branchstat: ${{ steps.brstat.outputs.stat}}
steps:
- name: Check code from PR branch
- name: Check code from PR branch
uses: actions/checkout@v6
with:
repository: ${{ github.event.pull_request.head.repo.full_name }}
ref: ${{ github.event.pull_request.head.sha }}
fetch-depth: 0
# PR code is only inspected with git rev-list, never executed
allow-unsafe-pr-checkout: true
persist-credentials: false
- name: Check whether branch is up-to-date
id: brstat
run: |
@@ -51,6 +54,9 @@ jobs:
repository: ${{ github.event.pull_request.head.repo.full_name }}
ref: ${{ github.event.pull_request.head.sha }}
path: pr
# PR code is only line-counted by master's sz.py, never executed
allow-unsafe-pr-checkout: true
persist-credentials: false
# the base default to tinygrad master and cannot be other fork branch for security purpose
- name: Checkout code from tinygrad master
uses: actions/checkout@v6
+25 -2
View File
@@ -171,7 +171,7 @@ jobs:
llvm: 'true'
amd: 'true'
- name: Run NULL backend tests
run: DEV=NULL python -m pytest -n=auto test/null/ --durations=20
run: SPEC=2 DEV=NULL python -m pytest -n=auto test/null/ --durations=20
- name: Run targeted tests on NULL backend
run: |
DEV=NULL python3 -m unittest test.backend.test_multitensor.TestMultiTensor.test_data_parallel_resnet_train_step
@@ -500,6 +500,29 @@ jobs:
- name: Run LLVM test
run: DEV=MOCKKFD+AMD:LLVM python test/device/test_amd_llvm.py
hcq2:
name: hcq2
runs-on: *linux
timeout-minutes: 5
steps:
- name: Checkout Code
uses: actions/checkout@v6
- name: Setup Environment
uses: ./.github/actions/setup-tinygrad
with:
key: hcq2
deps: testing_unit
amd: 'true'
- name: Run HCQ2 tests
run: HCQ_RUNTIME_DEV=PYTHON HCQ2=1 DEV=MOCKKFD+AMD FORWARD_ONLY=1 PYTHONPATH=. python test/test_tiny.py
- name: Run HCQ2 multi-device tests
run: |
HCQ_RUNTIME_DEV=PYTHON HCQ2=1 DEV=MOCKKFD+AMD FORWARD_ONLY=1 PYTHONPATH=. python test/unit/test_multitensor.py \
TestMultiTensor.test_simple_add TestMultiTensor.test_shard_reduce \
TestMultiTensor.test_backward_sum TestMultiTensor.test_matmul_shard_0_0
- name: Run HCQ2 JIT tests
run: HCQ_RUNTIME_DEV=PYTHON HCQ2=1 DEV=MOCKKFD+AMD FORWARD_ONLY=1 PYTHONPATH=. python test/unit/test_jit.py
testmockam:
name: Linux (am)
runs-on: *linux
@@ -620,7 +643,7 @@ jobs:
- name: Run unit tests
run: DEV=METAL python -m pytest -n=auto test/unit/ --durations=20
- name: Run NULL backend tests
run: DEV=NULL python -m pytest -n=auto test/null/ --durations=20
run: SPEC=2 DEV=NULL python -m pytest -n=auto test/null/ --durations=20
- name: Test tensor core ops (fake)
run: DEV=METAL DEBUG=3 TC=2 python test/backend/test_ops.py TestOps.test_gemm
- name: Test tensor core ops (real)
+5 -2
View File
@@ -1462,6 +1462,8 @@ def train_llama3():
@TinyJit
def minibatch(tokens:Tensor):
for nxt in fp8_next_amax: nxt.assign(0)
for nxt in fp8_next_grad_amax: nxt.assign(0)
if is_dp: tokens = tokens.to(None).shard(device, 0)
if is_mp: tokens = tokens.shard(device)
if not is_sharding: tokens = tokens.to(None)
@@ -1753,11 +1755,12 @@ def train_gptoss():
scheduler = CosineAnnealingLRWithWarmup(optim, opt_base_learning_rate, opt_end_learning_rate, opt_learning_rate_warmup_steps, opt_learning_rate_decay_steps)
# realize everything here
if optim.master_params: Tensor.realize(*optim.master_params)
if optim.master_params:
for m in optim.master_params: m.realize()
Tensor.realize(*optim.params, *fp8_inv_scales)
@TinyJit
@Context(TRAINING=1)
def minibatch(tokens:Tensor):
if is_dp: tokens = tokens.to(None).shard(device, 0)
if not is_sharding: tokens = tokens.to(None)
+71 -68
View File
@@ -37,8 +37,8 @@ def quantize_fp8(x:Tensor, amax_state:Tensor|None=None):
return x_clamped.cast(FP8_DTYPE), scale.float().reciprocal(), new_amax
def matmul(x:Tensor, w:Tensor, fp8:bool=True, amax_x:Tensor|None=None, w_inv_scale:Tensor|None=None,
x_fp8:Tensor|None=None, x_new_amax:Tensor|None=None,
grad_amax_state:Tensor|None=None, next_grad_amax_state:Tensor|None=None, x_prequant_mx:tuple|None=None) -> tuple[Tensor,...]:
x_fp8:Tensor|None=None, grad_amax_state:Tensor|None=None, next_grad_amax_state:Tensor|None=None, x_prequant_mx:tuple|None=None,
next_amax_x:Tensor|None=None) -> tuple[Tensor,...]:
if not fp8:
if ASM_GEMM:
from extra.gemm.cdna_asm_gemm import can_use_asm_gemm, asm_gemm
@@ -56,13 +56,14 @@ def matmul(x:Tensor, w:Tensor, fp8:bool=True, amax_x:Tensor|None=None, w_inv_sca
else:
x_phys = (x_q.cast(dtypes.bfloat16) * _mx_block_scale(x_e8)).reshape(*l_shape, x_q.shape[-1])
out = x_phys @ (w.cast(dtypes.bfloat16) * _mx_block_scale(w_inv_scale)).T
return out, (amax_x.detach() if amax_x is not None else None), x_q
return out, x_q
if x_fp8 is None:
if FUSED_INPUT_QUANTIZE and amax_x is not None:
if FUSED_INPUT_QUANTIZE:
from extra.llama_kernels.quantize_fp8_delayed import quantize_fp8_delayed
x_fp8, _, x_new_amax, _ = quantize_fp8_delayed(x, amax_x, FP8_DTYPE)
x_fp8, _ = quantize_fp8_delayed(x, amax_x, next_amax_x, FP8_DTYPE)
else:
x_fp8, _, x_new_amax = quantize_fp8(x, amax_state=amax_x)
x_fp8, _, new_amax_x = quantize_fp8(x, amax_state=amax_x)
next_amax_x.assign(new_amax_x)
if ASM_GEMM:
from extra.gemm.cdna_asm_gemm import can_use_asm_gemm, asm_gemm
if can_use_asm_gemm(x_fp8, w.T):
@@ -73,51 +74,51 @@ def matmul(x:Tensor, w:Tensor, fp8:bool=True, amax_x:Tensor|None=None, w_inv_sca
else:
out = asm_gemm(x_fp8, w.T, x_scale=amax_x, w_scale=w_inv_scale, grad_amax_state=grad_amax_state,
next_grad_amax_state=next_grad_amax_state)
return out, x_new_amax, x_fp8
return (x_fp8.dot(w.T, dtype=dtypes.float) * ((amax_x.float() + 1e-8) / FP8_MAX) * w_inv_scale).cast(dtypes.bfloat16), x_new_amax, x_fp8
return out, x_fp8
return (x_fp8.dot(w.T, dtype=dtypes.float) * ((amax_x.float() + 1e-8) / FP8_MAX) * w_inv_scale).cast(dtypes.bfloat16), x_fp8
def norm_quantize_matmul(x:Tensor, norm:Tensor, w:Tensor, w_inv_scale:Tensor, eps:float, amax_x:Tensor,
grad_amax_state:Tensor, next_grad_amax_state:Tensor):
next_amax_x:Tensor, grad_amax_state:Tensor, next_grad_amax_state:Tensor):
if FUSED_ADD_NORM_MUL_QUANTIZE:
from extra.llama_kernels.fused_rmsnorm_mul_quantize_fp8 import fused_rmsnorm_mul_quantize_fp8
x_fp8, new_amax, x_normed, rrms = fused_rmsnorm_mul_quantize_fp8(x, norm, amax_x, eps, FP8_DTYPE)
out, *ret = matmul(None, w, w_inv_scale=w_inv_scale, x_fp8=x_fp8, amax_x=amax_x, x_new_amax=new_amax,
x_fp8, x_normed, rrms = fused_rmsnorm_mul_quantize_fp8(x, norm, amax_x, eps, FP8_DTYPE, next_amax_x)
out, *ret = matmul(None, w, w_inv_scale=w_inv_scale, x_fp8=x_fp8, amax_x=amax_x,
grad_amax_state=grad_amax_state, next_grad_amax_state=next_grad_amax_state)
return out, x_normed, rrms, ret
x_normed, rrms = rmsnorm(x, eps)
out, *ret = matmul(x_normed * norm, w, amax_x=amax_x, w_inv_scale=w_inv_scale, grad_amax_state=grad_amax_state,
next_grad_amax_state=next_grad_amax_state)
next_grad_amax_state=next_grad_amax_state, next_amax_x=next_amax_x)
return out, x_normed, rrms, ret
def add_norm_quantize_matmul(x:Tensor, residual:Tensor, norm:Tensor, w:Tensor, w_inv_scale:Tensor, eps:float, amax_x:Tensor,
grad_amax_state:Tensor|None=None, next_grad_amax_state:Tensor|None=None):
next_amax_x:Tensor, grad_amax_state:Tensor|None=None, next_grad_amax_state:Tensor|None=None):
if FUSED_ADD_NORM_MUL_QUANTIZE:
from extra.llama_kernels.fused_rmsnorm_mul_quantize_fp8 import fused_add_rmsnorm_mul_quantize_fp8
x_fp8, new_amax, h, x_normed, rrms = fused_add_rmsnorm_mul_quantize_fp8(x, residual, norm, amax_x, eps, FP8_DTYPE)
out, *ret = matmul(None, w, w_inv_scale=w_inv_scale, x_fp8=x_fp8, amax_x=amax_x, x_new_amax=new_amax,
x_fp8, h, x_normed, rrms = fused_add_rmsnorm_mul_quantize_fp8(x, residual, norm, amax_x, eps, FP8_DTYPE, next_amax_x)
out, *ret = matmul(None, w, w_inv_scale=w_inv_scale, x_fp8=x_fp8, amax_x=amax_x,
grad_amax_state=grad_amax_state, next_grad_amax_state=next_grad_amax_state)
return out, h, x_normed, rrms, ret
h = x + residual
x_normed, rrms = rmsnorm(h, eps)
out, *ret = matmul(x_normed * norm, w, amax_x=amax_x, w_inv_scale=w_inv_scale, grad_amax_state=grad_amax_state,
next_grad_amax_state=next_grad_amax_state)
next_grad_amax_state=next_grad_amax_state, next_amax_x=next_amax_x)
return out, h, x_normed, rrms, ret
def silu_w13_quantize_matmul(x_w13:Tensor, w2:Tensor, s_2:Tensor,
amax_x2:Tensor,
amax_x2:Tensor, next_amax_x2:Tensor,
grad_amax_xw13:Tensor, next_grad_amax_xw13:Tensor,
grad_amax_xout:Tensor, next_grad_amax_xout:Tensor):
if FUSED_SILU_W13:
from extra.llama_kernels.cast_amax import fused_quantize_fp8_w13
x2_fp8, new_amax_x2 = fused_quantize_fp8_w13(x_w13, amax_x2, FP8_DTYPE, grad_amax_state=grad_amax_xw13,
next_grad_amax_state=next_grad_amax_xw13)
out, *ret = matmul(None, w2, w_inv_scale=s_2, x_fp8=x2_fp8, amax_x=amax_x2, x_new_amax=new_amax_x2,
x2_fp8 = fused_quantize_fp8_w13(x_w13, amax_x2, FP8_DTYPE, grad_amax_state=grad_amax_xw13,
next_grad_amax_state=next_grad_amax_xw13, amax_out=next_amax_x2)
out, *ret = matmul(None, w2, w_inv_scale=s_2, x_fp8=x2_fp8, amax_x=amax_x2,
grad_amax_state=grad_amax_xout, next_grad_amax_state=next_grad_amax_xout)
return out, ret
hidden = x_w13.shape[-1] // 2
x_w1, x_w3 = x_w13[..., :hidden], x_w13[..., hidden:]
out, *ret = matmul(x_w1.silu() * x_w3, w2, amax_x=amax_x2, w_inv_scale=s_2, grad_amax_state=grad_amax_xout,
next_grad_amax_state=next_grad_amax_xout)
next_grad_amax_state=next_grad_amax_xout, next_amax_x=next_amax_x2)
return out, ret
class FlatTransformer:
@@ -154,7 +155,7 @@ class FlatTransformer:
self.tok_embeddings = nn.Embedding(vocab_size, dim)
self.tok_embeddings.weight = Tensor.normal(vocab_size, dim, mean=0.0, std=0.02, dtype=dtypes.bfloat16)
self.output = Tensor.normal(1, vocab_size, dim, mean=0.0, std=0.02, dtype=dtypes.bfloat16)
self.freqs_cis = precompute_freqs_cis(dim // n_heads, max_context * 2, rope_theta).contiguous().is_param_(False)
self.freqs_cis = precompute_freqs_cis(dim // n_heads, max_context * 2, rope_theta).clone().is_param_(False)
def _amax(): return Tensor.full((), FP8_MAX, dtype=dtypes.float32).contiguous().is_param_(False)
names = ["xqkv", "xo", "x2"]
@@ -186,89 +187,87 @@ class FlatTransformer:
def attention(self, x:Tensor, freqs_cis:Tensor, *, attention_norm:Tensor, wqkv:Tensor, wo:Tensor,
amax_xqkv:Tensor, amax_xo:Tensor, s_qkv:Tensor, s_o:Tensor,
next_amax_xqkv:Tensor, next_amax_xo:Tensor,
grad_amax_xqkv:Tensor, grad_amax_xo:Tensor, next_grad_amax_xqkv:Tensor, next_grad_amax_xo:Tensor):
bsz, seqlen, _ = x.shape
amaxs, saves = [], []
saves = []
xqkv, x_normed, rrms, (new_amax, *s) = norm_quantize_matmul(x, attention_norm, wqkv, s_qkv, self.norm_eps,
xqkv, x_normed, rrms, s = norm_quantize_matmul(x, attention_norm, wqkv, s_qkv, self.norm_eps,
amax_x=amax_xqkv, grad_amax_state=grad_amax_xqkv,
next_grad_amax_state=next_grad_amax_xqkv)
amaxs.append(new_amax)
next_grad_amax_state=next_grad_amax_xqkv, next_amax_x=next_amax_xqkv)
saves.extend([x_normed, rrms, *s, xqkv])
xqkv = xqkv.reshape(bsz, seqlen, self.n_kv_heads, self.n_rep + 2, self.head_dim)
xq = xqkv[:, :, :, :self.n_rep].reshape(bsz, seqlen, self.n_heads, self.head_dim)
xk = xqkv[:, :, :, self.n_rep].reshape(bsz, seqlen, self.n_kv_heads, self.head_dim)
xv = xqkv[:, :, :, self.n_rep+1].reshape(bsz, seqlen, self.n_kv_heads, self.head_dim)
xq, xk = apply_rotary_emb(xq, xk, freqs_cis)
xq, xk, xv = xq.cast(dtypes.bfloat16), xk.cast(dtypes.bfloat16), xv.cast(dtypes.bfloat16)
if getenv("HK_FLASH_ATTENTION"):
from extra.thunder.amd.fa import flash_attention
from extra.thunder.amd.fa import flash_attention, fused_qkv_rope
xq, xk, xv = fused_qkv_rope(xqkv, freqs_cis, self.n_heads, self.n_kv_heads, self.head_dim)
attn, *save = flash_attention(xq, xk, xv, is_causal=True, write_flat=True)
saves.extend(save)
else:
xqkv = xqkv.reshape(bsz, seqlen, self.n_kv_heads, self.n_rep + 2, self.head_dim)
xq = xqkv[:, :, :, :self.n_rep].reshape(bsz, seqlen, self.n_heads, self.head_dim)
xk = xqkv[:, :, :, self.n_rep].reshape(bsz, seqlen, self.n_kv_heads, self.head_dim)
xv = xqkv[:, :, :, self.n_rep+1].reshape(bsz, seqlen, self.n_kv_heads, self.head_dim)
xq, xk = apply_rotary_emb(xq, xk, freqs_cis)
xq, xk, xv = xq.cast(dtypes.bfloat16), xk.cast(dtypes.bfloat16), xv.cast(dtypes.bfloat16)
xq, xk, xv = xq.transpose(1, 2), xk.transpose(1, 2), xv.transpose(1, 2)
attn = xq.scaled_dot_product_attention(xk, xv, is_causal=True, enable_gqa=True).transpose(1, 2)
attn = attn.reshape(bsz, seqlen, -1)
out, new_amax, *s = matmul(attn, wo, amax_x=amax_xo, w_inv_scale=s_o, grad_amax_state=grad_amax_xo,
next_grad_amax_state=next_grad_amax_xo)
amaxs.append(new_amax)
out, *s = matmul(attn, wo, amax_x=amax_xo, w_inv_scale=s_o, grad_amax_state=grad_amax_xo,
next_grad_amax_state=next_grad_amax_xo, next_amax_x=next_amax_xo)
saves.extend([*s, out])
return out, amaxs, saves
return out, saves
def feed_forward(self, x:Tensor, residual:Tensor, **kwargs):
amaxs, saves = [], []
saves = []
if SPLIT_W13:
h = x + residual
x_normed, rrms = rmsnorm(h, self.norm_eps)
saves.extend([x_normed, rrms])
inp = x_normed * kwargs["ffn_norm"]
x_w1, new_amax, *s = matmul(inp, kwargs["w1"], amax_x=kwargs["amax_x1"], w_inv_scale=kwargs["s_1"],
grad_amax_state=kwargs["grad_amax_xw1"], next_grad_amax_state=kwargs["next_grad_amax_xw1"])
amaxs.append(new_amax)
x_w1, *s = matmul(inp, kwargs["w1"], amax_x=kwargs["amax_x1"], w_inv_scale=kwargs["s_1"],
grad_amax_state=kwargs["grad_amax_xw1"], next_grad_amax_state=kwargs["next_grad_amax_xw1"],
next_amax_x=kwargs["next_amax_x1"])
saves.extend([*s, x_w1])
x_w3, new_amax, *s = matmul(inp, kwargs["w3"], amax_x=kwargs["amax_x3"], w_inv_scale=kwargs["s_3"],
grad_amax_state=kwargs["grad_amax_xw3"], next_grad_amax_state=kwargs["next_grad_amax_xw3"])
amaxs.append(new_amax)
x_w3, *s = matmul(inp, kwargs["w3"], amax_x=kwargs["amax_x3"], w_inv_scale=kwargs["s_3"],
grad_amax_state=kwargs["grad_amax_xw3"], next_grad_amax_state=kwargs["next_grad_amax_xw3"],
next_amax_x=kwargs["next_amax_x3"])
saves.extend([*s, x_w3])
if FUSED_SILU_W13 and MXFP8:
from extra.llama_kernels.fused_silu_mul_quantize_mxfp8 import fused_silu_mul_quantize_mxfp8
aq, ae8, asi = fused_silu_mul_quantize_mxfp8(x_w1.reshape(-1, x_w1.shape[-1]), x_w3.reshape(-1, x_w3.shape[-1]))
out, new_amax, *s = matmul(None, kwargs["w2"], x_prequant_mx=(aq, ae8, asi), amax_x=kwargs["amax_x2"],
w_inv_scale=kwargs["s_2"], grad_amax_state=kwargs["grad_amax_xout"],
next_grad_amax_state=kwargs["next_grad_amax_xout"])
out, *s = matmul(None, kwargs["w2"], x_prequant_mx=(aq, ae8, asi), amax_x=kwargs["amax_x2"],
w_inv_scale=kwargs["s_2"], grad_amax_state=kwargs["grad_amax_xout"],
next_grad_amax_state=kwargs["next_grad_amax_xout"], next_amax_x=kwargs["next_amax_x2"])
out = out.reshape(*x_w1.shape[:-1], kwargs["w2"].shape[0])
else:
out, new_amax, *s = matmul(x_w1.silu() * x_w3, kwargs["w2"], amax_x=kwargs["amax_x2"], w_inv_scale=kwargs["s_2"],
grad_amax_state=kwargs["grad_amax_xout"], next_grad_amax_state=kwargs["next_grad_amax_xout"])
amaxs.append(new_amax)
out, *s = matmul(x_w1.silu() * x_w3, kwargs["w2"], amax_x=kwargs["amax_x2"], w_inv_scale=kwargs["s_2"],
grad_amax_state=kwargs["grad_amax_xout"], next_grad_amax_state=kwargs["next_grad_amax_xout"],
next_amax_x=kwargs["next_amax_x2"])
saves.extend([*s, out])
else:
x_w13, h, x_normed, rrms, (new_amax, *s) = add_norm_quantize_matmul(x, residual, kwargs["ffn_norm"], kwargs["w13"], kwargs["s_13"],
x_w13, h, x_normed, rrms, s = add_norm_quantize_matmul(x, residual, kwargs["ffn_norm"], kwargs["w13"], kwargs["s_13"],
self.norm_eps, amax_x=kwargs["amax_x13"],
next_amax_x=kwargs["next_amax_x13"],
grad_amax_state=kwargs["grad_amax_xw13"],
next_grad_amax_state=kwargs["next_grad_amax_xw13"])
amaxs.append(new_amax)
saves.extend([x_normed, rrms, *s, x_w13])
out, (new_amax, *s) = silu_w13_quantize_matmul(x_w13, kwargs["w2"], kwargs["s_2"], amax_x2=kwargs["amax_x2"],
out, s = silu_w13_quantize_matmul(x_w13, kwargs["w2"], kwargs["s_2"], amax_x2=kwargs["amax_x2"],
next_amax_x2=kwargs["next_amax_x2"],
grad_amax_xw13=kwargs["grad_amax_xw13"],
next_grad_amax_xw13=kwargs["next_grad_amax_xw13"],
grad_amax_xout=kwargs["grad_amax_xout"],
next_grad_amax_xout=kwargs["next_grad_amax_xout"])
amaxs.append(new_amax)
saves.extend([*s, out])
return out, h, amaxs, saves
return out, h, saves
@function(precompile=True, precompile_backward=True)
def run_layer(self, x:Tensor, freqs_cis:Tensor, attn_kwargs:dict, ffn_kwargs:dict, save:bool=True):
attn, attn_amaxs, attn_saves = self.attention(x, freqs_cis, **attn_kwargs)
ffn, h, ffn_amaxs, ffn_saves = self.feed_forward(x, attn, **ffn_kwargs)
attn, attn_saves = self.attention(x, freqs_cis, **attn_kwargs)
ffn, h, ffn_saves = self.feed_forward(x, attn, **ffn_kwargs)
h = h + ffn
amaxs = tuple(a.detach() for a in (*attn_amaxs, *ffn_amaxs))
if save: return (h, *amaxs, *attn_saves, *ffn_saves)
else: return (h, *amaxs)
if save: return (h, *attn_saves, *ffn_saves)
else: return (h,)
def shard(self, device:tuple[str, ...], mp:bool=False):
from tinygrad.nn.state import get_parameters
@@ -313,26 +312,27 @@ class FlatTransformer:
def __call__(self, tokens:Tensor, save:bool=True):
h = self.tok_embeddings(tokens)
freqs_cis = self.freqs_cis.cast(h.dtype)[:, :tokens.shape[1], :, :, :]
freqs_cis = self.freqs_cis.cast(h.dtype)
if not getenv("HK_FLASH_ATTENTION"): freqs_cis = freqs_cis[:, :tokens.shape[1], :, :, :]
a, na, ga, nga, s = self._fp8_amax, self._fp8_next_amax, self._fp8_grad_amax, self._fp8_next_grad_amax, self._fp8_inv_scale
for i in range(self.n_layers):
attn_kwargs = dict(attention_norm=self.attention_norm[i], wqkv=self.wqkv[i], wo=self.wo[i],
amax_xqkv=a["xqkv"][i], amax_xo=a["xo"][i], s_qkv=s["wqkv"][i], s_o=s["wo"][i],
next_amax_xqkv=na["xqkv"][i], next_amax_xo=na["xo"][i],
grad_amax_xqkv=ga["xqkv"][i], grad_amax_xo=ga["xo"][i],
next_grad_amax_xqkv=nga["xqkv"][i], next_grad_amax_xo=nga["xo"][i])
ffn_kwargs = dict(ffn_norm=self.ffn_norm[i], w2=self.w2[i],
amax_x2=a["x2"][i], s_2=s["w2"][i], grad_amax_xout=ga["xout"][i], next_grad_amax_xout=nga["xout"][i])
amax_x2=a["x2"][i], s_2=s["w2"][i], grad_amax_xout=ga["xout"][i], next_grad_amax_xout=nga["xout"][i],
next_amax_x2=na["x2"][i])
if SPLIT_W13:
ffn_kwargs.update(w1=self.w1[i], w3=self.w3[i], amax_x1=a["x1"][i], amax_x3=a["x3"][i],
next_amax_x1=na["x1"][i], next_amax_x3=na["x3"][i],
s_1=s["w1"][i], s_3=s["w3"][i], grad_amax_xw1=ga["xw1"][i], grad_amax_xw3=ga["xw3"][i],
next_grad_amax_xw1=nga["xw1"][i], next_grad_amax_xw3=nga["xw3"][i])
else:
ffn_kwargs.update(w13=self.w13[i], amax_x13=a["x13"][i], s_13=s["w13"][i], grad_amax_xw13=ga["xw13"][i],
next_grad_amax_xw13=nga["xw13"][i])
h, *ret = self.run_layer(h, freqs_cis, attn_kwargs, ffn_kwargs, save=save)
amax_names = ["xqkv", "xo"] + (["x1", "x3"] if SPLIT_W13 else ["x13"]) + ["x2"]
for name, new_val in zip(amax_names, ret[:len(amax_names)]):
na[name][i].assign(new_val)
next_grad_amax_xw13=nga["xw13"][i], next_amax_x13=na["x13"][i])
h, *_ = self.run_layer(h, freqs_cis, attn_kwargs, ffn_kwargs, save=save)
logits = matmul(self.norm(h), self.output[0], fp8=False)[0]
return logits
@@ -415,6 +415,9 @@ if __name__ == "__main__":
@TinyJit
def fwd_bwd(tokens:Tensor):
with Timing("python forward: "):
for amax_dict in (model._fp8_next_amax, model._fp8_next_grad_amax):
for ts in amax_dict.values():
for nxt in ts: nxt.assign(0)
logits = model(tokens[:, :-1], save=llama_size=="8B")
loss = vocab_mask.where(-1e9, logits).sparse_categorical_crossentropy(tokens[:, 1:])
with Timing("python backward: "):
+6 -3
View File
@@ -12,7 +12,7 @@ from tinygrad.helpers import Timing, colored, GlobalCounters, profile_marker
from tinygrad.uop.ops import Ops, UOp
from extra.models.llama import apply_rotary_emb, precompute_freqs_cis
from extra.llama_kernels.rmsnorm import rmsnorm
from extra.gemm.cdna_asm_gemm import _mx_block_scale, quantize_mxfp8
from extra.gemm.cdna_asm_gemm import _mx_block_scale, _mx_block_scale_3d, quantize_mxfp8
FP8_DTYPE = dtypes.fp8e4m3
FP8_MAX = 448.0
@@ -39,8 +39,11 @@ def quant_dequant_mx(x:Tensor) -> Tensor:
fxn = _quant_dequant_fwd_fxn(x.as_param(0).uop, x.device)
return Tensor(UOp.maketuple(fxn.uop).call(x.uop, grad_fxn=_quant_dequant_bwd).gettuple(0))
def _mx_scale(e8:Tensor) -> Tensor:
return _mx_block_scale(e8) if e8.ndim == 2 else _mx_block_scale_3d(e8)
def _dequant_fwd(w_q:Tensor, w_scale:Tensor) -> Tensor:
return w_q.cast(dtypes.bfloat16) * _mx_block_scale(w_scale)
return w_q.cast(dtypes.bfloat16) * _mx_scale(w_scale)
@functools.cache
def _dequant_fwd_fxn(wq_p, ws_p, device):
@@ -48,7 +51,7 @@ def _dequant_fwd_fxn(wq_p, ws_p, device):
def _dequant_bwd(grad:UOp, call:UOp) -> tuple:
w_scale = Tensor(call.src[2])
return ((Tensor(grad).cast(dtypes.bfloat16) * _mx_block_scale(w_scale).cast(dtypes.bfloat16)).uop, None)
return ((Tensor(grad).cast(dtypes.bfloat16) * _mx_scale(w_scale).cast(dtypes.bfloat16)).uop, None)
def dequant_weight(w_q:Tensor, w_scale:Tensor) -> Tensor:
fxn = _dequant_fwd_fxn(w_q.as_param(0).uop, w_scale.as_param(1).uop, w_q.device)
+2 -1
View File
@@ -96,7 +96,7 @@ class GradAccClipAdamW(Optimizer):
up = up.float().shard_like(w) + self.lr.to(w.device) * wd * w.detach()
new_w = w.detach() - up
if master is not None: master.assign(new_w)
if self.zero: new_w = self._zero_gather(new_w)
if self.zero and not (MXFP8 and t.dtype in dtypes.fp8s): new_w = self._zero_gather(new_w)
# when master is offloaded to a different device than the param, results are resharded back onto the param's (sharded) device
offloaded = master is not None and master.device != t.device
if STOCHASTIC_ROUND and t.dtype == dtypes.bfloat16:
@@ -106,6 +106,7 @@ class GradAccClipAdamW(Optimizer):
if MXFP8:
from extra.gemm.cdna_asm_gemm import quantize_mxfp8
w_q, w_e8, _ = quantize_mxfp8(new_w.reshape(-1, new_w.shape[-1]))
if self.zero: w_q, w_e8 = self._zero_gather(w_q), self._zero_gather(w_e8)
new_e8 = w_e8.reshape(t._inv_scale.shape)
t._inv_scale.assign(new_e8.shard_like(t._inv_scale) if offloaded else new_e8)
ret = w_q.reshape(new_w.shape)
@@ -20,7 +20,7 @@ export ZERO_OPTIM=${ZERO_OPTIM:-1}
export OFFLOAD_OPTIM=${OFFLOAD_OPTIM:-0}
export DEFAULT_FLOAT="bfloat16" OPTIM_DTYPE="bfloat16"
export DP=${DP:-8} BS=${BS:-16} EVAL_BS=${EVAL_BS:-8} GRADIENT_ACC_STEPS=${GRADIENT_ACC_STEPS:-2}
export DP=${DP:-8} BS=${BS:-16} EVAL_BS=${EVAL_BS:-8} GRADIENT_ACC_STEPS=${GRADIENT_ACC_STEPS:-1}
export GBS=$((BS * GRADIENT_ACC_STEPS))
export MODEL="gptoss"
@@ -34,7 +34,7 @@ export SEED=${SEED:-5760}
export DATA_SEED=${DATA_SEED:-5760}
export JITBEAM=${JITBEAM:-3}
export BEAM_UOPS_MAX=6000 BEAM_UPCAST_MAX=256 BEAM_LOCAL_MAX=1024 BEAM_MIN_PROGRESS=5 BEAM_PADTO=1
export BEAM_UOPS_MAX=6000 BEAM_UPCAST_MAX=256 BEAM_LOCAL_MAX=1024 BEAM_MIN_PROGRESS=5 BEAM_PADTO=0
export FAKEDATA=${FAKEDATA:-1} BENCHMARK=${BENCHMARK:-10}
if [ -z "$FULL_LAYERS" ]; then
@@ -20,7 +20,7 @@ export ZERO_OPTIM=${ZERO_OPTIM:-1}
export OFFLOAD_OPTIM=${OFFLOAD_OPTIM:-0}
export DEFAULT_FLOAT="bfloat16" OPTIM_DTYPE="bfloat16"
export DP=${DP:-8} BS=${BS:-16} EVAL_BS=${EVAL_BS:-8} GRADIENT_ACC_STEPS=${GRADIENT_ACC_STEPS:-2}
export DP=${DP:-8} BS=${BS:-16} EVAL_BS=${EVAL_BS:-8} GRADIENT_ACC_STEPS=${GRADIENT_ACC_STEPS:-1}
export GBS=$((BS * GRADIENT_ACC_STEPS))
export MODEL="gptoss"
@@ -34,6 +34,6 @@ export SEED=${SEED:-$RANDOM}
export DATA_SEED=${DATA_SEED:-5760}
export JITBEAM=${JITBEAM:-3}
export BEAM_UOPS_MAX=6000 BEAM_UPCAST_MAX=256 BEAM_LOCAL_MAX=1024 BEAM_MIN_PROGRESS=5 BEAM_PADTO=1
export BEAM_UOPS_MAX=6000 BEAM_UPCAST_MAX=256 BEAM_LOCAL_MAX=1024 BEAM_MIN_PROGRESS=5 BEAM_PADTO=0
python3 examples/mlperf/model_train.py
@@ -2,5 +2,4 @@
export BENCHMARK=${BENCHMARK:-5}
export EVAL_BS=0
VIZ=${VIZ:--1} FULL_LAYERS=1 DEBUG=${DEBUG:--0} examples/mlperf/training_submission_v6.0/tinycorp/benchmarks/llama31_8b/implementations/tinybox_8xMI350X/dev_beam.sh
SRC="AMD"; [[ $DEV == NULL* ]] && SRC="NULL"
[ "$BENCHMARK" -le 3 ] || python -m tinygrad.viz.cli -s "$SRC" -t --interval "train @ 2" "train @ 3"
[ "$BENCHMARK" -le 3 ] || [[ $DEV == NULL* ]] || python -m tinygrad.viz.cli -s AMD -t --interval "train @ 2" "train @ 3"
@@ -0,0 +1,44 @@
#!/usr/bin/env bash
export PYTHONPATH="."
export PATH="/opt/rocm-7.1.1/bin:$PATH"
export ROCM_PATH="/opt/rocm-7.1.1"
export DEV=${DEV:-AMD}
export CHECK_OOB=0
export REWRITE_STACK_LIMIT=5000000 HCQDEV_WAIT_TIMEOUT_MS=240000
export DEVICE_IN_FUNCTION_BUG=1
export DEBUG=${DEBUG:-2}
export HK_FLASH_ATTENTION=${HK_FLASH_ATTENTION:-1}
export ALL2ALL=${ALL2ALL:-1}
export LATE_ALLREDUCE=${LATE_ALLREDUCE:-0}
export ALLREDUCE_CAST=${ALLREDUCE_CAST:-1}
export USE_ATOMICS=${USE_ATOMICS:-1}
export MASTER_WEIGHTS=${MASTER_WEIGHTS:-1}
export MXFP8=${MXFP8:-1}
export ZERO_OPTIM=${ZERO_OPTIM:-1}
export OFFLOAD_OPTIM=${OFFLOAD_OPTIM:-0}
export DEFAULT_FLOAT="bfloat16" OPTIM_DTYPE="bfloat16"
export DP=${DP:-8} BS=${BS:-16} EVAL_BS=${EVAL_BS:-8} GRADIENT_ACC_STEPS=${GRADIENT_ACC_STEPS:-1}
export GBS=$((BS * GRADIENT_ACC_STEPS))
export MODEL="gptoss"
export BASEDIR="/raid/datasets/c4-8b/"
export EVAL_TARGET=3.34 EVAL_FREQ=12288
export END_LR="4e-5" WARMUP_STEPS=128 MAX_STEPS=1200000
export SAMPLES=$((MAX_STEPS * GBS))
export SEQLEN=${SEQLEN:-8192}
export SEED=${SEED:-5760}
export DATA_SEED=${DATA_SEED:-5760}
export JITBEAM=${JITBEAM:-3}
export BEAM_UOPS_MAX=6000 BEAM_UPCAST_MAX=256 BEAM_LOCAL_MAX=1024 BEAM_MIN_PROGRESS=5 BEAM_PADTO=0
export FAKEDATA=${FAKEDATA:-1} BENCHMARK=${BENCHMARK:-10}
if [ -z "$FULL_LAYERS" ]; then
export LAYERS=${LAYERS:-2}
fi
python3 examples/mlperf/model_train.py
@@ -0,0 +1,39 @@
#!/usr/bin/env bash
export PYTHONPATH="."
export PATH="/opt/rocm-7.1.1/bin:$PATH"
export ROCM_PATH="/opt/rocm-7.1.1"
export DEV=${DEV:-AMD}
export CHECK_OOB=0
export REWRITE_STACK_LIMIT=5000000 HCQDEV_WAIT_TIMEOUT_MS=240000
export DEVICE_IN_FUNCTION_BUG=1
export DEBUG=${DEBUG:-0}
export HK_FLASH_ATTENTION=${HK_FLASH_ATTENTION:-1}
export ALL2ALL=${ALL2ALL:-1}
export LATE_ALLREDUCE=${LATE_ALLREDUCE:-0}
export ALLREDUCE_CAST=${ALLREDUCE_CAST:-1}
export USE_ATOMICS=${USE_ATOMICS:-1}
export MASTER_WEIGHTS=${MASTER_WEIGHTS:-1}
export MXFP8=${MXFP8:-1}
export ZERO_OPTIM=${ZERO_OPTIM:-1}
export OFFLOAD_OPTIM=${OFFLOAD_OPTIM:-0}
export DEFAULT_FLOAT="bfloat16" OPTIM_DTYPE="bfloat16"
export DP=${DP:-8} BS=${BS:-16} EVAL_BS=${EVAL_BS:-8} GRADIENT_ACC_STEPS=${GRADIENT_ACC_STEPS:-1}
export GBS=$((BS * GRADIENT_ACC_STEPS))
export MODEL="gptoss"
export BASEDIR="/raid/datasets/c4-8b/"
export EVAL_TARGET=3.34 EVAL_FREQ=12288
export END_LR="4e-5" WARMUP_STEPS=128 MAX_STEPS=1200000
export SAMPLES=$((MAX_STEPS * GBS))
export SEQLEN=${SEQLEN:-8192}
export SEED=${SEED:-$RANDOM}
export DATA_SEED=${DATA_SEED:-5760}
export JITBEAM=${JITBEAM:-3}
export BEAM_UOPS_MAX=6000 BEAM_UPCAST_MAX=256 BEAM_LOCAL_MAX=1024 BEAM_MIN_PROGRESS=5 BEAM_PADTO=0
python3 examples/mlperf/model_train.py
@@ -0,0 +1,54 @@
#!/usr/bin/env bash
export PYTHONPATH="."
export PATH="/opt/rocm-7.1.1/bin:$PATH"
export ROCM_PATH="/opt/rocm-7.1.1"
export DEV=${DEV:-AMD}
export CHECK_OOB=0
export REWRITE_STACK_LIMIT=5000000 HCQDEV_WAIT_TIMEOUT_MS=240000
export DEVICE_IN_FUNCTION_BUG=1
export DEBUG=${DEBUG:-2}
export HK_FLASH_ATTENTION=${HK_FLASH_ATTENTION:-1}
export ALL2ALL=${ALL2ALL:-1}
export LATE_ALLREDUCE=${LATE_ALLREDUCE:-0}
export USE_ATOMICS=${USE_ATOMICS:-1}
export ASM_GEMM=${ASM_GEMM:-1}
export WQKV=${WQKV:-1}
export MASTER_WEIGHTS=${MASTER_WEIGHTS:-1}
export FP8=${FP8:-1}
export ALLREDUCE_CAST=${ALLREDUCE_CAST:-1}
export FAST_CE=${FAST_CE:-1}
export FUSED_INPUT_QUANTIZE=${FUSED_INPUT_QUANTIZE:-1}
export FUSED_GRAD_QUANTIZE=${FUSED_GRAD_QUANTIZE:-1}
export FUSED_ADD_NORM_MUL_QUANTIZE=${FUSED_ADD_NORM_MUL_QUANTIZE:-1}
export FUSED_SILU_W13=${FUSED_SILU_W13:-1}
export SPLIT_W13=${SPLIT_W13:-0}
export OFFLOAD_OPTIM=${OFFLOAD_OPTIM:-0}
export DEFAULT_FLOAT="bfloat16" OPTIM_DTYPE="bfloat16"
export DP=${DP:-8} MP=${MP:-1} BS=${BS:-16} EVAL_BS=${EVAL_BS:-8} GRADIENT_ACC_STEPS=${GRADIENT_ACC_STEPS:-2}
export GBS=$((BS * GRADIENT_ACC_STEPS))
export MODEL="llama3"
export BASEDIR="/raid/datasets/c4-8b/"
export SMALL=1
export LLAMA3_SIZE=${LLAMA3_SIZE:-"8B"}
export EVAL_TARGET=3.3 EVAL_FREQ=12288
export LR="1e-3" END_LR="1e-4" WARMUP_SAMPLES=4096 MAX_STEPS=1200000
export WARMUP_STEPS=$((WARMUP_SAMPLES / GBS))
export SAMPLES=$((MAX_STEPS * GBS))
export SEQLEN=${SEQLEN:-8192}
export SEED=${SEED:-5760}
export DATA_SEED=${DATA_SEED:-5760}
export JITBEAM=${JITBEAM:-3}
export BEAM_UOPS_MAX=6000 BEAM_UPCAST_MAX=256 BEAM_LOCAL_MAX=1024 BEAM_MIN_PROGRESS=5 BEAM_PADTO=1
export FAKEDATA=${FAKEDATA:-1} BENCHMARK=${BENCHMARK:-10}
if [ -z "$FULL_LAYERS" ]; then
export LLAMA_LAYERS=${LLAMA_LAYERS:-2}
fi
python3 examples/mlperf/model_train.py
@@ -0,0 +1,54 @@
#!/usr/bin/env bash
export PYTHONPATH="."
export PATH="/opt/rocm-7.1.1/bin:$PATH"
export ROCM_PATH="/opt/rocm-7.1.1"
export DEV=${DEV:-AMD}
export CHECK_OOB=0
export REWRITE_STACK_LIMIT=5000000 HCQDEV_WAIT_TIMEOUT_MS=240000
export DEVICE_IN_FUNCTION_BUG=1
export DEBUG=${DEBUG:-2}
export HK_FLASH_ATTENTION=${HK_FLASH_ATTENTION:-1}
export ALL2ALL=${ALL2ALL:-1}
export LATE_ALLREDUCE=${LATE_ALLREDUCE:-1}
export USE_ATOMICS=${USE_ATOMICS:-1}
export ASM_GEMM=${ASM_GEMM:-1}
export WQKV=${WQKV:-1}
export MASTER_WEIGHTS=${MASTER_WEIGHTS:-1}
export FP8=${FP8:-1}
export ALLREDUCE_CAST=${ALLREDUCE_CAST:-1}
export FAST_CE=${FAST_CE:-0}
export FUSED_INPUT_QUANTIZE=${FUSED_INPUT_QUANTIZE:-0}
export FUSED_GRAD_QUANTIZE=${FUSED_GRAD_QUANTIZE:-0}
export FUSED_ADD_NORM_MUL_QUANTIZE=${FUSED_ADD_NORM_MUL_QUANTIZE:-0}
export FUSED_SILU_W13=${FUSED_SILU_W13:-0}
export SPLIT_W13=${SPLIT_W13:-1}
export OFFLOAD_OPTIM=${OFFLOAD_OPTIM:-1}
export DEFAULT_FLOAT="bfloat16" OPTIM_DTYPE="bfloat16"
export DP=${DP:-1} MP=${MP:-8} BS=${BS:-1} EVAL_BS=${EVAL_BS:-1} GRADIENT_ACC_STEPS=${GRADIENT_ACC_STEPS:-2}
export GBS=$((BS * GRADIENT_ACC_STEPS))
export MODEL="llama3"
export BASEDIR="/raid/datasets/c4-8b/"
export SMALL=1
export LLAMA3_SIZE=${LLAMA3_SIZE:-"8B"}
export EVAL_TARGET=3.3 EVAL_FREQ=12288
export LR="1e-3" END_LR="1e-4" WARMUP_SAMPLES=4096 MAX_STEPS=1200000
export WARMUP_STEPS=$((WARMUP_SAMPLES / GBS))
export SAMPLES=$((MAX_STEPS * GBS))
export SEQLEN=${SEQLEN:-8192}
export SEED=${SEED:-5760}
export DATA_SEED=${DATA_SEED:-5760}
export JITBEAM=${JITBEAM:-3}
export BEAM_UOPS_MAX=6000 BEAM_UPCAST_MAX=256 BEAM_LOCAL_MAX=1024 BEAM_MIN_PROGRESS=5 BEAM_PADTO=1
export FAKEDATA=${FAKEDATA:-1} BENCHMARK=${BENCHMARK:-10}
if [ -z "$FULL_LAYERS" ]; then
export LLAMA_LAYERS=${LLAMA_LAYERS:-2}
fi
python3 examples/mlperf/model_train.py
@@ -0,0 +1,49 @@
#!/usr/bin/env bash
export PYTHONPATH="."
export PATH="/opt/rocm-7.1.1/bin:$PATH"
export ROCM_PATH="/opt/rocm-7.1.1"
export DEV=${DEV:-AMD}
export CHECK_OOB=0
export REWRITE_STACK_LIMIT=5000000 HCQDEV_WAIT_TIMEOUT_MS=240000
export DEVICE_IN_FUNCTION_BUG=1
export DEBUG=${DEBUG:-0}
export HK_FLASH_ATTENTION=${HK_FLASH_ATTENTION:-1}
export ALL2ALL=${ALL2ALL:-1}
export LATE_ALLREDUCE=${LATE_ALLREDUCE:-0}
export USE_ATOMICS=${USE_ATOMICS:-1}
export ASM_GEMM=${ASM_GEMM:-1}
export WQKV=${WQKV:-1}
export MASTER_WEIGHTS=${MASTER_WEIGHTS:-1}
export FP8=${FP8:-1}
export ALLREDUCE_CAST=${ALLREDUCE_CAST:-1}
export FAST_CE=${FAST_CE:-1}
export FUSED_INPUT_QUANTIZE=${FUSED_INPUT_QUANTIZE:-1}
export FUSED_GRAD_QUANTIZE=${FUSED_GRAD_QUANTIZE:-1}
export FUSED_ADD_NORM_MUL_QUANTIZE=${FUSED_ADD_NORM_MUL_QUANTIZE:-1}
export FUSED_SILU_W13=${FUSED_SILU_W13:-1}
export SPLIT_W13=${SPLIT_W13:-0}
export OFFLOAD_OPTIM=${OFFLOAD_OPTIM:-0}
export DEFAULT_FLOAT="bfloat16" OPTIM_DTYPE="bfloat16"
export DP=${DP:-8} MP=${MP:-1} BS=${BS:-16} EVAL_BS=${EVAL_BS:-8} GRADIENT_ACC_STEPS=${GRADIENT_ACC_STEPS:-2}
export GBS=$((BS * GRADIENT_ACC_STEPS))
export MODEL="llama3"
export BASEDIR="/raid/datasets/c4-8b/"
export SMALL=1
export LLAMA3_SIZE=${LLAMA3_SIZE:-"8B"}
export EVAL_TARGET=3.3 EVAL_FREQ=12288
export LR="1e-3" END_LR="1e-4" WARMUP_SAMPLES=4096 MAX_STEPS=1200000
export WARMUP_STEPS=$((WARMUP_SAMPLES / GBS))
export SAMPLES=$((MAX_STEPS * GBS))
export SEQLEN=${SEQLEN:-8192}
export SEED=${SEED:-$RANDOM}
export DATA_SEED=${DATA_SEED:-5760}
export JITBEAM=${JITBEAM:-3}
export BEAM_UOPS_MAX=6000 BEAM_UPCAST_MAX=256 BEAM_LOCAL_MAX=1024 BEAM_MIN_PROGRESS=5 BEAM_PADTO=1
python3 examples/mlperf/model_train.py
@@ -0,0 +1,49 @@
#!/usr/bin/env bash
export PYTHONPATH="."
export PATH="/opt/rocm-7.1.1/bin:$PATH"
export ROCM_PATH="/opt/rocm-7.1.1"
export DEV=${DEV:-AMD}
export CHECK_OOB=0
export REWRITE_STACK_LIMIT=5000000 HCQDEV_WAIT_TIMEOUT_MS=240000
export DEVICE_IN_FUNCTION_BUG=1
export DEBUG=${DEBUG:-0}
export HK_FLASH_ATTENTION=${HK_FLASH_ATTENTION:-1}
export ALL2ALL=${ALL2ALL:-1}
export LATE_ALLREDUCE=${LATE_ALLREDUCE:-1}
export USE_ATOMICS=${USE_ATOMICS:-1}
export ASM_GEMM=${ASM_GEMM:-1}
export WQKV=${WQKV:-1}
export MASTER_WEIGHTS=${MASTER_WEIGHTS:-1}
export FP8=${FP8:-1}
export ALLREDUCE_CAST=${ALLREDUCE_CAST:-1}
export FAST_CE=${FAST_CE:-0}
export FUSED_INPUT_QUANTIZE=${FUSED_INPUT_QUANTIZE:-0}
export FUSED_GRAD_QUANTIZE=${FUSED_GRAD_QUANTIZE:-0}
export FUSED_ADD_NORM_MUL_QUANTIZE=${FUSED_ADD_NORM_MUL_QUANTIZE:-0}
export FUSED_SILU_W13=${FUSED_SILU_W13:-0}
export SPLIT_W13=${SPLIT_W13:-1}
export OFFLOAD_OPTIM=${OFFLOAD_OPTIM:-1}
export DEFAULT_FLOAT="bfloat16" OPTIM_DTYPE="bfloat16"
export DP=${DP:-1} MP=${MP:-8} BS=${BS:-1} EVAL_BS=${EVAL_BS:-1} GRADIENT_ACC_STEPS=${GRADIENT_ACC_STEPS:-32}
export GBS=$((BS * GRADIENT_ACC_STEPS))
export MODEL="llama3"
export BASEDIR="/raid/datasets/c4-8b/"
export SMALL=1
export LLAMA3_SIZE=${LLAMA3_SIZE:-"8B"}
export EVAL_TARGET=3.3 EVAL_FREQ=12288
export LR="1e-3" END_LR="1e-4" WARMUP_SAMPLES=4096 MAX_STEPS=1200000
export WARMUP_STEPS=$((WARMUP_SAMPLES / GBS))
export SAMPLES=$((MAX_STEPS * GBS))
export SEQLEN=${SEQLEN:-8192}
export SEED=${SEED:-$RANDOM}
export DATA_SEED=${DATA_SEED:-5760}
export JITBEAM=${JITBEAM:-3}
export BEAM_UOPS_MAX=6000 BEAM_UPCAST_MAX=256 BEAM_LOCAL_MAX=1024 BEAM_MIN_PROGRESS=5 BEAM_PADTO=1
python3 examples/mlperf/model_train.py
@@ -0,0 +1,5 @@
#!/bin/bash
export BENCHMARK=${BENCHMARK:-5}
export EVAL_BS=0
VIZ=${VIZ:--1} FULL_LAYERS=1 DEBUG=${DEBUG:--0} examples/mlperf/training_submission_v6.0/tinycorp/benchmarks/llama31_8b/implementations/tinybox_8xMI350X/dev_beam.sh
[ "$BENCHMARK" -le 3 ] || [[ $DEV == NULL* ]] || python -m tinygrad.viz.cli -s AMD -t --interval "train @ 2" "train @ 3"
@@ -0,0 +1,58 @@
#!/usr/bin/env bash
set -e # Exit on any error
set -o pipefail # Make pipeline fail if any command fails
export PYTHONPATH="."
export PATH="/opt/rocm-7.1.1/bin:$PATH"
export ROCM_PATH="/opt/rocm-7.1.1"
export DEV=AMD
export CHECK_OOB=0
export REWRITE_STACK_LIMIT=5000000 HCQDEV_WAIT_TIMEOUT_MS=240000
export DEVICE_IN_FUNCTION_BUG=1
export HK_FLASH_ATTENTION=1
export ALL2ALL=1
export LATE_ALLREDUCE=0
export USE_ATOMICS=1
export ASM_GEMM=1
export WQKV=1
export MASTER_WEIGHTS=1
export FP8=1
export ALLREDUCE_CAST=1
export FAST_CE=1
export FUSED_INPUT_QUANTIZE=1
export FUSED_GRAD_QUANTIZE=1
export FUSED_ADD_NORM_MUL_QUANTIZE=1
export FUSED_SILU_W13=1
export SPLIT_W13=0
export DEFAULT_FLOAT="bfloat16" OPTIM_DTYPE="bfloat16"
export DP=8 MP=1 BS=16 EVAL_BS=8 GRADIENT_ACC_STEPS=2
export GBS=$((BS * GRADIENT_ACC_STEPS))
export MODEL="llama3"
export BASEDIR="/raid/datasets/c4-8b/"
export SMALL=1
export LLAMA3_SIZE=8B
export EVAL_TARGET=3.3 EVAL_FREQ=12288
export LR="1e-3" END_LR="1e-4" WARMUP_SAMPLES=4096 MAX_STEPS=1200000
export WARMUP_STEPS=$((WARMUP_SAMPLES / GBS))
export SAMPLES=$((MAX_STEPS * GBS))
export SEQLEN=8192
export SEED=$RANDOM
export DATA_SEED=$SEED
export JITBEAM=3
export BEAM_UOPS_MAX=6000 BEAM_UPCAST_MAX=256 BEAM_LOCAL_MAX=1024 BEAM_MIN_PROGRESS=5 BEAM_PADTO=1
export LOGMLPERF=1
DATETIME=$(date "+%m%d%H%M")
LOGFILE="llama31_8b_8xMI350x_${DATETIME}_${SEED}.log"
# beam
FAKEDATA=1 BENCHMARK=10 INITMLPERF=1 LLAMA_LAYERS=2 python3 examples/mlperf/model_train.py | tee "$LOGFILE"
# run
RUNMLPERF=1 python3 examples/mlperf/model_train.py | tee -a "$LOGFILE"
@@ -0,0 +1,10 @@
#!/bin/bash
export BENCHMARK=5
export EVAL_BS=0
export FAKEDATA=1
export NULL_ALLOW_COPYOUT=1
export HIP_VISIBLE_DEVICES=""
export DEV=NULL:HIP:gfx950
export JITBEAM=0
export LLAMA_LAYERS=${LLAMA_LAYERS:-"2"}
time examples/mlperf/training_submission_v6.0/tinycorp/benchmarks/llama31_8b/implementations/tinybox_8xMI350X/dev_run.sh
@@ -0,0 +1,38 @@
{
"submitter": "tinycorp",
"division": "closed",
"status": "Available on-premise",
"system_name": "tinybox 8xMI300X",
"number_of_nodes": "1",
"host_processors_per_node": "2",
"host_processor_model_name": "AMD EPYC 9354",
"host_processor_core_count": "32",
"host_processor_vcpu_count": "64",
"host_processor_frequency": "",
"host_processor_caches": "",
"host_processor_interconnect": "",
"host_memory_capacity": "2304GB",
"host_storage_type": "NVMe SSD",
"host_storage_capacity": "3x 4TB raid array",
"host_networking": "",
"host_networking_topology": "",
"host_memory_configuration": "24x 96GB DDR5",
"accelerators_per_node": "8",
"accelerator_model_name": "AMD Instinct MI300X 192GB HBM3",
"accelerator_host_interconnect": "PCIe 5.0 x16",
"accelerator_frequency": "",
"accelerator_on-chip_memories": "",
"accelerator_memory_configuration": "HBM3",
"accelerator_memory_capacity": "192GB",
"accelerator_interconnect": "",
"accelerator_interconnect_topology": "",
"cooling": "air",
"hw_notes": "",
"framework": "tinygrad, branch mlperf_training_v5.0",
"other_software_stack": {
"python": "3.10.16",
"ROCm": "3.0.0+94441cb"
},
"operating_system": "Ubuntu 24.04.1 LTS",
"sw_notes": ""
}
@@ -0,0 +1,38 @@
{
"submitter": "tinycorp",
"division": "closed",
"status": "Available on-premise",
"system_name": "tinybox 8xMI350X",
"number_of_nodes": "1",
"host_processors_per_node": "2",
"host_processor_model_name": "AMD EPYC 9575F",
"host_processor_core_count": "32",
"host_processor_vcpu_count": "64",
"host_processor_frequency": "",
"host_processor_caches": "",
"host_processor_interconnect": "",
"host_memory_capacity": "3072 GiB",
"host_storage_type": "NVMe SSD",
"host_storage_capacity": "4TB",
"host_networking": "",
"host_networking_topology": "",
"host_memory_configuration": "24x 128GB DDR5",
"accelerators_per_node": "8",
"accelerator_model_name": "AMD Instinct MI350X 288GB HBM3e",
"accelerator_host_interconnect": "PCIe 5.0 x16",
"accelerator_frequency": "",
"accelerator_on-chip_memories": "",
"accelerator_memory_configuration": "HBM3",
"accelerator_memory_capacity": "288GB",
"accelerator_interconnect": "",
"accelerator_interconnect_topology": "",
"cooling": "air",
"hw_notes": "",
"framework": "tinygrad, branch mlperf_training_v6.0",
"other_software_stack": {
"python": "3.12.3",
"ROCm": "7.1.1"
},
"operating_system": "Ubuntu 24.04.3 LTS",
"sw_notes": ""
}
@@ -0,0 +1,38 @@
{
"submitter": "tinycorp",
"division": "closed",
"status": "Available on-premise",
"system_name": "tinybox green",
"number_of_nodes": "1",
"host_processors_per_node": "1",
"host_processor_model_name": "AMD EPYC 7532",
"host_processor_core_count": "32",
"host_processor_vcpu_count": "64",
"host_processor_frequency": "",
"host_processor_caches": "",
"host_processor_interconnect": "",
"host_memory_capacity": "128GB",
"host_storage_type": "NVMe SSD",
"host_storage_capacity": "4 TB raid array + 1 TB boot",
"host_networking": "",
"host_networking_topology": "",
"host_memory_configuration": "8x 16GB DDR4",
"accelerators_per_node": "6",
"accelerator_model_name": "NVIDIA GeForce RTX 4090",
"accelerator_host_interconnect": "PCIe 4.0 x16",
"accelerator_frequency": "",
"accelerator_on-chip_memories": "",
"accelerator_memory_configuration": "GDDR6X",
"accelerator_memory_capacity": "24GB",
"accelerator_interconnect": "",
"accelerator_interconnect_topology": "",
"cooling": "air",
"hw_notes": "",
"framework": "tinygrad, branch mlperf_training_v5.0",
"other_software_stack": {
"python": "3.10.12",
"CUDA": "12.4"
},
"operating_system": "Ubuntu 22.04.4",
"sw_notes": ""
}
@@ -0,0 +1,37 @@
{
"submitter": "tinycorp",
"division": "closed",
"status": "Available on-premise",
"system_name": "tinybox red",
"number_of_nodes": "1",
"host_processors_per_node": "1",
"host_processor_model_name": "AMD EPYC 7532",
"host_processor_core_count": "32",
"host_processor_vcpu_count": "64",
"host_processor_frequency": "",
"host_processor_caches": "",
"host_processor_interconnect": "",
"host_memory_capacity": "128GB",
"host_storage_type": "NVMe SSD",
"host_storage_capacity": "4 TB raid array + 1 TB boot",
"host_networking": "",
"host_networking_topology": "",
"host_memory_configuration": "8x 16GB DDR4",
"accelerators_per_node": "6",
"accelerator_model_name": "AMD Radeon RX 7900 XTX",
"accelerator_host_interconnect": "PCIe 4.0 x16",
"accelerator_frequency": "",
"accelerator_on-chip_memories": "",
"accelerator_memory_configuration": "GDDR6",
"accelerator_memory_capacity": "24GB",
"accelerator_interconnect": "",
"accelerator_interconnect_topology": "",
"cooling": "air",
"hw_notes": "",
"framework": "tinygrad, branch mlperf_training_v5.0",
"other_software_stack": {
"python": "3.10.12"
},
"operating_system": "Ubuntu 22.04.4",
"sw_notes": ""
}
+3 -28
View File
@@ -65,30 +65,6 @@ def compile(onnx_file):
if (allowed_gated_read_image:=getenv("ALLOWED_GATED_READ_IMAGE", -1)) != -1:
assert gated_read_image_count == allowed_gated_read_image, f"different gated read_image! {gated_read_image_count=}, {allowed_gated_read_image=}"
if Device.DEFAULT.startswith("QCOM") and getenv("OPENPILOT_QCOM_RPT", 1):
from extra.gemm.qcom_openpilot_vision_fp16 import patch_fp32_rpt
if (patched:=patch_fp32_rpt(run_onnx_jit)): print(f"repeat-packed {patched} QCOM vision kernels")
if Device.DEFAULT.startswith("QCOM") and getenv("OPENPILOT_QCOM_SCHEDULE", 1):
from extra.gemm.qcom_openpilot_schedule_projection import patch_projection
if (patched:=patch_projection(run_onnx_jit)): print(f"rescheduled {patched} QCOM vision kernels")
if Device.DEFAULT.startswith("QCOM") and getenv("OPENPILOT_QCOM_FULL_RPT", 1):
from extra.gemm.qcom_openpilot_inverse_full_rpt import patch_model as patch_full_rpt
if (patched:=patch_full_rpt(run_onnx_jit)): print(f"fully repeat-packed {patched} QCOM vision kernels")
if Device.DEFAULT.startswith("QCOM") and getenv("OPENPILOT_QCOM_DEDUPE", 1):
from extra.gemm.qcom_openpilot_dedupe_head import dedupe_identical_calls
if (removed:=dedupe_identical_calls(run_onnx_jit)): print(f"deduplicated {len(removed)} QCOM kernels")
if Device.DEFAULT.startswith("QCOM") and getenv("OPENPILOT_QCOM_PACK_CONV", 1):
from extra.gemm.qcom_openpilot_pack_conv_weights import patch_conv
if (patched:=patch_conv(run_onnx_jit)): print(f"packed weights for {patched} QCOM convolution kernels")
if Device.DEFAULT.startswith("QCOM") and getenv("OPENPILOT_QCOM_LEVEL_SCHEDULE", 1):
from extra.gemm.qcom_openpilot_level_schedule import schedule_levels
if (moved:=schedule_levels(run_onnx_jit)): print(f"rescheduled {moved} QCOM kernels by dependency level")
if Device.DEFAULT.startswith("QCOM") and getenv("OPENPILOT_QCOM_BATCH_HEAD", 1):
from extra.gemm.qcom_openpilot_batch_head import batch_head
if (combined:=batch_head(run_onnx_jit)): print(f"batched {combined} groups of QCOM head kernels")
if Device.DEFAULT.startswith("QCOM") and getenv("OPENPILOT_QCOM_INPUT_PACK", 1):
from extra.gemm.qcom_openpilot_input_pack import patch_input_pack
if (patched:=patch_input_pack(run_onnx_jit)): print(f"vectorized {patched} QCOM input kernel")
with open(OUTPUT, "wb") as f:
pickle.dump(run_onnx_jit, f)
mdl_sz = os.path.getsize(onnx_file)
@@ -96,7 +72,7 @@ def compile(onnx_file):
print(f"mdl size is {mdl_sz/1e6:.2f}M")
print(f"pkl size is {pkl_sz/1e6:.2f}M")
print("**** compile done ****")
return run_onnx_jit, inputs, test_val
return inputs, test_val
def test_vs_compile(run, inputs, test_val=None):
@@ -166,10 +142,9 @@ if __name__ == "__main__":
test_vs_compile(pickle_loaded, inputs)
else:
onnx_file = fetch(OPENPILOT_MODEL)
pickle_loaded, inputs, outputs = compile(onnx_file)
inputs, outputs = compile(onnx_file)
if OUTPUT != os.devnull:
with open(OUTPUT, "rb") as f: pickle_loaded = pickle.load(f)
with open(OUTPUT, "rb") as f: pickle_loaded = pickle.load(f)
test_vs_compile(pickle_loaded, inputs, outputs)
if getenv("SELFTEST"):
+1 -1
View File
@@ -80,7 +80,7 @@ def block_128x128_gemm(c:UOp, a:UOp, b:UOp) -> UOp:
# NOTE: since this is part of K, these 2 can be anywhere in the frags and long as a and b match
a_frag = a_frag.reshape(2, 8)[lane_m, :]
b_frag = b_frag.reshape(2, 8)[lane_m, :]
wmma = UOp.wmma(a_frag, b_frag, acc_frag.after(k), ((16, 16, 16), 'AMD', 32))
wmma = UOp.wmma(a_frag, b_frag, acc_frag.after(k), (16, 16, 16), 'AMD', 32)
acc_store = acc_frag.store(wmma).end(tile_m, tile_n)
else:
# registers for LOCAL -> REG
+3 -3
View File
@@ -13,7 +13,7 @@ WMMA_ACC = WMMA_M // LANES_PER_WAVE_M
THREADS_PER_BLOCK = WARP_SIZE * WAVES_M * WAVES_N
LDS_PAD = 4 # pad LDS rows to reduce bank conflicts
WMMA_ARG = ((WMMA_M, WMMA_N, WMMA_K), 'AMD', 32)
WMMA_ARG = (WMMA_M, WMMA_N, WMMA_K), 'AMD', 32
LOG2E = math.log2(math.e)
def warp_shfl_xor(val, offset, lane):
@@ -97,7 +97,7 @@ def amd_flash_attention(o:UOp, q:UOp, k:UOp, v:UOp) -> UOp:
S_frag = S_reg.reshape(TM // WMMA_ACC, WMMA_ACC, TN).permute(0, 2, 1)[tm1, tn1]
q_frag = Q_lds.reshape(WAVES_M, TM // WMMA_ACC, WMMA_M, D // WMMA_K, WMMA_K)[wave_m, tm1, lane_n, k_qk]
k_frag = KV_lds_k.reshape(WAVES_N, TN, WMMA_N, D // WMMA_K, WMMA_K)[wave_n, tn1, lane_n, k_qk]
qk = UOp.wmma(q_frag, k_frag, S_frag.after(k_qk), WMMA_ARG)
qk = UOp.wmma(q_frag, k_frag, S_frag.after(k_qk), *WMMA_ARG)
qk_done = S_frag.store(qk).end(tm1, tn1).end(k_qk)
S_reg = S_reg.after(qk_done)
@@ -158,7 +158,7 @@ def amd_flash_attention(o:UOp, q:UOp, k:UOp, v:UOp) -> UOp:
acc_frag = acc.reshape(TM // WMMA_ACC, WMMA_ACC, TD).permute(0, 2, 1)[tm2, tn2]
p_frag = P_lds.reshape(WAVES_M, TM // WMMA_ACC, WMMA_M, BLOCK_N // WMMA_K, WMMA_K)[wave_m, tm2, lane_n, k_pv]
v_frag = KV_lds_v.reshape(WAVES_N, TD, WMMA_N, BLOCK_N // WMMA_K, WMMA_K)[wave_n, tn2, lane_n, k_pv]
pv = UOp.wmma(p_frag, v_frag, acc_frag.after(k_pv), WMMA_ARG)
pv = UOp.wmma(p_frag, v_frag, acc_frag.after(k_pv), *WMMA_ARG)
# end KV tile loop
n_tile_end = acc_frag.store(pv).end(tm2, tn2).end(k_pv).barrier().end(n_tile)
+9 -7
View File
@@ -128,6 +128,11 @@ def _mx_block_scale(e8:Tensor) -> Tensor:
rows, scale_K = e8.shape
return (e8.cast(dtypes.float32) - 127.0).exp2().reshape(rows, scale_K, 1).expand(rows, scale_K, 32).reshape(rows, scale_K*32)
def _mx_block_scale_3d(e8:Tensor) -> Tensor:
# batched (E, rows, scale_K) dequant scale 2^(e8-127) broadcast to (E, rows, scale_K*32)
E, rows, scale_K = e8.shape
return (e8.cast(dtypes.float32) - 127.0).exp2().reshape(E, rows, scale_K, 1).expand(E, rows, scale_K, 32).reshape(E, rows, scale_K*32)
counters = {"used":0, "todos":[]}
def todo(msg:str) -> bool: counters["todos"].append(msg); return False
def _asm_gemm_report():
@@ -169,10 +174,10 @@ def custom_uop_gemm(C:UOp, A:UOp, B:UOp) -> UOp:
m = UOp.range(M, 1, AxisType.LOOP)
n = UOp.range(N, 2, AxisType.LOOP)
k = UOp.range(K, 0, AxisType.REDUCE)
mul = (A.flatten().index((m*UOp.const(dtypes.index, K)+k))*
B.flatten().index((k*UOp.const(dtypes.index, N)+n))).cast(dtypes.float32)
mul = (A.flatten().index((m*UOp.const(dtypes.weakint, K)+k))*
B.flatten().index((k*UOp.const(dtypes.weakint, N)+n))).cast(dtypes.float32)
red = mul.reduce(k, arg=Ops.ADD, dtype=dtypes.float32).cast(C.dtype)
store = C.flatten().index((m*UOp.const(dtypes.index, N)+n)).store(red).end(m, n)
store = C.flatten().index((m*UOp.const(dtypes.weakint, N)+n)).store(red).end(m, n)
return store.sink(arg=KernelInfo(name=f'uop_gemm_{M}_{N}_{K}'))
# ** bf16 A @ B.T kernel in C
@@ -275,10 +280,7 @@ def custom_gemm_bw(gradient:UOp, kernel:UOp, n_scales:int=2, has_grad_amax:bool=
elif getenv("FUSED_GRAD_QUANTIZE", 0):
grad_amax_t = Tensor(grad_amax_state, device=a.device)
g_amax = grad_amax_t
g_fp8, _, new_grad_amax, _ = quantize_fp8_delayed(g_t, g_amax)
store_effect = next_grad_amax_state.store(new_grad_amax.uop)
assert g_fp8.uop.op is Ops.AFTER, f"expected AFTER, got {g_fp8.uop.op}"
g_fp8 = Tensor(g_fp8.uop.replace(src=g_fp8.uop.src + (store_effect,)), device=a.device)
g_fp8, _ = quantize_fp8_delayed(g_t, g_amax, Tensor(next_grad_amax_state, device=a.device))
else:
grad_amax_t = Tensor(grad_amax_state, device=a.device)
g_amax = grad_amax_t
+2 -5
View File
@@ -1,5 +1,5 @@
from tinygrad import UOp, dtypes
from tinygrad.uop.ops import AxisType, Ops, KernelInfo, AddrSpace
from tinygrad.uop.ops import AxisType, KernelInfo, AddrSpace
from extra.gemm.amd_uop_matmul import test_matmul
N = 2048
@@ -27,11 +27,8 @@ def hand_spec_tc_cores():
acc = acc[0].set(0.0)
acc = acc[1].set(0.0)
# TODO: make this simple
wmma_arg = ('WMMA_8_8_8_float_float', (8, 8, 8), dtypes.float, dtypes.float, 'METAL', 32, (((3, 2),), ((3, 2),), ((3, 2),)), ())
acc_load = UOp.stack(acc.after(gk)[0], acc.after(gk)[1])
out = UOp(Ops.WMMA, dtypes.float, (a_tc, b_tc, acc_load), arg=wmma_arg)
out = UOp.wmma(a_tc, b_tc, acc_load, (8, 8, 8), 'METAL', 32)
end_loop = UOp.group(*[acc[i].store(out.index(i)) for i in range(2)]).end(gk)
+3 -5
View File
@@ -6,7 +6,7 @@ os.environ["AMD_LLVM"] = "0"
from tinygrad import Tensor, Context, dtypes, UOp, GlobalCounters
from tinygrad.helpers import DEBUG, getenv
from tinygrad.dtype import AddrSpace
from tinygrad.uop.ops import AxisType, KernelInfo, Ops
from tinygrad.uop.ops import AxisType, KernelInfo
WARP_SIZE = 64
@@ -137,8 +137,7 @@ def custom_gemm(C:UOp, A:UOp, B:UOp) -> UOp:
acc_load = acc_after[N_inner_loop, M_inner_loop]
# do WMMA
wmma_arg = ('WMMA_16_16_32_half_float', (16, 16, 32), dtypes.half, dtypes.float, 'AMD', 64, ((), (), ((3, 2), (2, 2))), ())
out = UOp(Ops.WMMA, dtypes.float, (Ar[M_inner_loop], Br[N_inner_loop], acc_load), arg=wmma_arg)
out = UOp.wmma(Ar[M_inner_loop], Br[N_inner_loop], acc_load, (16, 16, 32), 'AMD', 64)
# store back the acc
acc_store = acc[N_inner_loop, M_inner_loop].store(out)
@@ -193,8 +192,7 @@ acc = acc[init_l:=UOp.range(4, 1)].set(0.0, end=init_l)
# do the wmma
acc_load = UOp.stack(*[acc.after(K_loop)[i] for i in range(4)])
wmma_arg = ('WMMA_16_16_32_half_float', (16, 16, 32), dtypes.half, dtypes.float, 'AMD', 64, ((), (), ((3, 2), (2, 2))), ())
out = UOp(Ops.WMMA, dtypes.float, (A_in, B_in, acc_load), arg=wmma_arg)
out = UOp.wmma(A_in, B_in, acc_load, (16, 16, 32), 'AMD', 64)
# store back the acc
acc = acc.after(UOp.group(*[acc[i].store(out.index(i)) for i in range(4)]).end(K_loop))
+2 -3
View File
@@ -6,7 +6,7 @@ os.environ["AMD_LLVM"] = "0"
from tinygrad import Tensor, Context, dtypes, UOp, GlobalCounters
from tinygrad.helpers import DEBUG, getenv
from tinygrad.dtype import AddrSpace
from tinygrad.uop.ops import sint, AxisType, KernelInfo, Ops
from tinygrad.uop.ops import AxisType, KernelInfo
WARP_SIZE = 64
@@ -60,8 +60,7 @@ def compute_on_locals(acc:UOp, Asl:UOp, Bsl:UOp, rng:int, afters:tuple[UOp, ...]
acc_load = acc_after[N_inner_loop, M_inner_loop]
# do WMMA
wmma_arg = ('WMMA_16_16_32_half_float', (16, 16, 32), dtypes.half, dtypes.float, 'AMD', 64, ((), (), ((3, 2), (2, 2))), ())
out = UOp(Ops.WMMA, dtypes.float, (Ar[M_inner_loop], Br[N_inner_loop], acc_load), arg=wmma_arg)
out = UOp.wmma(Ar[M_inner_loop], Br[N_inner_loop], acc_load, (16, 16, 32), 'AMD', 64)
# store back the acc
acc_store = acc[N_inner_loop, M_inner_loop].store(out)
-110
View File
@@ -1,110 +0,0 @@
#!/usr/bin/env python3
"""Batch adjacent independent openpilot head kernels into one QCOM launch."""
import argparse, pickle, re
from dataclasses import replace
from tinygrad import Device
from tinygrad.engine.jit import create_graph_call
from tinygrad.uop.ops import Ops
from extra.gemm.qcom_openpilot_ir3 import plain_name
MAX_BATCH={"r_256_4_128_4":4,"r_128_16_4_16_4":4,"r_128_16_4_32_4":4,
"r_8_16_4_8_4":4,"r_8_4_8_4":4,"r_8_4_8_4n1":4}
MAX_BATCH.update({"r_16_16_4_8_4":4,"r_4_16_4_8_4":4,"r_16_16_4_4":4,
"r_4_4_4_4":4,"r_16_16_4_4n1":4,"r_4_4_4_4n1":4})
def batched_source(source:str, name:str, batch_count:int) -> str:
match=re.search(r"__kernel void \w+\((.*?)\) \{",source,re.S)
if match is None: raise RuntimeError("kernel signature not found")
declarations=[x.strip() for x in match.group(1).split(",")]
arg_names=[x.rsplit(" ",1)[1] for x in declarations]
renamed=[]
bodies=[]
body=source[match.end():source.rfind("}")]
local_decls=re.findall(r"__attribute__\s*\(\(aligned \(\d+\)\)\)\s*__local\s+[^;]+;",body)
hoisted=[]
for batch in range(batch_count):
mapping={arg:f"{arg}_{batch}" for arg in arg_names}
renamed.extend(decl.rsplit(" ",1)[0]+" "+mapping[arg] for decl,arg in zip(declarations,arg_names))
branch=body
for declaration in local_decls:
local_match=re.search(r"(\w+)(\[[^;]+;)$",declaration)
if local_match is None: raise RuntimeError(f"local declaration not understood: {declaration}")
old=local_match.group(1)
new=f"{old}_{batch}"
hoisted.append(declaration[:local_match.start(1)]+new+local_match.group(2))
branch=branch.replace(declaration,"")
branch=re.sub(rf"\b{re.escape(old)}\b",new,branch)
for old,new in mapping.items(): branch=re.sub(rf"\b{re.escape(old)}\b",new,branch)
bodies.append(branch)
prefix=source[:match.start()]
count=len(declarations)
order=tuple(batch*count for batch in range(batch_count))+tuple(
batch*count+arg for batch in range(batch_count) for arg in range(1,count))
branches=" else ".join((f"if (get_group_id(1)=={batch}) " if batch < batch_count-1 else "")+f"{{{body}}}"
for batch,body in enumerate(bodies))
return f"{prefix}__kernel void {name}_batch{batch_count}({','.join(renamed[i] for i in order)}) {{\n" \
f"{''.join(hoisted)}\n{branches}\n}}"
def independent(calls:list) -> bool:
outputs={call.src[out+1] for call in calls for out in call.src[0].arg.outs}
return not any(arg in outputs for call in calls for i,arg in enumerate(call.src[1:]) if i not in call.src[0].arg.outs)
def batch_head(model) -> int:
outer=model.captured.linear.src[0]
batch=list(outer.src[0].src[0].src)
new_batch=[]
combined=0
index=0
cache={}
while index < len(batch):
first=batch[index]
name=plain_name(first.src[0].arg.name) if first.op is Ops.CALL and first.src[0].op is Ops.PROGRAM else ""
if index+1 < len(batch) and name in MAX_BATCH:
calls=[first]
while index+len(calls) < len(batch) and len(calls) < MAX_BATCH[name]:
candidate=batch[index+len(calls)]
candidate_name=plain_name(candidate.src[0].arg.name) if candidate.op is Ops.CALL and candidate.src[0].op is Ops.PROGRAM else ""
if candidate_name != name or first.src[0].src[3].arg != candidate.src[0].src[3].arg: break
calls.append(candidate)
if len(calls) > 1 and independent(calls):
batch_count=len(calls)
program=first.src[0]
source=batched_source(program.src[2].arg,name,batch_count)
if source not in cache: cache[source]=Device["QCOM"].compiler.compile_cached(source)
aux0=program.arg.aux[0]
count=len(aux0)
ordered_aux=tuple(aux0[0] for _ in calls)+tuple(entry for _ in calls for entry in aux0[1:])
combined_aux=tuple(tuple((new_index,dtype,shape) for _old_index,dtype,shape in entry)
for new_index,entry in enumerate(ordered_aux))
info=replace(program.arg,name=f"{name}_batch{batch_count}",global_size=(program.arg.global_size[0],batch_count,1),
globals=tuple(range(count*batch_count)),outs=tuple(range(batch_count)),
ins=tuple(range(batch_count,count*batch_count)),aux=(combined_aux,))
program=program.replace(arg=info,src=program.src[:2]+
(program.src[2].replace(arg=source),program.src[3].replace(arg=cache[source])))
new_batch.append(first.replace(src=(program,*[call.src[1] for call in calls],
*[arg for call in calls for arg in call.src[2:]])))
combined+=1
index+=batch_count
continue
new_batch.append(first)
index+=1
model.captured._linear=model.captured.linear.substitute({outer:create_graph_call(new_batch)},walk=True)
model.captured.__dict__.pop("linear",None)
return combined
def main() -> None:
parser=argparse.ArgumentParser()
parser.add_argument("input")
parser.add_argument("output")
args=parser.parse_args()
with open(args.input,"rb") as f:model=pickle.load(f)
print("combined",batch_head(model))
with open(args.output,"wb") as f:pickle.dump(model,f)
if __name__ == "__main__":main()
-79
View File
@@ -1,79 +0,0 @@
#!/usr/bin/env python3
"""Remove byte-identical duplicate linear chains in the driving-vision head."""
import argparse, hashlib, pickle
from tinygrad.engine.jit import create_graph_call
from tinygrad.uop.ops import Ops, UOp
from extra.gemm.qcom_openpilot_ir3 import plain_name
TARGETS = {"r_128_16_4_32_4", "r_256_4_128_4", "r_128_16_4_16_4"}
def dedupe_identical_calls(model, all_calls:bool=True) -> list[tuple[int, str]]:
"""Alias calls with identical programs, inputs, and byte-identical constants."""
outer = model.captured.linear.src[0]
batch = outer.src[0].src[0].src
produced:dict[UOp, UOp] = {}
static_hash:dict[UOp, str] = {}
seen:dict[tuple, tuple[UOp, ...]] = {}
new_batch, removed = [], []
def representative(buf:UOp) -> UOp:
while buf in produced and produced[buf] is not buf: buf = produced[buf]
return buf
def content_hash(buf:UOp) -> str:
if buf not in static_hash:
static_hash[buf] = hashlib.sha256(memoryview(buf.buffer.numpy()).cast("B")).hexdigest()
return static_hash[buf]
for index, original in enumerate(batch):
call = original.replace(src=tuple(representative(x) if x in produced else x for x in original.src))
if call.op is not Ops.CALL or call.src[0].op is not Ops.PROGRAM or (not all_calls and plain_name(call.src[0].arg.name) not in TARGETS):
new_batch.append(call)
if call.op is Ops.CALL and call.src[0].op is Ops.PROGRAM:
for out_index in call.src[0].arg.outs: produced[original.src[out_index+1]] = call.src[out_index+1]
continue
program = call.src[0]
output_indices = set(program.arg.outs)
signature_args = []
for arg_index, (before, after) in enumerate(zip(original.src[1:], call.src[1:])):
if arg_index in output_indices: continue
if before.op is Ops.PARAM:
signature_args.append(("param", before.arg))
elif before in produced:
signature_args.append(("dynamic", representative(before)))
else:
signature_args.append((str(after.dtype), after.buffer.size, content_hash(after)))
signature = (plain_name(program.arg.name), program.src[3].arg, tuple(signature_args))
outputs = tuple(original.src[i+1] for i in program.arg.outs)
if signature in seen:
canonical_outputs = seen[signature]
for output, canonical in zip(outputs, canonical_outputs): produced[output] = representative(canonical)
removed.append((index, plain_name(program.arg.name)))
else:
new_batch.append(call)
canonical_outputs = tuple(call.src[i+1] for i in program.arg.outs)
seen[signature] = canonical_outputs
for output, canonical in zip(outputs, canonical_outputs): produced[output] = canonical
# Apply aliases to consumers which occur after the duplicate chains.
new_batch = [call.replace(src=tuple(representative(x) if x in produced else x for x in call.src)) for call in new_batch]
model.captured._linear = model.captured.linear.substitute({outer:create_graph_call(new_batch)}, walk=True)
model.captured.__dict__.pop("linear", None)
return removed
def main() -> None:
parser = argparse.ArgumentParser()
parser.add_argument("input")
parser.add_argument("output")
parser.add_argument("--all", action="store_true", help="deduplicate every program family, not only the head linears")
args = parser.parse_args()
with open(args.input, "rb") as f: model = pickle.load(f)
removed = dedupe_identical_calls(model, args.all)
with open(args.output, "wb") as f: pickle.dump(model, f)
print(f"removed {len(removed)} duplicate head calls: {removed}")
if __name__ == "__main__": main()
-49
View File
@@ -1,49 +0,0 @@
"""Vectorize the driving-vision uint8 input normalization kernel on QCOM."""
from dataclasses import replace
from tinygrad import Device
from tinygrad.engine.jit import create_graph_call
from tinygrad.uop.ops import Ops
from extra.gemm.qcom_openpilot_ir3 import plain_name
TARGET = "E_8192_3_4_2_4"
SOURCE = r"""#pragma OPENCL EXTENSION cl_khr_fp16 : enable
__kernel void E_8192_3_4_2_4(write_only image2d_t O,__global uchar *A,__global uchar *B,
__global half *MEAN,__global half *STD) {
int c=get_global_id(0),i=get_global_id(1),off=(c<<17)+(i<<2),mc=c<<2;
uchar4 a0=vload4(0,A+off),a1=vload4(0,A+off+32768);
uchar4 a2=vload4(0,A+off+65536),a3=vload4(0,A+off+98304);
uchar4 b0=vload4(0,B+off),b1=vload4(0,B+off+32768);
uchar4 b2=vload4(0,B+off+65536),b3=vload4(0,B+off+98304);
half4 ma=vload4(0,MEAN+mc),mb=vload4(0,MEAN+mc+12);
half4 ia=(half4)(1)/vload4(0,STD+mc),ib=(half4)(1)/vload4(0,STD+mc+12);
int x=c+(i&7)*24,y=i>>3;
write_imagef(O,(int2)(x,y),convert_float4(((half4)(a0.x,a1.x,a2.x,a3.x)-ma)*ia));
write_imagef(O,(int2)(x+3,y),convert_float4(((half4)(b0.x,b1.x,b2.x,b3.x)-mb)*ib));
write_imagef(O,(int2)(x+6,y),convert_float4(((half4)(a0.y,a1.y,a2.y,a3.y)-ma)*ia));
write_imagef(O,(int2)(x+9,y),convert_float4(((half4)(b0.y,b1.y,b2.y,b3.y)-mb)*ib));
write_imagef(O,(int2)(x+12,y),convert_float4(((half4)(a0.z,a1.z,a2.z,a3.z)-ma)*ia));
write_imagef(O,(int2)(x+15,y),convert_float4(((half4)(b0.z,b1.z,b2.z,b3.z)-mb)*ib));
write_imagef(O,(int2)(x+18,y),convert_float4(((half4)(a0.w,a1.w,a2.w,a3.w)-ma)*ia));
write_imagef(O,(int2)(x+21,y),convert_float4(((half4)(b0.w,b1.w,b2.w,b3.w)-mb)*ib));
}"""
def patch_input_pack(jit) -> int:
outer = jit.captured.linear.src[0]
batch = outer.src[0].src[0].src
lib, new_batch, replaced = None, [], 0
for call in batch:
name = plain_name(call.src[0].arg.name) if call.op is Ops.CALL and call.src[0].op is Ops.PROGRAM else ""
if name == TARGET:
if lib is None: lib = Device["QCOM"].compiler.compile_cached(SOURCE)
program = call.src[0]
program = program.replace(arg=replace(program.arg, global_size=(1, 64, 1), local_size=(3, 128, 1)),
src=program.src[:2] +
(program.src[2].replace(arg=SOURCE), program.src[3].replace(arg=lib)))
call, replaced = call.replace(src=(program, *call.src[1:])), replaced+1
new_batch.append(call)
if replaced:
jit.captured._linear = jit.captured.linear.substitute({outer:create_graph_call(new_batch)}, walk=True)
jit.captured.__dict__.pop("linear", None)
return replaced
@@ -1,123 +0,0 @@
#!/usr/bin/env python3
"""Pack all four FP32 accumulator vectors in the OpenPilot inverse projection."""
import argparse, hashlib, pickle, struct
from tinygrad.engine.jit import create_graph_call
from tinygrad.uop.ops import Ops
from extra.gemm.qcom_openpilot_ir3 import branch as BR, mad_f32 as MAD_F32, mov_f32 as MOV_F32, nop as NOP, plain_name
TARGET="r_32_64_4_4_192_4"
INVERSE_W_TARGETS={"r_512_16_4_4_48_4","r_128_32_4_4_96_4"}
SAFE_DONORS={"d4c281a1","1fe26758","e34e7e58"}
def replace_src2(ins:bytes, src2:int) -> bytes:
lo,hi=struct.unpack("<II",ins)
return struct.pack("<II",(lo&0xff00ffff)|(src2<<16),hi)
def replace_low_src(ins:bytes, src:int) -> bytes:
lo,hi=struct.unpack("<II",ins)
return struct.pack("<II",(lo&0xffffff00)|src,hi)
def pack_inverse_full(lib:bytes) -> bytes:
off,size=struct.unpack_from("<I",lib,0xc0)[0],struct.unpack_from("<I",lib,0x100)[0]
instrs=[lib[i:i+8] for i in range(off,off+size,8)]
if len(instrs)!=175: raise RuntimeError(f"expected 175 inverse instructions, got {len(instrs)}")
# Move loop control from r13.x into the existing r12.w zero register. This
# makes r13-r16 four contiguous accumulator vectors without growing the
# shader's declared register file.
out=instrs[:11]
for acc in ("r13.x","r14.x","r15.x","r16.x"): out.append(MOV_F32(acc,"r12.w",rpt=3))
loop_start=len(out)
body=list(instrs[16:32])
body[0]=replace_low_src(body[0],51) # add r8.z, r12.w, 192
body[1]=replace_low_src(body[1],51) # add r9.x, r12.w, 384
body[2]=replace_src2(body[2],51) # add r9.z, c28.y, r12.w
body[3]=replace_low_src(body[3],51) # mov r10.x, r12.w
out+=body
for component,weight in zip("xyzw",("r5.x","r2.x","r3.x","r4.x")):
out.append(MAD_F32("r13.x",f"r7.{component}",weight,"r13.x",rpt=3,r=True,sy=component=="x"))
out+=instrs[48:60]
control=list(instrs[60:65])
control[0]=replace_low_src(control[0],51) # increment r12.w
control[2]=replace_low_src(control[2],51) # compare r12.w
control[3]=MOV_F32("r12.w","r0.x")
out+=control
out.append(BR(loop_start-len(out)))
tail=list(instrs[66:131])
# The first residual moved from r12.w to r13.x; y/z/w were already in r13.
tail[90-66]=replace_src2(tail[90-66],52)
out+=tail
out += [NOP()]*(len(instrs)-len(out))
if len(out)!=len(instrs): raise RuntimeError(f"packed image has {len(out)} instructions")
return lib[:off]+b"".join(out)+lib[off+size:]
def with_fregs(lib:bytes, count:int) -> bytes:
out=bytearray(lib)
regoff=struct.unpack_from("<I",out,0x34)[0]+0x14
regs=struct.unpack_from("<I",out,regoff)[0]
struct.pack_into("<I",out,regoff,(regs&0x80000000)|max(regs&0x7fffffff,count))
return bytes(out)
def pack_inverse_w_full(lib:bytes) -> bytes:
off,size=struct.unpack_from("<I",lib,0xc0)[0],struct.unpack_from("<I",lib,0x100)[0]
ins=[lib[i:i+8] for i in range(off,off+size,8)]
if len(ins) not in (174,178): raise RuntimeError(f"expected 174/178 inverse-W instructions, got {len(ins)}")
out=ins[:19]+[MOV_F32("r17.x","r12.z",rpt=3)]
loop_start=len(out)
out+=ins[19:35]
for component,weight in zip("xyzw",("r5.x","r2.x","r3.x","r4.x")):
out.append(MAD_F32("r17.x",f"r7.{component}",weight,"r17.x",rpt=3,r=True,sy=component=="x"))
out+=ins[51:68]
out.append(BR(loop_start-len(out)))
tail=list(ins[69:])
mapping={50:68,52:69,53:70,54:71}
first_store=next(i for i,x in enumerate(tail) if struct.unpack_from("<I",x,4)[0]>>24==0xc0)
for i in range(first_store):
lo,_=struct.unpack("<II",tail[i])
src2=(lo>>16)&0xff
if src2 in mapping: tail[i]=replace_src2(tail[i],mapping[src2])
out+=tail
out += [NOP()]*(len(ins)-len(out))
return with_fregs(lib[:off]+b"".join(out)+lib[off+size:],18)
def patch_model(model,names:set[str]|None=None) -> int:
outer=model.captured.linear.src[0]
batch=list(outer.src[0].src[0].src)
cache,patched={},0
for index,call in enumerate(batch):
if call.op is not Ops.CALL or call.src[0].op is not Ops.PROGRAM: continue
name=plain_name(call.src[0].arg.name)
if name not in (names if names is not None else {TARGET}|INVERSE_W_TARGETS): continue
program=call.src[0]
old=program.src[3].arg
# These transforms relocate fixed compiler registers. A different QCOM compiler allocation can
# have the same instruction count but different live values and must not be patched by index.
if hashlib.sha1(old).hexdigest()[:8] not in SAFE_DONORS: continue
if old not in cache:
cache[old]=pack_inverse_full(old) if name==TARGET else pack_inverse_w_full(old)
program=program.replace(src=program.src[:3]+(program.src[3].replace(arg=cache[old]),))
batch[index]=call.replace(src=(program,*call.src[1:]))
patched+=1
model.captured._linear=model.captured.linear.substitute({outer:create_graph_call(batch)},walk=True)
model.captured.__dict__.pop("linear",None)
return patched
def main() -> None:
ap=argparse.ArgumentParser()
ap.add_argument("input")
ap.add_argument("output")
ap.add_argument("--names",help="comma-separated program families")
args=ap.parse_args()
with open(args.input,"rb") as f:model=pickle.load(f)
print("patched",patch_model(model,set(args.names.split(",")) if args.names else None))
with open(args.output,"wb") as f:pickle.dump(model,f)
if __name__=="__main__":main()
-34
View File
@@ -1,34 +0,0 @@
"""Small IR3 encoding helpers used by the OpenPilot QCOM graph patches."""
import re, struct
ANSI_RE = re.compile(r"\x1b\[[0-9;]*m")
def plain_name(name:str) -> str: return ANSI_RE.sub("", name)
def _freg(name:str|int) -> int:
if isinstance(name, int): return name
register, component = name.replace("r", "").split(".")
return int(register) * 4 + "xyzw".index(component)
def _pack(lo:int, hi:int) -> bytes: return struct.pack("<II", lo & 0xffffffff, hi & 0xffffffff)
def nop() -> bytes: return _pack(0, 0)
def branch(offset:int) -> bytes: return struct.pack("<iI", offset, 0x00900000)
def mov_f32(dst:str, src:str, rpt:int=0, sy:bool=False, ss:bool=False, r:bool=False) -> bytes:
return _pack(_freg(src), (0x30044000 if sy else 0x20044000) | (0x1000 if ss else 0) |
(0x800 if r else 0) | ((rpt & 0x7f) << 8) | _freg(dst))
def add_s(dst:str, src:str, imm:int, ss:bool=False) -> bytes:
lo = ((0x27 if imm < 0 else 0x20) << 24) | ((imm & 0xff) << 16) | _freg(src)
return _pack(lo, 0x42300000 | (0x1000 if ss else 0) | _freg(dst))
def mad_f32(dst:str, src1:str, src2:str, src3:str, rpt:int=0, sy:bool=False, r:bool=False) -> bytes:
d, s1, s2, s3 = _freg(dst), _freg(src1), _freg(src2), _freg(src3)
hi = ((0x73 if sy else 0x63) << 24) | (0x80 << 16) | ((s2 >> 1) << 16) | (((s2 & 1) << 7 | (rpt & 0x7f)) << 8) | d
lo = (0x20000000 if r else 0) | (s3 << 16) | (0x8000 if r else 0) | s1
return _pack(lo, hi)
def isam_f32(dst:str, coord:str, tex:int=0, samp:int=0) -> bytes:
return _pack((tex * 2) << 24 | ((samp & 7) << 21) | (_freg(coord) * 2 + 1), 0xa0001f00 | _freg(dst))
@@ -1,47 +0,0 @@
#!/usr/bin/env python3
"""Group ready OpenPilot graph calls by program family without crossing dependency levels."""
import argparse, pickle
from collections import defaultdict
from tinygrad.engine.jit import create_graph_call
from tinygrad.uop.ops import Ops
from extra.gemm.qcom_openpilot_ir3 import plain_name
def schedule_levels(model) -> int:
outer=model.captured.linear.src[0]
batch=list(outer.src[0].src[0].src)
writer,levels={},{}
grouped=defaultdict(list)
for sequence,call in enumerate(batch):
if call.op is not Ops.CALL or call.src[0].op is not Ops.PROGRAM:
grouped[sequence].append((sequence,call))
continue
deps={writer[arg] for arg in call.src[1:] if arg in writer}
level=1+max((levels[dep] for dep in deps),default=-1)
levels[sequence]=level
grouped[level].append((sequence,call))
for output in call.src[0].arg.outs: writer[call.src[output+1]]=sequence
scheduled=[]
moved=0
for entries in grouped.values():
ordered=sorted(entries,key=lambda item:(plain_name(item[1].src[0].arg.name),item[0]))
scheduled.extend(call for _index,call in ordered)
moved+=sum(old_index!=entries[new_index][0] for new_index,(old_index,_call) in enumerate(ordered))
if scheduled != batch:
model.captured._linear=model.captured.linear.substitute({outer:create_graph_call(scheduled)},walk=True)
model.captured.__dict__.pop("linear",None)
return moved
def main() -> None:
ap=argparse.ArgumentParser()
ap.add_argument("input")
ap.add_argument("output")
args=ap.parse_args()
with open(args.input,"rb") as f:model=pickle.load(f)
print("moved",schedule_levels(model))
with open(args.output,"wb") as f:pickle.dump(model,f)
if __name__=="__main__":main()
@@ -1,79 +0,0 @@
#!/usr/bin/env python3
"""Prepack static weights for the slow stride-2 openpilot 7x7 convolutions."""
import argparse, pickle, re
import numpy as np
from tinygrad import Device
from tinygrad.device import Buffer
from tinygrad.engine.jit import create_graph_call
from tinygrad.uop.ops import Ops, UOp
from extra.gemm.qcom_openpilot_ir3 import plain_name
TARGETS={"r_16_8_16_2_4_4_7_7", "r_8_4_32_2_4_4_7_7", "r_4_2_64_2_4_4_7_7"}
def packed_source(source:str) -> str:
start=source.index(" half val0")
end=source.index(" int alu20",start)
weight_name=re.search(r"__global half\* (data2_\d+)",source).group(1) # type: ignore[union-attr]
block=" int wp=(((alu0*2+alu2)*7+Ridx0)*28);\n"+"\n".join(
f" float4 w{i}=convert_float4(vload4(0,{weight_name}+wp+{i*4}));" for i in range(7))+"\n"
source=source[:start]+block+source[end:]
accum_start=source.index(" *(buf0+0)",start)
loop_prefix=source[start:accum_start]
casts=re.findall(r" float (cast\d+) = \(\(float\)\(val(\d+)\)\);\n",loop_prefix)
assert len(casts) == 28
source=source[:start]+re.sub(r" float cast\d+ = \(\(float\)\(val\d+\)\);\n", "", loop_prefix)+source[accum_start:]
mapping={cast:("w0.x" if int(val) == 27 else f"w{int(val)%7+1}.x" if int(val) < 6 else
f"w{(int(val)-6)%7}.{'yzw'[(int(val)-6)//7]}") for cast,val in casts}
for old,new in sorted(mapping.items(),key=lambda item:-len(item[0])):
source=re.sub(rf"\b{old}\b",new,source)
return source
def pack_weight_buffer(weight:UOp) -> UOp:
original=np.asarray(weight.buffer.numpy()).reshape(-1)
outputs=original.size//896
assert outputs*896 == original.size
packed=np.empty((outputs,2,7,7,4),dtype=np.float16)
for output in range(outputs):
for parity in range(2):
for row in range(7):
base=output*896+parity+row*28
for tap in range(7):
for component in range(4): packed[output,parity,row,tap,component]=original[base+component*224+tap*4]
buf=Buffer("QCOM",packed.size,weight.dtype,initial_value=bytearray(packed.tobytes()))
return UOp.from_buffer(buf)
def patch_conv(model) -> int:
outer=model.captured.linear.src[0]
batch=outer.src[0].src[0].src
new_batch=[]
replaced=0
for call in batch:
name=plain_name(call.src[0].arg.name) if call.op is Ops.CALL and call.src[0].op is Ops.PROGRAM else ""
if name in TARGETS:
program=call.src[0]
source=packed_source(program.src[2].arg)
lib=Device["QCOM"].compiler.compile_cached(source)
program=program.replace(src=program.src[:2]+(program.src[2].replace(arg=source),program.src[3].replace(arg=lib)))
call=call.replace(src=(program,call.src[1],call.src[2],pack_weight_buffer(call.src[3]),*call.src[4:]))
replaced+=1
new_batch.append(call)
model.captured._linear=model.captured.linear.substitute({outer:create_graph_call(new_batch)},walk=True)
model.captured.__dict__.pop("linear",None)
return replaced
def main() -> None:
parser=argparse.ArgumentParser()
parser.add_argument("input")
parser.add_argument("output")
args=parser.parse_args()
with open(args.input,"rb") as f:model=pickle.load(f)
print("patched",patch_conv(model))
with open(args.output,"wb") as f:pickle.dump(model,f)
if __name__ == "__main__":main()
@@ -1,111 +0,0 @@
#!/usr/bin/env python3
"""Reschedule independent texture addresses in the dominant vision projection."""
import argparse, pickle, struct
from tinygrad.engine.jit import create_graph_call
from tinygrad.uop.ops import Ops
from extra.gemm.qcom_openpilot_ir3 import add_s as ADD_S, branch as BR, isam_f32 as ISAM_F32, mov_f32 as MOV_F32, nop as NOP, plain_name
TARGET="r_32_192_4_4_64_4"
FIRST_CONV_TARGET="r_64_32_16_4_4_6_3_3_4"
FORWARD_STYLE={TARGET,"r_32_64_4_4_64_4","r_8_384_4_4_128_4"}
GAP_STYLE={"r_512_16_4_4_16_4","r_512_48_4_4_16_4","r_128_32_4_4_32_4","r_128_96_4_4_32_4"}
INVERSE_W_STYLE={"r_512_16_4_4_48_4","r_128_32_4_4_96_4"}
INVERSE_STYLE={"r_32_64_4_4_192_4"}
TARGETS=FORWARD_STYLE|GAP_STYLE|INVERSE_W_STYLE|INVERSE_STYLE|{FIRST_CONV_TARGET}
def schedule_first_conv(lib:bytes) -> bytes:
image_offset, image_size=struct.unpack_from("<I",lib,0xc0)[0],struct.unpack_from("<I",lib,0x100)[0]
instrs=[lib[i:i+8] for i in range(image_offset,image_offset+image_size,8)]
if len(instrs) != 262: raise RuntimeError(f"expected 262 first-conv instructions, got {len(instrs)}")
# Use registers that the subsequent texture loads overwrite, allowing all
# eight independent input/weight coordinates to precede the texture reads.
addresses=[]
for index,(register,offset) in enumerate(zip(("r0","r2","r3","r4"),(-36,-24,-12,0))):
addresses.append(MOV_F32(f"{register}.x","r16.w",ss=index > 0) if offset == 0 else
ADD_S(f"{register}.x","r16.w",offset,ss=index > 0))
addresses.append(MOV_F32(f"{register}.y","r16.z"))
addresses.extend(instrs[i] for i in (48,51,54,57))
loads=[ISAM_F32(dst,f"{coord}.x",tex=0) for dst,coord in zip(("r7.x","r6.x","r1.x","r0.x"),("r0","r2","r3","r4"))]
loads.extend(ISAM_F32(dst,coord,tex=1) for dst,coord in zip(("r2.x","r3.x","r4.x","r5.x"),("r8.x","r8.z","r9.x","r9.z")))
out=instrs[:32]+addresses+loads+instrs[60:86]
out.append(BR(31-len(out)))
out.extend(instrs[87:92])
out.append(BR(26-len(out)))
out.extend(instrs[93:99])
out.append(BR(24-len(out)))
out.extend(instrs[100:])
out.extend([NOP()]*(len(instrs)-len(out)))
if len(out) != len(instrs): raise RuntimeError(f"scheduled first conv has {len(out)} instructions")
return lib[:image_offset]+b"".join(out)+lib[image_offset+image_size:]
def schedule_native_f16(lib:bytes) -> bytes:
image_offset, image_size=struct.unpack_from("<I",lib,0xc0)[0],struct.unpack_from("<I",lib,0x100)[0]
instrs=[lib[i:i+8] for i in range(image_offset,image_offset+image_size,8)]
if len(instrs) != 222: raise RuntimeError(f"expected 222 native-FP16 instructions, got {len(instrs)}")
addresses=(21,24,27,30,33,40,47,54)
loads=(23,26,29,32,35,42,49,56)
mads=tuple(range(36,40))+tuple(range(43,47))+tuple(range(50,54))+tuple(range(57,61))
out=instrs[:21]+[instrs[i] for i in addresses]+[instrs[i] for i in loads]+[instrs[i] for i in mads]+instrs[61:67]
out.append(BR(21-len(out)))
out.extend(instrs[68:])
out.extend([NOP()]*(len(instrs)-len(out)))
if len(out) != len(instrs): raise RuntimeError(f"scheduled native FP16 has {len(out)} instructions")
return lib[:image_offset]+b"".join(out)+lib[image_offset+image_size:]
def schedule_loads(lib:bytes, name:str) -> bytes:
image_offset, image_size=struct.unpack_from("<I",lib,0xc0)[0],struct.unpack_from("<I",lib,0x100)[0]
instrs=[lib[i:i+8] for i in range(image_offset,image_offset+image_size,8)]
if name == TARGET and len(instrs) == 222: return schedule_native_f16(lib)
if len(instrs) < 160: raise RuntimeError(f"expected projection shader, got {len(instrs)} instructions")
# The compiler emits address, rpt5 nop, texture-read eight times. Calculate
# every independent address first, then issue the reads as one contiguous run.
if name in FORWARD_STYLE: start,body_end=20,66
elif name in GAP_STYLE: start,body_end=26,84
elif name in INVERSE_W_STYLE: start,body_end=19,76
elif name in INVERSE_STYLE: start,body_end=16,73
else: raise RuntimeError(f"unsupported projection {name}")
address_indices=tuple(start+3*i for i in range(8))
load_indices=tuple(start+3*i+2 for i in range(8))
out=instrs[:start]+[instrs[i] for i in address_indices]+[instrs[i] for i in load_indices]+instrs[start+24:body_end]
branch_index=len(out)
out.append(BR(start-branch_index))
out.extend(instrs[body_end+1:])
out.extend([NOP()]*(len(instrs)-len(out)))
if len(out) != len(instrs): raise RuntimeError(f"scheduled image has {len(out)} instructions")
return lib[:image_offset]+b"".join(out)+lib[image_offset+image_size:]
def patch_projection(model) -> int:
outer=model.captured.linear.src[0]
batch=outer.src[0].src[0].src
new_batch=[]
cache={}
replaced=0
for call in batch:
name=plain_name(call.src[0].arg.name) if call.op is Ops.CALL and call.src[0].op is Ops.PROGRAM else ""
if name in TARGETS:
program=call.src[0]
old=program.src[3].arg
if old not in cache: cache[old]=schedule_first_conv(old) if name == FIRST_CONV_TARGET else schedule_loads(old,name)
program=program.replace(src=program.src[:3]+(program.src[3].replace(arg=cache[old]),))
call=call.replace(src=(program,*call.src[1:]))
replaced+=1
new_batch.append(call)
model.captured._linear=model.captured.linear.substitute({outer:create_graph_call(new_batch)},walk=True)
model.captured.__dict__.pop("linear",None)
return replaced
def main() -> None:
parser=argparse.ArgumentParser()
parser.add_argument("input")
parser.add_argument("output")
args=parser.parse_args()
with open(args.input,"rb") as f:model=pickle.load(f)
print("patched",patch_projection(model))
with open(args.output,"wb") as f:pickle.dump(model,f)
if __name__ == "__main__":main()
-165
View File
@@ -1,165 +0,0 @@
#!/usr/bin/env python3
"""Replace selected driving_vision 1x1 convolutions with vector FP16-acc kernels."""
import struct
from dataclasses import replace
from tinygrad.engine.jit import create_graph_call
from tinygrad.uop.ops import Ops
from extra.gemm.qcom_openpilot_ir3 import branch as BR, mad_f32 as MAD_F32, nop as NOP, plain_name
TARGET = "r_32_192_4_4_64_4"
INVERSE_TARGET = "r_32_64_4_4_192_4"
FIRST_CONV_TARGET = "r_64_32_16_4_4_6_3_3_4"
FULL_Y_TARGETS = {TARGET, "r_32_64_4_4_64_4"}
FULL_Z_TARGETS = {"r_8_384_4_4_128_4"}
GAP_Y_TARGETS = {"r_512_16_4_4_16_4", "r_512_48_4_4_16_4", "r_128_32_4_4_32_4", "r_128_96_4_4_32_4"}
INVERSE_W_TARGETS = {"r_512_16_4_4_48_4", "r_128_32_4_4_96_4"}
OTHER_INVERSE_TARGETS: set[str] = set()
FP32_TARGETS = FULL_Y_TARGETS | FULL_Z_TARGETS | GAP_Y_TARGETS | INVERSE_W_TARGETS | OTHER_INVERSE_TARGETS | {INVERSE_TARGET, FIRST_CONV_TARGET}
def pack_fp32_mads(lib:bytes, component:str="y") -> bytes:
image_offset, image_size = struct.unpack_from("<I", lib, 0xc0)[0], struct.unpack_from("<I", lib, 0x100)[0]
image = lib[image_offset:image_offset+image_size]
instrs = [image[i:i+8] for i in range(0, len(image), 8)]
if len(instrs) < 116: raise RuntimeError(f"expected FP32 matmul loop through instruction 115, got {len(instrs)}")
out = instrs[:44]
for k_component, weight in zip("xyzw", ("r5.x", "r2.x", "r3.x", "r4.x")):
for acc, activation in zip(tuple(f"r{reg}.{component}" for reg in range(13, 17)), ("r7", "r6", "r1", "r0")):
out.append(MAD_F32(acc, f"{activation}.{k_component}", weight, acc, rpt=3,
sy=len(out) == 44, r=True))
out += instrs[108:114]
branch_index = len(out)
out.append(BR(20-branch_index))
out += instrs[115:]
out += [NOP()] * (len(instrs)-len(out))
if len(out) != len(instrs): raise RuntimeError(f"packed FP32 image has {len(out)} instructions")
return lib[:image_offset] + b"".join(out) + lib[image_offset+image_size:]
def pack_gap_y_fp32_mads(lib:bytes) -> bytes:
image_offset, image_size = struct.unpack_from("<I", lib, 0xc0)[0], struct.unpack_from("<I", lib, 0x100)[0]
image = lib[image_offset:image_offset+image_size]
instrs = [image[i:i+8] for i in range(0, len(image), 8)]
if len(instrs) < 237: raise RuntimeError(f"expected at least 237 gap-Y instructions, got {len(instrs)}")
out = instrs[:66]
for component, weight in zip("xyzw", ("r5.x", "r2.x", "r3.x", "r4.x")):
for acc, activation in zip(("r14.y", "r15.y", "r16.y"), ("r6", "r1", "r0")):
out.append(MAD_F32(acc, f"{activation}.{component}", weight, acc, rpt=3, r=True))
out += instrs[114:120]
branch_index = len(out)
out.append(BR(26-branch_index))
out += instrs[121:]
out += [NOP()] * (len(instrs)-len(out))
if len(out) != len(instrs): raise RuntimeError(f"packed gap-Y image has {len(out)} instructions")
return lib[:image_offset] + b"".join(out) + lib[image_offset+image_size:]
def pack_inverse_w_fp32_mads(lib:bytes) -> bytes:
image_offset, image_size = struct.unpack_from("<I", lib, 0xc0)[0], struct.unpack_from("<I", lib, 0x100)[0]
image = lib[image_offset:image_offset+image_size]
instrs = [image[i:i+8] for i in range(0, len(image), 8)]
if len(instrs) < 174: raise RuntimeError(f"expected at least 174 inverse-W instructions, got {len(instrs)}")
out = instrs[:59]
for component, weight in zip("xyzw", ("r5.x", "r2.x", "r3.x", "r4.x")):
for acc, activation in zip(("r13.w", "r14.w", "r15.w"), ("r6", "r1", "r0")):
out.append(MAD_F32(acc, f"{activation}.{component}", weight, acc, rpt=3, r=True))
out += instrs[107:112]
branch_index = len(out)
out.append(BR(19-branch_index))
out += instrs[113:]
out += [NOP()] * (len(instrs)-len(out))
if len(out) != len(instrs): raise RuntimeError(f"packed inverse-W image has {len(out)} instructions")
return lib[:image_offset] + b"".join(out) + lib[image_offset+image_size:]
def pack_other_inverse_fp32_mads(lib:bytes) -> bytes:
image_offset, image_size = struct.unpack_from("<I", lib, 0xc0)[0], struct.unpack_from("<I", lib, 0x100)[0]
image = lib[image_offset:image_offset+image_size]
instrs = [image[i:i+8] for i in range(0, len(image), 8)]
if len(instrs) != 179: raise RuntimeError(f"expected 179 other-inverse instructions, got {len(instrs)}")
# The first output vector is split by loop-control registers. Keep it scalar;
# the remaining three vectors are contiguous from r14.y through r17.x.
out = instrs[:60]
for component, weight in zip("xyzw", ("r5.x", "r2.x", "r3.x", "r4.x")):
for acc, activation in zip(("r14.y", "r15.y", "r16.y"), ("r6", "r1", "r0")):
out.append(MAD_F32(acc, f"{activation}.{component}", weight, acc, rpt=3, r=True))
out += instrs[108:114]
branch_index = len(out)
out.append(BR(20-branch_index))
out += instrs[115:]
out += [NOP()] * (len(instrs)-len(out))
if len(out) != len(instrs): raise RuntimeError(f"packed other-inverse image has {len(out)} instructions")
return lib[:image_offset] + b"".join(out) + lib[image_offset+image_size:]
def pack_first_conv_fp32_mads(lib:bytes) -> bytes:
image_offset, image_size = struct.unpack_from("<I", lib, 0xc0)[0], struct.unpack_from("<I", lib, 0x100)[0]
image = lib[image_offset:image_offset+image_size]
instrs = [image[i:i+8] for i in range(0, len(image), 8)]
if len(instrs) != 262: raise RuntimeError(f"expected 262 first-conv instructions, got {len(instrs)}")
out = instrs[:60]
first = True
for component, weight in zip("xyzw", ("r5", "r2", "r3", "r4")):
out.append(MAD_F32("r11.w", f"r7.{component}", f"{weight}.x", "r11.w", sy=first, r=True))
first = False
out.append(MAD_F32("r12.y", f"r7.{component}", f"{weight}.y", "r12.y", rpt=2, r=True))
for component, weight in zip("xyzw", ("r5.x", "r2.x", "r3.x", "r4.x")):
for acc, activation in zip(("r13.x", "r14.x", "r15.x"), ("r6", "r1", "r0")):
out.append(MAD_F32(acc, f"{activation}.{component}", weight, acc, rpt=3, r=True))
out += instrs[124:130]
branch_index = len(out)
out.append(BR(31-branch_index))
out += instrs[131:]
# Compacting the innermost loop also relocates the two enclosing-loop branches.
# Their targets remain in the untouched prologue, so rebuild their relative offsets.
out[92] = BR(26-92)
out[99] = BR(24-99)
out += [NOP()] * (len(instrs)-len(out))
if len(out) != len(instrs): raise RuntimeError(f"packed first-conv image has {len(out)} instructions")
return lib[:image_offset] + b"".join(out) + lib[image_offset+image_size:]
def pack_inverse_fp32_mads(lib:bytes) -> bytes:
image_offset, image_size = struct.unpack_from("<I", lib, 0xc0)[0], struct.unpack_from("<I", lib, 0x100)[0]
image = lib[image_offset:image_offset+image_size]
instrs = [image[i:i+8] for i in range(0, len(image), 8)]
if len(instrs) != 175: raise RuntimeError(f"expected 175 inverse instructions, got {len(instrs)}")
# The first output vector straddles loop-control r13.x, so retain its scalar
# instructions. The remaining r14/r15/r16 accumulator vectors are contiguous.
out = instrs[:56]
for component, weight in zip("xyzw", ("r5.x", "r2.x", "r3.x", "r4.x")):
for acc, activation in zip(("r14.x", "r15.x", "r16.x"), ("r6", "r1", "r0")):
out.append(MAD_F32(acc, f"{activation}.{component}", weight, acc, rpt=3, r=True))
out += instrs[104:109]
branch_index = len(out)
out.append(BR(16-branch_index))
out += instrs[110:]
out += [NOP()] * (len(instrs)-len(out))
if len(out) != len(instrs): raise RuntimeError(f"packed inverse image has {len(out)} instructions")
return lib[:image_offset] + b"".join(out) + lib[image_offset+image_size:]
def patch_fp32_rpt(jit, names:set[str]|None=None) -> int:
"""Apply the verified FP32-accumulate repeat packing to a captured vision JIT."""
outer = jit.captured.linear.src[0]
batch = outer.src[0].src[0].src
new_batch, replaced = [], 0
for call in batch:
name = plain_name(call.src[0].arg.name) if call.op is Ops.CALL and call.src[0].op is Ops.PROGRAM else ""
if name in FP32_TARGETS and (names is None or name in names):
program = call.src[0]
patch = (pack_fp32_mads if name in FULL_Y_TARGETS else
(lambda lib:pack_fp32_mads(lib, "z")) if name in FULL_Z_TARGETS else
pack_gap_y_fp32_mads if name in GAP_Y_TARGETS else
pack_inverse_w_fp32_mads if name in INVERSE_W_TARGETS else
pack_other_inverse_fp32_mads if name in OTHER_INVERSE_TARGETS else
pack_first_conv_fp32_mads if name == FIRST_CONV_TARGET else pack_inverse_fp32_mads)
program = program.replace(src=program.src[:3] + (program.src[3].replace(arg=patch(program.src[3].arg)),))
if name == INVERSE_TARGET:
program = program.replace(arg=replace(program.arg, global_size=(8, 2, 1), local_size=(8, 16, 1)))
call, replaced = call.replace(src=(program, *call.src[1:])), replaced+1
new_batch.append(call)
if replaced:
jit.captured._linear = jit.captured.linear.substitute({outer:create_graph_call(new_batch)}, walk=True)
jit.captured.__dict__.pop("linear", None)
return replaced
+140 -181
View File
@@ -2,8 +2,8 @@ from __future__ import annotations
from typing import cast, Callable, TypeVar, Generic, Any
import struct, functools, time, collections, itertools
from dataclasses import replace, dataclass
from tinygrad.helpers import DEV, getenv, select_first_inited, select_by_name, suppress_finalizing, dedup, pluralize
from tinygrad.helpers import to_tuple, round_up, partition, data64_le, panic
from tinygrad.helpers import DEV, getenv, select_first_inited, select_by_name, suppress_finalizing, dedup, pluralize, JIT_BATCH_SIZE
from tinygrad.helpers import to_tuple, round_up, partition, data64_le, panic, ContextVar
from tinygrad.device import Device, Buffer, BufferSpec, Compiled, LRUAllocator, MultiBuffer
from tinygrad.uop.ops import Ops, sint, UOp, UPat, PatternMatcher, KernelInfo, graph_rewrite, track_rewrites, GroupOp
from tinygrad.uop.symbolic import symbolic
@@ -19,6 +19,8 @@ HCQDeviceType = TypeVar('HCQDeviceType', bound='HCQ2Compiled')
# *****************
# 0. helpers
HCQ_RUNTIME_DEV = ContextVar("HCQ_RUNTIME_DEV", "CPU")
HCQ_DEVS = frozenset(("AMD",))
HCQ_P2P_DEVS = HCQ_DEVS | frozenset(("CPU",))
HCQ_CACHE_TAGS = frozenset(("program", "systems", "template"))
@@ -56,7 +58,7 @@ def make_patch(buf:UOp, off:sint, val:UOp, dtype=None) -> UOp:
return buf.index(UOp.const(dtypes.int, off // buf.dtype.itemsize)).store(val.simplify().cast(dtype or buf.dtype))
def make_binary_patch(buf:UOp, blob:bytes) -> UOp:
data = UOp(Ops.BITCAST, buf.dtype, (UOp(Ops.BINARY, dtypes.uint8, src=(), arg=blob),))
data = UOp(Ops.BINARY, src=(), arg=blob).bitcast(buf.dtype)
r = UOp.range(len(blob) // buf.dtype.itemsize, 0, dtype=dtypes.int, src=(buf, data))
return buf.index(r).store(data.index(r).load()).end(r)
@@ -100,159 +102,138 @@ def stage_copy(dst:UOp, src:UOp) -> UOp|None:
pm_insert_copy_staging = PatternMatcher([(UPat(Ops.CALL, src=(UPat(Ops.COPY), UPat(name="dst"), UPat(name="src"))), stage_copy)])
# *****************
# 2.1. tag hcq calls
def tag_hcq_call(ctx:itertools.count, call:UOp) -> UOp:
if (hcq_devs:=next((b.device for b in call.src[1:] if all_devices_in(b.device, HCQ_DEVS)), None)) is None: return call
queue = "COMPUTE:0" if call.src[0].op is Ops.PROGRAM else "COPY:0"
info = HCQInfo(get_call_name(call, get_call_arg_uops(call)), estimate_uop(call), to_tuple(hcq_devs), queue)
return call.replace(arg=replace(call.arg, aux=info)).rtag(next(ctx))
pm_tag_hcq_calls = PatternMatcher([(UPat(Ops.LINEAR, name="l"), lambda ctx, l: l.replace(src=tuple(tag_hcq_call(ctx, s) for s in l.src)))])
# *****************
# 2.2. deps tracking
# device.timeline_signal/value are the per-device schedule epoch. Before a schedule queue accesses memory owned by device N for the first time,
# it waits for device[N].timeline_signal >= device[N].timeline_value - 1. This orders the schedule after all prior schedules that touched device N.
#
# queue.timeline_signal/value are per-queue progress counters used only inside a schedule.
# Only the owner queue signals its queue.timeline_signal. Values are monotonic.
#
# At schedule end, one finalizer queue per touched device[N] waits for every active queue on device[N] to reach its schedule-local
# final queue.timeline value, then signals device[N].timeline_signal with the schedule's reserved device epoch. After that, buffers/transients
# for device N from this schedule are safe for the next schedule
#
# C programs reserve and bump timeline values, then patch command buffers with the concrete wait/signal values.
# 2. deps
class HCQDepsTracker(DepsTracker):
@staticmethod
def _key(buf:Any) -> tuple[Any, int, int]:
return (buf.arg.slot, 0, buf.max_numel() * buf.dtype.itemsize) if isinstance(buf, UOp) else DepsTracker._key(buf)
def make_deps(u:UOp, dep_lanes:list[tuple[UOp, int, int]], nlanes:int) -> UOp:
deps:dict[UOp, list[int|None]] = collections.defaultdict(lambda: [None]*nlanes)
for dep, dlane, lane in dep_lanes: deps[dep][lane] = dlane
return u.after(*deps, arg=tuple(tuple(v) for v in deps.values()))
def sched_sync(ctx:DepsTracker, call:UOp) -> UOp|None:
if not isinstance(call.arg.aux, HCQInfo): return None
def _get_call_bufs_by_lane(call:UOp, devices:tuple[str, ...]) -> list[list[Any]]:
refs = get_call_arg_uops(call)
outs, _ = get_call_outs_ins(call)
devices, queue = call.arg.aux.device, call.arg.aux.queue
return [[b if b.op is Ops.PARAM else mb.bufs[lane] if isinstance(mb:=b.buffer, MultiBuffer) else mb for b in refs] for lane in range(len(devices))]
dep_lanes:list[tuple[UOp, int, int]] = []
for lane, d in enumerate(devices):
lane_refs = [b if b.op is Ops.PARAM else mb.bufs[lane] if isinstance(mb:=b.buffer, MultiBuffer) else mb for b in refs]
for dep, dlane in ctx.access_resources(lane_refs, outs, (call, lane)): dep_lanes.append((dep, dlane, lane))
def _get_deps(ctx:DepsTracker, bufs_by_lane:list[list[Any]], write, key:tuple[tuple[str, ...], str, int]) -> list[tuple[tuple, int, int]]:
dep_lanes:list[tuple[tuple, int, int]] = []
for lane, bufs in enumerate(bufs_by_lane):
dep_lanes += [(dep, dlane, lane) for dep, dlane in ctx.access_resources(bufs, write if write is not None else range(len(bufs)), (key, lane))]
return dep_lanes
def _build_wait_cmds(dep_lanes:list[tuple[tuple, int, int]], devices:tuple[str, ...], queue:str) -> tuple[list[UOp], set[int]]:
# opt1: same-queue ops are fifo-ordered
if devices[0].split(":")[0] in {"AMD", "QCOM"} or queue.startswith("COPY"):
dep_lanes = [(dep, dlane, lane) for dep, dlane, lane in dep_lanes if (dep.arg.aux.device[dlane], dep.arg.aux.queue) != (devices[lane], queue)]
dep_lanes = [(dep, dlane, lane) for dep, dlane, lane in dep_lanes if (dep[0][dlane], dep[1]) != (devices[lane], queue)]
# keep latest dep per (dep device, queue, cur lane)
latest = {((dep.arg.aux.device[dlane], dep.arg.aux.queue), lane): (dep, dlane) for dep, dlane, lane in sorted(dep_lanes, key=lambda x: x[0].tag)}
return make_deps(call, [(dep, dlane, lane) for (_, lane), (dep, dlane) in latest.items()], len(devices))
pm_sched_sync = PatternMatcher([(UPat(Ops.CALL, name="call"), sched_sync)])
# opt2: keep latest dep per (dep device, queue, cur lane)
latest = {((dep[0][dlane], dep[1]), lane): (dep, dlane) for dep, dlane, lane in sorted(dep_lanes, key=lambda x: x[0][2])}
deps:dict[tuple, list[int|None]] = collections.defaultdict(lambda: [None]*len(devices))
for (_, lane), (dep, dlane) in latest.items(): deps[dep][lane] = dlane
waits = []
for (ddevs, dqueue, dtag), lanes in deps.items():
sig = make_mstack([make_signal(d if dl is None else ddevs[dl], queue=dqueue, sentinel=dl is None) for dl, d in zip(lanes, devices)])
val = make_mstack([make_signal_value(d if dl is None else ddevs[dl], queue=dqueue) for dl, d in zip(lanes, devices)])
waits.append((sig.index(zero:=UOp.const(dtypes.int, 0)).load() >= val.index(zero) + dtag).wait())
return waits, {dtag for _, _, dtag in deps}
def _build_finalizers(batch:list[tuple[UOp, tuple[str, ...]]], batch_info:list[tuple[tuple[str, ...], str]],
tracker:HCQDepsTracker) -> tuple[list[UOp], set[int]]:
# collect all buffers which belong to devices
dev_bufs:dict[str, dict[int, Any]] = collections.defaultdict(dict)
for call, devices in batch:
for b in itertools.chain.from_iterable(_get_call_bufs_by_lane(call, devices)):
for bd in to_tuple(b.device): dev_bufs[bd][id(b)] = b
zero, n, finalizers, waited = UOp.const(dtypes.int, 0), len(batch_info), [], set()
for _, devgroup in itertools.groupby(sorted(dedup([d for devs, _ in batch_info for d in devs])), key=lambda d: d.split(":")[0]):
devs = tuple(devgroup)
# to finalize the batch, sync all accesses from other devices to buffers that belong to this device
fin_deps = [dl for dl in _get_deps(tracker, [list(dev_bufs[d].values()) for d in devs], None, key=(devs, "COMPUTE:0", n)) if dl[0][2] < n]
waits, cur_waited = _build_wait_cmds(fin_deps, devs, "COMPUTE:0")
waited |= cur_waited
# wait the syncs, store the device epoch; value bumps are a separate call: no lane may bump until every lane has patched its waits
submit = make_submit(*waits, make_signal(devs).store((tl:=make_signal_value(devs)).index(zero)), devs=devs, queue="COMPUTE:0")
upd = [(tl, 1)] + [(make_signal_value(devs, queue=qn), n) for qn in dedup([qn for bdevs, qn in batch_info if set(bdevs) & set(devs)])]
bump = UOp.barrier(*[s.index(zero, dtype=s.dtype).store(s.index(zero) + inc) for s, inc in upd])
finalizers += [UOp.custom_function("hcq", b.sink()).call(aux=HCQInfo("hcq_finalizer", Estimates(), devs, "COMPUTE:0")) for b in (submit, bump)]
return finalizers, waited
def _finalize_batch(batch:list[tuple[UOp, tuple[str, ...]]]) -> list[UOp]:
batch_info = [(devices, "COMPUTE:0" if call.src[0].op is Ops.PROGRAM else "COPY:0") for call, devices in batch]
# schedule deps
waited:set[int] = set()
deps_tracker = HCQDepsTracker()
call_waits:list[list[UOp]] = []
for tag, ((call, _), (devices, queue)) in enumerate(zip(batch, batch_info)):
deps = _get_deps(deps_tracker, _get_call_bufs_by_lane(call, devices), get_call_outs_ins(call)[0], key=(devices, queue, tag))
cmds, cur_waited = _build_wait_cmds(deps, devices, queue)
call_waits.append(cmds)
waited |= cur_waited
# build finalizers
finalizers, finalizer_waited = _build_finalizers(batch, batch_info, deps_tracker)
waited |= finalizer_waited
src = []
for tag, ((call, _), (devices, queue), cmds) in enumerate(zip(batch, batch_info, call_waits)):
# first queue use, sync prior device work with main signal
if batch_info.index((devices, queue)) == tag:
epoch = (make_signal(devices).index(0).load() >= make_signal_value(devices).index(0) - 1).wait()
cmds = [UOp(Ops.BARRIER), epoch] + cmds
# signal queue timeline if someone waits for us
store = make_signal(devices, queue=queue).store(make_signal_value(devices, queue=queue).index(0) + tag) if tag in waited else None
# and make hcq call
info = HCQInfo(get_call_name(call, get_call_arg_uops(call)), estimate_uop(call), devices, queue)
cmds = [*cmds, call.replace(arg=replace(call.arg, aux=info))] + ([store] if store is not None else [])
src.append(UOp.custom_function("hcq", make_submit(*cmds, devs=devices, queue=queue).sink()).call(name="hcq", aux=info))
return src + finalizers
def sched_hcq_batches(l:UOp) -> UOp:
srcs:list[UOp] = []
batch:list[tuple[UOp, tuple[str, ...]]] = []
for call in l.src:
if (devs:=next((b.device for b in call.src[1:] if all_devices_in(b.device, HCQ_DEVS)), None)) is not None: batch.append((call, to_tuple(devs)))
else: srcs, batch = srcs + _finalize_batch(batch) + [call], []
return l.replace(src=tuple(srcs + _finalize_batch(batch)))
pm_sched_hcq_batches = PatternMatcher([(UPat(Ops.LINEAR, name="l"), sched_hcq_batches)])
# *****************
# 2.3. merge into queues
# 3. merge into queues
def _merged_hcq_call(calls:list[UOp]):
info = replace(unwrap_after(calls[0]).arg.aux, estimates=sum((unwrap_after(c).arg.aux.estimates for c in calls), start=Estimates()))
cmdbuf = make_submit(*calls, devs=info.device, queue=info.queue)
return UOp.custom_function("hcq", cmdbuf.sink()).call(name="hcq", aux=info)
def _merged_hcq_call(calls:list[UOp]) -> UOp: # TODO: simplify?
if len(calls) == 1: return calls[0]
info = replace(calls[0].arg.aux, name=f"submit {calls[0].arg.aux.queue} ({len(calls)})",
estimates=sum((c.arg.aux.estimates for c in calls), start=Estimates()))
cmds = [cmd for c in calls for cmd in get_submit(c).src[0].src]
return UOp.custom_function("hcq", make_submit(*cmds, devs=info.device, queue=info.queue).sink()).call(name="hcq", aux=info)
def merge_queues(linear:UOp) -> UOp:
new_src:list[UOp] = []
opened_qs:dict[tuple[tuple[str, ...], str], list[UOp]] = {} # (devs, queue) -> list of calls, kept in submit order
opened_qs:dict[tuple[tuple[str, ...], str], list[UOp]] = {} # (devs, queue) -> list of hcq calls, kept in submit order
limits = collections.defaultdict(lambda: JIT_BATCH_SIZE.value)
for call in linear.src:
if not isinstance(unwrap_after(call).arg.aux, HCQInfo):
if not isinstance(info:=call.arg.aux, HCQInfo) or info.name == "hcq_finalizer": # non-hcq call or finalizer: close all open queues
new_src += [_merged_hcq_call(opened_qs.pop(k)) for k in list(opened_qs)] + [call]
continue
devices, queue = unwrap_after(call).arg.aux.device, unwrap_after(call).arg.aux.queue
if (old:=opened_qs.pop((devices, queue), None)) is not None: new_rec = old + [call]
if (old:=opened_qs.pop(key:=(info.device, info.queue), None)) is not None:
if limits[key] and len(old) >= limits[key]: new_src, old, limits[key] = new_src + [_merged_hcq_call(old)], [], limits[key] * 2
new_rec = old + [call]
else:
# no such queue opened: close every open submit on this queue that shares a device, so submit order is kept
closing = [k for k in opened_qs if k[1] == queue and set(k[0]) & set(devices)]
closing = [k for k in opened_qs if k[1] == info.queue and set(k[0]) & set(info.device)]
new_src += [_merged_hcq_call(opened_qs.pop(k)) for k in closing]
new_rec = [call]
opened_qs[(devices, queue)] = new_rec
opened_qs[(info.device, info.queue)] = new_rec
return linear.replace(src=tuple(new_src + [_merged_hcq_call(c) for c in opened_qs.values()]))
pm_merge_queues = PatternMatcher([(UPat(Ops.LINEAR, name="linear"), merge_queues)])
# *****************
# 2.4. finalizer
def add_finalizer(ctx:itertools.count, linear:UOp) -> UOp:
# collect by device type
parts:dict[str, list[UOp]] = collections.defaultdict(list)
for call in linear.src:
if (c:=unwrap_after(call)).src[0].op is not Ops.CUSTOM_FUNCTION or c.src[0].arg != "hcq": continue
parts[c.arg.aux.device[0].split(':')[0]].append(unwrap_after(get_submit(call).src[0].src[0]))
nbump = next(ctx)
finalizers = []
for calls in parts.values():
devs = tuple(dedup(d for call in calls for d in unwrap_after(call).arg.aux.device))
zero = UOp.const(dtypes.int, 0)
tl = make_signal_value(devs)
# split each (multi-device) call into per-device deps, then store the device timeline value into the device signal after them
dep_lanes = [(call, dlane, devs.index(d)) for call in calls for dlane, d in enumerate(unwrap_after(call).arg.aux.device)]
store = make_deps(make_signal(devs).store(tl.index(zero)), dep_lanes, len(devs))
submit = make_submit(store, devs=devs, queue="COMPUTE:0")
upd = [(tl, 1)] + [(make_signal_value(devs, queue=qn), nbump) for qn in dedup([unwrap_after(call).arg.aux.queue for call in calls])]
patches = [s.after(submit).index(zero, dtype=s.dtype).store(s.index(zero) + inc) for s, inc in upd]
finalizers.append(UOp.custom_function("hcq", UOp.barrier(*patches).sink()).call(aux=HCQInfo("hcq finalizer", Estimates(), devs, "COMPUTE:0")))
return linear.replace(src=linear.src + tuple(finalizers))
pm_add_finalizer = PatternMatcher([(UPat(Ops.LINEAR, name="linear"), add_finalizer)])
# *****************
# 2.5. global sync
def add_global_sync(ctx:set[tuple[str, ...]], submit:UOp, q:UOp) -> UOp|None:
if (devs:=q.arg[0]) in ctx: return None
ctx.add(devs)
# some devices from a command buffer might be used for the first time this schedule, so we wait for their global timeline epoch.
wait = make_signal(devs).wait(make_signal_value(devs).index(UOp.const(dtypes.int, 0)) - 1)
return submit.replace(src=(q.replace(src=(UOp(Ops.BARRIER, dtypes.void), wait, *q.src)),))
pm_add_global_sync = PatternMatcher([(UPat(Ops.CUSTOM_FUNCTION, arg="submit_cmdbuf", src=(UPat(Ops.LINEAR, name="q"),), name="submit"), add_global_sync)])
# *****************
# 3.1. lower loads/stores
def add_loads(ctx:set[int], submit:UOp, q:UOp) -> UOp|None:
cur_devs = q.arg[0]
new_src:list[UOp] = []
for s in q.src:
if s.op is Ops.AFTER:
for lanes, dep in zip(s.arg, s.src[1:]):
devs, queue = dep.arg.aux.device, dep.arg.aux.queue
ctx.add(dep.tag) # mark op to update signal.
sig = make_mstack([make_signal(d if dl is None else devs[dl], queue=queue, sentinel=dl is None) for dl, d in zip(lanes, cur_devs)])
val = make_mstack([make_signal_value(d if dl is None else devs[dl], queue=queue) for dl, d in zip(lanes, cur_devs)]).index(UOp.const(dtypes.int, 0))
new_src.append(sig.wait(val + dep.tag))
s = s.src[0]
new_src.append(s)
return submit.replace(src=(q.replace(src=tuple(new_src)),))
pm_add_inner_loads = PatternMatcher([(UPat(Ops.CUSTOM_FUNCTION, arg="submit_cmdbuf", src=(UPat(Ops.LINEAR, name="q"),), name="submit"), add_loads)])
def add_stores(ctx:set[int], submit:UOp, q:UOp) -> UOp|None:
devs, queue = q.arg
new_src:list[UOp] = []
for op in q.src:
new_src.append(op)
if (sigval:=unwrap_after(op).tag) in ctx:
new_src.append(make_signal(devs, queue=queue).store(make_signal_value(devs, queue=queue).index(UOp.const(dtypes.int, 0)) + sigval))
return submit.replace(src=(q.replace(src=tuple(new_src)),))
pm_add_inner_stores = PatternMatcher([(UPat(Ops.CUSTOM_FUNCTION, arg="submit_cmdbuf", src=(UPat(Ops.LINEAR, name="q"),), name="submit"), add_stores)])
# *****************
# 4.1. hcq lowering: programs
@@ -277,7 +258,7 @@ def is_value_known_at_link(val:UOp) -> bool:
addressed_bufs = [b for g in val.toposort() if g.op is Ops.GETADDR for b in unwrap_mstack(g.buf_uop)]
# addr of input params is not known at link time
return not runtime_reads and all(b.op is not Ops.PARAM or b.tag is not None for b in addressed_bufs)
return not val.variables() and not runtime_reads and all(b.op is not Ops.PARAM or b.tag is not None for b in addressed_bufs)
def is_link_patch(p:UOp, jit:bool) -> bool:
store = p.src[0] if (is_binary_patch:=p.op is Ops.END) else p
@@ -303,35 +284,32 @@ pm_split_patches = PatternMatcher([(UPat(Ops.CALL, src=(UPat(Ops.CUSTOM_FUNCTION
# *****************
def make_addr_table(call:UOp, gaddrs:list[UOp], name:str):
def make_addr_table(call:UOp, gaddrs:list[UOp], name:str) -> tuple[dict[UOp, UOp], tuple[UOp, ...]]:
bare = {g: g.replace(src=(unwrap_after(g.src[0]),)) for g in gaddrs}
order = sorted(dedup(bare.values()), key=lambda g: ((b:=unwrap_mstack(g.buf_uop)[0]).arg.slot, to_tuple(b.tag)))
order = sorted(dedup(bare.values()), key=lambda g: ((b:=unwrap_mstack(g.buf_uop)[0]).arg.slot, repr(b.tag)))
slots, table = {g:i for i,g in enumerate(order)}, make_placeholder(call.arg.aux.device, len(order), dtypes.uint64, name)
reads = {g: table.after(*g.src[0].src[1:] if g.src[0].op is Ops.AFTER else ()).index(UOp.const(dtypes.int, slots[bare[g]])).load() for g in gaddrs}
return reads, (table.after(*[make_patch(table, i * table.dtype.itemsize, addr) for addr, i in slots.items()]),) if slots else ()
def rm_rt_getaddrs(call:UOp) -> UOp|None:
if not (gaddrs:=[u for u in call.src[0].toposort() if u.op is Ops.GETADDR]): return None
def make_blob_bufs(call:UOp, blobs:list[UOp]) -> tuple[dict[UOp, UOp], tuple[UOp, ...]]:
bufs = {b: make_placeholder(call.arg.aux.device, b.max_numel(), b.dtype, "template") for b in blobs}
return bufs, tuple(buf.after(make_binary_patch(buf, b.src[0].arg)) for b,buf in bufs.items())
def rm_rt_uops(call:UOp) -> UOp|None:
if not (rt_uops:=[u for u in call.src[0].toposort() if u.op is Ops.GETADDR or (u.op is Ops.BITCAST and u.src[0].op is Ops.BINARY)]): return None
gaddrs, blobs = partition(rt_uops, lambda u: u.op is Ops.GETADDR)
inputs, internals = partition(gaddrs, lambda g: all(x.op is Ops.PARAM and x.tag is None for x in unwrap_mstack(g.buf_uop)))
runtimes, systems = partition(internals, lambda g: any(x.tag in {"program", "kernargs", "cmdbuf"} for x in unwrap_mstack(g.buf_uop)))
# exec fills the inputs table with the input addresses every run, so it has no fill patches
(input_reads, _), (rt_reads, rt_fills), (sys_reads, sys_fills) = (make_addr_table(call, gs, name) for gs, name in
((inputs, "inputs"), (runtimes, "runtime"), (systems, "systems")))
return call.replace(src=(call.src[0].substitute(input_reads | rt_reads | sys_reads), *call.src[1:], *rt_fills, *sys_fills),
(reads, _), *tables = [make_addr_table(call, gs, n) for gs,n in ((inputs, "inputs"), (runtimes, "runtime"), (systems, "systems"))] + \
[make_blob_bufs(call, blobs)]
reads, fills = reads | {k:v for r,_ in tables for k,v in r.items()}, [f for _,fs in tables for f in fs]
return call.replace(src=(call.src[0].substitute(reads), *call.src[1:], *fills),
arg=replace(call.arg, aux=replace(call.arg.aux, input_idxs=tuple(sorted(dedup(g.buf_uop.arg.slot for g in inputs))))))
pm_rm_rt_getaddrs = PatternMatcher([(UPat(Ops.CALL, src=(UPat(Ops.CUSTOM_FUNCTION, arg="hcq"),), name="call", allow_any_len=True), rm_rt_getaddrs)])
# *****************
def rm_rt_binaries(call:UOp) -> UOp|None:
if not (blobs:=[u for u in call.src[0].toposort() if u.op is Ops.BITCAST and u.src[0].op is Ops.BINARY]): return None
blob_bufs = {blob: make_placeholder(call.arg.aux.device, blob.max_numel(), blob.dtype, "template") for blob in blobs}
fills = [buf.after(make_binary_patch(buf, blob.src[0].arg)) for blob, buf in blob_bufs.items()]
return call.replace(src=(call.src[0].substitute(blob_bufs), *call.src[1:], *fills))
pm_rm_rt_binaries = PatternMatcher([(UPat(Ops.CALL, src=(UPat(Ops.CUSTOM_FUNCTION, arg="hcq"),), name="call", allow_any_len=True), rm_rt_binaries)])
pm_rm_rt_uops = PatternMatcher([(UPat(Ops.CALL, src=(UPat(Ops.CUSTOM_FUNCTION, arg="hcq"),), name="call", allow_any_len=True), rm_rt_uops)])
# *****************
@@ -375,7 +353,7 @@ def pack_hcq_placeholders(call:UOp) -> UOp|None:
sizes[b.tag] = offs[b] + b.max_numel()
counts = collections.Counter(b.tag for b in bufs)
bases = {b.tag:make_placeholder(b.device, sizes[b.tag], b.dtype, b.tag) for b in bufs if counts[b.tag] > 1}
subs = {b:UOp(Ops.SLICE, b.dtype, (bases[b.tag], UOp.const(dtypes.index, offs.get(b, 0))), b.max_numel()) for b in bufs if b.tag in bases}
subs = {b:UOp(Ops.SLICE, b.dtype, (bases[b.tag], UOp.const(dtypes.weakint, offs.get(b, 0))), b.max_numel()) for b in bufs if b.tag in bases}
return call.replace(src=(call.src[0].substitute(subs, walk=True), *call.src[1:])) if subs else None
pm_pack_placeholders = PatternMatcher([(UPat(Ops.CALL, src=(UPat(Ops.CUSTOM_FUNCTION, arg="hcq"),), name="call", allow_any_len=True), pack_hcq_placeholders)])
@@ -383,7 +361,7 @@ pm_pack_placeholders = PatternMatcher([(UPat(Ops.CALL, src=(UPat(Ops.CUSTOM_FUNC
# 8. callify hcq programs
pm_callify_hcq = PatternMatcher([(UPat(Ops.CUSTOM_FUNCTION, arg="hcq", src=(UPat(Ops.SINK),), name="cf"),
lambda cf: cf.replace(src=(to_program(cf.src[0].replace(arg=KernelInfo("hcq_submit"), tag=1), Device["CPU"].renderer),)))])
lambda cf: cf.replace(src=(to_program(cf.src[0].replace(arg=KernelInfo("hcq_submit"), tag=1), Device[HCQ_RUNTIME_DEV.value].renderer),)))])
hcq_compile_cache:dict[tuple[bytes, bool], UOp] = {}
@@ -392,27 +370,23 @@ def hcq_compile(linear:UOp, input_uops:list[UOp]|None=None, jit=False) -> UOp:
if input_uops is not None: linear = graph_rewrite(linear, pm_replace_buffers, ctx=input_uops, walk=True, enter_calls=True, name="replace buffer")
if (final_linear:=(hcq_compile_cache.get(cache_key:=(linear.key, jit)))) is None:
# schedule
# prep
linear = linear.substitute(back_map:={s.param_like(i): s for i,s in enumerate(input_uops)} if input_uops is not None else {}, walk=True)
linear = graph_rewrite(linear, pm_insert_copy_staging + pm_flatten_linear, name="insert copy staging")
linear = graph_rewrite(linear, pm_tag_hcq_calls, ctx=(enumerator:=itertools.count(0)), walk=True, name="tag hcq calls")
linear = graph_rewrite(linear, pm_sched_sync, ctx=HCQDepsTracker(), walk=True, name="schedule sync")
linear = linear.substitute({s: p for p, s in back_map.items()}, walk=True)
# schedule
linear = graph_rewrite(linear, pm_sched_hcq_batches, walk=True, name="schedule hcq batches")
linear = linear.substitute({s: p for p, s in back_map.items()}, walk=True, enter_calls=True)
linear = graph_rewrite(linear, pm_merge_queues, walk=True, name="merge queues")
linear = graph_rewrite(linear, pm_add_finalizer, ctx=enumerator, walk=True, name="add finalizer")
linear = graph_rewrite(linear, pm_add_global_sync, ctx=set(), walk=True, name="add global sync", enter_calls=True)
# lowering to hcq ir
linear = graph_rewrite(linear, pm_add_inner_loads, ctx=(waited:=set()), walk=True, name="add loads", enter_calls=True)
linear = graph_rewrite(linear, pm_add_inner_stores, ctx=waited, walk=True, name="add stores", enter_calls=True)
linear = graph_rewrite(linear, pm_encode_cmdbufs, walk=True, name="encode cmdbufs", enter_calls=True)
linear = graph_rewrite(linear, pm_pack_placeholders, walk=True, name="pack placeholders")
# pie
linear = graph_rewrite(linear, pm_split_patches, ctx=jit, walk=True, name="split rt/lt patches")
linear = graph_rewrite(linear, pm_early_simplify + symbolic, bottom_up=False, name="simplify packed placeholders", enter_calls=True)
linear = graph_rewrite(linear, pm_rm_rt_getaddrs, walk=True, name="replace rt getaddrs")
linear = graph_rewrite(linear, pm_rm_rt_binaries, walk=True, name="replace rt binaries")
linear = graph_rewrite(linear, pm_rm_rt_uops, walk=True, name="replace rt uops")
linear = graph_rewrite(linear, pm_replace_params, walk=True, name="replace with args")
# and compile it
@@ -425,7 +399,7 @@ def hcq_compile(linear:UOp, input_uops:list[UOp]|None=None, jit=False) -> UOp:
def bufferize_buf(ctx:bool, buf:UOp) -> UOp|None:
if buf.tag is None: return None
return make_mstack(tuple(UOp.from_buffer((dv:=Device[dev]).pm_bufferize.rewrite(buf, ctx=(dv, ctx)), "CPU") for dev in to_tuple(buf.device)))
return make_mstack(tuple(UOp.from_buffer((dv:=Device[dev]).pm_bufferize.rewrite(buf, ctx=(dv, ctx)), HCQ_RUNTIME_DEV.value) for dev in to_tuple(buf.device)))
pm_bufferize = PatternMatcher([(UPat(Ops.PARAM, name="buf"), bufferize_buf)])
# *****************
@@ -440,7 +414,8 @@ def fold_binary(buf:UOp, blob:UOp) -> UOp:
def fold_const_store(buf:UOp, off:UOp, val:UOp) -> UOp:
for b, v in zip((bs:=mb.bufs if isinstance((mb:=buf.buffer), MultiBuffer) else (mb,)), val.src if val.op is Ops.STACK else (val,)*len(bs)):
struct.pack_into(f'<{v.dtype.fmt}', b.ensure_allocated()._buf.cpu_view().mv.cast('B'), off.arg * buf.dtype.itemsize, truncate[v.dtype](v.arg))
data = struct.pack(f'<{v.dtype.fmt}', truncate[v.dtype](v.arg))
b.ensure_allocated()._buf.cpu_view().view(offset=off.arg * buf.dtype.itemsize, size=len(data), fmt='B')[:] = data
return UOp(Ops.NOOP)
def resolve_getaddr(buf:UOp, g:UOp) -> UOp:
@@ -509,7 +484,8 @@ class HCQ2Compiled(Compiled):
self.rt_allocator = BumpAllocator(64 << 20, wrap=False)
def new_buffer(self, b:UOp, jit:bool) -> Buffer:
if jit or b.tag in HCQ_CACHE_TAGS: return Buffer(self.device, b.max_numel(), b.dtype, options=BufferSpec(cpu_access=True, nolru=True))
if jit or b.tag in HCQ_CACHE_TAGS:
return Buffer(self.device, b.max_numel(), b.dtype, options=BufferSpec(uncached=True, cpu_access=True, nolru=True))
return self.rt_buffer.view(b.max_numel(), b.dtype, self.rt_allocator.alloc(b.max_numel() * b.dtype.itemsize, alignment=128))
@functools.cache
@@ -571,6 +547,10 @@ class HCQ2Buffer:
def base(self) -> HCQ2Buffer: return self._base or self
class HCQAllocator(LRUAllocator[HCQDeviceType], Generic[HCQDeviceType]):
def _as_buffer(self, buf:HCQ2Buffer) -> memoryview:
self.dev.synchronize()
return buf.cpu_view().mv
def _map(self, buf:HCQ2Buffer) -> HCQ2Buffer:
if not hasattr(self, '_do_map'): raise NotImplementedError("map failed: no method implemented")
return self._do_map(buf)
@@ -586,24 +566,3 @@ class HCQAllocator(LRUAllocator[HCQDeviceType], Generic[HCQDeviceType]):
self.dev.iface.free(mb)
def _offset(self, buf, size:int, offset:int) -> HCQ2Buffer: return buf.offset(offset=offset, size=size)
def _wrap(self, dev:str, sz:int, opaque:HCQ2Buffer) -> Buffer:
return Buffer(dev, sz, dtypes.uint8, opaque=opaque, options=BufferSpec(external_ptr=1))
def _copy(self, dst:Buffer, src:Buffer):
from tinygrad.engine.realize import run_linear
du, su = UOp.from_buffer(dst), UOp.from_buffer(src)
run_linear(UOp(Ops.LINEAR, src=(su.param_like(1).copy_to_device(dst.device).call(du, su),)), update_stats=True)
def _copyin(self, dest:HCQ2Buffer, src:memoryview):
s = Buffer(self.dev.device, len(src), dtypes.uint8, options=BufferSpec(host=True), preallocate=True)
s._buf.cpu_view()[:len(src)] = src
self._copy(self._wrap(self.dev.device, len(src), dest), s)
def _copyout(self, dest:memoryview, src:HCQ2Buffer):
d = Buffer(self.dev.device, len(dest), dtypes.uint8, options=BufferSpec(host=True), preallocate=True)
self._copy(d, self._wrap(self.dev.device, len(dest), src))
self.dev.synchronize()
dest[:] = d._buf.cpu_view()[:len(dest)]
# def _as_buffer(self, buf): return buf.cpu_view().mv
+6 -6
View File
@@ -90,7 +90,7 @@ def memory_barrier(ctx):
reg_done=getattr(ctx.nbio, f'regBIF_BX_PF{pf}_GPU_HDP_FLUSH_DONE').addr[0], value=0xffffffff),
acquire_mem(ctx)))
def pm4_wait(ctx, dst, val): return wait_reg_mem(ctx, val, mem=make_getaddr(dst, ctx.devs))
def pm4_wait(ctx, x, y): return wait_reg_mem(ctx, y, mem=make_getaddr(x.buf_uop, ctx.devs))
def pm4_barrier(ctx): return memory_barrier(ctx)
@@ -138,7 +138,7 @@ def pm4_program(ctx, call, prg):
pm_pm4_opsel = PatternMatcher([
(UPat(Ops.CALL, src=(UPat(Ops.PROGRAM, name="prg"),), name="call", allow_any_len=True), pm4_program),
(UPat(Ops.WAIT, src=(UPat(name="dst"), UPat(name="val"))), pm4_wait),
(UPat(Ops.WAIT, src=(UPat.var("x") >= UPat.var("y"),)), pm4_wait),
(UPat(Ops.BARRIER), pm4_barrier),
(UPat(Ops.CUSTOM_FUNCTION, arg="timestamp", src=(UPat(name="dst"),)), pm4_timestamp),
(UPat(Ops.STORE, src=(UPat((Ops.BUFFER, Ops.PARAM), name="dst"), UPat(name="val"))), pm4_store),
@@ -184,10 +184,10 @@ def sdma_copy(ctx, call):
ctx.sdma.SDMA_PKT_COPY_LINEAR_COUNT_COUNT(min(sz - off, ctx.max_copy_size) - 1), 0,
*data64_le(src_addr + off), *data64_le(dst_addr + off)) for off in range(0, sz, ctx.max_copy_size)]))
def sdma_wait(ctx, dst, val):
def sdma_wait(ctx, x, y):
op = ctx.sdma.SDMA_OP_POLL_REGMEM | ctx.sdma.SDMA_PKT_POLL_REGMEM_HEADER_FUNC(WAIT_REG_MEM_FUNCTION_GEQ) \
| ctx.sdma.SDMA_PKT_POLL_REGMEM_HEADER_MEM_POLL(1)
return make_ins(SDMAOps.POLL_REGMEM, op, *data64_le(make_getaddr(dst, ctx.devs)), val, 0xffffffff,
return make_ins(SDMAOps.POLL_REGMEM, op, *data64_le(make_getaddr(x.buf_uop, ctx.devs)), y, 0xffffffff,
ctx.sdma.SDMA_PKT_POLL_REGMEM_DW5_INTERVAL(0x04) | ctx.sdma.SDMA_PKT_POLL_REGMEM_DW5_RETRY_COUNT(0xfff))
def sdma_store(ctx, dst, val):
@@ -203,7 +203,7 @@ pm_sdma_opsel = PatternMatcher([
(UPat(Ops.CALL, src=(UPat(Ops.COPY),), name="call", allow_any_len=True), sdma_copy),
(UPat(Ops.BARRIER), lambda: UOp(Ops.NOOP, dtypes.void, ())),
(UPat(Ops.WAIT, src=(UPat(name="dst"), UPat(name="val"))), sdma_wait),
(UPat(Ops.WAIT, src=(UPat.var("x") >= UPat.var("y"),)), sdma_wait),
(UPat(Ops.CUSTOM_FUNCTION, arg="timestamp", src=(UPat(name="dst"),)), sdma_timestamp),
(UPat(Ops.STORE, src=(UPat((Ops.BUFFER, Ops.PARAM), name="dst"), UPat(name="val"))), sdma_store),
])
@@ -536,7 +536,7 @@ class AMDDevice(HCQ2Compiled):
timestamp_divider = 100.0 # AMD GPU clock: ticks/us
ifaces = [KFDIface, PCIIface]
ifaces = [KFDIface, PCIIface, _mock(KFDIface, "MOCKIface"), _mock(KFDIface), _mock(PCIIface)]
def is_am(self) -> bool: return isinstance(self.iface, (PCIIface,))
def is_usb(self) -> bool: return False
-5
View File
@@ -20,11 +20,6 @@ def local_abs_max(x:Tensor) -> Tensor:
fxn = _local_abs_max_fxn(param.uop, x.device)
return Tensor(fxn[0].uop.call(x.uop).gettuple(0))
def scalar_amax(amax_buf:Tensor) -> Tensor:
if isinstance(amax_buf.device, tuple):
return local_abs_max(amax_buf).detach()
return amax_buf.max().detach()
def shard_shape(shape:tuple, axis:int, ndev:int) -> list:
s = list(shape)
s[axis] //= ndev
+15 -19
View File
@@ -3,7 +3,7 @@ import functools, pathlib
from tinygrad import Tensor, dtypes
from tinygrad.uop.ops import UOp, Ops, KernelInfo
from tinygrad.renderer import Estimates
from extra.llama_kernels import NUM_WG, THREADS_PER_WG, compile_cpp, alloc_like, alloc_local, scalar_amax, dname_of
from extra.llama_kernels import NUM_WG, THREADS_PER_WG, compile_cpp, alloc_like, dname_of
# module-level mailbox: grad_xw13 UOp -> (grad_xw13_fp8 UOp, delayed amax UOp)
# lets cdna_asm_gemm's bwd reuse the fp8 companion produced by the fused silu_mul bwd kernel
@@ -11,13 +11,13 @@ from extra.llama_kernels import NUM_WG, THREADS_PER_WG, compile_cpp, alloc_like,
_grad_fp8_mailbox:dict[UOp, tuple[UOp, UOp]] = {}
@functools.cache
def _custom_fused_bwd_w13(grad_xw13_fp8:UOp, grad_amax_buf:UOp, grad_amax:UOp,
def _custom_fused_bwd_w13(grad_xw13_fp8:UOp, grad_amax_next:UOp, grad_amax:UOp,
xw13:UOp, grad_x2:UOp, amax_state:UOp, grad_amax_state:UOp, dname:str) -> UOp:
hidden = xw13.shape[2] // 2
n_elems = xw13.shape[0] * xw13.shape[1] * hidden
threads, workgroups = UOp.special(THREADS_PER_WG, "lidx0"), UOp.special(NUM_WG, "gidx0")
mem = n_elems * 2 * 3 + n_elems * 2 + NUM_WG * 4 + 4
sink = UOp.sink(grad_xw13_fp8.base, grad_amax_buf.base, grad_amax.base,
mem = n_elems * 2 * 3 + n_elems * 2 + 4 + 4
sink = UOp.sink(grad_xw13_fp8.base, grad_amax_next.base, grad_amax.base,
xw13.base, grad_x2.base, amax_state.base, grad_amax_state.base, threads, workgroups,
arg=KernelInfo(f"fused_silu_mul_bwd_w13_{n_elems}", estimates=Estimates(ops=10*n_elems, mem=mem)))
src, lib = compile_cpp(pathlib.Path(__file__).parent, "cast_amax_bwd_w13.cpp", n_elems, hidden)
@@ -25,14 +25,14 @@ def _custom_fused_bwd_w13(grad_xw13_fp8:UOp, grad_amax_buf:UOp, grad_amax:UOp,
UOp(Ops.SOURCE, arg=src), UOp(Ops.BINARY, arg=lib)))
@functools.cache
def _custom_fused_cast_amax_w13(fp8_out:UOp, amax_buf:UOp, xw13:UOp, amax_state:UOp, grad_amax_state:UOp,
def _custom_fused_cast_amax_w13(fp8_out:UOp, amax_out:UOp, xw13:UOp, amax_state:UOp, grad_amax_state:UOp,
next_grad_amax_state:UOp, dname:str) -> UOp:
# NOTE: grad_amax_state is plumbed through as an unused fwd input so the bwd kernel can read it via kernel.src
hidden = xw13.shape[2] // 2
n_elems = xw13.shape[0] * xw13.shape[1] * hidden
threads, workgroups = UOp.special(THREADS_PER_WG, "lidx0"), UOp.special(NUM_WG, "gidx0")
mem = n_elems * 2 * 2 + n_elems + NUM_WG * 4
sink = UOp.sink(fp8_out.base, amax_buf.base, xw13.base, amax_state.base, threads, workgroups,
mem = n_elems * 2 * 2 + n_elems + 4
sink = UOp.sink(fp8_out.base, amax_out.base, xw13.base, amax_state.base, threads, workgroups,
arg=KernelInfo(f"fused_silu_mul_cast_amax_w13_{n_elems}", estimates=Estimates(ops=5*n_elems, mem=mem)))
src, lib = compile_cpp(pathlib.Path(__file__).parent, "cast_amax_fwd_w13.cpp", n_elems, hidden)
return UOp(Ops.PROGRAM, src=(sink, UOp(Ops.LINEAR, src=(*sink.src, sink)),
@@ -43,26 +43,23 @@ def _fused_quantize_bwd_w13(gradient:UOp, kernel:UOp):
device = xw13.device
axis = xw13.axis if isinstance(device, tuple) else None
grad_xw13_fp8 = alloc_like(xw13.shape, dtypes.fp8e4m3, device, axis)
grad_amax_buf = alloc_local((NUM_WG,), dtypes.float32, device, axis)
grad_amax_next = Tensor(next_grad_amax_state, device=device)
grad_amax_state_t = Tensor(grad_amax_state, device=device)
fxn = functools.partial(_custom_fused_bwd_w13, dname=dname_of(device))
grad_amax = grad_amax_state_t.empty_like()
grad_xw13_fp8, grad_amax_buf, grad_amax, *_ = Tensor.custom_kernel(
grad_xw13_fp8, grad_amax_buf, grad_amax,
grad_xw13_fp8, grad_amax_next, grad_amax, *_ = Tensor.custom_kernel(
grad_xw13_fp8, grad_amax_next, grad_amax,
Tensor(xw13, device=device), Tensor(gradient, device=device).cast(dtypes.bfloat16),
Tensor(amax_state, device=device), grad_amax_state_t, fxn=fxn)
grad_xw13_uop = grad_xw13_fp8.uop.cast(dtypes.bfloat16)
new_grad_amax = scalar_amax(grad_amax_buf)
store_effect = next_grad_amax_state.store(new_grad_amax.uop)
assert grad_xw13_fp8.uop.op is Ops.AFTER, f"expected AFTER, got {grad_xw13_fp8.uop.op}"
grad_xw13_fp8_uop = grad_xw13_fp8.uop.replace(src=grad_xw13_fp8.uop.src + (store_effect,))
# Stash fp8 companion for cdna_asm_gemm's bwd to attach to grad_a.
_grad_fp8_mailbox[grad_xw13_uop] = (grad_xw13_fp8_uop, grad_amax_state_t.uop)
_grad_fp8_mailbox[grad_xw13_uop] = (grad_xw13_fp8.uop, grad_amax_state_t.uop)
return (None, None, grad_xw13_uop, None, None, None)
def fused_quantize_fp8_w13(xw13:Tensor, amax_state:Tensor, fp8_dtype, grad_amax_state:Tensor,
next_grad_amax_state:Tensor) -> tuple[Tensor, Tensor]:
# NOTE: silu(xw1)*xw3 -> fp8 + amax over fused xw13 layout. Returns (fp8, new_amax)
next_grad_amax_state:Tensor, amax_out:Tensor) -> Tensor:
# NOTE: silu(xw1)*xw3 -> fp8 + amax over fused xw13 layout. Returns fp8.
# grad_amax_state: delayed amax for grad_xw13 fp8 quantization in the backward.
assert xw13.dtype == dtypes.bfloat16, f"expected bf16, got {xw13.dtype}"
MBS, SEQ, H2 = xw13.shape
@@ -70,8 +67,7 @@ def fused_quantize_fp8_w13(xw13:Tensor, amax_state:Tensor, fp8_dtype, grad_amax_
HIDDEN = H2 // 2
axis = xw13.uop.axis if isinstance(xw13.device, tuple) else None
fp8_out = alloc_like((MBS, SEQ, HIDDEN), fp8_dtype, xw13.device, axis)
amax_buf = alloc_local((NUM_WG,), dtypes.float32, xw13.device, axis)
fxn = functools.partial(_custom_fused_cast_amax_w13, dname=dname_of(xw13.device))
fp8_out, amax_buf, *_ = Tensor.custom_kernel(fp8_out, amax_buf, xw13, amax_state, grad_amax_state, next_grad_amax_state,
fp8_out, amax_out, *_ = Tensor.custom_kernel(fp8_out, amax_out, xw13, amax_state, grad_amax_state, next_grad_amax_state,
fxn=fxn, grad_fxn=_fused_quantize_bwd_w13)
return fp8_out, scalar_amax(amax_buf)
return fp8_out
@@ -23,14 +23,14 @@ static_assert(HIDDEN % VEC == 0, "HIDDEN must be divisible by VEC");
// fused silu*mul backward, three outputs in a single HBM pass:
// 1) fp8 grad_xw13_fp8 — delayed-scale quantize using grad_amax_state (mailbox to matmul bwd)
// 2) fp32 grad_amax_buf — per-WG partial |grad_xw13|, reduced into next step's grad_amax_state
// 2) fp32 grad_amax_next — scalar |grad_xw13| via global atomic max
// 3) fp32 grad_amax_out — delayed grad amax used for quantize/GEMM epilogue scale
// grad_amax_state is read for the fp8 scale. The store of new_grad_amax into grad_amax_state's
// buffer is built in Python as a separate effect and threaded into grad_a via .after(store).
extern "C" __global__ __launch_bounds__(THREADS_PER_WG) void
fused_silu_mul_bwd_w13(
__hip_fp8_storage_t* __restrict__ grad_xw13_fp8_out, // fp8, 2*N_ELEMS
float* __restrict__ grad_amax_buf, // fp32, NUM_WG per-WG partials
float* __restrict__ grad_amax_next, // fp32 scalar, initialized to 0 before launch
float* __restrict__ grad_amax_out, // fp32 scalar delayed grad amax
const __hip_bfloat16* __restrict__ xw13, // bf16, 2*N_ELEMS
const __hip_bfloat16* __restrict__ grad_x2, // bf16, N_ELEMS
@@ -92,5 +92,6 @@ fused_silu_mul_bwd_w13(
if (tid < s) sdata[tid] = fmaxf(sdata[tid], sdata[tid + s]);
__syncthreads();
}
if (tid == 0) grad_amax_buf[wg] = sdata[0];
if (tid == 0 && sdata[0] > *grad_amax_next)
atomicMax(reinterpret_cast<int32_t*>(grad_amax_next), __float_as_int(sdata[0]));
}
@@ -24,7 +24,7 @@ static_assert(HIDDEN % VEC == 0, "HIDDEN must be divisible by VEC (so VEC loads
extern "C" __global__ __launch_bounds__(THREADS_PER_WG) void
fused_silu_mul_cast_amax_w13(
__hip_fp8_storage_t* __restrict__ fp8_out, // fp8, N_ELEMS
float* __restrict__ amax_buf, // fp32, NUM_WG (per-WG amaxes)
float* __restrict__ amax_out, // fp32 scalar, initialized to 0 before launch
const __hip_bfloat16* __restrict__ xw13, // bf16, 2*N_ELEMS
const float* __restrict__ amax_state) // fp32 scalar
{
@@ -67,7 +67,7 @@ fused_silu_mul_cast_amax_w13(
*reinterpret_cast<uint64_t*>(&fp8_out[base]) = *reinterpret_cast<uint64_t*>(out);
}
// LDS tree reduction: per-workgroup amax
// LDS tree reduction: per-workgroup amax, then global atomic into the scalar.
sdata[tid] = local_max;
__syncthreads();
for (int s = THREADS_PER_WG / 2; s > 0; s >>= 1) {
@@ -75,5 +75,5 @@ fused_silu_mul_cast_amax_w13(
__syncthreads();
}
if (tid == 0) amax_buf[wg] = sdata[0];
if (tid == 0 && sdata[0] > *amax_out) atomicMax(reinterpret_cast<int32_t*>(amax_out), __float_as_int(sdata[0]));
}
+2 -2
View File
@@ -16,7 +16,7 @@ def _custom_fused_ce_loss_fwd(loss_out:UOp, max_out:UOp, lse_out:UOp, logits:UOp
row_lse = (logits[b, s, v_lse].cast(dtypes.float) - row_max).exp().reduce(v_lse, arg=Ops.ADD).log() + row_max
v_smooth = UOp.range(vocab, 3, axis_type=AxisType.REDUCE)
target = logits[b, s, targets[row].cast(dtypes.index)].cast(dtypes.float)
target = logits[b, s, targets[row].cast(dtypes.weakint)].cast(dtypes.float)
mean_logits = logits[b, s, v_smooth].cast(dtypes.float).reduce(v_smooth, arg=Ops.ADD) / vocab
loss = row_lse - (1.0 - label_smoothing) * target - label_smoothing * mean_logits
stores = UOp.group(loss_out[row].store(loss), max_out[row].store(row_max), lse_out[row].store(row_lse))
@@ -32,7 +32,7 @@ def _custom_fused_ce_loss_bwd(d_logits:UOp, logits:UOp, lse:UOp, targets:UOp, sc
s = row % seq
prob = (logits[b, s, v].cast(dtypes.float) - lse[row]).exp()
target = v.eq(targets[row].cast(dtypes.index)).where(1.0 - label_smoothing, 0.0)
target = v.eq(targets[row].cast(dtypes.weakint)).where(1.0 - label_smoothing, 0.0)
smooth = label_smoothing / vocab
grad = (prob - target - smooth) * scale[0]
@@ -3,19 +3,19 @@ import functools, pathlib
from tinygrad import Tensor, dtypes
from tinygrad.uop.ops import UOp, Ops, KernelInfo
from tinygrad.renderer import Estimates
from extra.llama_kernels import FP8_MAX, NUM_WG, THREADS_PER_WG, alloc_like, alloc_local, scalar_amax, dname_of, compile_hip
from extra.llama_kernels import NUM_WG, THREADS_PER_WG, alloc_like, alloc_local, dname_of, compile_hip
def _src() -> str: return (pathlib.Path(__file__).parent/"fused_rmsnorm_mul_quantize_fp8.cpp").read_text()
def _src_bwd() -> str: return (pathlib.Path(__file__).parent/"fused_rmsnorm_mul_quantize_fp8_bwd.cpp").read_text()
@functools.cache
def _custom_fwd(fp8_out:UOp, x_normed_out:UOp, rrms_out:UOp, amax_buf:UOp,
def _custom_fwd(fp8_out:UOp, x_normed_out:UOp, rrms_out:UOp, amax_out:UOp,
x:UOp, weight:UOp, amax_state:UOp, dname:str, eps_val:float) -> UOp:
MBS, SEQ, HIDDEN = x.shape
n_elems = MBS * SEQ * HIDDEN
threads, workgroups = UOp.special(THREADS_PER_WG, "lidx0"), UOp.special(NUM_WG, "gidx0")
mem = n_elems * 2 + n_elems + MBS * SEQ * 4 + n_elems + HIDDEN * 2 + NUM_WG * 4 + 4
sink = UOp.sink(fp8_out.base, x_normed_out.base, rrms_out.base, amax_buf.base,
mem = n_elems * 2 + n_elems + MBS * SEQ * 4 + n_elems + HIDDEN * 2 + 4 + 4
sink = UOp.sink(fp8_out.base, x_normed_out.base, rrms_out.base, amax_out.base,
x.base, weight.base, amax_state.base, threads, workgroups,
arg=KernelInfo(f"fused_rmsnorm_mul_quantize_fp8_{n_elems}_h{HIDDEN}_eps{eps_val:.0e}",
estimates=Estimates(ops=6*n_elems, mem=mem)))
@@ -26,13 +26,13 @@ def _custom_fwd(fp8_out:UOp, x_normed_out:UOp, rrms_out:UOp, amax_buf:UOp,
UOp(Ops.SOURCE, arg=src), UOp(Ops.BINARY, arg=compile_hip(src, defines))))
@functools.cache
def _custom_fwd_add(fp8_out:UOp, h_out:UOp, x_normed_out:UOp, rrms_out:UOp, amax_buf:UOp,
def _custom_fwd_add(fp8_out:UOp, h_out:UOp, x_normed_out:UOp, rrms_out:UOp, amax_out:UOp,
x:UOp, residual:UOp, weight:UOp, amax_state:UOp, dname:str, eps_val:float) -> UOp:
MBS, SEQ, HIDDEN = x.shape
n_elems = MBS * SEQ * HIDDEN
threads, workgroups = UOp.special(THREADS_PER_WG, "lidx0"), UOp.special(NUM_WG, "gidx0")
mem = n_elems * 2 * 4 + MBS * SEQ * 4 + HIDDEN * 2 + NUM_WG * 4 + 4
sink = UOp.sink(fp8_out.base, h_out.base, x_normed_out.base, rrms_out.base, amax_buf.base,
mem = n_elems * 2 * 4 + MBS * SEQ * 4 + HIDDEN * 2 + 4 + 4
sink = UOp.sink(fp8_out.base, h_out.base, x_normed_out.base, rrms_out.base, amax_out.base,
x.base, residual.base, weight.base, amax_state.base, threads, workgroups,
arg=KernelInfo(f"fused_add_rmsnorm_mul_quantize_fp8_{n_elems}_h{HIDDEN}_eps{eps_val:.0e}",
estimates=Estimates(ops=7*n_elems, mem=mem)))
@@ -85,7 +85,7 @@ def _bwd_common(fp8_grad_u, h_grad_u, x_u, x_normed_u, rrms_u, weight_u, amax_st
return grad_total.uop, grad_weight_uop
def _fused_bwd(gradient:UOp, kernel:UOp):
# NOTE: fwd inputs (fp8_out, x_normed_out, rrms_out, amax_buf, x, weight, amax_state)
# NOTE: fwd inputs (fp8_out, x_normed_out, rrms_out, amax_out, x, weight, amax_state)
_, x_normed_u, rrms_u, _, x_u, weight_u, amax_state_u = kernel.src[1:]
grad_x, grad_w = _bwd_common(gradient, None, x_u, x_normed_u, rrms_u, weight_u, amax_state_u, kernel)
return (None, None, None, None, grad_x, grad_w, None)
@@ -112,8 +112,9 @@ def _fused_add_bwd(*args, **kwargs):
grad_h, grad_w = _bwd_common(fp8_grad_u, h_grad_u, x_u, x_normed_u, rrms_u, weight_u, amax_state_u, kernel)
return (None, None, None, None, None, grad_h, grad_h, grad_w, None)
def fused_rmsnorm_mul_quantize_fp8(x:Tensor, weight:Tensor, amax_state:Tensor, eps:float, fp8_dtype) -> tuple[Tensor, Tensor, Tensor, Tensor]:
# NOTE: rmsnorm(x) * weight -> fp8 + amax. Returns (fp8, new_amax, x_normed, rrms).
def fused_rmsnorm_mul_quantize_fp8(x:Tensor, weight:Tensor, amax_state:Tensor, eps:float, fp8_dtype,
amax_out:Tensor) -> tuple[Tensor, Tensor, Tensor]:
# NOTE: rmsnorm(x) * weight -> fp8 + amax. Returns (fp8, x_normed, rrms).
# x_normed + rrms are saved for the rmsnorm backward (also recomputed here from x regs).
assert x.dtype == dtypes.bfloat16 and weight.dtype == dtypes.bfloat16
assert x.shape[-1] == weight.shape[-1], f"HIDDEN mismatch: x={x.shape}, weight={weight.shape}"
@@ -123,16 +124,15 @@ def fused_rmsnorm_mul_quantize_fp8(x:Tensor, weight:Tensor, amax_state:Tensor, e
fp8_out = alloc_like((MBS, SEQ, HIDDEN), fp8_dtype, x.device, axis)
x_normed_out = alloc_like((MBS, SEQ, HIDDEN), dtypes.bfloat16, x.device, axis)
rrms_out = alloc_like((MBS, SEQ), dtypes.float32, x.device, axis)
amax_buf = alloc_local((NUM_WG,), dtypes.float32, x.device, axis)
fxn = functools.partial(_custom_fwd, dname=dname_of(x.device), eps_val=eps)
fp8_out, x_normed_out, rrms_out, amax_buf, *_ = Tensor.custom_kernel(
fp8_out, x_normed_out, rrms_out, amax_buf, x, weight, amax_state, fxn=fxn, grad_fxn=_fused_bwd)
return fp8_out, scalar_amax(amax_buf), x_normed_out, rrms_out
fp8_out, x_normed_out, rrms_out, amax_out, *_ = Tensor.custom_kernel(
fp8_out, x_normed_out, rrms_out, amax_out, x, weight, amax_state, fxn=fxn, grad_fxn=_fused_bwd)
return fp8_out, x_normed_out, rrms_out
def fused_add_rmsnorm_mul_quantize_fp8(x:Tensor, residual:Tensor, weight:Tensor, amax_state:Tensor,
eps:float, fp8_dtype) -> tuple[Tensor, Tensor, Tensor, Tensor, Tensor]:
eps:float, fp8_dtype, amax_out:Tensor) -> tuple[Tensor, Tensor, Tensor, Tensor]:
# NOTE: h = x + residual; y_normed = rmsnorm(h); fp8 = quantize(y_normed * weight).
# Returns (fp8, new_amax, h, x_normed, rrms). h is also written so downstream can
# Returns (fp8, h, x_normed, rrms). h is also written so downstream can
# reuse it without recomputing x+residual — eliminates the separate residual-add kernel.
assert x.dtype == dtypes.bfloat16 and residual.dtype == dtypes.bfloat16 and weight.dtype == dtypes.bfloat16
assert x.shape == residual.shape
@@ -143,9 +143,8 @@ def fused_add_rmsnorm_mul_quantize_fp8(x:Tensor, residual:Tensor, weight:Tensor,
h_out = alloc_like((MBS, SEQ, HIDDEN), dtypes.bfloat16, x.device, axis)
x_normed_out = alloc_like((MBS, SEQ, HIDDEN), dtypes.bfloat16, x.device, axis)
rrms_out = alloc_like((MBS, SEQ), dtypes.float32, x.device, axis)
amax_buf = alloc_local((NUM_WG,), dtypes.float32, x.device, axis)
fxn = functools.partial(_custom_fwd_add, dname=dname_of(x.device), eps_val=eps)
fp8_out, h_out, x_normed_out, rrms_out, amax_buf, *_ = Tensor.custom_kernel(
fp8_out, h_out, x_normed_out, rrms_out, amax_buf, x, residual, weight, amax_state,
fp8_out, h_out, x_normed_out, rrms_out, amax_out, *_ = Tensor.custom_kernel(
fp8_out, h_out, x_normed_out, rrms_out, amax_out, x, residual, weight, amax_state,
fxn=fxn, grad_fxn=_fused_add_bwd)
return fp8_out, scalar_amax(amax_buf), h_out, x_normed_out, rrms_out
return fp8_out, h_out, x_normed_out, rrms_out
@@ -7,7 +7,7 @@
// fp8 = fp8_sat(y * (FP8_MAX / amax_state))
// Also writes:
// rrms[row] — saved for the rmsnorm backward
// amax_buf[wg] — per-WG |y| partials, reduced later to update amax_state
// amax_out — scalar |y| via global atomic max
//
// Layout: one WG per row, ROWS_PER_WG rows per WG via grid-stride (ROWS = N_ELEMS / HIDDEN).
// Each thread handles HIDDEN / THREADS_PER_WG elements per row.
@@ -48,7 +48,7 @@ fused_add_rmsnorm_mul_quantize_fp8(
__hip_bfloat16* __restrict__ h_out, // bf16, ROWS*HIDDEN — x + residual (saved for downstream)
__hip_bfloat16* __restrict__ x_normed_out, // bf16, ROWS*HIDDEN
float* __restrict__ rrms_out, // fp32, ROWS
float* __restrict__ amax_buf, // fp32, NUM_WG
float* __restrict__ amax_out, // fp32 scalar, initialized to 0 before launch
const __hip_bfloat16* __restrict__ x, // bf16, ROWS*HIDDEN
const __hip_bfloat16* __restrict__ residual, // bf16, ROWS*HIDDEN — added into x before rmsnorm
const __hip_bfloat16* __restrict__ weight, // bf16, HIDDEN
@@ -60,7 +60,7 @@ fused_rmsnorm_mul_quantize_fp8(
__hip_fp8_storage_t* __restrict__ fp8_out, // fp8, ROWS*HIDDEN
__hip_bfloat16* __restrict__ x_normed_out, // bf16, ROWS*HIDDEN (saved for rmsnorm bwd)
float* __restrict__ rrms_out, // fp32, ROWS (fp32 to match rmsnorm_bwd.cpp expectation)
float* __restrict__ amax_buf, // fp32, NUM_WG per-WG partials
float* __restrict__ amax_out, // fp32 scalar, initialized to 0 before launch
const __hip_bfloat16* __restrict__ x, // bf16, ROWS*HIDDEN
const __hip_bfloat16* __restrict__ weight, // bf16, HIDDEN (per-hidden scale)
const float* __restrict__ amax_state) // fp32 scalar
@@ -144,12 +144,12 @@ fused_rmsnorm_mul_quantize_fp8(
__syncthreads(); // before next row's sum_sq reduce reuses sdata
}
// Final per-WG amax reduce.
// Final per-WG amax reduce, then global atomic into the scalar.
sdata[tid] = local_max;
__syncthreads();
for (int s = THREADS_PER_WG / 2; s > 0; s >>= 1) {
if (tid < s) sdata[tid] = fmaxf(sdata[tid], sdata[tid + s]);
__syncthreads();
}
if (tid == 0) amax_buf[wg] = sdata[0];
if (tid == 0 && sdata[0] > *amax_out) atomicMax(reinterpret_cast<int32_t*>(amax_out), __float_as_int(sdata[0]));
}
@@ -3,14 +3,13 @@ from tinygrad import Tensor, dtypes
from tinygrad.dtype import AddrSpace
from tinygrad.helpers import prod
from tinygrad.uop.ops import UOp, Ops, KernelInfo, AxisType
from extra.llama_kernels import FP8_MAX, NUM_WG, THREADS_PER_WG, alloc_like, alloc_local, scalar_amax
from extra.llama_kernels import FP8_MAX, NUM_WG, THREADS_PER_WG, alloc_like
@functools.cache
def _custom_quantize_fp8_with_amax(fp8_out:UOp, amax_partial:UOp, x:UOp, amax_state:UOp) -> UOp:
def _custom_quantize_fp8_with_amax(fp8_out:UOp, amax_out:UOp, x:UOp, amax_state:UOp, device=None) -> UOp:
VEC = 8
n_elems = prod(x.shape)
assert n_elems % (NUM_WG * THREADS_PER_WG * VEC) == 0
assert amax_partial.shape[0] == NUM_WG
x = x.reshape(n_elems)
fp8_out = fp8_out.reshape(n_elems)
@@ -46,8 +45,13 @@ def _custom_quantize_fp8_with_amax(fp8_out:UOp, amax_partial:UOp, x:UOp, amax_st
lds = lds.after(lds[tid.valid(active)].store(lds[tid].maximum(other)).barrier())
step //= 2
amax_store = amax_partial[tid.eq(0).where(wg, UOp.invalid())].store(lds[0])
return amax_store.end(tid, wg).sink(arg=KernelInfo(f"quantize_fp8_with_amax_{n_elems}", opts_to_apply=()))
device = device[0].split(":")[0] if isinstance(device, tuple) else device.split(":")[0]
if device in {"AMD", "NULL"}: atomic_arg = "if ({2} > {3}) __hip_atomic_fetch_max((int*){0}, {1}, __ATOMIC_RELAXED, __HIP_MEMORY_SCOPE_AGENT);"
else: raise NotImplementedError(f"no atomic max for device {device}")
amax_idx = amax_out.reshape((1,)).index(UOp.const(dtypes.weakint, 0))
max_val = lds[0].load()
atomic = UOp(Ops.CUSTOM, dtypes.void, (amax_idx, max_val.bitcast(dtypes.int32), max_val, amax_idx.load()), arg=atomic_arg)
return atomic.end(tid, wg).sink(arg=KernelInfo(f"quantize_fp8_with_amax_{n_elems}", opts_to_apply=()))
@functools.cache
def _custom_quantize_fp8_scalar(fp8_out:UOp, x:UOp, amax_state:UOp) -> UOp:
@@ -69,25 +73,19 @@ def _quantize_fp8_delayed_bwd(gradient:UOp, kernel:UOp):
grad_x = (Tensor(gradient, device=device).float() * scale).cast(dtypes.bfloat16)
return (None, None, grad_x.uop, None)
def quantize_fp8_delayed(x:Tensor, amax_state:Tensor, fp8_dtype=dtypes.fp8e4m3) -> tuple[Tensor, Tensor, Tensor, UOp]:
# NOTE: one-pass bf16 -> fp8 quantize with delayed scaling. Returns (fp8, inv_scale, new_amax, store_effect).
# Fused kernel reads x once and writes fp8 + per-WG |x| partials (then a small reduce produces scalar new_amax).
# store_effect writes new_amax into amax_state's buffer — the caller must thread it into a realized
# output via `.after(store_effect)`. Calling `amax_state.assign(new_amax)` inside a grad_fxn does
# NOT work because .assign mutates only the temp Tensor's .uop, not the original layer-owned buffer.
def quantize_fp8_delayed(x:Tensor, amax_state:Tensor, amax_out:Tensor, fp8_dtype=dtypes.fp8e4m3) -> tuple[Tensor, Tensor]:
# NOTE: one-pass bf16 -> fp8 quantize with delayed scaling.
# Fused kernel reads x once and writes fp8 + scalar amax via global atomic max.
assert x.dtype == dtypes.bfloat16, f"expected bf16, got {x.dtype}"
axis = x.uop.axis if isinstance(x.device, tuple) else None
fp8_out = alloc_like(x.shape, fp8_dtype, x.device, axis)
n_elems = prod(x.uop.shard_shape)
assert n_elems % NUM_WG == 0, f"{n_elems=} must divide over {NUM_WG=}"
amax_partial = alloc_local((NUM_WG,), dtypes.float32, x.device, axis)
fxn = _custom_quantize_fp8_with_amax
fp8_out, amax_partial, *_ = Tensor.custom_kernel(fp8_out, amax_partial, x, amax_state,
fxn=fxn, grad_fxn=_quantize_fp8_delayed_bwd)
new_amax = scalar_amax(amax_partial)
fxn = functools.partial(_custom_quantize_fp8_with_amax, device=x.device)
fp8_out, amax_out, *_ = Tensor.custom_kernel(fp8_out, amax_out, x, amax_state,
fxn=fxn, grad_fxn=_quantize_fp8_delayed_bwd)
inv_scale = (amax_state.float() + 1e-8) / FP8_MAX
store_effect = amax_state.uop.store(new_amax.uop)
return fp8_out, inv_scale, new_amax, store_effect
return fp8_out, inv_scale
def quantize_fp8_scalar(x:Tensor, amax_state:Tensor, fp8_dtype=dtypes.fp8e4m3) -> Tensor:
# NOTE: pure one-pass bf16 -> fp8 quantize with delayed scalar scale. No amax computation.
+98 -5
View File
@@ -16,6 +16,96 @@ def _sharded_empty(shape:Tensor, ref:Tensor, axis:int|None, dtype:DTypeLike|None
axis = ref.uop.axis if axis is None else axis
return Tensor(Tensor.invalids(*shape, dtype=dtype, device=ref.device).uop.multi(axis), dtype=dtype, device=ref.device)
@functools.cache
def custom_fused_qkv_rope_forward(q:UOp, k:UOp, v:UOp, xqkv:UOp, freqs_cis:UOp,
device:str, arch:str, B:int, N:int, H:int, H_KV:int, D:int):
code = (pathlib.Path(__file__).parent / "fused_qkv_rope.cpp").read_text()
threads = 256
thread_idx = UOp.special(threads, "lidx0")
block_idx_x, block_idx_y = UOp.special(B, "gidx0"), UOp.special(N, "gidx1")
sink = UOp.sink(q.base, k.base, v.base, xqkv.base, freqs_cis.base, thread_idx, block_idx_x, block_idx_y,
arg=KernelInfo(name="fused_qkv_rope_forward"))
compile_args = ["-std=c++20", "-ffast-math", f"-DATTN_B={B}", f"-DATTN_N={N}", f"-DATTN_H={H}",
f"-DATTN_H_KV={H_KV}", f"-DATTN_D={D}", f"-DTHREADS_PER_BLOCK={threads}"]
lib = HIPCCCompiler(arch, compile_args).compile_cached(code)
return UOp(Ops.PROGRAM, src=(sink, UOp(Ops.LINEAR, src=(*sink.src, sink)), UOp(Ops.SOURCE, arg=code), UOp(Ops.BINARY, arg=lib)))
@functools.cache
def custom_fused_qkv_rope_backward(dxqkv:UOp, dq:UOp, dk:UOp, dv:UOp, freqs_cis:UOp,
device:str, arch:str, B:int, N:int, H:int, H_KV:int, D:int):
assert (B, N, H, H_KV, D) == (2, 8192, 32, 8, 128)
code = (pathlib.Path(__file__).parent / "fused_qkv_rope_bwd.cpp").read_text()
threads = 256
thread_idx = UOp.special(threads, "lidx0")
gsz = (B, N // 64, H + 2 * H_KV)
block_idx_x, block_idx_y, block_idx_z = (UOp.special(x, f"gidx{i}") for i, x in enumerate(gsz))
sink = UOp.sink(dxqkv.base, dq.base, dk.base, dv.base, freqs_cis.base, thread_idx, block_idx_x, block_idx_y, block_idx_z,
arg=KernelInfo(name="fused_qkv_rope_backward"))
compile_args = [f"-I{(pathlib.Path(__file__).parent / 'include').as_posix()}", "-std=c++20", "-DKITTENS_CDNA4", "-DHIP_ENABLE_WARP_SYNC_BUILTINS", "-ffast-math", f"-DATTN_B={B}", f"-DATTN_N={N}", f"-DATTN_H={H}",
f"-DATTN_H_KV={H_KV}", f"-DATTN_D={D}", f"-DTHREADS_PER_BLOCK={threads}"]
lib = HIPCCCompiler(arch, compile_args).compile_cached(code)
return UOp(Ops.PROGRAM, src=(sink, UOp(Ops.LINEAR, src=(*sink.src, sink)), UOp(Ops.SOURCE, arg=code), UOp(Ops.BINARY, arg=lib)))
def _fa_native_grads(dq:UOp, dk:UOp, dv:UOp) -> tuple[UOp, UOp, UOp]|None:
def unwrap_partial(x:UOp) -> UOp|None:
expected = (Ops.CAST, Ops.REDUCE, Ops.PERMUTE, Ops.CAST, Ops.RESHAPE, Ops.AFTER)
for op in expected:
if x.op is not op: return None
if op is not Ops.AFTER: x = x.src[0]
return x
dq_native, dk_partial, dv_partial = dq.base, unwrap_partial(dk), unwrap_partial(dv)
if dq_native.op is not Ops.AFTER or dk_partial is None or dv_partial is None: return None
B, N, H, D, H_KV = dq.shape[0], dq.shape[1], dq.shape[2], dq.shape[3], dk.shape[2]
heads_per_wg = 2 if D == 128 and (H // H_KV) % 2 == 0 else 1
partials = (H // H_KV) // heads_per_wg
if dq_native.shape != (B, H, N, D) or dk_partial.shape != (B * partials, N, H_KV, D) or dv_partial.shape != dk_partial.shape: return None
return dq_native, dk_partial, dv_partial
def _fused_qkv_rope_grad(dq_u:UOp, dk_u:UOp, dv_u:UOp, call:UOp) -> tuple[None, None, None, UOp, None]:
dq, dk, dv = Tensor(dq_u, device=dq_u.device), Tensor(dk_u, device=dk_u.device), Tensor(dv_u, device=dv_u.device)
xqkv_u, freqs_u = call.src[4], call.src[5]
xqkv, freqs_cis = Tensor(xqkv_u, device=xqkv_u.device), Tensor(freqs_u, device=freqs_u.device)
B, N, _ = xqkv.shape
H, H_KV, D = dq.shape[2], dk.shape[2], dq.shape[3]
num_devices = len(xqkv.device) if isinstance(xqkv.device, tuple) else 1
is_dp, is_mp = xqkv.uop.axis == 0, xqkv.uop.axis == 2
B_local = B // num_devices if is_dp else B
H_local = H // num_devices if is_mp else H
H_KV_local = H_KV // num_devices if is_mp else H_KV
single_device = xqkv.device[0] if isinstance(xqkv.device, tuple) else xqkv.device
arch = Device[single_device].renderer.target.arch
fa_native = _fa_native_grads(dq_u, dk_u, dv_u)
assert fa_native is not None, "fused QKV RoPE backward requires native Flash Attention gradients"
dq, dk, dv = (Tensor(x, device=x.device) for x in fa_native)
dxqkv = _sharded_empty_like(xqkv, axis=xqkv.uop.axis if isinstance(xqkv.device, tuple) else None)
fxn = functools.partial(custom_fused_qkv_rope_backward, device=single_device, arch=arch,
B=B_local, N=N, H=H_local, H_KV=H_KV_local, D=D)
dxqkv = Tensor.custom_kernel(dxqkv, dq, dk, dv, freqs_cis, fxn=fxn)[0]
return None, None, None, dxqkv.uop, None
def fused_qkv_rope(xqkv:Tensor, freqs_cis:Tensor, n_heads:int, n_kv_heads:int, head_dim:int) -> tuple[Tensor, Tensor, Tensor]:
B, N, packed_dim = xqkv.shape
assert packed_dim == n_kv_heads * (n_heads // n_kv_heads + 2) * head_dim
assert freqs_cis.dtype == dtypes.bfloat16, f"fused QKV RoPE requires bfloat16 frequencies, got {freqs_cis.dtype}"
assert freqs_cis.shape == (1, freqs_cis.shape[1], 1, head_dim // 2, 2) and freqs_cis.shape[1] >= N, \
f"invalid RoPE frequency shape {freqs_cis.shape} for sequence length {N} and head dimension {head_dim}"
num_devices = len(xqkv.device) if isinstance(xqkv.device, tuple) else 1
is_dp, is_mp = xqkv.uop.axis == 0, xqkv.uop.axis == 2
B_local = B // num_devices if is_dp else B
H_local = n_heads // num_devices if is_mp else n_heads
H_KV_local = n_kv_heads // num_devices if is_mp else n_kv_heads
assert H_local % H_KV_local == 0 and head_dim % 2 == 0 and head_dim <= 512
single_device = xqkv.device[0] if isinstance(xqkv.device, tuple) else xqkv.device
arch = Device[single_device].renderer.target.arch
axis = 0 if is_dp else 2 if is_mp else None
q = _sharded_empty((B, N, n_heads, head_dim), xqkv, axis=axis, dtype=dtypes.bfloat16)
k = _sharded_empty((B, N, n_kv_heads, head_dim), xqkv, axis=axis, dtype=dtypes.bfloat16)
v = _sharded_empty((B, N, n_kv_heads, head_dim), xqkv, axis=axis, dtype=dtypes.bfloat16)
fxn = functools.partial(custom_fused_qkv_rope_forward, device=single_device, arch=arch,
B=B_local, N=N, H=H_local, H_KV=H_KV_local, D=head_dim)
q, k, v, *_ = Tensor.custom_kernel(q, k, v, xqkv, freqs_cis, fxn=fxn, grad_fxn=_fused_qkv_rope_grad)
return q, k, v
def _sharded_empty_like(ref:Tensor, axis:int|None=None) -> Tensor:
return _sharded_empty(ref.shape, ref, axis)
@@ -31,8 +121,9 @@ def _fa_grad_fxn(B, H, N, D, H_local, H_KV_local, H_KV, B_local, shard_axis, sha
dq = _sharded_empty((B, H, N, D), xq, axis=shard_axis_t)
GROUP_SIZE = H_local // H_KV_local
dk_partial = _sharded_empty((B * GROUP_SIZE, N, H_KV, D), xk, axis=shard_axis)
dv_partial = _sharded_empty((B * GROUP_SIZE, N, H_KV, D), xv, axis=shard_axis)
HEADS_PER_WG = 2 if D == 128 and GROUP_SIZE % 2 == 0 else 1
dk_partial = _sharded_empty((B * GROUP_SIZE // HEADS_PER_WG, N, H_KV, D), xk, axis=shard_axis)
dv_partial = _sharded_empty((B * GROUP_SIZE // HEADS_PER_WG, N, H_KV, D), xv, axis=shard_axis)
# delta_vec = (do * attn).sum(-1, dtype=dtypes.float32).transpose(1, 2).unsqueeze(-2).detach()
delta_vec = _sharded_empty((B, H, 1, N), xq, dtype=dtypes.float32, axis=shard_axis_t)
@@ -46,8 +137,8 @@ def _fa_grad_fxn(B, H, N, D, H_local, H_KV_local, H_KV, B_local, shard_axis, sha
dq = dq.reshape(B, H, N//16, 4, 2, 2, D//32, 4, 4, 2).permute(0, 1, 2, 7, 8, 3, 4, 6, 5, 9).reshape(B, H, N, D).transpose(1, 2)
# reduce partial dK/dV across GROUP_SIZE query heads
dk = dk_partial.reshape(B, GROUP_SIZE, N, H_KV, D).sum(1)
dv = dv_partial.reshape(B, GROUP_SIZE, N, H_KV, D).sum(1)
dk = dk_partial.reshape(B, GROUP_SIZE // HEADS_PER_WG, N, H_KV, D).sum(1)
dv = dv_partial.reshape(B, GROUP_SIZE // HEADS_PER_WG, N, H_KV, D).sum(1)
if not has_sink: return None, None, dq.uop, dk.uop, dv.uop
sinks = Tensor(ker.src[6], device=ker.src[6].device)
@@ -160,9 +251,11 @@ def custom_fa_backward(dq:UOp, dk:UOp, dv:UOp, do:UOp, q:UOp, k:UOp, v:UOp, l_ve
f"-DATTN_B={B}", f"-DATTN_N={N}", f"-DATTN_H={H}", f"-DATTN_H_KV={H_KV}", f"-DATTN_D={D}"]
BLOCK_SIZE_KV = 256
GROUP_SIZE = H // H_KV
HEADS_PER_WG = 2 if D == 128 and GROUP_SIZE % 2 == 0 else 1
NUM_WARPS = 4
NUM_THREADS = 64 * NUM_WARPS
gsz = (H, N // BLOCK_SIZE_KV, B)
gsz = (H // HEADS_PER_WG, N // BLOCK_SIZE_KV, B)
lsz = (NUM_THREADS, 1, 1)
threadIdx_x = UOp.special(lsz[0], "lidx0")
blockIdx_x, blockIdx_y, blockIdx_z = UOp.special(gsz[0], "gidx0"), UOp.special(gsz[1], "gidx1"), UOp.special(gsz[2], "gidx2")
+24 -23
View File
@@ -28,6 +28,7 @@ constexpr int ATTN_H_KV = 8; // number of key/value heads (for GQA)
#endif
constexpr int GROUP_SIZE = ATTN_H / ATTN_H_KV; // queries per KV head group
constexpr int HEADS_PER_WG = (ATTN_D == 128 && GROUP_SIZE % 2 == 0) ? 2 : 1;
#ifndef ATTN_N
constexpr int ATTN_N = 1024; // sequence length
@@ -53,7 +54,7 @@ using namespace kittens;
using _gl_QdO = gl<bf16, ATTN_B, ATTN_N, ATTN_H, ATTN_D>;
using _gl_KV = gl<bf16, ATTN_B, ATTN_N, ATTN_H_KV, ATTN_D>;
using _gl_dQ = gl<bf16, ATTN_B, ATTN_H, ATTN_N, ATTN_D>;
using _gl_dKV = gl<bf16, ATTN_B * GROUP_SIZE, ATTN_N, ATTN_H_KV, ATTN_D>;
using _gl_dKV = gl<bf16, ATTN_B * (GROUP_SIZE / HEADS_PER_WG), ATTN_N, ATTN_H_KV, ATTN_D>;
using _gl_Lvec = gl<float, ATTN_B, ATTN_H, 1, ATTN_N>;
template<int D> struct attn_bwd_combined_globals {
@@ -63,7 +64,7 @@ template<int D> struct attn_bwd_combined_globals {
_gl_dQ dQg;
_gl_dKV dKg, dVg;
_gl_Lvec L_vec, delta_vec;
dim3 grid() { return dim3(ATTN_H, (ATTN_N / BLOCK_SIZE_KV), ATTN_B); }
dim3 grid() { return dim3(ATTN_H / HEADS_PER_WG, (ATTN_N / BLOCK_SIZE_KV), ATTN_B); }
dim3 block() { return dim3(NUM_THREADS); }
size_t dynamic_shared_memory() { return MAX_SHARED_MEMORY; }
};
@@ -71,7 +72,7 @@ template<int D> struct attn_bwd_combined_globals {
template<int D> __launch_bounds__(NUM_THREADS, 1)
__global__ void attend_bwd_combined_ker(bf16 *dQ_ptr, bf16 *dK_ptr, bf16 *dV_ptr, bf16 *dO_ptr, bf16 *Q_ptr, bf16 *K_ptr, bf16 *V_ptr, float *L_vec_ptr, float *delta_vec_ptr) {
const int q_head_idx_fixed = blockIdx.x; // This is the query head index [0, ATTN_H)
const int q_head_idx_fixed = blockIdx.x * HEADS_PER_WG; // First query head handled by this workgroup.
const int kv_head_idx = q_head_idx_fixed / GROUP_SIZE;
const int q_head_in_group = q_head_idx_fixed % GROUP_SIZE;
const int seq_idx = blockIdx.y;
@@ -88,7 +89,7 @@ __global__ void attend_bwd_combined_ker(bf16 *dQ_ptr, bf16 *dK_ptr, bf16 *dV_ptr
// first Q step that can overlap this K_span:
const int first_step = max(0, k_start_min / STEP_QO);
const int num_steps_per_head = total_steps_per_head - first_step;
const int num_steps = num_steps_per_head;
const int num_steps = num_steps_per_head * HEADS_PER_WG;
const int k_pos = j * WARP_SIZE_KV;
constexpr float L_SCALE_FACTOR = 1.44269504089f;
@@ -270,12 +271,12 @@ __global__ void attend_bwd_combined_ker(bf16 *dQ_ptr, bf16 *dK_ptr, bf16 *dV_ptr
if (num_steps > 1) {
// Prologue
{
const int q_head_idx = (0) / num_steps_per_head + first_q_head;
const int q_seq_idx = ((0) % num_steps_per_head) + first_step;
const int q_head_idx = first_q_head;
const int q_seq_idx = first_step;
const int q_pos = q_seq_idx * STEP_QO;
const int next_q_head_idx = (0 + 1) / num_steps_per_head + first_q_head;
const int next_q_seq_idx = ((0 + 1) % num_steps_per_head) + first_step;
const int next_q_head_idx = first_q_head;
const int next_q_seq_idx = first_step + 1;
// dot slice 0
{
@@ -1332,15 +1333,18 @@ __global__ void attend_bwd_combined_ker(bf16 *dQ_ptr, bf16 *dK_ptr, bf16 *dV_ptr
// 9. for 1 <= i <= T_r (1024 / 32 = 32)
for (int i = 1; i < num_steps - 1; ++i, tic ^= 1, toc ^= 1) {
const int last_q_head_idx = (i - 1) / num_steps_per_head + first_q_head;
const int last_q_seq_idx = ((i - 1) % num_steps_per_head) + first_step;
const int last_head_offset = (i - 1) >= num_steps_per_head;
const int last_q_head_idx = last_head_offset + first_q_head;
const int last_q_seq_idx = i - 1 - last_head_offset * num_steps_per_head + first_step;
const int q_head_idx = i / num_steps_per_head + first_q_head;
const int q_seq_idx = (i % num_steps_per_head) + first_step;
const int head_offset = i >= num_steps_per_head;
const int q_head_idx = head_offset + first_q_head;
const int q_seq_idx = i - head_offset * num_steps_per_head + first_step;
const int q_pos = q_seq_idx * STEP_QO;
const int next_q_head_idx = (i + 1) / num_steps_per_head + first_q_head;
const int next_q_seq_idx = ((i + 1) % num_steps_per_head) + first_step;
const int next_head_offset = (i + 1) >= num_steps_per_head;
const int next_q_head_idx = next_head_offset + first_q_head;
const int next_q_seq_idx = i + 1 - next_head_offset * num_steps_per_head + first_step;
// dot slice 0
{
@@ -2378,11 +2382,11 @@ __global__ void attend_bwd_combined_ker(bf16 *dQ_ptr, bf16 *dK_ptr, bf16 *dV_ptr
}
}
const int last_q_head_idx = (num_steps - 2) / num_steps_per_head + first_q_head;
const int last_q_seq_idx = ((num_steps - 2) % num_steps_per_head) + first_step;
const int last_q_head_idx = first_q_head + HEADS_PER_WG - 1;
const int last_q_seq_idx = first_step + num_steps_per_head - 2;
const int q_head_idx = (num_steps - 1) / num_steps_per_head + first_q_head;
const int q_seq_idx = ((num_steps - 1) % num_steps_per_head) + first_step;
const int q_head_idx = first_q_head + HEADS_PER_WG - 1;
const int q_seq_idx = first_step + num_steps_per_head - 1;
const int q_pos = q_seq_idx * STEP_QO;
// Epilogue
{
@@ -3407,14 +3411,14 @@ __global__ void attend_bwd_combined_ker(bf16 *dQ_ptr, bf16 *dK_ptr, bf16 *dV_ptr
}
}
store<1>(g.dVg, dV_j, {batch_idx * GROUP_SIZE + q_head_in_group, 0, kv_head_idx, 0}, {0, j, 0, 0});
store<1>(g.dVg, dV_j, {batch_idx * (GROUP_SIZE / HEADS_PER_WG) + q_head_in_group / HEADS_PER_WG, 0, kv_head_idx, 0}, {0, j, 0, 0});
__builtin_amdgcn_s_waitcnt(0);
__builtin_amdgcn_s_barrier();
// We first copy dV_j_T from accumulator GPRs to vector GPRs and then perform the store
accvgpr_read(dV_j_T, dK_j_T);
mul(dV_j_T, dV_j_T, dP_SCALE_FACTOR);
store<1>(g.dKg, dV_j, {batch_idx * GROUP_SIZE + q_head_in_group, 0, kv_head_idx, 0}, {0, j, 0, 0});
store<1>(g.dKg, dV_j, {batch_idx * (GROUP_SIZE / HEADS_PER_WG) + q_head_in_group / HEADS_PER_WG, 0, kv_head_idx, 0}, {0, j, 0, 0});
// Write out final dQ_i slice
mul(dQ_i_T, dQ_i_T, dP_SCALE_FACTOR);
@@ -3422,6 +3426,3 @@ __global__ void attend_bwd_combined_ker(bf16 *dQ_ptr, bf16 *dK_ptr, bf16 *dV_ptr
}
template __global__ void attend_bwd_combined_ker<ATTN_D>(bf16*, bf16*, bf16*, bf16*, bf16*, bf16*, bf16*, float*, float*);
+69
View File
@@ -0,0 +1,69 @@
#include <hip/hip_runtime.h>
#include <hip/hip_bf16.h>
#ifndef ATTN_B
#define ATTN_B 2
#endif
#ifndef ATTN_N
#define ATTN_N 8192
#endif
#ifndef ATTN_H
#define ATTN_H 32
#endif
#ifndef ATTN_H_KV
#define ATTN_H_KV 8
#endif
#ifndef ATTN_D
#define ATTN_D 128
#endif
#ifndef THREADS_PER_BLOCK
#define THREADS_PER_BLOCK 256
#endif
constexpr int GROUP_SIZE = ATTN_H / ATTN_H_KV;
constexpr int HALF_D = ATTN_D / 2;
constexpr int PACKED_D = (GROUP_SIZE + 2) * ATTN_D;
extern "C" __global__ __launch_bounds__(THREADS_PER_BLOCK) void
fused_qkv_rope_forward(
__hip_bfloat16* __restrict__ q,
__hip_bfloat16* __restrict__ k,
__hip_bfloat16* __restrict__ v,
const __hip_bfloat16* __restrict__ xqkv,
const __hip_bfloat16* __restrict__ freqs_cis) {
const int b = blockIdx.x;
const int n = blockIdx.y;
const int bn = b * ATTN_N + n;
const int packed_bn = bn * ATTN_H_KV * PACKED_D;
const int q_bn = bn * ATTN_H * ATTN_D;
const int kv_bn = bn * ATTN_H_KV * ATTN_D;
if (threadIdx.x < HALF_D) {
const int pair = threadIdx.x;
const int even = pair << 1;
const float c = static_cast<float>(freqs_cis[((n * HALF_D + pair) * 2) + 0]);
const float s = static_cast<float>(freqs_cis[((n * HALF_D + pair) * 2) + 1]);
for (int kvh = 0; kvh < ATTN_H_KV; kvh++) {
const int base = packed_bn + kvh * PACKED_D;
for (int rep = 0; rep < GROUP_SIZE; rep++) {
const int qbase = base + rep * ATTN_D;
const int h = kvh * GROUP_SIZE + rep;
const float a = static_cast<float>(xqkv[qbase + even]);
const float bb = static_cast<float>(xqkv[qbase + even + 1]);
const int out = q_bn + h * ATTN_D + even;
q[out] = static_cast<__hip_bfloat16>(a * c - bb * s);
q[out + 1] = static_cast<__hip_bfloat16>(a * s + bb * c);
}
const float a = static_cast<float>(xqkv[base + GROUP_SIZE * ATTN_D + even]);
const float bb = static_cast<float>(xqkv[base + GROUP_SIZE * ATTN_D + even + 1]);
const int out = kv_bn + kvh * ATTN_D + even;
k[out] = static_cast<__hip_bfloat16>(a * c - bb * s);
k[out + 1] = static_cast<__hip_bfloat16>(a * s + bb * c);
v[out] = xqkv[base + (GROUP_SIZE + 1) * ATTN_D + even];
v[out + 1] = xqkv[base + (GROUP_SIZE + 1) * ATTN_D + even + 1];
}
}
}
+161
View File
@@ -0,0 +1,161 @@
#include "kittens.cuh"
using namespace kittens;
#ifndef ATTN_B
#define ATTN_B 2
#endif
#ifndef ATTN_N
#define ATTN_N 8192
#endif
#ifndef ATTN_H
#define ATTN_H 32
#endif
#ifndef ATTN_H_KV
#define ATTN_H_KV 8
#endif
#ifndef ATTN_D
#define ATTN_D 128
#endif
#ifndef THREADS_PER_BLOCK
#define THREADS_PER_BLOCK 256
#endif
constexpr int GROUP_SIZE = ATTN_H / ATTN_H_KV;
constexpr int HALF_D = ATTN_D / 2;
constexpr int PACKED_H = ATTN_H_KV * (GROUP_SIZE + 2);
constexpr int HEADS_PER_WG = ATTN_D == 128 && GROUP_SIZE % 2 == 0 ? 2 : 1;
constexpr int KV_PARTIALS = GROUP_SIZE / HEADS_PER_WG;
constexpr int NUM_WARPS = 4;
constexpr int TILE_N = 16;
template<typename T> using grad_tile = rt<T, TILE_N, ATTN_D, row_l, rt_16x32_s>;
template<int axis, ducks::rt::row_layout RT, ducks::gl::all GL, ducks::coord::tile COORD=coord<RT>>
__device__ __forceinline__ void load_fa_shuffled(RT &dst, const GL &src, const COORD &idx) {
using U = typename GL::dtype;
using U2 = base_types::packing<U>::packed_type;
U *src_ptr = (U*)&src[(idx.template unit_coord<axis, 3>())];
const int row_stride = src.template stride<axis>();
const int lane = kittens::laneid();
const int tile_row_stride = row_stride * dst.base_tile_rows;
const int tile_stride = dst.base_tile_rows * dst.base_tile_cols;
const uint32_t buffer_size = src.batch() * src.depth() * src.rows() * src.cols() * sizeof(U);
const buffer_resource br = make_buffer_resource(reinterpret_cast<uintptr_t>(src_ptr), buffer_size, 0x00020000);
#pragma unroll
for (int i = 0; i < dst.height; i++) {
#pragma unroll
for (int j = 0; j < dst.width; j++) {
const float4 loaded = std::bit_cast<float4>(llvm_amdgcn_raw_buffer_load_b128(
std::bit_cast<i32x4>(br), (i * tile_row_stride + j * tile_stride + lane * 8) * sizeof(U), 0, 0));
const U2 *packed = reinterpret_cast<const U2*>(&loaded);
#pragma unroll
for (int k = 0; k < dst.packed_per_base_tile; k++) dst.tiles[i][j].data[k] = packed[k];
}
}
}
template<int axis, ducks::rt::row_layout RT, ducks::gl::all GL, ducks::coord::tile COORD=coord<RT>>
__device__ __forceinline__ void store_fa_shuffled(const GL &dst, const RT &src, const COORD &idx) {
using U = typename GL::dtype;
U *dst_ptr = (U*)&dst[(idx.template unit_coord<axis, 3>())];
const int row_stride = dst.template stride<axis>();
const int lane = kittens::laneid();
const int row_offset = (lane % 4) * 4;
const int col_offset = ((lane / 32) * 16) + (((lane % 32) / 16) * 2) + (((lane % 16) / 4) * 4);
const uint32_t buffer_size = dst.batch() * dst.depth() * dst.rows() * dst.cols() * sizeof(U);
const buffer_resource br = make_buffer_resource(reinterpret_cast<uintptr_t>(dst_ptr), buffer_size, 0x00020000);
#pragma unroll
for (int i = 0; i < src.height; i++) {
const int row = src.base_tile_rows * i + row_offset;
#pragma unroll
for (int j = 0; j < src.width; j++) {
const int col = src.base_tile_cols * j + col_offset;
#pragma unroll
for (int k = 0; k < src.packed_per_base_tile; k++) llvm_amdgcn_raw_buffer_store_b32(
*reinterpret_cast<const uint32_t*>(&src.tiles[i][j].data[k]), std::bit_cast<i32x4>(br),
((row + k) * row_stride + col) * sizeof(U), 0, 0);
}
}
}
template<ducks::rt::row_layout RT>
__device__ __forceinline__ void inverse_rope_fa(RT &tile, const bf16_2 *freqs, const int n_base) {
const int lane = kittens::laneid();
const int row_offset = (lane % 4) * 4;
const int col_offset = ((lane / 32) * 16) + (((lane % 32) / 16) * 2) + (((lane % 16) / 4) * 4);
#pragma unroll
for (int i = 0; i < tile.height; i++) {
#pragma unroll
for (int j = 0; j < tile.width; j++) {
const int col = tile.base_tile_cols * j + col_offset;
#pragma unroll
for (int k = 0; k < tile.packed_per_base_tile; k++) {
const int row = tile.base_tile_rows * i + row_offset + k;
const float2 cs = __bfloat1622float2(freqs[(n_base + row) * HALF_D + col / 2]);
const float2 g = __bfloat1622float2(tile.tiles[i][j].data[k]);
tile.tiles[i][j].data[k] = __float22bfloat162_rn(make_float2(g.x * cs.x + g.y * cs.y, -g.x * cs.y + g.y * cs.x));
}
}
}
}
template<ducks::rt::row_layout RT>
__device__ __forceinline__ void inverse_rope(RT &tile, const bf16_2 *freqs, const int n_base) {
const int lane = kittens::laneid();
#pragma unroll
for (int i = 0; i < tile.height; i++) {
const int row = tile.base_tile_rows * i + lane % tile.base_tile_rows;
#pragma unroll
for (int j = 0; j < tile.width; j++) {
#pragma unroll
for (int k = 0; k < tile.packed_per_base_tile; k++) {
const int col = tile.base_tile_cols * j + tile.base_tile_stride * (lane / tile.base_tile_rows) + 2 * k;
const float2 cs = __bfloat1622float2(freqs[(n_base + row) * HALF_D + col / 2]);
const float2 g = __bfloat1622float2(tile.tiles[i][j].data[k]);
tile.tiles[i][j].data[k] = __float22bfloat162_rn(make_float2(g.x * cs.x + g.y * cs.y, -g.x * cs.y + g.y * cs.x));
}
}
}
}
extern "C" __global__ __launch_bounds__(THREADS_PER_BLOCK) void
fused_qkv_rope_backward(
bf16* __restrict__ dxqkv,
const bf16* __restrict__ dq,
const bf16* __restrict__ dk,
const bf16* __restrict__ dv,
const bf16* __restrict__ freqs_cis) {
gl<bf16, -1, -1, -1, -1> out{dxqkv, ATTN_B, ATTN_N, PACKED_H, ATTN_D};
gl<bf16, -1, -1, -1, -1> dqg{const_cast<bf16*>(dq), ATTN_B, ATTN_H, ATTN_N, ATTN_D};
gl<bf16, -1, -1, -1, -1> dkg{const_cast<bf16*>(dk), ATTN_B * KV_PARTIALS, ATTN_N, ATTN_H_KV, ATTN_D};
gl<bf16, -1, -1, -1, -1> dvg{const_cast<bf16*>(dv), ATTN_B * KV_PARTIALS, ATTN_N, ATTN_H_KV, ATTN_D};
const int b = blockIdx.x, n_tile = blockIdx.y * NUM_WARPS + kittens::warpid(), n_base = n_tile * TILE_N;
const int field = blockIdx.z;
if (field < ATTN_H) {
grad_tile<bf16> tile;
load_fa_shuffled<2>(tile, dqg, {b, field, n_tile, 0});
inverse_rope_fa(tile, reinterpret_cast<const bf16_2*>(freqs_cis), n_base);
const int out_head = (field / GROUP_SIZE) * (GROUP_SIZE + 2) + field % GROUP_SIZE;
store_fa_shuffled<1>(out, tile, {b, n_tile, out_head, 0});
} else {
const bool is_k = field < ATTN_H + ATTN_H_KV;
const int kvh = field - ATTN_H - (is_k ? 0 : ATTN_H_KV);
const auto &src = is_k ? dkg : dvg;
grad_tile<bf16> partial, tile;
grad_tile<float> partial_f, sum;
zero(sum);
#pragma unroll
for (int p = 0; p < KV_PARTIALS; p++) {
load<1>(partial, src, {b * KV_PARTIALS + p, n_tile, kvh, 0});
copy(partial_f, partial);
add(sum, sum, partial_f);
}
copy(tile, sum);
if (is_k) inverse_rope(tile, reinterpret_cast<const bf16_2*>(freqs_cis), n_base);
const int out_head = kvh * (GROUP_SIZE + 2) + GROUP_SIZE + !is_k;
store<1>(out, tile, {b, n_tile, out_head, 0});
}
}
+22 -6
View File
@@ -148,12 +148,28 @@ __global__ __launch_bounds__(512, 2) void hk_fp8_gemm(bf16 *C_ptr, fp8e4m3 *A_pt
RT_C cC;
RT_C cD;
// Calculate which block this threadblock should work on
int global_block_id = blockIdx.x;
// Convert linear block ID to 2D coordinates
int block_row = global_block_id / blocks_per_col;
int block_col = global_block_id % blocks_per_col;
int block_row, block_col;
if constexpr (N > M) {
// Wide outputs repeatedly consume the same A rows. Keep a short strip
// resident on each XCD while walking N to improve local cache reuse.
int wgid = chiplet_transform_chunked(int(blockIdx.x), total_blocks_needed, NUM_XCDS, 64);
constexpr int WGM = 3;
const int num_wgid_in_group = WGM * blocks_per_col;
const int group_id = wgid / num_wgid_in_group;
const int first_block_row = group_id * WGM;
const int group_size_m = min(blocks_per_row - first_block_row, WGM);
block_row = first_block_row + ((wgid % num_wgid_in_group) % group_size_m);
block_col = (wgid % num_wgid_in_group) / group_size_m;
} else {
int wgid = chiplet_transform_chunked(int(blockIdx.x), total_blocks_needed, NUM_XCDS, 64);
constexpr int WGM = 8;
const int num_wgid_in_group = WGM * blocks_per_col;
const int group_id = wgid / num_wgid_in_group;
const int first_block_row = group_id * WGM;
const int group_size_m = min(blocks_per_row - first_block_row, WGM);
block_row = first_block_row + ((wgid % num_wgid_in_group) % group_size_m);
block_col = (wgid % num_wgid_in_group) / group_size_m;
}
int block_m = block_row * BLOCK_SIZE_ROW;
int block_n = block_col * BLOCK_SIZE_COL;
+1 -1
View File
@@ -134,7 +134,7 @@ __global__ __launch_bounds__(512, 2) void hk_fp8_atb_gemm(bf16 *C_ptr, fp8e4m3 *
int wgid = blockIdx.x;
const int WGM = 8;
wgid = chiplet_transform_chunked(wgid, total_blocks_needed, NUM_XCDS, 64);
wgid = chiplet_transform_chunked(wgid, total_blocks_needed, NUM_XCDS, 32);
const int num_wgid_in_group = WGM * blocks_per_col;
int group_id = wgid / num_wgid_in_group;
+13 -14
View File
@@ -1,7 +1,6 @@
import math
from typing import cast, Callable
from tinygrad import dtypes
from tinygrad.uop.ops import AxisType, UOp, Ops
from tinygrad.uop.ops import AxisType, UOp
from tinygrad.dtype import AddrSpace
from tinygrad.helpers import prod
@@ -75,9 +74,9 @@ class Group:
a_base_shape = cast(RT, a).base_shape
if a_base_shape.cols == 16:
wmma_arg = ('WMMA_16_16_16___bf16_float', (16, 16, 16), dtypes.bfloat16, dtypes.float, 'AMD', 64, (((4, 2), (3, 2)), ((4, 2), (3, 2)), ((4, 2), (3, 2))), ()) # type: ignore
wmma_dims = (16, 16, 16)
elif a_base_shape.cols == 32:
wmma_arg = ('WMMA_16_16_32___bf16_float', (16, 16, 32), dtypes.bfloat16, dtypes.float, 'AMD', 64, (((4, 2), (3, 2), (9, 2)), ((4, 2), (3, 2), (9, 2)), ((4, 2), (3, 2))), ()) # type: ignore
wmma_dims = (16, 16, 32)
else: raise NotImplementedError(f"mma_AB not implemented for {a_base_shape.cols=}")
for height in self.ker.range(c.shape[-3], track=False):
@@ -92,7 +91,7 @@ class Group:
else: raise NotImplementedError(f"mma_AB not implemented for {a_base_shape.cols=}")
d_in = UOp.stack(*[c[height, width, i] for i in range(4)])
out = UOp(Ops.WMMA, dtypes.float32, (a_in, b_in, d_in), arg=wmma_arg)
out = UOp.wmma(a_in, b_in, d_in, wmma_dims, 'AMD', 64)
c_i = [c[height, width, i].store(out.index(i)) for i in range(4)]
c_store = UOp.group(*c_i).end(height, width, inner)
@@ -105,9 +104,9 @@ class Group:
a_base_shape = cast(RT, a).base_shape
if a_base_shape.cols == 16:
wmma_arg = ('WMMA_16_16_16___bf16_float', (16, 16, 16), dtypes.bfloat16, dtypes.float, 'AMD', 64, (((4, 2), (3, 2)), ((4, 2), (3, 2)), ((4, 2), (3, 2))), ()) # type: ignore
wmma_dims = (16, 16, 16)
elif a_base_shape.cols == 32:
wmma_arg = ('WMMA_16_16_32___bf16_float', (16, 16, 32), dtypes.bfloat16, dtypes.float, 'AMD', 64, (((4, 2), (3, 2), (9, 2)), ((4, 2), (3, 2), (9, 2)), ((4, 2), (3, 2))), ()) # type: ignore
wmma_dims = (16, 16, 32)
else: raise NotImplementedError(f"mma_ABt not implemented for {a_base_shape.cols=}")
for height in self.ker.range(c.shape[-3], track=False):
@@ -122,7 +121,7 @@ class Group:
else: raise NotImplementedError(f"mma_ABt not implemented for {a_base_shape.cols=}")
d_in = UOp.stack(*[c[height, width, i] for i in range(4)])
out = UOp(Ops.WMMA, dtypes.float32, (a_in, b_in, d_in), arg=wmma_arg)
out = UOp.wmma(a_in, b_in, d_in, wmma_dims, 'AMD', 64)
c_i = [c[height, width, i].store(out.index(i)) for i in range(4)]
c_store = UOp.group(*c_i).end(height, width, inner)
@@ -135,9 +134,9 @@ class Group:
a_base_shape = cast(RT, a).base_shape
if a_base_shape.cols == 16:
wmma_arg = ('WMMA_16_16_16___bf16_float', (16, 16, 16), dtypes.bfloat16, dtypes.float, 'AMD', 64, (((4, 2), (3, 2)), ((4, 2), (3, 2)), ((4, 2), (3, 2))), ()) # type: ignore
wmma_dims = (16, 16, 16)
elif a_base_shape.cols == 32:
wmma_arg = ('WMMA_16_16_32___bf16_float', (16, 16, 32), dtypes.bfloat16, dtypes.float, 'AMD', 64, (((4, 2), (3, 2), (9, 2)), ((4, 2), (3, 2), (9, 2)), ((4, 2), (3, 2))), ()) # type: ignore
wmma_dims = (16, 16, 32)
else: raise NotImplementedError(f"mma_AtB not implemented for {a_base_shape.cols=}")
for height in self.ker.range(c.shape[-3], track=False):
@@ -152,7 +151,7 @@ class Group:
else: raise NotImplementedError(f"mma_AtB not implemented for {a_base_shape.cols=}")
d_in = UOp.stack(*[c[height, width, i] for i in range(4)])
out = UOp(Ops.WMMA, dtypes.float32, (a_in, b_in, d_in), arg=wmma_arg)
out = UOp.wmma(a_in, b_in, d_in, wmma_dims, 'AMD', 64)
c_i = [c[height, width, i].store(out.index(i)) for i in range(4)]
c_store = UOp.group(*c_i).end(height, width, inner)
@@ -165,9 +164,9 @@ class Group:
a_base_shape = cast(RT, a).base_shape
if a_base_shape.cols == 16:
wmma_arg = ('WMMA_16_16_16___bf16_float', (16, 16, 16), dtypes.bfloat16, dtypes.float, 'AMD', 64, (((4, 2), (3, 2)), ((4, 2), (3, 2)), ((4, 2), (3, 2))), ()) # type: ignore
wmma_dims = (16, 16, 16)
elif a_base_shape.cols == 32:
wmma_arg = ('WMMA_16_16_32___bf16_float', (16, 16, 32), dtypes.bfloat16, dtypes.float, 'AMD', 64, (((4, 2), (3, 2), (9, 2)), ((4, 2), (3, 2), (9, 2)), ((4, 2), (3, 2))), ()) # type: ignore
wmma_dims = (16, 16, 32)
else: raise NotImplementedError(f"mma_AtBt not implemented for {a_base_shape.cols=}")
for height in self.ker.range(c.shape[-3], track=False):
@@ -182,7 +181,7 @@ class Group:
else: raise NotImplementedError(f"mma_AtBt not implemented for {a_base_shape.cols=}")
d_in = UOp.stack(*[c[height, width, i] for i in range(4)])
out = UOp(Ops.WMMA, dtypes.float32, (a_in, b_in, d_in), arg=wmma_arg)
out = UOp.wmma(a_in, b_in, d_in, wmma_dims, 'AMD', 64)
c_i = [c[height, width, i].store(out.index(i)) for i in range(4)]
c_store = UOp.group(*c_i).end(height, width, inner)
+2
View File
@@ -565,8 +565,10 @@ tiny_backend = {**{k:wrap_out(v) for k,v in tiny_backend_out.items()}, **{
"aten.floor_divide": lambda x,y: x//y,
"aten.floor_divide_.Tensor": lambda x,y: x//y,
"aten.__lshift__.Scalar": lambda x,y: x<<y,
"aten.__lshift__.Tensor": lambda x,y: x<<y,
"aten.__ilshift__.Scalar": lambda x,y: x<<y,
"aten.__rshift__.Scalar": lambda x,y: x>>y,
"aten.__rshift__.Tensor": lambda x,y: x>>y,
"aten.__irshift__.Scalar": lambda x,y: x>>y,
# inplace ops using replace for fusion
"aten.zero_": lambda x: x.const_like(0),
-1
View File
@@ -161,7 +161,6 @@ norecursedirs = [
".git",
]
timeout = 300
timeout_method = "thread"
timeout_func_only = true
testpaths = ["test"]
filterwarnings = [
+1 -1
View File
@@ -36,7 +36,7 @@ def custom_add_var(A:UOp, B:UOp) -> UOp:
A,B = A.flatten(), B.flatten()
assert A.dtype == dtypes.uint32, f"buffer dtype must be uint32, got {A.dtype}"
threads = UOp.special(A.numel(), "lidx0")
var = UOp.param(2, dtypes.index, vmin_vmax=(0, 10), name="var", addrspace=AddrSpace.ALU)
var = UOp.param(2, dtypes.weakint, vmin_vmax=(0, 10), name="var", addrspace=AddrSpace.ALU)
insts = [
s_load_b128(s[4:7], s[0:1]),
s_load_b32(s[8], s[0:1], offset=0x10), # all threads load the same variable
+21 -10
View File
@@ -8,7 +8,7 @@ from tinygrad.renderer.ptx import PTXRenderer
from tinygrad.renderer.nir import NIRRenderer
from tinygrad import Context, Device, Tensor, dtypes
from hypothesis import given, settings, strategies as strat
from test.helpers import rand_for_dtype
from test.helpers import rand_for_dtype, min_normal
from test.unit.test_dtype_spec import _assert_eq, core_dtypes, dtype_ints, dtype_floats, FP8E4M3_MAX, FP8E5M2_MAX, FP8E4M3FNUZ_MAX, FP8E5M2FNUZ_MAX
import pytest
pytestmark = pytest.mark.filterwarnings("ignore")
@@ -25,10 +25,10 @@ def get_available_cast_dtypes(dtype: DType) -> List[DType]:
if dtype not in supported_dtypes and dtype not in dtypes.fp8s+(dtypes.half,dtypes.bfloat16): return []
return dts
def _to_torch_storage_type(dtype:DType):
if dtype == dtypes.bfloat16: return torch.float32
if dtype in dtypes.fp8s: return torch.float32
return _to_torch_dtype(dtype)
def _to_torch_storage(a:Tensor) -> torch.Tensor:
# tolist() of an fp8 Tensor gives floats, so convert and store in uint8
if a.dtype in dtypes.fp8s: return torch.tensor([float_to_fp8(x, a.dtype) for x in a.flatten().tolist()], dtype=torch.uint8).reshape(a.shape)
return torch.tensor(a.tolist(), dtype=_to_torch_dtype(a.dtype))
def _test_to_np(a:Tensor, np_dtype, target):
if DEBUG >= 2: print(a)
@@ -46,12 +46,15 @@ def _test_cast(a:Tensor, target_dtype:DType):
if a.is_floating_point() and dtypes.is_unsigned(target_dtype):
# converting negative float to unsigned integer is undefined
a = a.abs()
if a.is_floating_point() and dtypes.is_float(target_dtype) and (mn:=min_normal(target_dtype)) >= min_normal(a.dtype):
# subnormals are zero, so an input below the target's min normal casts to 0
a = (a.abs() < mn).where(0, a)
expected = list(a.numpy().astype(_to_np_dtype(target_dtype)))
if target_dtype in dtypes.fp8s: expected = [truncate[target_dtype](x) for x in expected]
_test_op(lambda: a.cast(target_dtype), target_dtype, expected)
def _test_bitcast(a:Tensor, target_dtype:DType, target=None):
expected = torch.tensor(a.tolist(), dtype=_to_torch_storage_type(a.dtype)).view(_to_torch_dtype(target_dtype)).tolist()
expected = _to_torch_storage(a).view(_to_torch_dtype(target_dtype)).tolist()
if target_dtype in dtypes.fp8s: expected = [fp8_to_float(x, target_dtype) for x in expected]
_test_op(lambda: a.bitcast(target_dtype), target_dtype, target or expected)
@@ -61,10 +64,12 @@ class TestDType(unittest.TestCase):
@classmethod
def setUpClass(cls):
if cls.DTYPE is None: raise unittest.SkipTest("base class")
cls.DATA = rand_for_dtype(cls.DTYPE, 0x10, allow_subnormal=cls.DTYPE in supported_dtypes)
cls.DATA = rand_for_dtype(cls.DTYPE, 0x10, allow_subnormal=cls.DTYPE in supported_dtypes and cls.DTYPE not in dtypes.fp8s)
def test_to_np(self):
_test_to_np(Tensor(self.DATA, dtype=self.DTYPE), _to_np_dtype(self.DTYPE), np.array(self.DATA, dtype=_to_np_dtype(self.DTYPE)))
a = Tensor(self.DATA, dtype=self.DTYPE)
self.assertEqual(a.dtype, self.DTYPE)
_test_to_np(a, _to_np_dtype(self.DTYPE), np.array(self.DATA, dtype=_to_np_dtype(self.DTYPE)))
def test_casts_to(self):
for dtype in get_available_cast_dtypes(self.DTYPE):
@@ -273,10 +278,11 @@ class TestBitCast(unittest.TestCase):
@given(strat.sampled_from(dtype_ints + dtype_floats), strat.sampled_from(dtype_ints + dtype_floats))
def test_shape_change_bitcast(self, dt1, dt2):
data = rand_for_dtype(dt1, 32).reshape(2, 2, 8)
expected = torch.tensor(data.tolist(), dtype=_to_torch_storage_type(dt1)).view(_to_torch_dtype(dt2))
a = Tensor(data, dtype=dt1)
expected = _to_torch_storage(a).view(_to_torch_dtype(dt2))
if dt2 in dtypes.fp8s:
expected = torch.tensor([fp8_to_float(x, dt2) for x in expected.view(-1).tolist()]).view_as(expected)
_test_op(lambda: Tensor(data, dtype=dt1).bitcast(dt2), dt2, expected.tolist())
_test_op(lambda: a.bitcast(dt2), dt2, expected.tolist())
def test_shape_change_bitcast_exceptions(self):
with self.assertRaises(RuntimeError):
@@ -293,6 +299,11 @@ class TestBitCast(unittest.TestCase):
b = a.bitcast(dtypes.float32)
assert b.numpy()[0,0] == 1.
def test_bitcast_bf16_from_cast(self):
# a bfloat16 from a cast holds bfloat16 bits. 1.0 is 0x3f80 in bfloat16, which is 1.875 in half
a = Tensor([1.0], dtype=dtypes.float32).cast(dtypes.bfloat16)
assert a.bitcast(dtypes.half).numpy()[0] == 1.875
class TestInt16DType(TestDType): DTYPE = dtypes.int16
class TestUint16DType(TestDType):
+27
View File
@@ -291,6 +291,33 @@ class TestDTypeALU(unittest.TestCase):
@Context(EMULATED_DTYPES="long")
def test_emulated_int64(self, a, b, op): universal_test(a, b, dtypes.int64, op)
def _test_shl(self):
for dtype, values, distances in ((dtypes.int64, [-0x1234, 0x80000001, -1, 0x1234, 1], [0, 5, 31, 32, 62]),
(dtypes.uint64, [0x80000001, 0x80000001, 1, 0xFEDC, 1], [0, 5, 31, 32, 62]),
(dtypes.int8, [-3, 1, 7, -2, 1], [0, 1, 3, 5, 6]),
(dtypes.uint16, [3, 1, 0xFF, 7, 1], [0, 1, 7, 12, 15])):
with self.subTest(dtype=dtype):
result = Tensor(values, dtype=dtype) << Tensor(distances, dtype=dtype)
np.testing.assert_equal(result.numpy(), [x << d for x, d in zip(values, distances)])
def _test_shr(self):
for dtype, values, distances in ((dtypes.int64, [-(2**40), -1, -(2**50), -(2**40), 0x123456789ABCDEF], [0, 5, 31, 32, 63]),
(dtypes.uint64, [0xFEDCBA9876543210] * 5, [0, 5, 31, 32, 63]),
(dtypes.int8, [-128, -1, 64, -37, 1], [0, 1, 3, 5, 7]),
(dtypes.uint16, [0xFFFF] * 5, [0, 1, 8, 13, 15])):
with self.subTest(dtype=dtype):
result = Tensor(values, dtype=dtype) >> Tensor(distances, dtype=dtype)
np.testing.assert_equal(result.numpy(), [x >> d for x, d in zip(values, distances)])
def test_shl(self): self._test_shl()
def test_shr(self): self._test_shr()
@Context(EMULATED_DTYPES="long")
def test_emulated_shl(self): self._test_shl()
@Context(EMULATED_DTYPES="long")
def test_emulated_shr(self): self._test_shr()
@given(ht.uint8, strat.sampled_from(integer_unary_operations))
def test_uint8_unary(self, a, op): universal_test_unary(a, dtypes.uint8, op)
+6 -9
View File
@@ -1,9 +1,9 @@
import numpy as np
import functools, unittest, ctypes
import functools, unittest
from tinygrad.device import Device, Buffer
from tinygrad.tensor import Tensor
from tinygrad.helpers import Context, from_mv
from tinygrad.helpers import Context
from tinygrad.dtype import dtypes
from tinygrad.engine.jit import MultiGraphRunner
from tinygrad.engine.realize import run_linear, compile_linear
@@ -31,7 +31,7 @@ def make_buffer(device, size=BUF_SIZE, fill=False):
buf = Buffer(device, size, dtypes.int).ensure_allocated()
if fill:
with Context(DEBUG=0):
buf.copyin(Tensor(np.random.randint(-10000, 10000, size=size, dtype=np.int32)).realize().uop.base.realized.as_memoryview())
buf.copy_from(Tensor(np.random.randint(-10000, 10000, size=size, dtype=np.int32)).realize().uop.base.realized)
return buf
def make_view(base, offset_elems, size_elems):
@@ -55,15 +55,12 @@ def run_schedule(calls:list[UOp]):
run_linear(UOp(Ops.LINEAR, src=tuple(calls)))
def zero_bufs(bufs):
for b in bufs:
mv = memoryview(bytearray(b.nbytes))
ctypes.memset(from_mv(mv), 0, len(mv))
b.copyin(mv)
for b in bufs: b.copy_from(Buffer("PYTHON", b.size, b.dtype, opaque=memoryview(bytearray(b.nbytes))))
@unittest.skipUnless(Device[Device.DEFAULT].graph is not None, "graph support required")
class TestGraph(unittest.TestCase):
def skip_if_no_offset(self):
if not hasattr(Device[Device.DEFAULT].allocator, "_offset"): self.skipTest("device does not support _offset")
if Device.DEFAULT in {"WEBGPU", "CL"}: self.skipTest("device does not support _offset")
def skip_if_not_multigraph(self):
graph = g.func if isinstance(g:=(d:=Device[Device.DEFAULT]).graph, functools.partial) else g
@@ -213,8 +210,8 @@ class TestGraph(unittest.TestCase):
def test_graph_offset_bufs(self):
self.skip_if_not_multigraph()
self.skip_if_no_offset()
d0 = Device.DEFAULT
if not hasattr(Device[d0].allocator, "_offset"): self.skipTest("device does not support _offset")
b0 = make_buffer(d0, fill=True)
b1 = make_view(b0, 0, b0.size)
+1 -1
View File
@@ -424,7 +424,7 @@ def copyout_outputs(outbufs:list[Buffer]) -> list[np.ndarray]:
return [np.frombuffer(x.as_memoryview(), _to_np_dtype(x.dtype)) for x in outbufs]
def reset_bufs(bufs:list[Buffer]):
for buf in bufs: buf.copyin(np.zeros((buf.size*buf.dtype.itemsize,), dtype=np.uint8).data)
for buf in bufs: buf.copy_from(Buffer("PYTHON", buf.size, buf.dtype, opaque=memoryview(bytearray(buf.nbytes))))
def _helper_linearizer_opt_ast(realized_ast:UOp, real_bufs:list[Buffer], opts=[],
apply_tc=False, atol=1e-4, rtol=1e-4, color_sizes=[], wanna_output=[]):
+9 -9
View File
@@ -12,18 +12,18 @@ class TestLinearizerFailure(unittest.TestCase):
@unittest.skipUnless(Device.DEFAULT == "METAL", "only tested on METAL")
def test_failure_beam_mnist(self):
c0 = UOp.param(0, dtypes.uchar, (4014080,))
c1 = UOp.range(UOp.const(dtypes.index, 512), 0, AxisType.GLOBAL)
c2 = UOp.range(UOp.const(dtypes.index, 784), 1, AxisType.GLOBAL)
c3 = UOp.range(UOp.const(dtypes.index, 10), 3, AxisType.GLOBAL)
c1 = UOp.range(UOp.const(dtypes.weakint, 512), 0, AxisType.GLOBAL)
c2 = UOp.range(UOp.const(dtypes.weakint, 784), 1, AxisType.GLOBAL)
c3 = UOp.range(UOp.const(dtypes.weakint, 10), 3, AxisType.GLOBAL)
c4 = UOp.param(1, dtypes.int, (512,))
c5 = c4.index(c1.valid(UOp.const(dtypes.bool, True)))
c6 = UOp.range(UOp.const(dtypes.index, 6000), 1004, AxisType.REDUCE)
c7 = UOp.range(UOp.const(dtypes.index, 3750), 2006, AxisType.REDUCE)
c8 = UOp.range(UOp.const(dtypes.index, 16), 2007, AxisType.GROUP_REDUCE)
c6 = UOp.range(UOp.const(dtypes.weakint, 6000), 1004, AxisType.REDUCE)
c7 = UOp.range(UOp.const(dtypes.weakint, 3750), 2006, AxisType.REDUCE)
c8 = UOp.range(UOp.const(dtypes.weakint, 16), 2007, AxisType.GROUP_REDUCE)
c9 = UOp.param(2, dtypes.uchar, (47040000,))
c10 = c9.index((((c3*UOp.const(dtypes.index, 4704000))+c2)+(c6*UOp.const(dtypes.index, 784))).valid(UOp.const(dtypes.bool, True)))
c11 = c5.alu(Ops.CMPNE, ((((c3*UOp.const(dtypes.index, 6000))+c6)+((c7*UOp.const(dtypes.index, 16))+c8)).alu(Ops.CMPLT, UOp.const(dtypes.index, 59999)).where(UOp.const(dtypes.int, 0), UOp.const(dtypes.int, 1)).reduce(c7, c8, arg=Ops.ADD)+UOp.const(dtypes.int, -1))).where(UOp.const(dtypes.uchar, 0), c10).reduce(c6, arg=Ops.ADD)
c12 = c0.index((((c1*UOp.const(dtypes.index, 7840))+(c2*UOp.const(dtypes.index, 10)))+c3).valid(UOp.const(dtypes.bool, True))).store(c11).end(c1, c2, c3)
c10 = c9.index((((c3*UOp.const(dtypes.weakint, 4704000))+c2)+(c6*UOp.const(dtypes.weakint, 784))).valid(UOp.const(dtypes.bool, True)))
c11 = c5.alu(Ops.CMPNE, ((((c3*UOp.const(dtypes.weakint, 6000))+c6)+((c7*UOp.const(dtypes.weakint, 16))+c8)).alu(Ops.CMPLT, UOp.const(dtypes.weakint, 59999)).where(UOp.const(dtypes.int, 0), UOp.const(dtypes.int, 1)).reduce(c7, c8, arg=Ops.ADD)+UOp.const(dtypes.int, -1))).where(UOp.const(dtypes.uchar, 0), c10).reduce(c6, arg=Ops.ADD)
c12 = c0.index((((c1*UOp.const(dtypes.weakint, 7840))+(c2*UOp.const(dtypes.weakint, 10)))+c3).valid(UOp.const(dtypes.bool, True))).store(c11).end(c1, c2, c3)
ast = c12.sink(arg=KernelInfo(name='test', axis_types=(), dont_use_locals=False, applied_opts=(Opt(op=OptOps.GROUP, axis=1, arg=16),), opts_to_apply=None))
_ = to_program(ast, Device["METAL"].renderer)
+77 -8
View File
@@ -1,11 +1,14 @@
import unittest
import unittest, functools
from tinygrad import Tensor, Device, dtypes, Context, GlobalCounters
from tinygrad.helpers import getenv
from examples.mlperf.models.flat_llama import FP8_DTYPE, quantize_fp8
from extra.llama_kernels.fused_ce import fused_ce_loss
from extra.llama_kernels import local_abs_max
from extra.llama_kernels.quantize_fp8_delayed import quantize_fp8_delayed, quantize_fp8_scalar
from extra.models.llama import apply_rotary_emb, precompute_freqs_cis
from extra.thunder.amd.fa import custom_fused_qkv_rope_backward, fused_qkv_rope
from test.helpers import needs_second_gpu
from test.backend.test_asm_gemm import has_hipcc
def run_fused_ce(bs:int, seqlen:int, vocab:int, label_smoothing:float=0.0) -> None:
Tensor.manual_seed(0)
@@ -46,9 +49,10 @@ def run_quantize_fp8(shape:tuple[int, ...], delayed:bool=True) -> None:
with Context(DEBUG=0): Tensor.realize(x, amax_state)
if delayed:
fp8, inv_scale, new_amax, _ = quantize_fp8_delayed(x, amax_state, FP8_DTYPE)
amax_out = Tensor.zeros((), dtype=dtypes.float32, device=x.device).realize()
fp8, inv_scale = quantize_fp8_delayed(x, amax_state, amax_out, FP8_DTYPE)
ref_fp8, ref_inv_scale, ref_new_amax = quantize_fp8(x, amax_state=amax_state)
Tensor.realize(fp8, inv_scale, new_amax)
Tensor.realize(fp8, inv_scale)
Tensor.realize(ref_fp8, ref_inv_scale, ref_new_amax)
else:
fp8 = quantize_fp8_scalar(x, amax_state, FP8_DTYPE)
@@ -60,9 +64,10 @@ def run_quantize_fp8(shape:tuple[int, ...], delayed:bool=True) -> None:
assert fp8.cast(dtypes.float).allclose(ref_fp8.cast(dtypes.float), atol=0, rtol=0).item(), "fp8 mismatch"
if delayed:
assert inv_scale.allclose(ref_inv_scale, atol=0, rtol=0).item(), "inv_scale mismatch"
assert new_amax.allclose(ref_new_amax, atol=0, rtol=0).item(), \
f"amax mismatch: got={new_amax.item()} ref={ref_new_amax.item()} diff={abs(new_amax.item()-ref_new_amax.item())}"
assert amax_out.allclose(ref_new_amax, atol=0, rtol=0).item(), \
f"amax mismatch: got={amax_out.item()} ref={ref_new_amax.item()} diff={abs(amax_out.item()-ref_new_amax.item())}"
@unittest.skipUnless(Device.DEFAULT == "AMD", "requires atomic max")
class TestQuantizeFP8(unittest.TestCase):
def setUp(self):
ren = Device[Device.DEFAULT].renderer
@@ -78,10 +83,11 @@ class TestQuantizeFP8(unittest.TestCase):
x = Tensor.empty(2048*8, 1024, dtype=dtypes.bfloat16, device=devs).uop.multi(0)
x = Tensor(x, device=devs)
amax_state = Tensor.full((), 2.0, dtype=dtypes.float32, device=devs).contiguous()
fp8, _, new_amax, _ = quantize_fp8_delayed(x, amax_state, FP8_DTYPE)
Tensor.realize(fp8, new_amax)
amax_out = Tensor.zeros((), dtype=dtypes.float32, device=devs).realize()
fp8, _ = quantize_fp8_delayed(x, amax_state, amax_out, FP8_DTYPE)
Tensor.realize(fp8)
assert fp8.uop.shape == x.uop.shape
assert new_amax.shape == ()
assert amax_out.shape == ()
class TestLocalAmax(unittest.TestCase):
def test_multi_tensor_local_shard_amax(self):
@@ -92,5 +98,68 @@ class TestLocalAmax(unittest.TestCase):
self.assertEqual(GlobalCounters.kernel_count, 2)
self.assertEqual(out.tolist(), [[0., 7., 14., 21.], [28., 35., 42., 49.], [120., 135., 150., 165.], [180., 195., 210., 225.]])
@unittest.skipUnless(has_hipcc() and Device.DEFAULT == "AMD", "requires hipcc to compile and amd device to run")
class TestFusedQKVRoPE(unittest.TestCase):
SHAPE = (2, 8192, 32, 8, 128)
def rand_bf16(self, *shape:int) -> Tensor:
return (Tensor.randn(*shape) * 0.1).cast(dtypes.bfloat16).contiguous().realize()
def freqs_cis(self) -> Tensor:
_, N, _, _, D = self.SHAPE
return precompute_freqs_cis(D, N * 2).cast(dtypes.bfloat16).clone().realize()
def test_llama31_8b_forward(self):
Tensor.manual_seed(0)
B, N, H, H_KV, D = self.SHAPE
GROUP = H // H_KV
freqs_cis = self.freqs_cis()
x = self.rand_bf16(B, N, H_KV * (GROUP + 2) * D)
q, k, v = fused_qkv_rope(x, freqs_cis, H, H_KV, D)
Tensor.realize(q, k, v)
packed_ref = x.reshape(B, N, H_KV, GROUP + 2, D)
q_ref = packed_ref[:, :, :, :GROUP].reshape(B, N, H, D)
k_ref, v_ref = packed_ref[:, :, :, GROUP], packed_ref[:, :, :, GROUP+1]
q_ref, k_ref = apply_rotary_emb(q_ref, k_ref, freqs_cis[:, :N])
q_ref, k_ref, v_ref = q_ref.cast(dtypes.bfloat16), k_ref.cast(dtypes.bfloat16), v_ref.cast(dtypes.bfloat16)
Tensor.realize(q_ref, k_ref, v_ref)
with Context(DEBUG=0):
self.assertTrue(q.allclose(q_ref, atol=2e-2, rtol=0).item(), "Q forward mismatch")
self.assertTrue(k.allclose(k_ref, atol=2e-2, rtol=0).item(), "K forward mismatch")
self.assertTrue(v.allclose(v_ref, atol=0, rtol=0).item(), "V forward mismatch")
def test_llama31_8b_backward(self):
Tensor.manual_seed(1)
B, N, H, H_KV, D = self.SHAPE
PARTIALS = 2
GROUP = H // H_KV
freqs_cis = self.freqs_cis()
dq = self.rand_bf16(B, N, H, D)
dk_partial = self.rand_bf16(B * PARTIALS, N, H_KV, D)
dv_partial = self.rand_bf16(B * PARTIALS, N, H_KV, D)
# Invert Flash Attention's dQ layout transform to reproduce its native buffer.
dq_native = dq.transpose(1, 2).reshape(B, H, N//16, 4, 4, 4, 2, D//32, 2, 2) \
.permute(0, 1, 2, 5, 6, 8, 7, 3, 4, 9).reshape(B, H, N, D).contiguous().realize()
dx = Tensor.empty(B, N, H_KV * (GROUP + 2) * D, dtype=dtypes.bfloat16)
arch = Device[Device.DEFAULT].renderer.target.arch
fxn = functools.partial(custom_fused_qkv_rope_backward, device=Device.DEFAULT, arch=arch,
B=B, N=N, H=H, H_KV=H_KV, D=D)
dx = Tensor.custom_kernel(dx, dq_native, dk_partial, dv_partial, freqs_cis, fxn=fxn)[0].realize()
def inverse_rope(x:Tensor) -> Tensor:
x = x.reshape(*x.shape[:-1], D//2, 2).float()
cs = freqs_cis[:, :N].float()
return Tensor.stack(x[..., 0] * cs[..., 0] + x[..., 1] * cs[..., 1],
-x[..., 0] * cs[..., 1] + x[..., 1] * cs[..., 0], dim=-1).flatten(-2).cast(dtypes.bfloat16)
dq_ref = inverse_rope(dq).reshape(B, N, H_KV, GROUP, D)
dk_ref = inverse_rope(dk_partial.float().reshape(B, PARTIALS, N, H_KV, D).sum(1).cast(dtypes.bfloat16)).unsqueeze(3)
dv_ref = dv_partial.float().reshape(B, PARTIALS, N, H_KV, D).sum(1).cast(dtypes.bfloat16).unsqueeze(3)
ref = Tensor.cat(dq_ref, dk_ref, dv_ref, dim=3).reshape(*dx.shape).realize()
with Context(DEBUG=0): self.assertTrue(dx.allclose(ref, atol=2e-2, rtol=2e-2).item(), "backward mismatch")
if __name__ == '__main__':
unittest.main()
+2 -2
View File
@@ -385,7 +385,7 @@ class TestMultiBufferView(unittest.TestCase):
b_ref = view_fn(a_ref)
b_multi = view_fn(a_multi).contiguous()
linear, var_vals = b_multi.linear_with_vars()
if all(hasattr(Device[d].allocator, "_offset") for d in b_multi.device):
if all(not d.startswith(("WEBGPU", "CL")) for d in b_multi.device):
compiled = [call for call in linear.src if call.src[0].op is Ops.SINK]
self.assertEqual(len(compiled), 0, f"expected zero compiled kernels, got {len(compiled)}")
run_linear(linear, var_vals)
@@ -417,7 +417,7 @@ class TestMultiBufferView(unittest.TestCase):
a = Tensor.arange(8*12).reshape(8, 12).clone().shard(devices_4, axis=1).realize()
out = a[5].contiguous()
linear, var_vals = out.linear_with_vars()
if all(hasattr(Device[d].allocator, "_offset") for d in out.device):
if all(not d.startswith(("WEBGPU", "CL")) for d in out.device):
compiled = [call for call in linear.src if call.src[0].op is Ops.SINK]
self.assertEqual(len(compiled), 0)
run_linear(linear, var_vals)
+6
View File
@@ -848,6 +848,8 @@ class TestOps(unittest.TestCase):
helper_test_op([], lambda: tor << 0, lambda: (ten << 0).cast(dtypes.int32), forward_only=True)
helper_test_op([], lambda: tor << 2, lambda: (ten << 2).cast(dtypes.int32), forward_only=True)
helper_test_op([], lambda: tor << 31, lambda: (ten << 31).cast(dtypes.int32), forward_only=True)
helper_test_op([], lambda: tor << torch.tensor([0,2,4]).int(),
lambda: (ten << Tensor([0,2,4], dtype=dtypes.uint32)).cast(dtypes.int32), forward_only=True)
helper_test_op([], lambda: tor.__lshift__(2), lambda: ten.__lshift__(2).cast(dtypes.int32), forward_only=True)
helper_test_op([], lambda: tor.bitwise_left_shift(2), lambda: ten.lshift(2).cast(dtypes.int32), forward_only=True)
@@ -859,6 +861,8 @@ class TestOps(unittest.TestCase):
helper_test_op([], lambda: tor >> 0, lambda: (ten >> 0).cast(dtypes.int32), forward_only=True)
helper_test_op([], lambda: tor >> 2, lambda: (ten >> 2).cast(dtypes.int32), forward_only=True)
helper_test_op([], lambda: tor >> 31, lambda: (ten >> 31).cast(dtypes.int32), forward_only=True)
helper_test_op([], lambda: tor >> torch.tensor([0,2,4]).int(),
lambda: (ten >> Tensor([0,2,4], dtype=dtypes.uint32)).cast(dtypes.int32), forward_only=True)
helper_test_op([], lambda: tor.__rshift__(2), lambda: ten.__rshift__(2).cast(dtypes.int32), forward_only=True)
helper_test_op([], lambda: tor.bitwise_right_shift(2), lambda: ten.rshift(2).cast(dtypes.int32), forward_only=True)
@@ -870,6 +874,7 @@ class TestOps(unittest.TestCase):
helper_test_op([], lambda: tor << 2, lambda: ten << 2, forward_only=True)
helper_test_op([], lambda: tor << 8, lambda: ten << 8, forward_only=True)
helper_test_op([], lambda: tor << 31, lambda: ten << 31, forward_only=True)
helper_test_op([], lambda: tor << torch.tensor([0,2,8,31]).int(), lambda: ten << Tensor([0,2,8,31], dtype=dtypes.int), forward_only=True)
def test_rshift_signed(self):
data = [[-1, -3, 1, 7], [0, -2147483648, 2147483647, -1]]
@@ -879,6 +884,7 @@ class TestOps(unittest.TestCase):
helper_test_op([], lambda: tor >> 2, lambda: ten >> 2, forward_only=True)
helper_test_op([], lambda: tor >> 8, lambda: ten >> 8, forward_only=True)
helper_test_op([], lambda: tor >> 31, lambda: ten >> 31, forward_only=True)
helper_test_op([], lambda: tor >> torch.tensor([0,2,8,31]).int(), lambda: ten >> Tensor([0,2,8,31], dtype=dtypes.int), forward_only=True)
def test_idiv_shift_rewrite_negative(self):
a = Tensor(-5).div(2, rounding_mode="trunc").item()
+8 -8
View File
@@ -70,9 +70,9 @@ class TestProfiler(unittest.TestCase):
buf1 = Buffer(Device.DEFAULT, 2, dtypes.float, options=BufferSpec(nolru=True)).ensure_allocated()
with helper_collect_profile(TestProfiler.d0) as profile:
buf1.copyin(memoryview(bytearray(struct.pack("ff", 0, 1))))
buf1.copy_from(Buffer("PYTHON", 2, dtypes.float, opaque=memoryview(bytearray(struct.pack("ff", 0, 1)))))
kernel_runs = [x for x in profile if isinstance(x, ProfileRangeEvent) and x.device.startswith(TestProfiler.d0.device)]
kernel_runs = [x for x in profile if isinstance(x, ProfileRangeEvent) and x.device.startswith((TestProfiler.d0.device, "PYTHON"))]
assert len(kernel_runs) == 1, "one kernel run is expected"
def test_profile_multiops(self):
@@ -80,12 +80,12 @@ class TestProfiler(unittest.TestCase):
buf1 = Buffer(Device.DEFAULT, 2, dtypes.float, options=BufferSpec(nolru=True)).ensure_allocated()
with helper_collect_profile(TestProfiler.d0) as profile:
buf1.copyin(memoryview(bytearray(struct.pack("ff", 0, 1))))
buf1.copy_from(Buffer("PYTHON", 2, dtypes.float, opaque=memoryview(bytearray(struct.pack("ff", 0, 1)))))
gs, ls = TestProfiler.prg.arg.launch_dims({})
TestProfiler.runtime(buf1._buf, TestProfiler.a.uop.buffer._buf, global_size=gs, local_size=ls)
buf1.copyout(memoryview(bytearray(buf1.nbytes)))
buf1.as_memoryview()
evs = [x for x in profile if isinstance(x, ProfileRangeEvent) and x.device.startswith(TestProfiler.d0.device)]
evs = [x for x in profile if isinstance(x, ProfileRangeEvent) and x.device.startswith((TestProfiler.d0.device, "PYTHON"))]
assert len(evs) == 3, "3 kernel runs are expected"
# NOTE: order of events does not matter, the tool is responsible for sorting them
@@ -103,12 +103,12 @@ class TestProfiler(unittest.TestCase):
buf2 = Buffer(f"{Device.DEFAULT}:1", 2, dtypes.float, options=BufferSpec(nolru=True)).ensure_allocated()
with helper_collect_profile(TestProfiler.d0, d1) as profile:
buf1.copyin(memoryview(bytearray(struct.pack("ff", 0, 1))))
buf2.copyin(memoryview(bytearray(struct.pack("ff", 0, 1))))
buf1.copy_from(Buffer("PYTHON", 2, dtypes.float, opaque=memoryview(bytearray(struct.pack("ff", 0, 1)))))
buf2.copy_from(Buffer("PYTHON", 2, dtypes.float, opaque=memoryview(bytearray(struct.pack("ff", 0, 1)))))
for dev in [TestProfiler.d0.device, d1.device]:
evs = [x for x in profile if isinstance(x, ProfileRangeEvent) and _dev_base(x.device) == dev]
assert len(evs) == 1, "one kernel runs are expected"
assert len(evs) == (0 if hasattr(TestProfiler.d0.allocator, '_as_buffer') else 1), "one kernel runs are expected"
def test_profile_multidev_transfer(self):
try: d1 = Device[f"{Device.DEFAULT}:1"]
+3 -3
View File
@@ -1,6 +1,6 @@
import unittest
import numpy as np
from tinygrad.device import Device
from tinygrad.device import Device, Buffer
from tinygrad.dtype import dtypes, ConstType
from tinygrad.engine.realize import run_linear
from tinygrad.codegen import to_program
@@ -10,13 +10,13 @@ from tinygrad.renderer.ptx import PTXRenderer
from tinygrad.renderer.wgsl import WGSLRenderer
from tinygrad.runtime.ops_python import PythonRenderer
from tinygrad.uop.ops import UOp, Ops, KernelInfo, python_alu
from tinygrad.tensor import Tensor, _to_np_dtype
from tinygrad.tensor import Tensor
def _test_uop_result(inputs:list[Tensor], sink:UOp, local_size=None):
for x in inputs: x.realize()
sz = 1 if local_size is None else prod(local_size)
outs = [UOp.new_buffer(Device.DEFAULT, sz, u.src[1].dtype) for u in sink.src if u.op is Ops.STORE]
for u in outs: u.buffer.allocate().copyin(np.zeros(sz, dtype=_to_np_dtype(u.dtype)).data)
for u in outs: u.buffer.allocate().copy_from(Buffer("PYTHON", sz, u.dtype, opaque=memoryview(bytearray(u.buffer.nbytes))))
run_linear(UOp(Ops.LINEAR, src=(sink.call(*outs, *(x.uop.base for x in inputs)),)))
return [u.buffer.numpy() for u in outs]
+15 -1
View File
@@ -2,7 +2,7 @@
# schedule confirms the right things are capable of fusing
# NOTE: this has overlap with external_test_opt.py
import unittest
import unittest, time
import numpy as np
from tinygrad import nn, dtypes, Device, Tensor, Variable
@@ -197,6 +197,20 @@ class TestLimitBufs(unittest.TestCase):
base = (idx >= i).where(a + b, base)
assert all(x > 0 for x in base.tolist())
def test_limit_bufs_linear_scaling(self):
def sched_time(n):
with Context(TRACK_MATCH_STATS=0, DEBUG=0):
bufs = [Tensor.ones(16).contiguous().realize() for _ in range(4)]
root = bufs[0]
for i in range(n): root = root + bufs[i % 4]
with Context(MAX_KERNEL_BUFFERS=8, SCACHE=0):
st = time.perf_counter()
root.schedule_linear()
return time.perf_counter() - st
sched_time(400)
t1, t2 = min(sched_time(400) for _ in range(3)), min(sched_time(1600) for _ in range(3))
self.assertLess(t2/t1, 8, f"{t1*1e3:.1f}ms -> {t2*1e3:.1f}ms")
class TestSwizzle(unittest.TestCase):
def test_swizzle_simple(self):
Tensor.manual_seed(0)
+12 -13
View File
@@ -4,11 +4,10 @@ from tinygrad.device import Buffer
from tinygrad.helpers import Context, DEV
from test.helpers import needs_second_gpu
@unittest.skipUnless(hasattr(Device[Device.DEFAULT].allocator, "_offset"), "subbuffer not supported")
@unittest.skipIf(Device.DEFAULT in {"WEBGPU", "CL"}, "subbuffer not supported")
class TestSubBuffer(unittest.TestCase):
def setUp(self):
self.buf = Buffer(Device.DEFAULT, 10, dtypes.uint8).ensure_allocated()
self.buf.copyin(memoryview(bytearray(range(10))))
self.buf = Buffer(Device.DEFAULT, 10, dtypes.uint8, initial_value=bytes(range(10)))
self.buf_unalloc = Buffer(Device.DEFAULT, 10, dtypes.uint8)
def test_subbuffer(self):
@@ -59,7 +58,7 @@ class TestSubBuffer(unittest.TestCase):
_ = Buffer(Device.DEFAULT, 10, dtypes.uint8).ensure_allocated()
self.buf.ensure_allocated()
self.buf.copyin(memoryview(bytearray(range(10, 20))))
self.buf.copy_from(Buffer("PYTHON", 10, dtypes.uint8, opaque=memoryview(bytearray(range(10, 20)))))
vbuf.ensure_allocated()
@@ -109,15 +108,15 @@ class TestSubBuffer(unittest.TestCase):
def test_subbuffer_copy_in_out(self):
sub_buf = self.buf.view(3, dtypes.uint8, offset=3).ensure_allocated() # [3:6]
data_out_sub = bytearray([0]*3)
sub_buf.copyout(memoryview(data_out_sub))
data_out_sub[:] = sub_buf.as_memoryview()
assert data_out_sub == bytearray(range(3, 6))
sub_buf.copyin(memoryview(bytearray(range(3))))
sub_buf.copy_from(Buffer("PYTHON", 3, dtypes.uint8, opaque=memoryview(bytearray(range(3)))))
assert sub_buf.as_memoryview().tolist() == list(range(3))
assert self.buf.as_memoryview().tolist()[3:6] == list(range(3))
sub_buf.copyout(memoryview(data_out_sub))
data_out_sub[:] = sub_buf.as_memoryview()
assert data_out_sub == bytearray(range(3))
data_out_base = bytearray([0]*10)
self.buf.copyout(memoryview(data_out_base))
data_out_base[:] = self.buf.as_memoryview()
assert data_out_base[0:3] == bytearray(range(0, 3))
assert data_out_base[3:6] == data_out_sub
assert data_out_base[6:10] == bytearray(range(6, 10))
@@ -129,27 +128,27 @@ class TestSubBuffer(unittest.TestCase):
self.assertTrue(view2.is_allocated())
data_in = bytearray([7, 8, 9])
view2.copyin(memoryview(data_in))
view2.copy_from(Buffer("PYTHON", 3, view2.dtype, opaque=memoryview(data_in)))
data_out_v2 = bytearray([0]*3)
view2.copyout(memoryview(data_out_v2))
data_out_v2[:] = view2.as_memoryview()
assert data_in == data_out_v2
expected_base_data = memoryview(bytearray(range(10)))
expected_base_data[4:7] = data_in
data_out_base = bytearray([0]*10)
self.buf.copyout(memoryview(data_out_base))
data_out_base[:] = self.buf.as_memoryview()
assert expected_base_data == data_out_base
def test_subbuffer_alloc(self):
sub_buf = self.buf.view(4, dtypes.int8, offset=3)
sub_buf.allocate()
sub_buf.copyin(memoryview(bytearray(range(10, 14))))
sub_buf.copy_from(Buffer("PYTHON", 4, dtypes.int8, opaque=memoryview(bytearray(range(10, 14)))))
assert self.buf.as_memoryview().tolist()[3:7] == sub_buf.as_memoryview().tolist()
sub_buf = self.buf_unalloc.view(4, dtypes.int8, offset=3)
sub_buf.allocate()
sub_buf.copyin(memoryview(bytearray(range(10, 14))))
sub_buf.copy_from(Buffer("PYTHON", 4, dtypes.int8, opaque=memoryview(bytearray(range(10, 14)))))
assert self.buf_unalloc.as_memoryview().tolist()[3:7] == sub_buf.as_memoryview().tolist()
def test_subbuffer_dealloc(self):
+4 -10
View File
@@ -33,11 +33,9 @@ def _test_single_value(vals, op, dts):
alu = uop(uops, op, output_dtype, loads)
out = uop(uops, Ops.STORE, dtypes.void, (buf_store.index(uop(uops, Ops.CONST, dtypes.int32, (), 0)), alu))
buf = Buffer(Device.DEFAULT, 1, output_dtype).allocate()
buf2 = [Buffer(Device.DEFAULT, 1, dtype).allocate().copyin(np.array([a], dtype=_to_np_dtype(dtype)).data) for a,dtype in zip(vals, dts)]
buf2 = [Buffer(Device.DEFAULT, 1, dtype, initial_value=np.array([a], dtype=_to_np_dtype(dtype)).tobytes()) for a,dtype in zip(vals, dts)]
run_uops([out], [buf]+buf2)
ret = np.empty(1, _to_np_dtype(output_dtype))
buf.copyout(ret.data)
return ret[0]
return np.frombuffer(buf.as_memoryview(), _to_np_dtype(output_dtype))[0]
def _test_single_value_const(vals, op, dts):
uops = []
@@ -48,9 +46,7 @@ def _test_single_value_const(vals, op, dts):
out = buf_store[UOp.const(dtypes.int32, 0)].store(alu)
buf = Buffer(Device.DEFAULT, 1, output_dtype).allocate()
run_uops([out], [buf])
ret = np.empty(1, _to_np_dtype(output_dtype))
buf.copyout(ret.data)
return ret[0]
return np.frombuffer(buf.as_memoryview(), _to_np_dtype(output_dtype))[0]
def _test_uops_result(output_dtype, uops, res):
# uops = []
@@ -59,9 +55,7 @@ def _test_uops_result(output_dtype, uops, res):
out = uop(uops, Ops.STORE, dtypes.void, (buf_store.index(uop(uops, Ops.CONST, dtypes.int32, (), 0)), res))
buf = Buffer(Device.DEFAULT, 1, output_dtype).allocate()
run_uops([out], [buf])
ret = np.empty(1, _to_np_dtype(output_dtype))
buf.copyout(ret.data)
return ret[0]
return np.frombuffer(buf.as_memoryview(), _to_np_dtype(output_dtype))[0]
class TestUOps(unittest.TestCase):
def _equal(self, v1, v2):
+5 -5
View File
@@ -31,8 +31,8 @@ class TestHCQ(unittest.TestCase):
def setUp(self):
TestHCQ.d0.synchronize()
TestHCQ.a.uop.buffer.copyin(memoryview(bytearray(struct.pack("ff", 0, 1))))
TestHCQ.b.uop.buffer.copyin(memoryview(bytearray(struct.pack("ff", 0, 0))))
TestHCQ.a.uop.buffer.copy_from(Buffer("PYTHON", 2, dtypes.float, opaque=memoryview(bytearray(struct.pack("ff", 0, 1)))))
TestHCQ.b.uop.buffer.copy_from(Buffer("PYTHON", 2, dtypes.float, opaque=memoryview(bytearray(struct.pack("ff", 0, 0)))))
TestHCQ.d0.synchronize() # wait for copyins to complete
# Test signals
@@ -376,7 +376,7 @@ class TestHCQ(unittest.TestCase):
SZ = 200_000_000
b = Buffer(f"{Device.DEFAULT}:1", SZ, dtypes.uint8, options=BufferSpec(nolru=True)).allocate()
a = Buffer(Device.DEFAULT, SZ, dtypes.uint8, options=BufferSpec(nolru=True)).allocate()
TestHCQ.d0.allocator.map(b._buf)
TestHCQ.d0.allocator._map(b._buf)
sig_st, sig_en = TestHCQ.d0.new_signal(), TestHCQ.d0.new_signal()
TestHCQ.d0.hw_copy_queue_t().timestamp(sig_st) \
@@ -454,7 +454,7 @@ class TestHCQ(unittest.TestCase):
buf1 = Buffer(Device.DEFAULT, 1, dtypes.int8, options=BufferSpec(nolru=True)).ensure_allocated()
buf2 = Buffer(f"{Device.DEFAULT}:1", 1, dtypes.int8, options=BufferSpec(nolru=True)).ensure_allocated()
buf3 = Buffer(Device.DEFAULT, 1, dtypes.int8, options=BufferSpec(host=True, nolru=True)).ensure_allocated()
TestHCQ.d0.allocator.map(buf2._buf)
TestHCQ.d0.allocator._map(buf2._buf)
for i in range(256):
ctypes.memset(buf3._buf.va_addr, i, 1)
@@ -569,7 +569,7 @@ class TestHCQ(unittest.TestCase):
local_buf = Buffer(f"{Device.DEFAULT}:{devid}", sz, dtypes.uint8, options=BufferSpec(cpu_access=True)).ensure_allocated()
d.allocator.map(cpu_buffer._buf)
d.allocator._map(cpu_buffer._buf)
d.hw_copy_queue_t().wait(d.timeline_signal, d.timeline_value - 1) \
.copy(local_buf._buf, cpu_buffer._buf, sz) \
+2 -2
View File
@@ -34,7 +34,7 @@ class TestCLError(unittest.TestCase):
data = list(range(65))
unaligned = memoryview(bytearray(data))[1:]
buffer = Buffer("CL", 64, dtypes.uint8).allocate()
buffer.copyin(unaligned)
buffer.copy_from(Buffer("PYTHON", 64, dtypes.uint8, opaque=unaligned))
result = memoryview(bytearray(len(data) - 1))
buffer.copyout(result)
result[:] = buffer.as_memoryview()
assert unaligned == result, "Unaligned data copied in must be equal to data copied out."
+2 -2
View File
@@ -24,7 +24,7 @@ def vision_conv_143():
c32 = ((c27<3)!=True)&(c27<67)
c34 = UOp.param(1, dtypes.half, shape=(32, 1024, 4))
c38 = c5//2
c45 = (c32&c24).where((c27*64+c38+c17*4096+-12480), UOp.const(dtypes.index, Invalid))
c45 = (c32&c24).where((c27*64+c38+c17*4096+-12480), UOp.const(dtypes.weakint, Invalid))
c48 = (c24&c32).where(c34.index(c45), UOp.const(dtypes.float, 0.0))
c49 = UOp.param(2, dtypes.half, shape=(64, 49, 4))
c61 = c48*c49.index((c26*4+c5%2+c16*28+c38*196))
@@ -50,7 +50,7 @@ def vision_conv_153():
c32 = ((c27<3)!=True)&(c27<35)
c34 = UOp.param(1, dtypes.half, shape=(16, 1024, 4))
c38 = c5//2
c45 = (c32&c24).where((c27*128+c38+c17*4096+-12672), UOp.const(dtypes.index, Invalid))
c45 = (c32&c24).where((c27*128+c38+c17*4096+-12672), UOp.const(dtypes.weakint, Invalid))
c48 = (c24&c32).where(c34.index(c45), UOp.const(dtypes.float, 0.0))
c49 = UOp.param(2, dtypes.half, shape=(128, 49, 4))
c61 = c48*c49.index((c26*4+c5%2+c16*28+c38*196))
+2 -2
View File
@@ -47,8 +47,8 @@ class TestHCQ(unittest.TestCase):
def setUp(self):
TestHCQ.d0.synchronize()
TestHCQ.a.uop.buffer.copyin(memoryview(bytearray(struct.pack("ff", 0, 1))))
TestHCQ.b.uop.buffer.copyin(memoryview(bytearray(struct.pack("ff", 0, 0))))
TestHCQ.a.uop.buffer.copy_from(Buffer("PYTHON", 2, dtypes.float, opaque=memoryview(bytearray(struct.pack("ff", 0, 1)))))
TestHCQ.b.uop.buffer.copy_from(Buffer("PYTHON", 2, dtypes.float, opaque=memoryview(bytearray(struct.pack("ff", 0, 0)))))
TestHCQ.d0.synchronize() # wait for copyins to complete
def test_run_1000_times_one_submit(self):
+7 -42
View File
@@ -1,66 +1,31 @@
import unittest, time
from tinygrad.runtime.support.usb import ASM24Controller
from tinygrad.helpers import Timing
import unittest
from tinygrad.helpers import Timing, getenv
from tinygrad import Tensor, Device
import numpy as np
class TestASMController(unittest.TestCase):
@classmethod
def setUpClass(cls):
cls.ctrl = ASM24Controller()
def test_write_and_read(self):
base = 0xF000
data = b"hello!"
self.ctrl.write(base, data)
out = self.ctrl.read(base, len(data))
self.assertEqual(out, data)
def test_scsi_write_and_read_from_f000(self):
payload = bytes([0x5B]) * 4096
self.ctrl.scsi_write(payload, lba=0)
back = self.ctrl.read(0xF000, len(payload))
self.assertEqual(back, payload)
def test_scsi_write_speed_4k(self):
payload = bytes([0x5A]) * 4096
start = time.perf_counter()
self.ctrl.scsi_write(payload, lba=0)
dur_ms = (time.perf_counter() - start) * 1000
print(f"scsi_write 4K took {dur_ms:.3f} ms")
def test_read_speed_4k(self):
payload = bytes([0xA5]) * 4096
self.ctrl.write(0xF000, payload)
start = time.perf_counter()
out = self.ctrl.read(0xF000, 4096)
dur_ms = (time.perf_counter() - start) * 1000
print(f"read 4K took {dur_ms:.3f} ms")
self.assertEqual(out, payload)
class TestDevCopySpeeds(unittest.TestCase):
@classmethod
def setUpClass(cls):
cls.sz = 512
cls.sz = getenv("SIZE", 2e6)
cls.dev = Device["AMD"]
if not cls.dev.is_usb(): raise unittest.SkipTest("only test this on USB devices")
def testCopyCPUtoDefault(self):
for _ in range(10):
t = Tensor.ones(self.sz, self.sz, device="CPU").contiguous().realize()
t = Tensor.ones(self.sz, device="CPU", dtype='uchar').contiguous().realize()
with Timing(f"copyin of {t.nbytes()/1e6:.2f} MB: ", on_exit=lambda ns: f" @ {t.nbytes()/ns * 1e3:.2f} MB/s"): # noqa: F821
t.to(Device.DEFAULT).realize()
Device[Device.DEFAULT].synchronize()
del t
def testCopyDefaulttoCPU(self):
t = Tensor.ones(self.sz, self.sz).contiguous().realize()
t = Tensor.ones(self.sz, dtype='uchar').contiguous().realize()
for _ in range(10):
with Timing(f"copyout of {t.nbytes()/1e6:.2f} MB: ", on_exit=lambda ns: f" @ {t.nbytes()/ns * 1e3:.2f} MB/s"):
t.to('CPU').realize()
def testValidateCopies(self):
t = Tensor.randn(self.sz, self.sz, device="CPU").contiguous().realize()
t = Tensor.randn(self.sz, device="CPU", dtype='uchar').contiguous().realize()
x = t.to(Device.DEFAULT).realize()
Device[Device.DEFAULT].synchronize()
@@ -70,4 +35,4 @@ class TestDevCopySpeeds(unittest.TestCase):
del x, y, t
if __name__ == "__main__":
unittest.main()
unittest.main()
+4 -6
View File
@@ -1,7 +1,7 @@
import random, ctypes
import random
import numpy as np
from tinygrad.device import Buffer, Device
from tinygrad.helpers import Context, getenv, from_mv
from tinygrad.helpers import Context, getenv
from tinygrad.dtype import dtypes
from tinygrad.tensor import Tensor, _to_np_dtype
from tinygrad.engine.realize import BufferXfer, get_runner, ExecItem
@@ -29,7 +29,7 @@ def alloc_rawbuffer(device, fill=False):
if fill:
with Context(DEBUG=0):
data = np.random.randint(-10000, 10000, size=rawbuf.size, dtype=_to_np_dtype(rawbuf.dtype))
rawbuf.copyin(Tensor(data).realize().uop.base.realized.as_memoryview())
rawbuf.copy_from(Tensor(data).realize().uop.base.realized)
return rawbuf
def gen_kernel_ji(device, deps):
@@ -84,9 +84,7 @@ def run_jit(jis, all_buffers, input_buffers, var_vals):
with Context(DEBUG=0):
for rawbuf in all_buffers:
if rawbuf in input_buffers: continue
mv = memoryview(bytearray(rawbuf.nbytes))
ctypes.memset(from_mv(mv), 0, len(mv))
rawbuf.copyin(mv)
rawbuf.copy_from(Buffer("PYTHON", rawbuf.size, rawbuf.dtype, opaque=memoryview(bytearray(rawbuf.nbytes))))
for ei in jis: ei.run(var_vals, jit=True)
+3 -3
View File
@@ -40,7 +40,7 @@ def random_int_expr(depth=10):
def random_bool_expr(depth=10, expr1=None):
if depth == 0: return True
if expr1 is None: expr1 = random_int_expr(depth-1)
expr2 = random.choice([random_or_sub_expression_int(depth-1, expr1), UOp.const(dtypes.index, random.randint(-10, 10))])
expr2 = random.choice([random_or_sub_expression_int(depth-1, expr1), UOp.const(dtypes.weakint, random.randint(-10, 10))])
return random.choice(comp_ops)(expr1, expr2)
@@ -82,8 +82,8 @@ if __name__ == "__main__":
f"v2=Variable(\"{u2.arg[0]}\", {u2.arg[1]}, {u2.arg[2]})\n" +\
f"v3=Variable(\"{u3.arg[0]}\", {u3.arg[1]}, {u3.arg[2]})\n" +\
f"expr = {expr}\n" +\
f"v1_val, v2_val, v3_val = UOp.const(dtypes.index, {n1.as_long()}), UOp.const(dtypes.index, {n2.as_long()})," +\
f"UOp.const(dtypes.index, {n3.as_long()})\n" +\
f"v1_val, v2_val, v3_val = UOp.const(dtypes.weakint, {n1.as_long()}), UOp.const(dtypes.weakint, {n2.as_long()})," +\
f"UOp.const(dtypes.weakint, {n3.as_long()})\n" +\
"num = expr.simplify().substitute({v1:v1_val, v2:v2_val, v3:v3_val}).ssimplify()\n" +\
"rn = expr.substitute({v1:v1_val, v2:v2_val, v3:v3_val}).ssimplify()\n" +\
"assert num==rn, f\"{num} != {rn}\"\n"
+10 -4
View File
@@ -6,7 +6,7 @@ from tinygrad import Tensor, dtypes, Device
from tinygrad.uop.ops import UOp, Ops, KernelInfo
from tinygrad.tensor import _to_np_dtype
from tinygrad.codegen import to_program
from tinygrad.dtype import DType
from tinygrad.dtype import DType, truncate
from tinygrad.nn.state import get_parameters
from tinygrad.helpers import T, Target, DEV
from tinygrad.renderer import Renderer
@@ -38,6 +38,10 @@ def call_is_graph(call:UOp) -> bool:
ast = call.src[0]
return ast.op is Ops.CUSTOM_FUNCTION and ast.arg == "graph"
def call_is_hcq(call:UOp) -> bool:
ast = call.src[0]
return ast.op is Ops.CUSTOM_FUNCTION and ast.arg == "hcq"
def jit_cache_count(linear:UOp) -> int:
n = 0
for call in linear.src:
@@ -51,6 +55,7 @@ def assert_jit_cache_len(fxn, expected_len):
if linear is None or not linear.src:
assert expected_len == 0, expected_len
return
if expected_len and all(call_is_hcq(call) for call in linear.src): expected_len = 3 # HCQ2: merged same-queue calls + finalizer + bumps
if call_is_graph(linear.src[0]):
assert len(linear.src) == 1, len(linear.src)
inner = linear.src[0].src[0].src[0] # LINEAR UOp inside CUSTOM_FUNCTION
@@ -58,6 +63,8 @@ def assert_jit_cache_len(fxn, expected_len):
else:
assert len(linear.src) == expected_len, f"expected {expected_len}, got {len(linear.src)}"
def min_normal(dt:DType) -> float: return 2.0 ** (2 - (1 << (dtypes.finfo(dt)[0] - 1)))
def rand_for_dtype(dt:DType, size:int, allow_subnormal=True):
if dtypes.is_unsigned(dt):
return np.random.randint(0, 100, size=size, dtype=_to_np_dtype(dt))
@@ -66,9 +73,8 @@ def rand_for_dtype(dt:DType, size:int, allow_subnormal=True):
elif dt == dtypes.bool:
return np.random.choice([True, False], size=size)
ret = np.random.uniform(-10, 10, size=size).astype(_to_np_dtype(dt))
if not allow_subnormal:
min_normal = 2.0 ** (2 - (1 << (dtypes.finfo(dt)[0] - 1)))
ret = np.where(np.abs(ret) < min_normal, 0, ret)
if dt == dtypes.bfloat16 or dt in dtypes.fp8s: ret = np.array([truncate[dt](x) for x in ret], dtype=ret.dtype)
if not allow_subnormal: ret = np.where(np.abs(ret) < min_normal(dt), 0, ret)
return ret
def timeit(fxn:Callable[..., T], *args, **kwargs) -> tuple[T, float]:
+8 -1
View File
@@ -9,7 +9,7 @@ MOCKGPU_ARCH = "cdna4" if DEV.arch == "gfx950" else "rdna4" if DEV.arch.startswi
assert (ma:=getenv("MOCKGPU_ARCH", "")) == "", "MOCKGPU_ARCH is deprecated, use DEV=" + \
str(replace(DEV.value, arch={"cdna4":"gfx950", "rdna4":"gfx1201"}.get(ma, "gfx1100"))) # type: ignore
GFX_TARGET_VERSION = {"rdna3": 110000, "rdna4": 120000, "cdna4": 90500}[MOCKGPU_ARCH]
import tinygrad.runtime.autogen.amd_gpu as amd_gpu, tinygrad.runtime.autogen.am.pm4_nv as pm4
import tinygrad.runtime.autogen.amd_gpu as amd_gpu, tinygrad.runtime.autogen.am.pm4_nv as pm4, tinygrad.runtime.autogen.am.sdma_6_0_0 as sdma
SDMA_MAX_COPY_SIZE = 0x400000
@@ -275,6 +275,7 @@ class SDMAExecutor(AMDQueue):
elif op == amd_gpu.SDMA_OP_POLL_REGMEM: cont = self._execute_poll_regmem()
elif op == amd_gpu.SDMA_OP_GCR: self._execute_gcr()
elif op == amd_gpu.SDMA_OP_COPY: self._execute_copy()
elif op == sdma.SDMA_OP_WRITE: self._execute_write()
elif op == amd_gpu.SDMA_OP_TIMESTAMP: self._execute_timestamp()
elif op == 32: self.rptr[0] += 4 # SDMA_OP_DUMMY_TRAP: pipeline flush, no interrupt
else: raise RuntimeError(f"Unknown SDMA op {op}")
@@ -289,6 +290,12 @@ class SDMAExecutor(AMDQueue):
struct = sdma_pkts.trap.from_address(self.base + self.rptr[0] % self.size)
self.rptr[0] += ctypes.sizeof(struct)
def _execute_write(self):
packet = to_mv(self.base + self.rptr[0] % self.size, 16).cast('I')
addr, count = packet[1] | packet[2] << 32, packet[3] + 1
ctypes.memmove(self.gpu.translate_addr(addr), self.base + self.rptr[0] % self.size + 16, count * 4)
self.rptr[0] += (4 + count) * 4
def _execute_poll_regmem(self):
struct = sdma_pkts.poll_regmem.from_address(self.base + self.rptr[0] % self.size)
+2 -2
View File
@@ -803,7 +803,7 @@ def _compile_sopp(inst: ir3.SOPP | ir4.SOPP, ctx: _Ctx) -> UOp:
pcode = get_pcode(inst.op)
pc_bytes = ctx.rpc() # PC is already 64-bit byte address
vcc, exec_val = ctx.rmask(_c(VCC_LO.offset)), ctx.rexec()
srcs = {'PC': pc_bytes.cast(dtypes.int64), 'SIMM16': simm16, 'SCC': ctx.rsgpr_dyn(_c(SCC.offset)), 'VCC': vcc,
srcs: dict[str, UOp|int] = {'PC': pc_bytes.cast(dtypes.int64), 'SIMM16': simm16, 'SCC': ctx.rsgpr_dyn(_c(SCC.offset)), 'VCC': vcc,
'VCCZ': vcc.eq(UOp.const(vcc.dtype, 0)).cast(dtypes.uint32),
'EXECZ': exec_val.eq(UOp.const(exec_val.dtype, 0)).cast(dtypes.uint32)}
for dest, val in parse_pcode(pcode, srcs)[1]:
@@ -858,7 +858,7 @@ def _compile_sop(inst: ir3.SOP1|ir3.SOP2|ir3.SOPC|ir3.SOPK|ir4.SOP1|ir4.SOP2|ir4
if isinstance(inst, ir4.SOPK): s0 = simm16
elif isinstance(inst, irc.SOPK) and 'CMPK' not in op_name and 'SETREG' not in op_name: s0 = simm16_sext
else: s0 = ctx.rsgpr_dyn(sdst_off)
srcs = {'S0': s0, 'S1': simm16_sext, 'SIMM16': simm16_sext, 'D0': ctx.rsgpr_dyn(sdst_off)}
srcs: dict[str, UOp|int] = {'S0': s0, 'S1': simm16_sext, 'SIMM16': simm16_sext, 'D0': ctx.rsgpr_dyn(sdst_off)}
dst_off, dst_size = sdst_off, 1
# S_GETREG_B32: extract bits from HW register. Handle as special case since HW_REGISTERS is not a normal variable.
# HW register values are stored at SGPR[SGPR_COUNT-16 + hwRegId] by _init_wave.
+8 -8
View File
@@ -688,10 +688,10 @@ class Parser:
return _extract_bits(base, hi, lo)
# Dynamic bit slice: (base >> lo) & ((1 << (hi - lo + 1)) - 1)
dt = dtypes.uint64 if base.dtype in (dtypes.uint64, dtypes.int64) else dtypes.uint32
hi, lo = first.cast(dt), second.cast(dt)
width = hi - lo + _const(dt, 1)
hi_u, lo_u = first.cast(dt), second.cast(dt)
width = hi_u - lo_u + _const(dt, 1)
mask = (_const(dt, 1) << width) - _const(dt, 1)
return (base.cast(dt) >> lo) & mask
return (base.cast(dt) >> lo_u) & mask
self.eat('RBRACKET')
dt_suffix = None
if self.try_eat('DOT'):
@@ -1123,11 +1123,11 @@ def parse_block(lines: list[str], start: int, env: dict[str, VarVal], funcs: dic
val = parse_tokens(toks[j:], env, funcs)
lo_dt, hi_dt = DTYPES.get(lo_type, dtypes.uint64), DTYPES.get(hi_type, dtypes.uint32)
lo_bits = 64 if lo_dt in (dtypes.uint64, dtypes.int64) else 32
lo_val = val.cast(lo_dt) if val.dtype.itemsize * 8 <= lo_bits else (val & _const(val.dtype, (1 << lo_bits) - 1)).cast(lo_dt)
hi_val = (val >> _const(val.dtype, lo_bits)).cast(hi_dt)
block_assigns[lo_var] = env[lo_var] = lo_val
block_assigns[hi_var] = env[hi_var] = hi_val
if assigns is not None: assigns.extend([(f'{lo_var}.{lo_type}', lo_val), (f'{hi_var}.{hi_type}', hi_val)])
lo_u = val.cast(lo_dt) if val.dtype.itemsize * 8 <= lo_bits else (val & _const(val.dtype, (1 << lo_bits) - 1)).cast(lo_dt)
hi_u = (val >> _const(val.dtype, lo_bits)).cast(hi_dt)
block_assigns[lo_var] = env[lo_var] = lo_u
block_assigns[hi_var] = env[hi_var] = hi_u
if assigns is not None: assigns.extend([(f'{lo_var}.{lo_type}', lo_u), (f'{hi_var}.{hi_type}', hi_u)])
i += 1
continue
+2 -2
View File
@@ -1,4 +1,4 @@
import ctypes, time, os, builtins, fcntl
import ctypes, time, os, builtins, fcntl, typing
from tinygrad.helpers import DEV
from tinygrad.runtime.support.hcq import FileIOInterface
from tinygrad.runtime.autogen import libc
@@ -9,7 +9,7 @@ start = time.perf_counter()
drivers = [cls() for t in DEV.value if (cls:={"MOCKPCI+AMD": AMDriver, "MOCKKFD+AMD": AMDDriver, "MOCK+AMD": AMDDriver, "MOCKUSB+AMD": AMUSBDriver,
"MOCK+NV": NVDriver}.get(f"{t.interface}+{t.device}"))]
tracked_fds = {}
tracked_fds: dict[int, typing.Any] = {}
original_memoryview = builtins.memoryview
class TrackedMemoryView:
+96 -100
View File
@@ -6,38 +6,25 @@ class MockUSB:
def __init__(self, mem):
self.mem = mem
def read(self, address, size): return bytes(self.mem[address:address+size])
def write(self, address, data, ignore_cache=False): self.mem[address:address+len(data)] = data
def pcie_mem_req(self, address, value=None, size=1):
if value is None: return int.from_bytes(self.mem[address:address+size], "little")
else: self.mem[address:address+size] = value.to_bytes(size, "little")
def pcie_mem_write(self, address, values, size):
for i, value in enumerate(values): self.pcie_mem_req(address + i * size, value, size)
def write(self, address, data): self.mem[address:address+len(data)] = data
def pcie_mem_read(self, address, nbytes): return bytes(self.mem[address:address+nbytes])
def pcie_mem_write(self, address, data): self.mem[address:address+len(data)] = data
# *** ASM24 Controller Mock ***
_mock_usb_state: MockASM24State|None = None
class MockASM24State:
"""Mock ASM24 controller: XRAM memory map, DMA windows, TLP engine, PCI config space.
"""Mock custom ASM24 controller: XRAM, DMA windows, PCI config space, and GPU BARs.
Memory map (64KB XRAM):
0xA000-0xAFFF: DMA window -> sys 0x820000
0xB000-0xB1FF: DMA window -> sys 0x800000
0xB200-0xB7FF: PCI MMIO (TLP engine)
0xB200-0xB7FF: controller PCI MMIO
0xF000-0xFFFF: DMA window -> sys 0x200000 (512KB)
"""
XRAM_SIZE = 0x10000
TLP_FMT_TYPE = 0xB210
TLP_BYTE_EN = 0xB217
TLP_ADDR_LO = 0xB218
TLP_ADDR_HI = 0xB21C
TLP_DATA = 0xB220
TLP_COMPL = 0xB22A
TLP_TRIGGER = 0xB254
TLP_LINK_STATUS = 0xB284
TLP_STATUS = 0xB296
def __init__(self, gpu, driver, vram_size:int, doorbell_size:int, mmio_size:int):
self.gpu, self.driver = gpu, driver
self._xram = bytearray(self.XRAM_SIZE)
@@ -93,53 +80,7 @@ class MockASM24State:
if ctrl_addr <= addr < ctrl_addr + dma_size:
(ctypes.c_ubyte * 1).from_address(host_addr + (addr - ctrl_addr))[0] = value
return
if addr == self.TLP_STATUS:
self._xram[addr] &= ~value & 0xFF
return
self._xram[addr] = value
if addr == self.TLP_TRIGGER and value == 0x0F: self._process_tlp()
# --- TLP engine ---
def _process_tlp(self):
fmt_type, byte_en = self._xram[self.TLP_FMT_TYPE], self._xram[self.TLP_BYTE_EN]
addr_lo = int.from_bytes(self._xram[self.TLP_ADDR_LO:self.TLP_ADDR_LO+4], 'big')
addr_hi = int.from_bytes(self._xram[self.TLP_ADDR_HI:self.TLP_ADDR_HI+4], 'big')
address = addr_lo | (addr_hi << 32)
size, offset, tmp = 0, 0, byte_en
while tmp and not (tmp & 1):
offset += 1
tmp >>= 1
while tmp:
size += tmp & 1
tmp >>= 1
is_write, is_cfg = bool(fmt_type & 0x40), (fmt_type & 0xbe) == 0x04
if is_cfg:
bus, dev, fn, byte_addr = (address >> 24) & 0xFF, (address >> 19) & 0x1F, (address >> 16) & 0x7, address & 0xFFC
if is_write:
data = int.from_bytes(self._xram[self.TLP_DATA:self.TLP_DATA+4], 'big')
self._cfg_write(bus, dev, fn, byte_addr + offset, (data >> (8 * offset)) & ((1 << (8 * size)) - 1), size)
else:
self._xram[self.TLP_DATA:self.TLP_DATA+4] = int.from_bytes(self._get_cfg(bus, dev, fn)[byte_addr:byte_addr+4], 'little').to_bytes(4, 'big')
self._xram[self.TLP_COMPL:self.TLP_COMPL+2] = (4).to_bytes(2, 'big')
self._xram[self.TLP_LINK_STATUS] = 0x01 if not is_write else 0x00
self._xram[self.TLP_STATUS] = 0x02
return
if is_write:
data = int.from_bytes(self._xram[self.TLP_DATA:self.TLP_DATA+4], 'big')
self._pcie_dispatch(address + offset, (data >> (8 * offset)) & ((1 << (8 * size)) - 1), size)
else:
result = self._pcie_dispatch(address + offset, None, size)
if result is not None:
self._xram[self.TLP_DATA:self.TLP_DATA+4] = ((result << (8 * offset)) & 0xFFFFFFFF).to_bytes(4, 'big')
self._xram[self.TLP_COMPL:self.TLP_COMPL+2] = (size & 0xFFF).to_bytes(2, 'big')
self._xram[self.TLP_LINK_STATUS] = 0x01 if not is_write else 0x00
self._xram[self.TLP_STATUS] = 0x02
def _cfg_write(self, bus:int, dev:int, fn:int, byte_addr:int, val:int, size:int):
cfg = self._get_cfg(bus, dev, fn)
@@ -168,49 +109,104 @@ class MockASM24State:
# Generic config write
for i in range(size): cfg[byte_addr + i] = (val >> (8 * i)) & 0xFF
def _pcie_dispatch(self, address:int, value:int|None, size:int) -> int|None:
def _find_bar(self, address:int, size:int) -> tuple[int, int]:
for reg_off, (bar_addr, bar_size) in self._bar_addrs.items():
if bar_addr <= address < bar_addr + bar_size:
offset = address - bar_addr
if reg_off == 0x10: # BAR0 - VRAM
if value is None: return int.from_bytes(bytes(self.gpu.vram[offset:offset+size]), "little")
self.gpu.vram[offset:offset+size] = list(value.to_bytes(size, "little"))
return None
if reg_off == 0x18: # BAR2 - Doorbell
if value is None: return int.from_bytes(bytes(self._doorbell[offset:offset+size]), "little")
for i, b in enumerate(value.to_bytes(size, "little")): self._doorbell[offset + i] = b
self.driver._emulate_execute()
return None
if reg_off == 0x24: # BAR5 - MMIO
if value is None: return self.gpu.mmio[offset // 4]
self.gpu.mmio[offset // 4] = value
return None
raise ValueError(f"PCIe address {address:#x} not mapped to any BAR")
if bar_addr <= address and address + size <= bar_addr + bar_size: return reg_off, address - bar_addr
raise ValueError(f"PCIe range {address:#x}+{size:#x} not mapped to any BAR")
# --- CDB processing (called by MockUSB3.send_batch) ---
def _pcie_read(self, address:int, size:int) -> bytes:
reg_off, offset = self._find_bar(address, size)
if reg_off == 0x10: return bytes(self.gpu.vram[offset:offset+size])
if reg_off == 0x18: return bytes(self._doorbell[offset:offset+size])
if reg_off == 0x24: return bytes((self.gpu.mmio[(offset+i)//4] >> (8*((offset+i)&3))) & 0xFF for i in range(size))
raise RuntimeError(f"unsupported BAR register {reg_off:#x}")
def process_cdb(self, cdb:bytes, rlen:int, send_data:bytes|None) -> bytes|None:
op = cdb[0]
if op == 0xE5: # write byte
self._xram_write_byte(((cdb[2] << 16) | (cdb[3] << 8) | cdb[4]) & 0xFFFF, cdb[1])
return None
if op == 0xE4: # read
return self._xram_read(((cdb[2] << 16) | (cdb[3] << 8) | cdb[4]) & 0xFFFF, cdb[1])
if op == 0x8A and send_data is not None and 0xF000 in self._dma_regions: # SCSI write
host_addr, dma_size = self._dma_regions[0xF000]
ctypes.memmove(host_addr, send_data, min(len(send_data), dma_size))
def _pcie_write(self, address:int, data:bytes):
reg_off, offset = self._find_bar(address, len(data))
if reg_off == 0x10: self.gpu.vram[offset:offset+len(data)] = list(data)
elif reg_off == 0x18:
self._doorbell[offset:offset+len(data)] = list(data)
self.driver._emulate_execute()
elif reg_off == 0x24:
updates: dict[int, int] = {}
for i, byte in enumerate(data):
idx, shift = (offset+i)//4, 8*((offset+i)&3)
updates[idx] = (updates.get(idx, self.gpu.mmio[idx]) & ~(0xFF << shift)) | (byte << shift)
for idx, val in updates.items(): self.gpu.mmio[idx] = val
else: raise RuntimeError(f"unsupported BAR register {reg_off:#x}")
def _pcie_dispatch(self, address:int, value:int|None, size:int) -> int|None:
if value is None: return int.from_bytes(self._pcie_read(address, size), 'little')
self._pcie_write(address, value.to_bytes(size, 'little'))
return None
class MockUSB3:
@classmethod
def list_devices(cls, vendor, dev): return [(0, "usb:mock")]
def __init__(self, *args, **kwargs):
self.product, self.is_custom = "", False
def send_batch(self, cdbs:list[bytes], idata:list[int]|None=None, odata:list[bytes|None]|None=None) -> list[bytes|None]:
self.product = "custom mock"
self._bulk_read_op: tuple[str, int, int]|None = None
self._bulk_write_op: tuple[str, int, int]|None = None
self._f0_reply = bytes(8)
@property
def state(self) -> MockASM24State:
assert _mock_usb_state is not None
idata, odata = idata or [0] * len(cdbs), odata or [None] * len(cdbs)
results: list[bytes|None] = []
for cdb, rlen, sdata in zip(cdbs, idata, odata):
result = _mock_usb_state.process_cdb(cdb, rlen, sdata)
results.append(result if rlen > 0 else None)
return results
return _mock_usb_state
def control_write(self, request:int, value:int=0, index:int=0, data:bytes=b'', timeout:int=1000):
if request == 0xF3:
self.state._xram[0xB450] = 0x78 if value else 0
elif request == 0xE5:
self.state._xram_write_byte(value, index)
elif request == 0xF2:
op = ("sram_read" if value & 0x8000 else "sram_write", 0xF000, (value & 0x7FFF) * 512)
if value & 0x8000: self._bulk_read_op = op
else: self._bulk_write_op = op
elif request == 0xF0:
address_lo, address_hi, payload = struct.unpack('<III', data)
address, fmt_type, byte_en = address_lo | (address_hi << 32), value & 0xFF, value >> 8
if index == 1: self._bulk_write_op = ("pcie_write", address, payload * 4)
elif index == 2: self._bulk_read_op = ("pcie_read", address, payload * 4)
else:
assert index == 0 and byte_en
offset = (byte_en & -byte_en).bit_length() - 1
size, is_write, is_cfg = byte_en.bit_count(), bool(fmt_type & 0x40), (fmt_type & 0xBE) == 0x04
if is_cfg:
bus, dev, fn, byte_addr = (address >> 24) & 0xFF, (address >> 19) & 0x1F, (address >> 16) & 0x7, address & 0xFFC
if is_write: self.state._cfg_write(bus, dev, fn, byte_addr + offset, (payload >> (8 * offset)) & ((1 << (8 * size))-1), size)
else: payload = int.from_bytes(self.state._get_cfg(bus, dev, fn)[byte_addr:byte_addr+4], 'little')
elif is_write:
self.state._pcie_dispatch(address + offset, (payload >> (8 * offset)) & ((1 << (8 * size))-1), size)
else: payload = (self.state._pcie_dispatch(address + offset, None, size) or 0) << (8 * offset)
self._f0_reply = struct.pack('<I', payload & 0xFFFFFFFF) + bytes(4)
else: raise ValueError(f"unsupported control OUT request 0x{request:02X}")
def control_read(self, request:int, length:int, value:int=0, index:int=0, timeout:int=1000) -> memoryview:
if request == 0xE4: data = self.state._xram_read(value, length)
elif request == 0xF0: data = self._f0_reply
else: raise ValueError(f"unsupported control IN request 0x{request:02X}")
return memoryview(data[:length])
def bulk_write(self, data:bytes, timeout:int=1000):
assert self._bulk_write_op is not None
op, address, size = self._bulk_write_op
assert len(data) == size
if op == "sram_write":
host_addr, region_size = self.state._dma_regions[address]
ctypes.memmove(host_addr, data, min(len(data), region_size))
elif op == "pcie_write": self.state._pcie_write(address, data)
else: raise RuntimeError(f"cannot bulk write for {op}")
self._bulk_write_op = None
def bulk_read(self, length:int, timeout:int=1000) -> memoryview:
assert self._bulk_read_op is not None
op, address, size = self._bulk_read_op
assert length == size
if op == "sram_read":
host_addr, region_size = self.state._dma_regions[address]
data = bytes((ctypes.c_ubyte * min(length, region_size)).from_address(host_addr))
elif op == "pcie_read": data = self.state._pcie_read(address, length)
else: raise RuntimeError(f"cannot bulk read for {op}")
self._bulk_read_op = None
return memoryview(data)
+19 -1
View File
@@ -1,6 +1,6 @@
import unittest, itertools, math
from tinygrad import Tensor, dtypes, Context
from tinygrad.dtype import DType, ConstType
from tinygrad.dtype import DType, ConstType, Invalid
from tinygrad.uop.ops import Ops, UOp
from test.helpers import full_rewrite
import numpy as np
@@ -34,6 +34,24 @@ class TestUnaryOpsConstFolding(unittest.TestCase):
x = x.clip(0, 1).realize()
_check_ast_count(1, x.neg())
class TestWeakConstFolding(unittest.TestCase):
def test_weakint_math(self):
out = (UOp.const(dtypes.weakint, 2**40) + UOp.const(dtypes.weakint, 2**40)).simplify()
self.assertEqual((out.op, out.dtype, out.arg), (Ops.CONST, dtypes.weakint, 2**41))
def test_float_unaries(self):
for dtype in (dtypes.weakfloat,):
for op in (Ops.SIN, Ops.LOG2, Ops.EXP2, Ops.SQRT, Ops.RECIPROCAL):
out = UOp.const(dtype, 4).alu(op).simplify()
self.assertEqual((out.op, out.dtype), (Ops.CONST, dtypes.weakfloat))
def test_weakfloat_math(self):
out = (UOp.const(dtypes.weakfloat, 1.25) + UOp.const(dtypes.weakfloat, 2.5)).simplify()
self.assertEqual((out.op, out.dtype, out.arg), (Ops.CONST, dtypes.weakfloat, 3.75))
def test_invalid_poison(self):
self.assertIs(UOp.const(dtypes.weakint, Invalid).alu(Ops.CDIV, UOp.const(dtypes.weakint, 0)).simplify().arg, Invalid)
class TestBinaryOpsConstFolding(unittest.TestCase):
def test_add_literal_zero(self):
_check_ast_count(0, Tensor([1.0, 2, 3, 4]) + 0)
+6 -19
View File
@@ -1,6 +1,6 @@
import unittest, math, struct, operator
from tinygrad import Tensor, Device
from tinygrad.dtype import DTYPES_DICT, dtypes, truncate, float_to_fp16, float_to_bf16, _to_np_dtype, least_upper_dtype, least_upper_float
from tinygrad.dtype import DTYPES_DICT, dtypes, Invalid, truncate, float_to_fp16, float_to_bf16, _to_np_dtype, least_upper_dtype, least_upper_float
from tinygrad.helpers import getenv
from hypothesis import given, settings, strategies as strat
@@ -57,6 +57,7 @@ class TestHelpers(unittest.TestCase):
def test_from_py(self):
assert dtypes.from_py(True) == dtypes.bool
assert dtypes.from_py(Invalid) == dtypes.bool
assert dtypes.from_py(2) == dtypes.default_int
assert dtypes.from_py(3.0) == dtypes.default_float
assert dtypes.from_py([]) == dtypes.default_float
@@ -223,30 +224,16 @@ class TestTypePromotion(unittest.TestCase):
assert least_upper_dtype(dtypes.fp8e5m2, dtypes.uint64) == dtypes.fp8e5m2
def test_weakint_promo(self):
# weakint with itself is weakint
assert least_upper_dtype(dtypes.weakint, dtypes.weakint) == dtypes.weakint
# weakint is above bool
assert least_upper_dtype(dtypes.weakint, dtypes.bool) == dtypes.weakint
# weakint defers to any concrete int type
assert least_upper_dtype(dtypes.weakint, dtypes.int8) == dtypes.int8
assert least_upper_dtype(dtypes.weakint, dtypes.uint8) == dtypes.uint8
assert least_upper_dtype(dtypes.weakint, dtypes.int16) == dtypes.int16
assert least_upper_dtype(dtypes.weakint, dtypes.int32) == dtypes.int32
assert least_upper_dtype(dtypes.weakint, dtypes.int64) == dtypes.int64
assert least_upper_dtype(dtypes.weakint, dtypes.uint64) == dtypes.uint64
# weakint defers to any float type
assert least_upper_dtype(dtypes.weakint, dtypes.float16) == dtypes.float16
assert least_upper_dtype(dtypes.weakint, dtypes.float32) == dtypes.float32
assert least_upper_dtype(dtypes.weakint, dtypes.float64) == dtypes.float64
with self.assertRaises(KeyError): least_upper_dtype(dtypes.weakint, dtypes.weakint)
with self.assertRaises(KeyError): least_upper_dtype(dtypes.weakint, dtypes.int8)
def test_weakfloat_promo(self):
# weakfloat is a float, but like weakint it is not one of dtypes.floats
# weakfloat is a float, but is not one of dtypes.floats
assert dtypes.is_float(dtypes.weakfloat) and dtypes.weakfloat not in dtypes.floats
# weakfloat with itself is weakfloat
assert least_upper_dtype(dtypes.weakfloat, dtypes.weakfloat) == dtypes.weakfloat
# weakfloat is above bool, weakint and any concrete int (they defer up to it)
# weakfloat is above bool and any concrete int (they defer up to it)
assert least_upper_dtype(dtypes.weakfloat, dtypes.bool) == dtypes.weakfloat
assert least_upper_dtype(dtypes.weakfloat, dtypes.weakint) == dtypes.weakfloat
assert least_upper_dtype(dtypes.weakfloat, dtypes.int32) == dtypes.weakfloat
assert least_upper_dtype(dtypes.weakfloat, dtypes.uint64) == dtypes.weakfloat
# weakfloat defers to any concrete float type
+1 -1
View File
@@ -24,7 +24,7 @@ class TestGroupedDims(unittest.TestCase):
total = math.prod(dims)
specials = sorted(dedup(flatten([[y for y in x.toposort() if y.op is Ops.SPECIAL] for x in idxs])), key=lambda u: u.arg)
# build flat index and primed flat (same expression with renamed SPECIALs)
flat = UOp.const(dtypes.index, 0)
flat = UOp.const(dtypes.weakint, 0)
for i, idx in enumerate(idxs):
flat = flat + idx * int(math.prod(dims[i+1:]))
flat_p = flat.substitute({s: UOp(Ops.SPECIAL, src=s.src, arg=s.arg+"_p") for s in specials})
+4 -4
View File
@@ -107,21 +107,21 @@ class TestFoldingAndReduction(unittest.TestCase):
class TestModuloAndDivisionFolding(unittest.TestCase):
def test_full_graph_rewrite_modulo_folding_with_define_var(self):
# index dtype because div-mod rules only work on index
x_var_uop = UOp.variable('x', 0, 100).cast(dtypes.index)
x_var_uop = UOp.variable('x', 0, 100).cast(dtypes.weakint)
optimized_mod_uop = apply_rewrite(((x_var_uop * 4) + 2) % 4)
self.assertEqual(optimized_mod_uop.op, Ops.CONST)
self.assertEqual(optimized_mod_uop.arg, 2)
def test_full_graph_rewrite_division_folding_with_define_var(self):
# index dtype because div-mod rules only work on index
n_var_uop = UOp.variable('n', 1, 1000).cast(dtypes.index)
n_var_uop = UOp.variable('n', 1, 1000).cast(dtypes.weakint)
optimized_div_uop = apply_rewrite((n_var_uop * 6) // 3)
self.assertEqual(optimized_div_uop.op, Ops.MUL)
self.assertEqual(optimized_div_uop.src[1].arg, 2)
def test_full_graph_rewrite_complex_mod_div_folding(self):
# index dtype because div-mod rules only work on index
k_var_uop = UOp.variable('k', 0, 50).cast(dtypes.index)
k_var_uop = UOp.variable('k', 0, 50).cast(dtypes.weakint)
optimized_div_uop = apply_rewrite(((k_var_uop * 12 + 8) % 6) // 2)
self.assertEqual(optimized_div_uop.op, Ops.CONST)
self.assertEqual(optimized_div_uop.arg, 1)
@@ -140,7 +140,7 @@ class TestModuloAndDivisionFolding(unittest.TestCase):
def test_full_graph_rewrite_modulo_large_divisor(self):
# index dtype because div-mod rules only work on index
x_var_uop = UOp.variable('x', 1, 5)
self.assertIs(apply_rewrite(x_var_uop.cast(dtypes.index) % 10).render(simplify=False), x_var_uop.render(simplify=False))
self.assertIs(apply_rewrite(x_var_uop.cast(dtypes.weakint) % 10).render(simplify=False), x_var_uop.render(simplify=False))
def test_full_graph_rewrite_division_with_remainder(self):
x_var_uop = UOp.variable('x', 7, 9)
+9 -9
View File
@@ -84,22 +84,22 @@ class TestUSBMMIOInterface(unittest.TestCase):
self.mmio[2] = 0xFE
self.assertEqual(full_view[2], 0xFE)
def test_pcimem_byte(self):
def test_pcimem_dword(self):
usb2 = MockUSB(bytearray(self.size))
mmio_pci = USBMMIOInterface(usb2, 0, self.size, fmt='B', pcimem=True)
mmio_pci[3] = 0x11
self.assertEqual(mmio_pci[3], 0x11)
self.assertEqual(usb2.mem[3], 0x11)
mmio_pci = USBMMIOInterface(usb2, 0, self.size, fmt='I', pcimem=True)
mmio_pci[3] = 0x11223344
self.assertEqual(mmio_pci[3], 0x11223344)
self.assertEqual(usb2.mem[12:16], b'\x44\x33\x22\x11')
def test_pcimem_slice(self):
usb3 = MockUSB(bytearray(self.size))
mmio_pci = USBMMIOInterface(usb3, 0, self.size, fmt='B', pcimem=True)
values = [2, 3, 4]
mmio_pci[4:7] = values
raw = mmio_pci[4:7]
values = [2, 3, 4, 5]
mmio_pci[4:8] = values
raw = mmio_pci[4:8]
self.assertIsInstance(raw, bytes)
self.assertEqual(list(raw), values)
self.assertEqual([mmio_pci[i] for i in range(4, 7)], values)
self.assertEqual(list(usb3.mem[4:8]), values)
if __name__ == "__main__":
unittest.main()
+2 -1
View File
@@ -1,6 +1,7 @@
import ctypes, gzip, unittest, timeit, pickle
from tinygrad import Variable
from tinygrad.helpers import Context, ContextVar, argfix, colored, word_wrap, is_numpy_ndarray, mv_address, count, all_same
from tinygrad.helpers import Context, ContextVar, argfix, colored, word_wrap, mv_address, count, all_same
from tinygrad.tensor import is_numpy_ndarray
from tinygrad.helpers import merge_dicts, strip_parens, prod, round_up, fetch, fully_flatten, from_mv, to_mv, polyN, time_to_str, cdiv, cmod, getbits
from tinygrad.helpers import ceildiv, ansistrip, get_shape
from tinygrad.tensor import Tensor
+5 -5
View File
@@ -8,12 +8,12 @@ from tinygrad.codegen import to_program
class TestLinearizerFailures(unittest.TestCase):
def test_fail_1(self):
c0 = UOp.param(0, dtypes.float, (64,))
c1 = UOp.range(UOp.const(dtypes.index, 2), 1, AxisType.LOOP)
c2 = UOp.range(UOp.const(dtypes.index, 32), 2, AxisType.LOOP)
c3 = ((c1*UOp.const(dtypes.index, 32))+c2)
c1 = UOp.range(UOp.const(dtypes.weakint, 2), 1, AxisType.LOOP)
c2 = UOp.range(UOp.const(dtypes.weakint, 32), 2, AxisType.LOOP)
c3 = ((c1*UOp.const(dtypes.weakint, 32))+c2)
c4 = UOp.param(1, dtypes.float, (163840,))
c5 = UOp.range(UOp.const(dtypes.index, 2560), 0, AxisType.REDUCE)
c6 = c4.index(((((((c5//UOp.const(dtypes.index, 8))%UOp.const(dtypes.index, 8))*UOp.const(dtypes.index, 8))+(c5%UOp.const(dtypes.index, 8)))+(((c2*UOp.const(dtypes.index, 40))+(c5//UOp.const(dtypes.index, 64)))*UOp.const(dtypes.index, 64)))+(c1*UOp.const(dtypes.index, 81920))))
c5 = UOp.range(UOp.const(dtypes.weakint, 2560), 0, AxisType.REDUCE)
c6 = c4.index(((((((c5//UOp.const(dtypes.weakint, 8))%UOp.const(dtypes.weakint, 8))*UOp.const(dtypes.weakint, 8))+(c5%UOp.const(dtypes.weakint, 8)))+(((c2*UOp.const(dtypes.weakint, 40))+(c5//UOp.const(dtypes.weakint, 64)))*UOp.const(dtypes.weakint, 64)))+(c1*UOp.const(dtypes.weakint, 81920))))
c7 = UOp.param(2, dtypes.float, (64,))
c8 = c7.index(c3)
c9 = ((((c6+(c8*UOp.const(dtypes.float, -1.0)))*(c6+(c8*UOp.const(dtypes.float, -1.0)))).reduce(c5, arg=Ops.ADD)*UOp.const(dtypes.float, 0.000390625))+UOp.const(dtypes.float, 1e-05)).sqrt().reciprocal()
+113 -50
View File
@@ -1,4 +1,4 @@
import unittest, threading, time
import unittest, threading, time, json
from unittest.mock import Mock
class TestLLMServer(unittest.TestCase):
@@ -7,12 +7,9 @@ class TestLLMServer(unittest.TestCase):
@classmethod
def setUpClass(cls):
cls.mock_tok = Mock()
cls.mock_tok.role = Mock(return_value=[100, 101])
cls.mock_tok.encode = Mock(return_value=[200, 201, 202])
cls.mock_tok.decode = Mock(return_value="Hello")
cls.mock_tok.stream_decoder = Mock(return_value=lambda tid=None: "Hello" if tid is not None else "")
cls.mock_tok.end_turn = Mock(return_value=[998])
cls.mock_tok.prefix = Mock(return_value=[1])
cls.mock_tok.preset = "llama3"
cls.mock_tok.bos_id = 1
cls.mock_tok.eos_id = 999
@@ -20,12 +17,14 @@ class TestLLMServer(unittest.TestCase):
cls.mock_tok.is_end = Mock(side_effect=lambda tid: tid in (999,))
cls.mock_model = Mock()
cls.mock_model.max_context = 4
cls.mock_model.generate = Mock(side_effect=lambda ids, **kwargs: iter([300, 301, 999]))
cls.mock_model.get_start_pos = Mock(return_value=0)
from tinygrad.llm.cli import LLMServer
from tinygrad.llm.cli import FallbackTemplate
from tinygrad.llm.serve import LLMServer
cls.server = LLMServer(('127.0.0.1', 0), cls.mock_model, "test-model", cls.mock_tok)
cls.server = LLMServer(('127.0.0.1', 0), cls.mock_model, "test-model", cls.mock_tok, FallbackTemplate(cls.mock_tok))
cls.port = cls.server.server_address[1]
cls.server_thread = threading.Thread(target=cls.server.serve_forever, daemon=True)
cls.server_thread.start()
@@ -131,6 +130,16 @@ class TestLLMServer(unittest.TestCase):
self.assertIsNotNone(resp.usage.prompt_tokens)
self.assertIsNotNone(resp.usage.completion_tokens)
def test_context_length_error(self):
from openai import BadRequestError
self.mock_tok.encode.return_value = [200, 201, 202, 203]
try:
with self.assertRaises(BadRequestError) as err:
self.client.chat.completions.create(model="test-model", messages=[{"role":"user", "content":"too long"}])
self.assertEqual(err.exception.code, "context_length_exceeded")
finally:
self.mock_tok.encode.return_value = [200, 201, 202]
def test_max_tokens_streaming(self):
self.mock_model.generate = Mock(side_effect=lambda ids, **kwargs: iter([300, 301, 302, 303, 999]))
stream = self.client.chat.completions.create(
@@ -149,50 +158,6 @@ class TestLLMServer(unittest.TestCase):
self.assertEqual(resp.choices[0].finish_reason, "length")
self.assertEqual(resp.usage.completion_tokens, 2)
def test_assistant_prefill(self):
"""Last assistant message should be treated as prefill (not a completed turn)."""
self.mock_model.generate = Mock(side_effect=lambda ids, **kwargs: iter([300, 999]))
captured_ids = []
orig_generate = self.mock_model.generate.side_effect
def capture_generate(ids, **kwargs):
captured_ids.extend(ids)
return orig_generate(ids, **kwargs)
self.mock_model.generate = Mock(side_effect=capture_generate)
resp = self.client.chat.completions.create(
model="test", messages=[
{"role": "user", "content": "Hello"},
{"role": "assistant", "content": "Sure"}
], stream=False
)
# prefill tokens should be in ids: role("assistant") + encode("Sure") but NO end_turn after it
# and NO extra role("assistant") appended
role_tokens = self.mock_tok.role.call_args_list
# last role() call should be for "assistant" (the prefill message), not an extra one
self.assertEqual(role_tokens[-1], unittest.mock.call("assistant"))
# end_turn should be called once less than role() — the prefill assistant msg doesn't get end_turn
# NOTE: this is flaky in random order
#self.assertEqual(self.mock_tok.end_turn.call_count, self.mock_tok.role.call_count - 1)
self.assertIsNotNone(resp.choices[0].message.content)
def test_assistant_prefill_not_last(self):
"""Assistant message that's NOT last should be a normal completed turn."""
self.mock_model.generate = Mock(side_effect=lambda ids, **kwargs: iter([300, 999]))
self.mock_tok.role.reset_mock()
self.mock_tok.end_turn.reset_mock()
self.client.chat.completions.create(
model="test", messages=[
{"role": "user", "content": "Hello"},
{"role": "assistant", "content": "Sure"},
{"role": "user", "content": "Continue"}
], stream=False
)
# all messages get end_turn, plus an extra role("assistant") at the end
# roles: user, assistant, user, assistant(generation prompt) = 4 role calls
# end_turns: user, assistant, user = 3 end_turn calls (one per message)
self.assertEqual(self.mock_tok.end_turn.call_count, 3)
self.assertEqual(self.mock_tok.role.call_count, 4)
def test_models_endpoint(self):
import requests as req
resp = req.get(f"http://127.0.0.1:{self.port}/v1/models")
@@ -203,5 +168,103 @@ class TestLLMServer(unittest.TestCase):
self.assertEqual(data["data"][0]["id"], "test-model")
self.assertEqual(data["data"][0]["object"], "model")
class TestLLMToolCalls(unittest.TestCase):
"""Tool calling through the OpenAI-compatible HTTP API."""
@classmethod
def setUpClass(cls):
cls.mock_tok = Mock()
cls.mock_tok.encode = Mock(return_value=[200, 201, 202])
cls.mock_tok.decode = Mock(return_value="")
cls.mock_tok.preset = "qwen2"
cls.mock_tok.bos_id, cls.mock_tok.eos_id, cls.mock_tok.eot_id = None, 999, None
cls.mock_tok.is_end = Mock(return_value=False)
cls.mock_model = Mock()
cls.mock_model.max_context = 4
cls.mock_model.get_start_pos = Mock(return_value=0)
from tinygrad.llm.serve import LLMServer
import jinja2
# .items() matches tool-aware templates and ensures OpenAI JSON argument strings are normalized before rendering the next turn.
template = jinja2.Template("""{% for m in messages %}{{ m.content or '' }}{% for tc in m.tool_calls or [] %}
{% for key, value in tc.function.arguments.items() %}{{ key }}={{ value }}{% endfor %}{% endfor %}{% endfor %}""")
cls.server = LLMServer(('127.0.0.1', 0), cls.mock_model, "tool-model", cls.mock_tok, template)
cls.port = cls.server.server_address[1]
cls.server_thread = threading.Thread(target=cls.server.serve_forever, daemon=True)
cls.server_thread.start()
time.sleep(0.1)
from openai import OpenAI
cls.client = OpenAI(base_url=f"http://127.0.0.1:{cls.port}/v1", api_key="test")
@classmethod
def tearDownClass(cls):
cls.server.shutdown()
cls.server.server_close()
def set_output(self, text:str):
pieces = dict(enumerate(text, 1))
self.mock_tok.stream_decoder = Mock(return_value=lambda tid=None: pieces[tid] if tid is not None else "")
self.mock_model.generate = Mock(side_effect=lambda ids, **kwargs: iter(pieces))
@staticmethod
def tools():
return [{"type":"function", "function":{"name":"read", "description":"Read a file",
"parameters":{"type":"object", "properties":{"path":{"type":"string"}}, "required":["path"]}}}]
def test_streaming_tool_call(self):
self.set_output('before<tool_call>{"name":"read","arguments":{"path":"README.md"}}</tool_call>')
chunks = list(self.client.chat.completions.create(model="tool-model", messages=[{"role":"user", "content":"Read README.md"}],
tools=self.tools(), stream=True))
self.assertEqual("".join(c.choices[0].delta.content or "" for c in chunks if c.choices), "before")
calls = [tc for c in chunks if c.choices for tc in c.choices[0].delta.tool_calls or []]
self.assertEqual(len(calls), 1)
self.assertEqual(calls[0].function.name, "read")
self.assertEqual(json.loads(calls[0].function.arguments), {"path":"README.md"})
self.assertEqual(chunks[-1].choices[0].finish_reason, "tool_calls")
def test_multiple_xml_tool_calls(self):
self.set_output("<tool_call><function=read><parameter=path>\"a\"</parameter></function></tool_call>"
"<tool_call><function=read><parameter=path>\"b\"</parameter></function></tool_call>")
response = self.client.chat.completions.create(model="tool-model", messages=[{"role":"user", "content":"Read a and b"}],
tools=self.tools())
self.assertEqual([json.loads(tc.function.arguments)["path"] for tc in response.choices[0].message.tool_calls], ["a", "b"])
self.assertEqual(response.choices[0].finish_reason, "tool_calls")
def test_multiline_tool_argument_preserves_trailing_newline(self):
self.set_output("<tool_call>\n<function=write>\n<parameter=content>\nfirst\nsecond\n\n</parameter>\n"
"<parameter=filePath>\nout.txt\n</parameter>\n</function>\n</tool_call>")
response = self.client.chat.completions.create(model="tool-model", messages=[{"role":"user", "content":"Write out.txt"}], tools=self.tools())
args = json.loads(response.choices[0].message.tool_calls[0].function.arguments)
self.assertEqual(args, {"content":"first\nsecond\n", "filePath":"out.txt"})
def test_invalid_tool_call_becomes_content(self):
self.set_output("<tool_call>not a call</tool_call>")
response = self.client.chat.completions.create(model="tool-model", messages=[{"role":"user", "content":"Hello"}], tools=self.tools())
self.assertEqual(response.choices[0].message.content, "<tool_call>not a call</tool_call>")
self.assertIsNone(response.choices[0].message.tool_calls)
self.assertEqual(response.choices[0].finish_reason, "stop")
def test_tool_call_in_reasoning_is_not_executed(self):
self.set_output('<think>draft <tool_call>{"name":"wrong","arguments":{}}</tool_call></think>answer')
response = self.client.chat.completions.create(model="tool-model", messages=[{"role":"user", "content":"Hello"}], tools=self.tools())
self.assertEqual(response.choices[0].message.content, "answer")
self.assertIsNone(response.choices[0].message.tool_calls)
self.assertEqual(response.choices[0].finish_reason, "stop")
def test_tool_result_round_trip(self):
self.set_output('<tool_call>{"name":"read","arguments":{"path":"README.md"}}</tool_call>')
first = self.client.chat.completions.create(model="tool-model", messages=[{"role":"user", "content":"Read README.md"}], tools=self.tools())
call = first.choices[0].message.tool_calls[0]
self.set_output("done")
second = self.client.chat.completions.create(model="tool-model", messages=[
{"role":"user", "content":"Read README.md"},
{"role":"assistant", "content":None, "tool_calls":[call.model_dump()]},
{"role":"tool", "tool_call_id":call.id, "content":"file contents"},
], tools=self.tools())
self.assertEqual(second.choices[0].message.content, "done")
self.assertEqual(second.choices[0].finish_reason, "stop")
if __name__ == '__main__':
unittest.main()
+41 -5
View File
@@ -1,5 +1,5 @@
import unittest, base64, functools, sys
from tinygrad.llm.cli import SimpleTokenizer
import unittest, base64, functools, re, sys, time, unicodedata
from tinygrad.llm.cli import SimpleTokenizer, FallbackTemplate
from tinygrad.helpers import fetch
@unittest.skipIf(sys.platform == 'win32', "fetch race condition on Windows")
@@ -46,6 +46,41 @@ class TestLLMTokenizer(unittest.TestCase):
def test_llama_repeat(self): self._test_coding(self.llama_tok, "00000000000000000", [ 931, 931, 931, 931, 931, 410 ])
def test_llama_pat(self): self._test_coding(self.llama_tok, "today\n \n", [ 31213, 14211 ])
def test_split_regex_matches_naive_listing(self):
# the compacted codepoint ranges must match the same text as listing every codepoint
def naive(pre): return "".join(re.escape(chr(cp)) for cp in range(0x323b0) if unicodedata.category(chr(cp)).startswith(pre))
r_ws, r_p_N, r_p_L = r"\t\n\x0b\x0c\r\x85" + naive("Z"), naive("N"), naive("L")
naive_re = re.compile("(?i:'s|'t|'re|'ve|'m|'ll|'d)|" +
f"[^\\r\\n{r_p_N}{r_p_L}]?[{r_p_L}]+|[{r_p_N}]{{1,3}}| ?[^{r_ws}{r_p_N}{r_p_L}]+[\\r\\n]*|[{r_ws}]*[\\r\\n]+|[{r_ws}]+(?![^{r_ws}])|[{r_ws}]+")
sample = "hello world 한국어 中文 текст ١٢٣ 123 😊\n \ttoday\n'équivalent ²³№ "
self.assertEqual(SimpleTokenizer({}, {})._split_to_word.findall(sample), naive_re.findall(sample))
def test_split_regex_speed(self):
# the naive listing compiles a 429KB pattern that takes 10+s to match a 225KB prompt; ranges keep it small and fast
tok = SimpleTokenizer({}, {})
self.assertLess(len(tok._split_to_word.pattern), 100_000)
text = "The quick brown fox jumps over the lazy dog. " * 5000
tok._split_to_word.findall(text) # warmup
tms = []
for _ in range(5):
st = time.perf_counter()
words = tok._split_to_word.findall(text)
tms.append(time.perf_counter() - st)
self.assertLess(min(tms), 4) # best-of-5 is robust to CI scheduling pauses; new code takes ~60ms
self.assertEqual(len(words), 50001)
def test_llama_continued_conversation(self):
self._test_coding(self.llama_tok, "hello <|eot_id|>world", [15339, 220, 128009, 14957])
self._test_coding(self.llama_tok, "hello <|eot_id|>world again", [15339, 220, 128009, 14957, 1578])
self._test_coding(self.llama_tok, "hello changed <|eot_id|>world again", [15339, 5614, 220, 128009, 14957, 1578])
def test_long_cached_prompt_matches_fresh_tokenization(self):
prefix = "system tools\n" * 700 + "<|eot_id|>"
first, changed = prefix + "run tower of hanoi", prefix + "run ls /"
expected = self.llama_tok.encode(changed)
self.llama_tok.encode(first)
self.assertEqual(self.llama_tok.encode(changed), expected)
def test_tekken_from_gguf_kv(self):
kv = {
"tokenizer.ggml.tokens": ["<unk>", "<s>", "</s>", "[INST]", "[/INST]", "hello"],
@@ -54,10 +89,11 @@ class TestLLMTokenizer(unittest.TestCase):
"tokenizer.ggml.eos_token_id": 2,
}
tok = SimpleTokenizer.from_gguf_kv(kv)
self.assertEqual(tok.role("user"), [3])
template = FallbackTemplate(tok)
self.assertEqual(template.role("user"), "[INST]")
self.assertEqual(tok.encode("hello"), [5])
self.assertEqual(tok.end_turn(), [4])
self.assertEqual(tok.role("assistant"), [])
self.assertEqual(template.end_turn(), "[/INST]")
self.assertEqual(template.role("assistant"), "")
def test_stream_decoder(self):
"""stream_decoder buffers incomplete UTF-8: token 25677 has 3/4 of emoji, token 138 completes it."""
+2 -2
View File
@@ -51,8 +51,8 @@ class TestPatternMatcher(unittest.TestCase):
ctx.append(True)
assert len(x.src) == 0
return x.replace(src=(UOp(Ops.NOOP),))
matcher = PatternMatcher([(UPat(Ops.CONST, src=(), name="x"), fxn)])
c1 = UOp.const(dtypes.float, 1.0)
matcher = PatternMatcher([(UPat(Ops.NOOP, src=(), name="x"), fxn)])
c1 = UOp(Ops.NOOP)
# second rewrite shouldn't match anything
ctx = []
c1 = matcher.rewrite(c1, ctx)
+12
View File
@@ -1472,6 +1472,18 @@ class TestSchedule(unittest.TestCase):
x.softmax().sum().backward()
run_linear(*check_schedule(x.grad, 4))
def test_logsumexp_backward(self):
Tensor.manual_seed(0)
x = Tensor.randn(4, 12, 64, 64).realize()
x.logsumexp(-1).sum().backward()
run_linear(*check_schedule(x.grad, 3))
def test_logcumsumexp_backward(self):
Tensor.manual_seed(0)
x = Tensor.randn(4, 512).realize()
x.logcumsumexp(-1).sum().backward()
run_linear(*check_schedule(x.grad, 3))
def test_scaled_dot_product_attention_fusion(self):
x, y, z, m = (Tensor.empty(32, 8, 16, 16) for _ in range(4))
out = Tensor.scaled_dot_product_attention(x, y, z, attn_mask=m)
+14 -14
View File
@@ -1,6 +1,6 @@
import unittest, itertools
from tinygrad.codegen.late.coalese import indexing_simplify
from tinygrad.codegen.late.coalesce import indexing_simplify
from tinygrad.dtype import dtypes
from tinygrad.uop.ops import UOp, Ops, graph_rewrite
from tinygrad.uop.symbolic import simplify_valid, sym, pm_move_where_on_load
@@ -23,7 +23,7 @@ def get_load_image_uop(image_shape:tuple[int, ...], valid:UOp, idx:tuple[UOp, UO
UOp.param(0, dtypes.float, image_shape).index(idx[1].valid(valid), idx[0].valid(valid)),
))
def Special(expr, nmax): return UOp(Ops.SPECIAL, src=(UOp.const(dtypes.index, nmax),), arg=expr)
def Special(expr, nmax): return UOp(Ops.SPECIAL, src=(UOp.const(dtypes.weakint, nmax),), arg=expr)
def Variable(expr, nmin, nmax): return UOp.variable(expr, nmin, nmax)
def Range(n, nmax): return UOp.range(nmax, n)
@@ -455,7 +455,7 @@ class TestImageSimplification(unittest.TestCase):
A1 = lidx0*32 + r0*32 + lidx1*4 - 99
valid = ((lidx1 < 1).ne(True)) & ((lidx0 + r0) < 3).ne(True) & ((lidx0 + r0) < 19)
alu0 = gidx0 + (A1 % 32)*32 + (A1 // 32 % 16)*1024
load = get_load_image_uop((1, 16384, 4), valid, (alu0, UOp.const(dtypes.index, 0)))
load = get_load_image_uop((1, 16384, 4), valid, (alu0, UOp.const(dtypes.weakint, 0)))
try:
self.check(load, None, "(gidx0+lidx0*1024+r0*1024+lidx1*128+-3168)", "0")
except AssertionError:
@@ -474,7 +474,7 @@ class TestImageSimplification(unittest.TestCase):
A1 = lidx0*16 + r0*16 + lidx1*4 - 51
valid = ((lidx1 < 1).ne(True)) & ((lidx0 + r0) < 3).ne(True) & ((lidx0 + r0) < 11)
alu0 = lidx2 + gidx0*4 + (A1 % 16)*64 + (A1 // 16 % 8)*1024
load = get_load_image_uop((1, 8192, 4), valid, (alu0, UOp.const(dtypes.index, 0)))
load = get_load_image_uop((1, 8192, 4), valid, (alu0, UOp.const(dtypes.weakint, 0)))
try:
self.check(load, None, "(lidx2+gidx0*4+lidx0*1024+r0*1024+lidx1*256+-3264)", "0")
except AssertionError:
@@ -488,18 +488,18 @@ class TestImageSimplification(unittest.TestCase):
gidx0 = Special("gidx0", 1064)
r12 = Range(12, 3)
valid = ((gidx0 < 645).ne(True)) & (gidx0 < 653)
idx = (r12*4 + (gidx0+3)%4 + (gidx0+3)//4*24 - 3888, UOp.const(dtypes.index, 0))
idx = (r12*4 + (gidx0+3)%4 + (gidx0+3)//4*24 - 3888, UOp.const(dtypes.weakint, 0))
load = get_load_image_uop((1, 48, 4), valid, idx)
self.check(load, None, "(r12*4+(gidx0+3)%4+(gidx0+3)//4*24+-3888)", "0")
class TestDropTrueGate(unittest.TestCase):
def test_drop_true_gate_on_index(self):
# test that INDEX with a constant True valid gets simplified to drop the valid
from tinygrad.codegen.late.coalese import indexing_simplify
from tinygrad.codegen.late.coalesce import indexing_simplify
from tinygrad.uop.ops import graph_rewrite
from tinygrad.uop.symbolic import sym
buf = UOp.param(0, dtypes.int, (1,))
idx = UOp.const(dtypes.index, 0)
idx = UOp.const(dtypes.weakint, 0)
true_gate = UOp.const(dtypes.bool, True)
index_with_gate = UOp(Ops.INDEX, src=(buf, idx.valid(true_gate)))
# apply the optimization
@@ -516,7 +516,7 @@ class TestRangeShrink(unittest.TestCase):
def test_range_shrink_single_guard(self):
# range 0..203 guarded by r < 4 everywhere -> shrink to 0..3
r = Range(0, 204)
load = get_gated_load_uop(r < UOp.const(dtypes.index, 4), r)
load = get_gated_load_uop(r < UOp.const(dtypes.weakint, 4), r)
ranges = self.get_ranges(load.sink())
self.assertEqual(len(ranges), 1)
self.assertEqual(ranges[0].src[0].arg, 4)
@@ -524,8 +524,8 @@ class TestRangeShrink(unittest.TestCase):
def test_range_shrink_picks_max_guard(self):
# two loads guard the same range with r < 4 and r < 8 -> shrink to max(4, 8) = 8
r = Range(0, 204)
load1 = get_gated_load_uop(r < UOp.const(dtypes.index, 4), r)
load2 = get_gated_load_uop(r < UOp.const(dtypes.index, 8), r)
load1 = get_gated_load_uop(r < UOp.const(dtypes.weakint, 4), r)
load2 = get_gated_load_uop(r < UOp.const(dtypes.weakint, 8), r)
ranges = self.get_ranges(UOp.sink(load1, load2))
self.assertEqual(len(ranges), 1)
self.assertEqual(ranges[0].src[0].arg, 8)
@@ -533,7 +533,7 @@ class TestRangeShrink(unittest.TestCase):
def test_range_no_shrink_guard_ge_max(self):
# guard r < 300 with range max 204 -> no shrink (guard doesn't constrain)
r = Range(0, 204)
load = get_gated_load_uop(r < UOp.const(dtypes.index, 300), r)
load = get_gated_load_uop(r < UOp.const(dtypes.weakint, 300), r)
ranges = self.get_ranges(load.sink())
self.assertEqual(len(ranges), 1)
self.assertEqual(ranges[0].src[0].arg, 204)
@@ -541,7 +541,7 @@ class TestRangeShrink(unittest.TestCase):
def test_range_no_shrink_when_unguarded_elsewhere(self):
# one load guards r < 4, but another load uses r without a gate -> no shrink
r = Range(0, 204)
load1 = get_gated_load_uop(r < UOp.const(dtypes.index, 4), r)
load1 = get_gated_load_uop(r < UOp.const(dtypes.weakint, 4), r)
load2 = UOp(Ops.LOAD, src=(UOp.param(1, dtypes.float, (204,)).index(r),))
ranges = self.get_ranges(UOp.sink(load1, load2))
self.assertEqual(len(ranges), 1)
@@ -550,7 +550,7 @@ class TestRangeShrink(unittest.TestCase):
def test_range_no_shrink_when_used_in_reduce(self):
# range used in both a gated load AND directly in the reduce expression -> no shrink
r = Range(0, 204)
gated_load = get_gated_load_uop(r < UOp.const(dtypes.index, 4), r)
gated_load = get_gated_load_uop(r < UOp.const(dtypes.weakint, 4), r)
red = (r.cast(dtypes.float) + gated_load).reduce(r, arg=Ops.ADD)
ranges = self.get_ranges(red.sink())
self.assertEqual(len(ranges), 1)
@@ -559,7 +559,7 @@ class TestRangeShrink(unittest.TestCase):
def test_range_shrink_to_single_iteration(self):
# guard r < 1 shrinks range to 1 -> single iteration, range eliminated entirely
r = Range(0, 204)
load = get_gated_load_uop(r < UOp.const(dtypes.index, 1), r)
load = get_gated_load_uop(r < UOp.const(dtypes.weakint, 1), r)
ranges = self.get_ranges(load.sink())
self.assertEqual(len(ranges), 0)
+1 -1
View File
@@ -382,7 +382,7 @@ class TestTensorUOpStack(unittest.TestCase):
self.assertIs(_t(2, 3).uop.stack(w.uop).dtype, dtypes.float32)
def test_stack_index_dtype(self):
# index is outside the promotion lattice, equal dtypes bypass promotion
self.assertEqual(UOp.const(dtypes.index, 1).stack(UOp.const(dtypes.index, 2)).shape, (2,))
self.assertEqual(UOp.const(dtypes.weakint, 1).stack(UOp.const(dtypes.weakint, 2)).shape, (2,))
class TestTensorUOpConv2d(unittest.TestCase):
def test_conv2d_basic(self):
+25 -74
View File
@@ -2,7 +2,7 @@ import unittest, pytest
from tinygrad import dtypes, Variable
from tinygrad.dtype import AddrSpace
from tinygrad.helpers import DEBUG, Context
from tinygrad.uop.ops import Ops, UOp, UPat, PatternMatcher, graph_rewrite, GroupOp, AxisType
from tinygrad.uop.ops import Ops, UOp, UPat, PatternMatcher, graph_rewrite, GroupOp, AxisType, broadcast_axes
from tinygrad.uop.symbolic import sym
from test.helpers import to_uops_list
@@ -202,7 +202,7 @@ class TestUOpGraph(unittest.TestCase):
def test_where_same_fold(self):
v = UOp.variable('tmp', 0, 1)
c0 = UOp.const(dtypes.index, 0)
c0 = UOp.const(dtypes.weakint, 0)
vc = v != c0
c1 = UOp.const(dtypes.float, 1.0)
out = vc.where(c1, c1)
@@ -309,66 +309,6 @@ class TestUOpGraph(unittest.TestCase):
for uop, const in zip(uops, consts):
self.assertEqual(uop, const)
@unittest.skip("no longer testable standalone")
def test_wmma_vectorize_fold(self):
for i in [2, 4, 8]:
vec = UOp(Ops.STACK, dtypes.half, tuple(UOp.const(dtypes.half, 0.0) for _ in range(i)))
var = UOp.variable("var", 0, 1, dtypes.half)
acc = UOp.variable('acc', 0, 1, dtypes.half)
wmma = UOp(Ops.WMMA, src=(vec, var, acc))
uops = to_uops_list([wmma])
self.assertEqual(uops[0], acc)
self.assertEqual(len(uops), 2) # +1 for SINK
for i in [2, 4, 8]:
var = UOp.variable("var", 0, 1, dtypes.half)
vec = UOp(Ops.STACK, dtypes.half, tuple(UOp.const(dtypes.half, 0.0) for _ in range(i)))
acc = UOp.variable('acc', 0, 1, dtypes.half)
wmma = UOp(Ops.WMMA, src=(var, vec, acc))
uops = to_uops_list([wmma])
self.assertEqual(uops[0], acc)
self.assertEqual(len(uops), 2) # +1 for SINK
@unittest.skip("wmma is wrong here, it needs an arg")
def test_wmma_vectorize_no_fold(self):
for i in [4, 8]:
vec = UOp(Ops.STACK, dtypes.half,
tuple(UOp.const(dtypes.half, 0.0) for _ in range(i//2)) +
tuple(UOp.variable(f'tmp{j}', 0, 1, dtypes.half) for j in range(i//2)))
var = UOp.variable(f'tmp{i}', 0, 1, dtypes.half)
acc = UOp.variable('acc', 0, 1, dtypes.half)
wmma = UOp(Ops.WMMA, src=(vec, var, acc))
uops = to_uops_list([wmma])
self.assertEqual(uops[-2], wmma) # -2 to skip SINK
for i in [4, 8]:
var = UOp.variable(f'tmp{i}', 0, 1, dtypes.half)
vec = UOp(Ops.STACK, dtypes.half,
tuple(UOp.const(dtypes.half, 0.0) for _ in range(i//2)) +
tuple(UOp.variable(f'tmp{j}', 0, 1, dtypes.half) for j in range(i//2)))
acc = UOp.variable('acc', 0, 1, dtypes.half)
wmma = UOp(Ops.WMMA, src=(var, vec, acc))
uops = to_uops_list([wmma])
self.assertEqual(uops[-2], wmma) # -2 to skip SINK
for i in [2, 4, 8]:
vec = UOp(Ops.STACK, dtypes.half,
tuple(UOp.const(dtypes.half, 1.0 if j == 0 else 0.0) for j in range(i)))
var = UOp.variable(f'tmp{i}', 0, 1, dtypes.half)
acc = UOp.variable('acc', 0, 1, dtypes.half)
wmma = UOp(Ops.WMMA, src=(vec, var, acc))
uops = to_uops_list([wmma])
self.assertEqual(uops[-2], wmma) # -2 to skip SINK
for i in [2, 4, 8]:
var = UOp.variable(f'tmp{i}', 0, 1, dtypes.half)
vec = UOp(Ops.STACK, dtypes.half,
tuple(UOp.const(dtypes.half, 1.0 if j == 0 else 0.0) for j in range(i)))
acc = UOp.variable('acc', 0, 1, dtypes.half)
wmma = UOp(Ops.WMMA, src=(var, vec, acc))
uops = to_uops_list([wmma])
self.assertEqual(uops[-2], wmma) # -2 to skip SINK
def test_cast_alu_fold(self):
d0 = UOp.param(0, dtypes.bool, (1,))
d1 = UOp.param(1, dtypes.int, (1,))
@@ -484,16 +424,16 @@ class TestUOpGraph(unittest.TestCase):
# mnist indexing with split reduceop
# Make sure we are not doign math on the loaded index, which would promote it to long
c0 = UOp.param(0, dtypes.uchar, (128000,))
c1 = UOp.range(UOp.const(dtypes.index, 512), 1, AxisType.LOOP)
c2 = UOp.range(UOp.const(dtypes.index, 250), 2, AxisType.LOOP)
c1 = UOp.range(UOp.const(dtypes.weakint, 512), 1, AxisType.LOOP)
c2 = UOp.range(UOp.const(dtypes.weakint, 250), 2, AxisType.LOOP)
c3 = UOp.param(1, dtypes.int, (512,))
c4 = c3.index(c1)
c5 = UOp.range(UOp.const(dtypes.index, 240), 0, AxisType.REDUCE)
c6 = ((c2*UOp.const(dtypes.index, 240))+c5)
c5 = UOp.range(UOp.const(dtypes.weakint, 240), 0, AxisType.REDUCE)
c6 = ((c2*UOp.const(dtypes.weakint, 240))+c5)
c7 = UOp.param(2, dtypes.uchar, (60000,))
c8 = c7.index(c6)
c9 = ((c4<0).where((c4+60000), c4)!=c6.cast(dtypes.int)).where(0, c8.cast(dtypes.uint).cast(dtypes.uchar)).reduce(c5, arg=Ops.ADD)
c10 = c0.index(((c1*UOp.const(dtypes.index, 250))+c2)).store(c9).end(c1, c2)
c10 = c0.index(((c1*UOp.const(dtypes.weakint, 250))+c2)).store(c9).end(c1, c2)
uops = to_uops_list([c10])
for u in uops:
self.assertNotEqual(u.dtype, dtypes.long)
@@ -501,19 +441,19 @@ class TestUOpGraph(unittest.TestCase):
def test_load_idx_no_math_on_loaded(self):
# test the (x+y)<c pattern where x has loads - we shouldn't do math on loaded indices
c0 = UOp.param(0, dtypes.uchar, (128000,))
c1 = UOp.range(UOp.const(dtypes.index, 512), 1, AxisType.LOOP)
c2 = UOp.range(UOp.const(dtypes.index, 250), 2, AxisType.LOOP)
c1 = UOp.range(UOp.const(dtypes.weakint, 512), 1, AxisType.LOOP)
c2 = UOp.range(UOp.const(dtypes.weakint, 250), 2, AxisType.LOOP)
c3 = UOp.param(1, dtypes.int, (512,))
c4 = c3.index(c1) # c4 is a load
c5 = UOp.range(UOp.const(dtypes.index, 240), 0, AxisType.REDUCE)
c6 = ((c2*UOp.const(dtypes.index, 240))+c5)
c5 = UOp.range(UOp.const(dtypes.weakint, 240), 0, AxisType.REDUCE)
c6 = ((c2*UOp.const(dtypes.weakint, 240))+c5)
c7 = UOp.param(2, dtypes.uchar, (60000,))
c8 = c7.index(c6)
# (loaded + range) < const pattern - loaded value shouldn't be promoted to long
loaded_idx = c4.cast(dtypes.index)
comparison = (loaded_idx + c5) < UOp.const(dtypes.index, 60000)
loaded_idx = c4.cast(dtypes.weakint)
comparison = (loaded_idx + c5) < UOp.const(dtypes.weakint, 60000)
c9 = comparison.where(c8.cast(dtypes.uint).cast(dtypes.uchar), 0).reduce(c5, arg=Ops.ADD)
c10 = c0.index(((c1*UOp.const(dtypes.index, 250))+c2)).store(c9).end(c1, c2)
c10 = c0.index(((c1*UOp.const(dtypes.weakint, 250))+c2)).store(c9).end(c1, c2)
uops = to_uops_list([c10])
for u in uops:
self.assertNotEqual(u.dtype, dtypes.long)
@@ -767,5 +707,16 @@ class TestUOpBroadcast(unittest.TestCase):
c = a + b
self.assertEqual(c.op, Ops.ADD)
def test_broadcast_axes(self):
t = Variable("t", 1, 10)
self.assertEqual(broadcast_axes((4, 8), (4, 8)), ())
self.assertEqual(broadcast_axes((8,), (4, 8)), (0,))
self.assertEqual(broadcast_axes((), (4, 8)), (0, 1))
self.assertEqual(broadcast_axes((3, 1), (4, 3, 8)), (0, 2))
self.assertEqual(broadcast_axes((1, 8), (1, 8)), ())
self.assertEqual(broadcast_axes((t, 8), (t, 8)), ())
self.assertEqual(broadcast_axes((1, 8), (t, 8)), (0,))
with self.assertRaises(RuntimeError): broadcast_axes((4, 8), (8,))
if __name__ == '__main__':
unittest.main(verbosity=2)
+12 -12
View File
@@ -11,13 +11,13 @@ from tinygrad.uop.validate import uops_to_z3
def check_uop_against_string(self, v:UOp, s:str):
sym_vars = {v.render():v for v in v.toposort() if v.op in (Ops.RANGE, Ops.SPECIAL, Ops.PARAM)}
s_eval = eval(s, sym_vars)
if isinstance(s_eval, int) and v.dtype==dtypes.index: s_eval = UOp.const(dtypes.index, s_eval)
if isinstance(s_eval, int) and v.dtype==dtypes.weakint: s_eval = UOp.const(dtypes.weakint, s_eval)
elif isinstance(s_eval, (bool, int, float)): s_eval = UOp.const(dtypes.from_py(s_eval), s_eval)
s_eval = graph_rewrite(s_eval, commutative, name="cannonicalize eval")
self.assertIs(s_eval, v, f"eval did not match simplified: {s_eval} != {v.render()} for {s}")
def Variable(name: str, min_val: ConstType, max_val: ConstType, dtype: DType=dtypes.index): return UOp.variable(name,min_val,max_val,dtype)
def uconst(val): return UOp.const(dtypes.index, val)
def Variable(name: str, min_val: ConstType, max_val: ConstType, dtype: DType=dtypes.weakint): return UOp.variable(name,min_val,max_val,dtype)
def uconst(val): return UOp.const(dtypes.weakint, val)
def usum(ops): return functools.reduce(lambda x,y: x+y, ops)
def uand(ops): return functools.reduce(lambda x,y: x*y, ops)
@@ -247,12 +247,12 @@ class TestSymbolic(unittest.TestCase):
self.assertEqual((Variable("x", -10, 0)%Variable("y", 1, 10))._min_max, (0, 9))
def test_range_div_its_symbolic_bound(self):
a = Variable("a", 1, 10, dtypes.index)
a = Variable("a", 1, 10, dtypes.weakint)
ridx0 = UOp.range(a+2, 0)
self.helper_test_variable(ridx0//(a+2), 0, 0, "0")
def test_range_mod_its_symbolic_bound(self):
a = Variable("a", 1, 10, dtypes.index)
a = Variable("a", 1, 10, dtypes.weakint)
ridx = UOp.range(a+2, 0)
self.helper_test_variable(ridx%(a+2), 0, 11, "r0")
@@ -919,8 +919,8 @@ class TestSymbolic(unittest.TestCase):
self.helper_test_variable(cond.cast(dtypes.int).ne(2), 1, 1, "True")
self.helper_test_variable(cond.cast(dtypes.int).ne(-1), 1, 1, "True")
# CAST(bool -> index) folds too
self.helper_test_variable(cond.cast(dtypes.index).ne(0), 0, 1, "(a<2)")
self.helper_test_variable(cond.cast(dtypes.index).ne(1), 0, 1, "((a<2)!=True)")
self.helper_test_variable(cond.cast(dtypes.weakint).ne(0), 0, 1, "(a<2)")
self.helper_test_variable(cond.cast(dtypes.weakint).ne(1), 0, 1, "((a<2)!=True)")
def test_where_removal(self):
cond = Variable("a", 0, 3) < 2
@@ -1021,7 +1021,7 @@ class TestSymbolic(unittest.TestCase):
self.helper_test_variable((numerator//denominator)<=0, 1, 1, "True")
def test_symbolic_range_doesnt_collapse(self):
r0 = UOp.range((Variable("a", 1, 10)<5).cast(dtypes.index), 0)
r0 = UOp.range((Variable("a", 1, 10)<5).cast(dtypes.weakint), 0)
self.helper_test_variable(r0, 0, 0, "r0")
def test_const_reciprocal(self):
@@ -1289,16 +1289,16 @@ class TestInvalidIndex(unittest.TestCase):
self.assertIs((UOp.invalid()<Variable("a",0,10)).simplify().dtype, dtypes.bool)
def test_alu_invalid_vconst(self):
c1 = UOp.const(dtypes.index, (1, 1, Invalid, Invalid))
c2 = UOp.const(dtypes.index, (1, Invalid, 1, 1))
self.assertIs((c1+c2).simplify(), UOp.const(dtypes.index, (2, Invalid, Invalid, Invalid)))
c1 = UOp.const(dtypes.weakint, (1, 1, Invalid, Invalid))
c2 = UOp.const(dtypes.weakint, (1, Invalid, 1, 1))
self.assertIs((c1+c2).simplify(), UOp.const(dtypes.weakint, (2, Invalid, Invalid, Invalid)))
class TestStoreLoadFolding(unittest.TestCase):
"""Tests for store(index, load(index)) -> NOOP rule. This rule matches patterns that EMERGE during simplification."""
def test_store_load_folding(self):
# store(idx, load(idx)) -> NOOP, including emergent patterns like store(idx, load(idx) + 0)
buf = UOp.param(0, dtypes.int, (1,))
index = buf.index(UOp.const(dtypes.index, 0))
index = buf.index(UOp.const(dtypes.weakint, 0))
# Direct: store(idx, load(idx)) -> NOOP
self.assertEqual(graph_rewrite(index.store(index.load()), sym).op, Ops.NOOP)
# Emergent: store(idx, load(idx) + 0) -> store(idx, load(idx)) -> NOOP

Some files were not shown because too many files have changed in this diff Show More