Compare commits

...
87 Commits
Author SHA1 Message Date
geohot 9d8a2bd212 toposort recursive_property is faster 2025-11-24 22:09:32 -08:00
George HotzandGitHub 8e8fec408e fix n^2 _apply_map_to_tensors [pr] (#13443)
* clean up slow rules

* fix rule

* non n^2 toposort

* topovisit

* state dict profile_marker
2025-11-24 18:59:16 -08:00
wozeparrotandGitHub 249553a119 tinyfs tweaks (#13444) 2025-11-24 18:07:32 -08:00
wozeparrotandGitHub f46bc31156 tk: start and step in range (#13442) 2025-11-24 15:43:24 -08:00
George HotzandGitHub cc5e6323ac stable diffusion profiling (#13441)
* stable diffusion profiling

Signed-off-by: George Hotz <[email protected]>

* profile_marker

* profile per step

* fix slow Context

* profile that

---------

Signed-off-by: George Hotz <[email protected]>
2025-11-24 15:25:45 -08:00
nimlgenandGitHub 18cfb54736 amd: a bit better se limiting (#13440)
* amd: a bit better se limiting

* SQTT_LIMIT_SE=0
2025-11-24 21:51:47 +03:00
C TandGitHub 2d53029be3 Whisper less flaky tests (#13435)
* use less flaky metric for whisper long transcription

* multiline long transcription 3 reference

* fix reference transcript

see https://homepage.ntu.edu.tw/~karchung/miniconversations/MC.htm
sanitized for whisper

* try lower wer threshold

* add test for wer metric

* extract TRANSCRIPTION_3_ALT

* rename test

* rename

* add tests for high WER difference

* move tests

* sync metric
2025-11-24 09:50:49 -08:00
qazalandGitHub 2a9bd12700 sqtt: add occupancy events to the timeline (#13430) 2025-11-24 22:28:05 +08:00
Sieds LyklesandGitHub 63a931ff76 Symbolic divisor fuzzer (#13433)
* render z3 range better

* working version

* rename

* add to workflow

* factor out variable_names

* smaller expressions

* smaller

* + back
2025-11-23 20:29:32 +01:00
nimlgenandGitHub 677db34eba nv: cleanup map flags (#13434) 2025-11-23 19:54:52 +03:00
qazalandGitHub 712c7a6448 sqtt loader cleanups from the occupancy branch (#13431)
* cleanup err handling

* from disasms

* s/wave_execs/wave_insts
2025-11-23 21:50:34 +08:00
George HotzandGitHub 9d7a17ee39 beautiful SQTT_PARSE=1 with color (#13428)
* beautiful SQTT_PARSE=1 with color

* linter

* linter 2

* a few more labels

* filter and or

* wave alloc

* a few more
2025-11-23 01:05:14 -08:00
qazalandGitHub 474a631877 viz: align left offset for nested items (#13420) 2025-11-23 14:22:51 +08:00
geohot da0aa57a3b add cu parsing to attempt_sqtt_parse 2025-11-22 22:09:05 -08:00
qazalandGitHub 320ed78803 can view wave timeline with SQTT_ITRACE_SE_MASK=0 (#13427) 2025-11-23 13:55:47 +08:00
PranilandGitHub c1838c71fc display service name typo (#13426)
its tinybox-display.service
2025-11-22 20:49:56 -08:00
George HotzandGitHub 5110409339 continue work on parse sqtt, enable with SQTT_PARSE (#13425)
* continue work on parse sqtt, enable with SQTT_PARSE

* fix timing

* delta is pre instruction

* hi8 values

* a few more

* a bit more

* let it crash if you enabled it

* figure out simd

* hide 0x11
2025-11-22 19:03:17 -08:00
George HotzandGitHub 92170d0ff1 lil op cleanup (#13424)
* track flag count and op count

* text

* more

* file count

* lil op cleanup

* cleanups

* move
2025-11-22 15:21:15 -08:00
George HotzandGitHub 423b76a852 improve sqtt format parser (saturday coffee shop project) (#13419)
* improve sqtt format parser

* actually read the trash code ChatGPT wrote

* cleanups

* hand written parser

* quality

* more

* was missing first packet

* maybe

* filt

* fixups

* label the waves

* progress
2025-11-22 15:04:10 -08:00
geohot 9d6cf3472e remove op/sentinel 2025-11-22 15:01:47 -08:00
sirhcmandGitHub 310da2a201 remove hashFiles in setup-tinygrad (#13423)
* fix hashFiles in setup-tinygrad on macos

* remove hashFiles altogether
2025-11-22 17:47:10 -05:00
qazalandGitHub c14033e10f viz: faster startup time with SQTT=1 (#13337)
* roc.py cleanups

* direct append

* viz index cleanup

* simd row details

* add kernel arg

* late instructions decode

* more instruction decode to sep server request

* 200ms startup, 6 second to waves timeline

* sort units

* creating new http paths is easy now

* instructions unpacker

* min diff, use hyphens

* summary table
2025-11-22 22:02:30 +08:00
qazalandGitHub 1655fdb6de viz: cleanup sqtt loader (#13417) 2025-11-22 20:10:23 +08:00
qazalandGitHub 903eec3754 fix sz.py tinygrad import in ci (#13418) 2025-11-22 19:20:26 +08:00
nimlgenandGitHub 3a42680e22 amd: pmc generic arch for gfx10+ (#13407) 2025-11-22 12:31:23 +03:00
George HotzandGitHub 1f8b24a6b9 track flag count and op count (#13416)
* track flag count and op count

* text

* more

* file count
2025-11-21 22:46:33 -08:00
George HotzandGitHub 4c0f4226b9 delete the PRECAST op [p] (#13415)
* don't use PRECAST in cstyle renderer [p]

* fix in metal

* fix opencl

* __builtin_bit_cast

* precast is unused

* cuda is c99?

* lambda_union_bitcast

* helper function

* delete precast op
2025-11-21 21:47:14 -08:00
wozeparrotandGitHub 1f648bb1ba feat: reenable mobilenetv2 dsp (#13320) 2025-11-21 15:21:49 -08:00
chenyuandGitHub 054477a44f remove full_symbolic in simplify (#13413)
only flip one schedule in winograd backward, no functional difference
2025-11-21 15:04:00 -05:00
chenyuandGitHub cb29265f23 add test that shows the validhack regression with bad rewrite order (#13411) 2025-11-21 13:48:30 -05:00
qazalandGitHub fdfe83880b viz: unique sqtt wave names (#13410)
* viz: unique sqtt wave names

* better name for the shape

* it's a per program counter now

* table view, refactor to wave:insts dict
2025-11-22 02:43:31 +08:00
chenyuandGitHub a6c9b4ff6a fix symbolic comments [pr] (#13408) 2025-11-21 09:18:50 -05:00
Sieds LyklesandGitHub 114bb94c55 Fix load collapse MAX to ADD (#13406)
* add Ops.ADD to pattern

* add test
2025-11-21 12:26:14 +01:00
qazalandGitHub 87c248eafa small cleanups from viz memory usage fixes (#13405)
* shape link cleanups

* cleanup findRectAtPosition
2025-11-21 17:05:08 +08:00
qazalandGitHub 0de1b24154 viz: SE : CU : SIMD : WAVE in sqtt timeline (#13404)
* wave id in device rows

* SE : CU : SIMD : WAVE

* automatic width

* better styling

* rm the blue

* sort
2025-11-21 15:42:29 +08:00
George HotzandGitHub dabb02767f set AMD profile mode with sudo on SQTT or PMC (#13403)
* require profile mode

* add mode setter

* cleanup

* not needed

* SQTT_LIMIT_SE
2025-11-20 23:19:11 -08:00
George HotzandGitHub e1051d00d7 multi like on full_like as well as rand_like (#13402)
* multi like on full_like as well as rand_like

* add test and fix bug

* mismatch, optim match

* one line
2025-11-20 20:46:48 -08:00
chenyuandGitHub fa3def2f12 call less simplify in simplify_valid_load [pr] (#13401) 2025-11-20 19:54:22 -05:00
qazalandGitHub 895ec7417e viz: enable mapping function names to colors (#13400) 2025-11-21 06:43:02 +08:00
George HotzandGitHub a74f6020d5 track apply map to tensors (#13399)
* track apply map to tensors

* sub
2025-11-20 14:24:55 -08:00
chenyuandGitHub 647fde64e6 no sym in pm_reduce [pr] (#13398)
* no sym in pm_reduce [pr]

* fix that
2025-11-20 16:49:09 -05:00
qazalandGitHub 1313250e0d viz: use system helper for llvm-mca (#13395) 2025-11-21 04:47:25 +08:00
sirhcmandGitHub de3593957f Revert "Revert "autogen: fix formatting on zero-argument function-like macros…" (#13388)
This reverts commit 0901a40685.
2025-11-20 15:36:13 -05:00
qazalandGitHub 1220072328 viz: refactor to generic steps api (#13393) 2025-11-21 04:33:23 +08:00
George HotzandGitHub 26ccbf7040 debufferize with symbolic in one pm (#13392) 2025-11-20 11:47:03 -08:00
George HotzandGitHub c46f608703 top down remove_bufferize (#13391)
* top down remove_bufferize

* removable if ALWAYS_CONTIGUOUS
2025-11-20 11:32:00 -08:00
sirhcmandGitHub 4043489803 set curl -f in setup-tinygrad (#13389)
* set curl -f in setup-tinygrad

* test bad redirect

* Revert "test bad redirect"

This reverts commit ad945e7ffc.
2025-11-20 13:45:47 -05:00
chenyuandGitHub 0251a8e628 parse_valid minor cleanup [pr] (#13385)
* stricter parse_valid [pr]

* not stricter

* no VCONST

* Revert "no VCONST"

This reverts commit 330dbdf4060562596febcbf970bda6051a35012f.
2025-11-20 13:15:06 -05:00
sirhcmandGitHub 0901a40685 Revert "autogen: fix formatting on zero-argument function-like macros (#13386)" (#13387)
This reverts commit 58d85d4bab.
2025-11-20 12:45:35 -05:00
91e289cb14 amd fp8 llvm (#13186)
* amd fp8 llvm support

* fix max

* clean

* add test_mi350.sh

---------

Co-authored-by: chenyu <[email protected]>
2025-11-20 12:35:57 -05:00
Roelof van DijkandGitHub 1058748440 torch backend: no aten.detach for torch 2.10 compat (#13381)
* this works, less cpp?

* simpler = better

* keep torch 2.9 working as well
2025-11-20 09:12:15 -08:00
sirhcmandGitHub 58d85d4bab autogen: fix formatting on zero-argument function-like macros (#13386)
* fix formatting on zero-argument function-like macros

* autogen tests should run

* ugh
2025-11-20 12:11:04 -05:00
qazalandGitHub 9dbc550692 roc: map disassembly to prog name (#13384) 2025-11-20 23:47:19 +08:00
qazalandGitHub ebcdf68bab viz: use content headers for profiler (#13383) 2025-11-20 23:33:16 +08:00
nimlgenandGitHub 0b0ea4981c hcq: unwrap signals (#13382) 2025-11-20 18:12:41 +03:00
qazalandGitHub 9dcd52287a add external_benchmark_pyrender (#13378)
* add external_benchmark_pyrender

* can ctrlc it

* cpu_profile exists
2025-11-20 17:38:28 +08:00
geohot cb38c704c3 delete nonfunctional ramp.py 2025-11-19 20:43:44 -08:00
George HotzandGitHub 8919c994b7 Revert "AxisType.PLACEHOLDER in reshape to do less graph_rewrite (#13373)" (#13375)
This reverts commit ac7559e33d.
2025-11-19 19:34:30 -08:00
George HotzandGitHub ac7559e33d AxisType.PLACEHOLDER in reshape to do less graph_rewrite (#13373)
* AxisType.PLACEHOLDER in reshape to do less graph_rewrite

* _apply_movement_op cache
2025-11-19 19:19:58 -08:00
chenyuandGitHub 050682ab40 use invalid_gate consistently [pr] (#13374) 2025-11-19 22:15:12 -05:00
0dc2ff431d fix: revive torch backend (#13280)
* fix: revive torch backend

* as_strided view vs copy

* Revert "as_strided view vs copy"

This reverts commit 82a61223f2.

* add extra tests (move inplace, add fusion tests)

* better fusion with inplace_op

* no optimizer hooks (break mnist training fusion)

* split off fusion tests in separate file, assert on resnet fusion

fix: remove comments

* cleanup, reduce diff

* reduce diff

* better fusion and identity checks

---------

Co-authored-by: George Hotz <[email protected]>
2025-11-19 15:26:50 -08:00
wozeparrotandGitHub 56b2540349 tk: keep extra tile data by replacing uop (#13370) 2025-11-19 15:11:43 -08:00
George HotzandGitHub ab7df42c78 bring back fold_divmod_general with bugfix and test [pr] (#13369)
* Revert "Revert "merge to fold_divmod_general [p] (#13359)""

This reverts commit 05ccc69248.

* Revert "Revert "actually merge to fold_divmod_general [pr] (#13363)""

This reverts commit 90e5752199.

* Revert "Revert "add cache to fold_divmod_general (#13365)""

This reverts commit 8e17bd6791.

* bring back fold_divmod_general with bugfix and test
2025-11-19 14:51:51 -08:00
George HotzandGitHub 986d113024 symbolic fuzz failure (#13367)
* symbolic fuzz failure

* skip flaky test
2025-11-19 14:21:08 -08:00
geohot 05ccc69248 Revert "merge to fold_divmod_general [p] (#13359)"
This reverts commit 7711bbac7f.
2025-11-19 14:18:09 -08:00
geohot 90e5752199 Revert "actually merge to fold_divmod_general [pr] (#13363)"
This reverts commit 3d82b83cec.
2025-11-19 14:18:08 -08:00
geohot 8e17bd6791 Revert "add cache to fold_divmod_general (#13365)"
This reverts commit b5309a5043.
2025-11-19 14:18:08 -08:00
George HotzandGitHub b5309a5043 add cache to fold_divmod_general (#13365) 2025-11-19 13:49:18 -08:00
George HotzandGitHub 3d82b83cec actually merge to fold_divmod_general [pr] (#13363)
* actually merge to fold_divmod_general [pr]

* one more merge

* Revert "one more merge"

This reverts commit aa79f6781c.

* avoid that case for speed

* faster and simpler
2025-11-19 13:17:56 -08:00
chenyuandGitHub a91f00925b remove VECTORIZE and WMMA rules from sym [pr] (#13362) 2025-11-19 14:51:21 -05:00
George HotzandGitHub 7711bbac7f merge to fold_divmod_general [p] (#13359)
* merge to fold_divmod_general [p]

* merge more

* merge more

* merge more
2025-11-19 11:37:45 -08:00
George HotzandGitHub 6fdbd03104 more divmod cleanup [p] (#13358)
* more divmod cleanup [p]

* lil cleanups, faster
2025-11-19 10:35:15 -08:00
George HotzandGitHub bd88a72149 div and mod to its own file, try 2 [p] (#13357) 2025-11-19 10:10:06 -08:00
George HotzandGitHub 957cf717e7 Python speed (#13355)
* skip process replay by default

* work on python speed

* fix names of rewrite rules

* fix that test
2025-11-19 09:03:00 -08:00
chenyuandGitHub fc19ea76b5 clean up threefry rules (#13354) 2025-11-19 11:48:07 -05:00
George HotzandGitHub 385618d45b skip process replay by default (#13353) 2025-11-19 08:25:34 -08:00
chenyuandGitHub fba4535289 remove hacks for threefry long removal when padded [pr] (#13352) 2025-11-19 11:11:39 -05:00
George HotzandGitHub 225eb1500f generic range changes that work for str + int (#13350)
* generic range changes that work for str + int

* opt range counts up
2025-11-19 08:07:49 -08:00
chenyuandGitHub 1a72ac16a6 move where same false branch rule to symbolic_simple [pr] (#13349) 2025-11-19 10:15:38 -05:00
chenyuandGitHub 79055ddb8b clean propagate_invalid more [pr] (#13347) 2025-11-19 09:47:50 -05:00
nimlgenandGitHub 0c9fbf87e1 nvioctl: classes (#13346) 2025-11-19 16:14:15 +03:00
qazalandGitHub f2221130bb viz: pick shape by event type (#13279) 2025-11-19 20:15:52 +08:00
wozeparrotandGitHub be72b78dcb tk: small fixes (#13345)
* fix: handle case where final uop isn't a tk wrapped one

* clean: remove after from mma
2025-11-19 00:58:50 -08:00
wozeparrotandGitHub e4fbde5b3b fix: extra options need to go on second step too (#13344) 2025-11-19 00:58:09 -08:00
George HotzandGitHub 1a332afa76 spec test on 3.14 (#12957) 2025-11-19 00:43:04 -08:00
sirhcmandGitHub a438c277de autogen tests for 3.14 (#13343) 2025-11-18 22:16:59 -05:00
chenyuandGitHub 722e7a16ed remove rule in propagate_invalid [pr] (#13342) 2025-11-18 21:38:33 -05:00
70 changed files with 2512 additions and 1476 deletions
+4 -4
View File
@@ -61,7 +61,7 @@ runs:
uses: actions/cache@v4 uses: actions/cache@v4
with: with:
path: ${{ github.workspace }}/.venv path: ${{ github.workspace }}/.venv
key: venv-${{ runner.os }}-python-${{ steps.setup-python.outputs.python-version }}-${{ inputs.deps }}-${{ inputs.pydeps }}-${{ hashFiles('**/pyproject.toml') }}-${{ env.CACHE_VERSION }} key: venv-${{ runner.os }}-python-${{ steps.setup-python.outputs.python-version }}-${{ inputs.deps }}-${{ inputs.pydeps }}-${{ env.CACHE_VERSION }}
# **** Caching downloads **** # **** Caching downloads ****
@@ -221,7 +221,7 @@ runs:
sudo mkdir -p /usr/local/lib sudo mkdir -p /usr/local/lib
curl -s -H "Authorization: token $GH_TOKEN" curl -s https://api.github.com/repos/nimlgen/amdcomgr_dylib/releases/latest | \ curl -s -H "Authorization: token $GH_TOKEN" curl -s https://api.github.com/repos/nimlgen/amdcomgr_dylib/releases/latest | \
jq -r '.assets[] | select(.name == "libamd_comgr.dylib").browser_download_url' | \ jq -r '.assets[] | select(.name == "libamd_comgr.dylib").browser_download_url' | \
sudo xargs curl -L -o /usr/local/lib/libamd_comgr.dylib sudo xargs curl -fL -o /usr/local/lib/libamd_comgr.dylib
cargo build --release --manifest-path ./extra/remu/Cargo.toml cargo build --release --manifest-path ./extra/remu/Cargo.toml
# **** gpuocelot **** # **** gpuocelot ****
@@ -278,7 +278,7 @@ runs:
if: inputs.webgpu == 'true' && runner.os == 'Linux' if: inputs.webgpu == 'true' && runner.os == 'Linux'
shell: bash shell: bash
run: | run: |
sudo curl -L https://github.com/wpmed92/pydawn/releases/download/v0.1.6/libwebgpu_dawn.so -o /usr/local/lib/libwebgpu_dawn.so sudo curl -fL https://github.com/wpmed92/pydawn/releases/download/v0.1.6/libwebgpu_dawn.so -o /usr/local/lib/libwebgpu_dawn.so
sudo ldconfig sudo ldconfig
- name: Install WebGPU dawn (macOS) - name: Install WebGPU dawn (macOS)
if: inputs.webgpu == 'true' && runner.os == 'macOS' if: inputs.webgpu == 'true' && runner.os == 'macOS'
@@ -298,7 +298,7 @@ runs:
- name: Install mesa (linux) - name: Install mesa (linux)
if: inputs.mesa == 'true' && runner.os == 'Linux' if: inputs.mesa == 'true' && runner.os == 'Linux'
shell: bash shell: bash
run: sudo curl -L https://github.com/sirhcm/tinymesa/releases/download/tinymesa-32dc66c/libtinymesa_cpu-mesa-25.2.4-linux-amd64.so -o /usr/lib/libtinymesa_cpu.so run: sudo curl -fL https://github.com/sirhcm/tinymesa/releases/download/tinymesa-32dc66c/libtinymesa_cpu-mesa-25.2.4-linux-amd64.so -o /usr/lib/libtinymesa_cpu.so
- name: Install mesa (macOS) - name: Install mesa (macOS)
if: inputs.mesa == 'true' && runner.os == 'macOS' if: inputs.mesa == 'true' && runner.os == 'macOS'
shell: bash shell: bash
+2
View File
@@ -13,9 +13,11 @@ on:
pull_request: pull_request:
paths: paths:
- 'tinygrad/runtime/autogen/**/*' - 'tinygrad/runtime/autogen/**/*'
- 'tinygrad/runtime/support/autogen.py'
workflow_dispatch: workflow_dispatch:
paths: paths:
- 'tinygrad/runtime/autogen/**/*' - 'tinygrad/runtime/autogen/**/*'
- 'tinygrad/runtime/support/autogen.py'
jobs: jobs:
autogen: autogen:
+8 -8
View File
@@ -643,14 +643,14 @@ jobs:
run: BENCHMARK_LOG=openpilot_0_10_1_policy PYTHONPATH="." ASSERT_MIN_STEP_TIME=4 DEV=QCOM FLOAT16=1 IMAGE=2 NOLOCALS=1 taskset -c 4-7 python3 examples/openpilot/compile3.py https://github.com/commaai/openpilot/raw/720392c9a5b986981fdbed1bb8c47a6c5573a50e/selfdrive/modeld/models/driving_policy.onnx run: BENCHMARK_LOG=openpilot_0_10_1_policy PYTHONPATH="." ASSERT_MIN_STEP_TIME=4 DEV=QCOM FLOAT16=1 IMAGE=2 NOLOCALS=1 taskset -c 4-7 python3 examples/openpilot/compile3.py https://github.com/commaai/openpilot/raw/720392c9a5b986981fdbed1bb8c47a6c5573a50e/selfdrive/modeld/models/driving_policy.onnx
- name: openpilot compile3 0.10.1 dmonitoring - name: openpilot compile3 0.10.1 dmonitoring
run: BENCHMARK_LOG=openpilot_0_10_1_dmonitoring PYTHONPATH="." ASSERT_MIN_STEP_TIME=10 DEV=QCOM FLOAT16=1 IMAGE=2 NOLOCALS=1 taskset -c 4-7 python3 examples/openpilot/compile3.py https://github.com/commaai/openpilot/raw/720392c9a5b986981fdbed1bb8c47a6c5573a50e/selfdrive/modeld/models/dmonitoring_model.onnx run: BENCHMARK_LOG=openpilot_0_10_1_dmonitoring PYTHONPATH="." ASSERT_MIN_STEP_TIME=10 DEV=QCOM FLOAT16=1 IMAGE=2 NOLOCALS=1 taskset -c 4-7 python3 examples/openpilot/compile3.py https://github.com/commaai/openpilot/raw/720392c9a5b986981fdbed1bb8c47a6c5573a50e/selfdrive/modeld/models/dmonitoring_model.onnx
# - name: benchmark MobileNetV2 on DSP - name: benchmark MobileNetV2 on DSP
# run: | run: |
# # generate quantized weights # generate quantized weights
# ln -s /data/home/tiny/tinygrad/extra/datasets/imagenet extra/datasets/imagenet ln -s /data/home/tiny/tinygrad/extra/datasets/imagenet extra/datasets/imagenet
# ln -s /data/home/tiny/tinygrad/testsig-*.so . ln -s /data/home/tiny/tinygrad/testsig-*.so .
# PYTHONPATH=. CC=clang-19 CPU=1 CPU_LLVM=0 QUANT=1 CNT=0 python3 examples/test_onnx_imagenet.py https://github.com/xamcat/mobcat-samples/raw/refs/heads/master/onnx_runtime/InferencingSample/InferencingSample/mobilenetv2-7.onnx /tmp/model.quant.onnx PYTHONPATH=. CC=clang-19 CPU=1 CPU_LLVM=0 QUANT=1 CNT=0 python3 examples/test_onnx_imagenet.py https://github.com/xamcat/mobcat-samples/raw/refs/heads/master/onnx_runtime/InferencingSample/InferencingSample/mobilenetv2-7.onnx /tmp/model.quant.onnx
# # benchmark on DSP with NOOPT=1, the devectorizer has issues # benchmark on DSP with NOOPT=1, the devectorizer has issues
# PYTHONPATH=. CC=clang-19 DSP=1 NOOPT=1 CNT=2 DEBUG=2 python3 examples/test_onnx_imagenet.py /tmp/model.quant.onnx PYTHONPATH=. CC=clang-19 DSP=1 NOOPT=1 CNT=2 DEBUG=2 python3 examples/test_onnx_imagenet.py /tmp/model.quant.onnx
- name: Run process replay tests - name: Run process replay tests
run: cp test/external/process_replay/process_replay.py ./process_replay.py && git fetch origin master && git -c advice.detachedHead=false checkout origin/master && PYTHONPATH=. python3 process_replay.py run: cp test/external/process_replay/process_replay.py ./process_replay.py && git fetch origin master && git -c advice.detachedHead=false checkout origin/master && PYTHONPATH=. python3 process_replay.py
+3 -3
View File
@@ -56,15 +56,15 @@ jobs:
uses: actions/checkout@v4 uses: actions/checkout@v4
with: with:
path: base path: base
- name: Set up Python 3.10 - name: Set up Python 3.12
uses: actions/setup-python@v5 uses: actions/setup-python@v5
with: with:
python-version: '3.10' python-version: '3.12'
- name: Count Line Diff - name: Count Line Diff
run: | run: |
pip install tabulate
BASE="$GITHUB_WORKSPACE/base" BASE="$GITHUB_WORKSPACE/base"
PR="$GITHUB_WORKSPACE/pr" PR="$GITHUB_WORKSPACE/pr"
pip install tabulate $BASE
cp "$BASE/sz.py" . cp "$BASE/sz.py" .
echo "loc_content<<EOF" >> "$GITHUB_ENV" echo "loc_content<<EOF" >> "$GITHUB_ENV"
python sz.py "$BASE" "$PR" >> "$GITHUB_ENV" python sz.py "$BASE" "$PR" >> "$GITHUB_ENV"
+63 -58
View File
@@ -86,65 +86,67 @@ jobs:
clang -O2 recognize.c -lm -o recognize clang -O2 recognize.c -lm -o recognize
cat test/models/efficientnet/Chicken.jpg | ./recognize | grep cock cat test/models/efficientnet/Chicken.jpg | ./recognize | grep cock
# TODO: fix the torch backend and reenable torchbackend:
# torchbackend: name: Torch Backend Tests
# name: Torch Backend Tests runs-on: ubuntu-latest
# runs-on: ubuntu-latest timeout-minutes: 15
# timeout-minutes: 15 steps:
# steps: - name: Checkout Code
# - name: Checkout Code uses: actions/checkout@v4
# uses: actions/checkout@v4 - name: Setup Environment
# - name: Setup Environment uses: ./.github/actions/setup-tinygrad
# uses: ./.github/actions/setup-tinygrad with:
# with: key: torch-backend-pillow-torchvision-et-pt
# key: torch-backend-pillow-torchvision-et-pt deps: testing_minimal
# deps: testing_minimal pydeps: "pillow torchvision expecttest"
# pydeps: "pillow torchvision expecttest" llvm: 'true'
# llvm: 'true' - name: Install ninja
# - name: Install ninja run: |
# run: | sudo apt update || true
# sudo apt update || true sudo apt install -y --no-install-recommends ninja-build
# sudo apt install -y --no-install-recommends ninja-build - name: Lint with ruff
# - name: Lint with ruff run: |
# run: | pip3 install --upgrade --force-reinstall ruff==0.11.0
# pip3 install --upgrade --force-reinstall ruff==0.11.0 python3 -m ruff check extra/torch_backend/backend.py
# python3 -m ruff check extra/torch_backend/backend.py - name: Test one op
# - name: Test one op run: FORWARD_ONLY=1 TINY_BACKEND=1 python3 test/test_ops.py TestOps.test_add
# run: FORWARD_ONLY=1 TINY_BACKEND=1 python3 test/test_ops.py TestOps.test_add - name: Test ResNet-18
# - name: Test ResNet-18 run: DEBUG=2 python3 extra/torch_backend/example.py
# run: DEBUG=2 python3 extra/torch_backend/example.py - name: My (custom) tests
# - name: My (custom) tests run: python3 extra/torch_backend/test.py
# run: python3 extra/torch_backend/test.py - name: Test one op in torch tests
# - name: Test one op in torch tests run: DEBUG=2 python3 extra/torch_backend/torch_tests.py TestTinyBackendPRIVATEUSE1.test_unary_log_tiny_float32
# run: DEBUG=2 python3 extra/torch_backend/torch_tests.py TestTinyBackendPRIVATEUSE1.test_unary_log_tiny_float32 - name: Test Ops with TINY_BACKEND
# - name: Test Ops with TINY_BACKEND run: CPU=1 CPU_LLVM=1 LLVMOPT=0 TINY_BACKEND=1 python3 -m pytest -n auto test/test_ops.py --durations=20
# run: CPU=1 CPU_LLVM=1 LLVMOPT=0 TINY_BACKEND=1 python3 -m pytest -n auto test/test_ops.py --durations=20 - name: Test in-place operations on views
# - name: Test in-place operations on views run: TORCH_DEBUG=1 python3 extra/torch_backend/test_inplace.py
# run: TORCH_DEBUG=1 python3 extra/torch_backend/test_inplace.py - name: Test multi-gpu
# - name: Test multi-gpu run: CPU=1 CPU_LLVM=1 GPUS=4 TORCH_DEBUG=1 python3 extra/torch_backend/test_multigpu.py
# run: CPU=1 CPU_LLVM=1 GPUS=4 TORCH_DEBUG=1 python3 extra/torch_backend/test_multigpu.py - name: Test kernel fusion
run: python3 extra/torch_backend/test_kernel_fusion.py
# torchbackendmore:
# name: Torch Backend Tests More torchbackendmore:
# runs-on: ubuntu-latest name: Torch Backend Tests More
# timeout-minutes: 15 runs-on: ubuntu-latest
# steps: timeout-minutes: 15
# - name: Checkout Code steps:
# uses: actions/checkout@v4 - name: Checkout Code
# - name: Setup Environment uses: actions/checkout@v4
# uses: ./.github/actions/setup-tinygrad - name: Setup Environment
# with: uses: ./.github/actions/setup-tinygrad
# key: torch-backend-pillow-torchvision-et-pt with:
# deps: testing_minimal key: torch-backend-pillow-torchvision-et-pt
# llvm: 'true' deps: testing_minimal
# - name: Install ninja llvm: 'true'
# run: | - name: Install ninja
# sudo apt update || true run: |
# sudo apt install -y --no-install-recommends ninja-build sudo apt update || true
# - name: Test beautiful_mnist in torch with TINY_BACKEND sudo apt install -y --no-install-recommends ninja-build
# run: STEPS=20 CPU=1 TARGET_EVAL_ACC_PCT=90.0 TINY_BACKEND=1 python3 examples/other_mnist/beautiful_mnist_torch.py - name: Test beautiful_mnist in torch with TINY_BACKEND
# - name: Test some torch tests (expect failure) run: STEPS=20 CPU=1 TARGET_EVAL_ACC_PCT=90.0 TINY_BACKEND=1 python3 examples/other_mnist/beautiful_mnist_torch.py
# run: python3 -m pytest extra/torch_backend/torch_tests.py -v --tb=no || true - name: Test some torch tests (expect failure)
run: python3 -m pytest extra/torch_backend/torch_tests.py -v --tb=no || true
bepython: bepython:
name: Python Backend name: Python Backend
@@ -306,6 +308,7 @@ jobs:
with: with:
key: spec-unit key: spec-unit
deps: testing_unit deps: testing_unit
python-version: '3.14'
- name: Test SPEC=2 - name: Test SPEC=2
run: IGNORE_OOB=0 SPEC=2 PYTHONPATH="." pytest --maxfail=10 -n auto --durations=30 --ignore=test/models --ignore test/unit/test_hashing.py --timeout 60 -k "not test_setitem_big" --splits 2 --group ${{ matrix.group }} run: IGNORE_OOB=0 SPEC=2 PYTHONPATH="." pytest --maxfail=10 -n auto --durations=30 --ignore=test/models --ignore test/unit/test_hashing.py --timeout 60 -k "not test_setitem_big" --splits 2 --group ${{ matrix.group }}
@@ -323,6 +326,8 @@ jobs:
deps: testing_unit deps: testing_unit
- name: Fuzz Test symbolic - name: Fuzz Test symbolic
run: python test/external/fuzz_symbolic.py run: python test/external/fuzz_symbolic.py
- name: Fuzz Test symbolic (symbolic divisors)
run: python test/external/fuzz_symbolic_symbolic_div.py
- name: Fuzz Test fast idiv - name: Fuzz Test fast idiv
run: python test/external/fuzz_fast_idiv.py run: python test/external/fuzz_fast_idiv.py
- name: Fuzz Test shape ops - name: Fuzz Test shape ops
-293
View File
@@ -1,293 +0,0 @@
#!/usr/bin/env python3
# this file is a "ramp" for people new to tinygrad to think about how to approach it
# it is runnable and editable.
# whenever you see stuff like DEBUG=2 or CPU=1 discussed, these are environment variables
# in a unix shell like bash `DEBUG=2 CPU=1 python docs/ramp.py`
# this pip installs tinygrad master for the system
# the -e allows you to edit the tinygrad folder and update system tinygrad
# tinygrad is pure Python, so you are encouraged to do this
# git pull in the tinygrad directory will also get you the latest
"""
git clone https://github.com/tinygrad/tinygrad.git
cd tinygrad
python3 -m pip install -e .
"""
# %% ********
print("******* PART 1 *******")
# we start with a Device.
# a Device is where Tensors are stored and compute is run
# tinygrad autodetects the best device on your system and makes it the DEFAULT
from tinygrad import Device
print(Device.DEFAULT) # on Mac, you can see this prints METAL
# now, lets create a Tensor
from tinygrad import Tensor, dtypes
t = Tensor([1,2,3,4])
# you can see this Tensor is on the DEFAULT device with int dtype and shape (4,)
assert t.device == Device.DEFAULT
assert t.dtype == dtypes.int
assert t.shape == (4,)
# unlike in torch, if we print it, it doesn't print the contents
# this is because tinygrad is lazy
# this Tensor has not been computed yet
print(t)
# <Tensor <UOp METAL (4,) int (<Ops.COPY: 7>, None)> on METAL with grad None>
# the ".uop" property on Tensor contains the specification of how to compute it
print(t.uop)
"""
UOp(Ops.COPY, dtypes.int, arg=None, src=(
UOp(Ops.BUFFER, dtypes.int, arg=4, src=(
UOp(Ops.UNIQUE, dtypes.void, arg=0, src=()),
UOp(Ops.DEVICE, dtypes.void, arg='PYTHON', src=()),)),
UOp(Ops.DEVICE, dtypes.void, arg='METAL', src=()),))
"""
# as you can see, it's specifying a copy from PYTHON device
# which is where the [1,2,3,4] array lives
# UOps are the specification language in tinygrad
# they are immutable and form a DAG
# they have a "Ops", a "dtype", a tuple of srcs (parents), and an arg
t.realize()
# if we want to "realize" a tensor, we can with the "realize" method
# now when we look at the uop, it's changed
print(t.uop)
"""
UOp(Ops.BUFFER, dtypes.int, arg=4, src=(
UOp(Ops.UNIQUE, dtypes.void, arg=1, src=()),
UOp(Ops.DEVICE, dtypes.void, arg='METAL', src=()),))
"""
# the copy was actually run, and now the "uop" of the Tensor is just a BUFFER
# if you run this script with DEBUG=2 in the environment, you can see the copy happen
# *** METAL 1 copy 16, METAL <- PYTHON ...
# now let's do some compute
# we look at the uop to see the specification of the compute
t_times_2 = t * 2
print(t_times_2.uop)
"""
UOp(Ops.MUL, dtypes.int, arg=None, src=(
UOp(Ops.BUFFER, dtypes.int, arg=4, src=(
UOp(Ops.UNIQUE, dtypes.void, arg=1, src=()),
x2:=UOp(Ops.DEVICE, dtypes.void, arg='METAL', src=()),)),
UOp(Ops.EXPAND, dtypes.int, arg=(4,), src=(
UOp(Ops.RESHAPE, dtypes.int, arg=(1,), src=(
UOp(Ops.CONST, dtypes.int, arg=2, src=(
UOp(Ops.VIEW, dtypes.void, arg=ShapeTracker(views=(View(shape=(), strides=(), offset=0, mask=None, contiguous=True),)), src=(
x2,)),)),)),)),))
"""
# the BUFFER from above is being multiplied by a CONST 2
# it's RESHAPEd and EXPANDed to broadcast the CONST to the BUFFER
# we can check the result with
assert t_times_2.tolist() == [2, 4, 6, 8]
# UOps are both immutable and globally unique
# if i multiply the Tensor by 4 twice, these result Tensors will have the same uop specification
t_times_4_try_1 = t * 4
t_times_4_try_2 = t * 4
assert t_times_4_try_1.uop is t_times_4_try_2.uop
# the specification isn't just the same, it's the exact same Python object
assert t_times_4_try_1 is not t_times_4_try_2
# the Tensor is a different Python object
# if we realize `t_times_4_try_1` ...
t_times_4_try_1.realize()
print(t_times_4_try_2.uop)
"""
UOp(Ops.BUFFER, dtypes.int, arg=4, src=(
UOp(Ops.UNIQUE, dtypes.void, arg=4, src=()),
UOp(Ops.DEVICE, dtypes.void, arg='METAL', src=()),))
"""
# ... `t_times_4_try_2` also becomes the same BUFFER
assert t_times_4_try_1.uop is t_times_4_try_2.uop
# so this print doesn't require any computation, just a copy back to the CPU so we can print it
print("** only the copy start")
print(t_times_4_try_2.tolist()) # [4, 8, 12, 16]
print("** only the copy end")
# you can confirm this with DEBUG=2, seeing what's printed in between the "**" prints
# tinygrad has an auto differentiation engine that operates according to these same principles
# the derivative of "log(x)" is "1/x", and you can see this on line 20 of gradient.py
t_float = Tensor([3.0])
t_log = t_float.log()
t_log_grad, = t_log.sum().gradient(t_float)
# due to how log is implemented, this gradient contains a lot of UOps
print(t_log_grad.uop)
# ...not shown here...
# but if you run with DEBUG=4 (CPU=1 used here for simpler code), you can see the generated code
"""
void E_(float* restrict data0, float* restrict data1) {
float val0 = *(data1+0);
*(data0+0) = (1/val0);
}
"""
# the derivative is close to 1/3
assert (t_log_grad.item() - 1/3) < 1e-6
# %% ********
print("******* PART 2 *******")
# we redefine the same t here so this cell can run on it's own
from tinygrad import Tensor
t = Tensor([1,2,3,4])
# what's above gives you enough of an understanding to go use tinygrad as a library
# however, a lot of the beauty of tinygrad is in how easy it is to interact with the internals
# NOTE: the APIs here are subject to change
t_plus_3_plus_4 = t + 3 + 4
print(t_plus_3_plus_4.uop)
"""
UOp(Ops.ADD, dtypes.int, arg=None, src=(
UOp(Ops.ADD, dtypes.int, arg=None, src=(
UOp(Ops.BUFFER, dtypes.int, arg=4, src=(
UOp(Ops.UNIQUE, dtypes.void, arg=1, src=()),
x3:=UOp(Ops.DEVICE, dtypes.void, arg='CPU', src=()),)),
UOp(Ops.EXPAND, dtypes.int, arg=(4,), src=(
UOp(Ops.RESHAPE, dtypes.int, arg=(1,), src=(
UOp(Ops.CONST, dtypes.int, arg=3, src=(
x7:=UOp(Ops.VIEW, dtypes.void, arg=ShapeTracker(views=(View(shape=(), strides=(), offset=0, mask=None, contiguous=True),)), src=(
x3,)),)),)),)),)),
UOp(Ops.EXPAND, dtypes.int, arg=(4,), src=(
UOp(Ops.RESHAPE, dtypes.int, arg=(1,), src=(
UOp(Ops.CONST, dtypes.int, arg=4, src=(
x7,)),)),)),))
"""
# you can see it's adding both 3 and 4
# but by the time we are actually running the code, it's adding 7
# `kernelize` will simplify and group the operations in the graph into kernels
t_plus_3_plus_4.kernelize()
print(t_plus_3_plus_4.uop)
"""
UOp(Ops.ASSIGN, dtypes.int, arg=None, src=(
x0:=UOp(Ops.BUFFER, dtypes.int, arg=4, src=(
UOp(Ops.UNIQUE, dtypes.void, arg=7, src=()),
x2:=UOp(Ops.DEVICE, dtypes.void, arg='CPU', src=()),)),
UOp(Ops.KERNEL, dtypes.void, arg=<Kernel 12 SINK(<Ops.STORE: 48>,) (__add__,)>, src=(
x0,
UOp(Ops.BUFFER, dtypes.int, arg=4, src=(
UOp(Ops.UNIQUE, dtypes.void, arg=1, src=()),
x2,)),)),))
"""
# ASSIGN has two srcs, src[0] is the BUFFER that's assigned to, and src[1] is the thing to assign
# src[1] is the GPU Kernel that's going to be run
# we can get the ast of the Kernel as follows
kernel_ast = t_plus_3_plus_4.uop.src[1].arg.ast
# almost everything in tinygrad functions as a rewrite of the UOps
# the codegen rewrites the ast to a simplified form ready for "rendering"
from tinygrad.codegen import full_rewrite_to_sink
rewritten_ast = full_rewrite_to_sink(kernel_ast)
print(rewritten_ast)
"""
UOp(Ops.SINK, dtypes.void, arg=None, src=(
UOp(Ops.STORE, dtypes.void, arg=None, src=(
UOp(Ops.INDEX, dtypes.int.ptr(4), arg=None, src=(
UOp(Ops.DEFINE_GLOBAL, dtypes.int.ptr(4), arg=0, src=()),
x3:=UOp(Ops.SPECIAL, dtypes.int, arg=('gidx0', 4), src=()),)),
UOp(Ops.ADD, dtypes.int, arg=None, src=(
UOp(Ops.LOAD, dtypes.int, arg=None, src=(
UOp(Ops.INDEX, dtypes.int.ptr(4), arg=None, src=(
UOp(Ops.DEFINE_GLOBAL, dtypes.int.ptr(4), arg=1, src=()),
x3,)),)),
UOp(Ops.CONST, dtypes.int, arg=7, src=()),)),)),))
"""
# you can see at this point we are adding 7, not 3 and 4
# with DEBUG=4, we can see the code.
# since optimizations are on, it UPCASTed the operation, explicitly writing out all 4 +7s
t_plus_3_plus_4.realize()
"""
void E_4n2(int* restrict data0, int* restrict data1) {
int val0 = *(data1+0);
int val1 = *(data1+1);
int val2 = *(data1+2);
int val3 = *(data1+3);
*(data0+0) = (val0+7);
*(data0+1) = (val1+7);
*(data0+2) = (val2+7);
*(data0+3) = (val3+7);
}
"""
# the function name E_4n2 is "E" for elementwise op (as opposed to "r" for reduce op)
# "4" for the size, and "n2" for name deduping (it's the 3rd function with the same E and 4 in this session)
# when you print the name with DEBUG=2, you'll see the 4 is yellow, meaning that it's upcasted
# if you run with NOOPT=1 ...
"""
void E_4n2(int* restrict data0, int* restrict data1) {
for (int ridx0 = 0; ridx0 < 4; ridx0++) {
int val0 = *(data1+ridx0);
*(data0+ridx0) = (val0+7);
}
}
"""
# ... you get this unoptimized code with a loop and the 4 is blue (for global). the color code is in kernel.py
# %% ********
print("******* PART 3 *******")
# now, we go even lower and understand UOps better and how the graph rewrite engine works.
# it's much simpler than what's in LLVM or MLIR
from tinygrad import dtypes
from tinygrad.uop.ops import UOp, Ops
# first, we'll construct some const UOps
a = UOp(Ops.CONST, dtypes.int, arg=2)
b = UOp(Ops.CONST, dtypes.int, arg=2)
# if you have been paying attention, you should know these are the same Python object
assert a is b
# UOps support normal Python math operations, so a_plus_b expresses the spec for 2 + 2
a_plus_b = a + b
print(a_plus_b)
"""
UOp(Ops.ADD, dtypes.int, arg=None, src=(
x0:=UOp(Ops.CONST, dtypes.int, arg=2, src=()),
x0,))
"""
# we could actually render this 2+2 into a language like c and run it
# or, we can use tinygrad's graph rewrite engine to "constant fold"
from tinygrad.uop.ops import graph_rewrite, UPat, PatternMatcher
# a `PatternMatcher` is a list of tuples. for each element in the list:
# [0] is the pattern to match, and [1] is the function to run.
# this function can return either a UOp to replace the pattern with, or None to not replace
simple_pm = PatternMatcher([
(UPat(Ops.ADD, src=(UPat(Ops.CONST, name="c1"), UPat(Ops.CONST, name="c2"))),
lambda c1,c2: UOp(Ops.CONST, dtype=c1.dtype, arg=c1.arg+c2.arg)),
])
# this pattern matches the addition of two CONST and rewrites it into a single CONST UOp
# to actually apply the pattern to a_plus_b, we use graph_rewrite
a_plus_b_simplified = graph_rewrite(a_plus_b, simple_pm)
print(a_plus_b_simplified)
"""
UOp(Ops.CONST, dtypes.int, arg=4, src=())
"""
# 2+2 is in fact, 4
# we can also use syntactic sugar to write the pattern nicer
simpler_pm = PatternMatcher([
(UPat.cvar("c1")+UPat.cvar("c2"), lambda c1,c2: c1.const_like(c1.arg+c2.arg))
])
assert graph_rewrite(a_plus_b, simple_pm) is graph_rewrite(a_plus_b, simpler_pm)
# note again the use of is, UOps are immutable and globally unique
# %% ********
# that brings you to an understanding of the most core concepts in tinygrad
# you can run this with VIZ=1 to use the web based graph rewrite explorer
# hopefully now you understand it. the nodes in the graph are just UOps
+1 -1
View File
@@ -41,7 +41,7 @@ The BMC also has a web interface you can use if you find that easier.
It is recommended that you change the BMC password after setting up the box, as the password on the screen is only the initial password. It is recommended that you change the BMC password after setting up the box, as the password on the screen is only the initial password.
If you do decide to change the BMC password and no longer want the initial password to be displayed, remove the `/root/.bmc_password` file. If you do decide to change the BMC password and no longer want the initial password to be displayed, remove the `/root/.bmc_password` file.
Reboot after making these changes or restart the `displayservice.service` service. Reboot after making these changes or restart the `tinybox-display.service` service.
## What do I use it for? ## What do I use it for?
+15 -8
View File
@@ -9,7 +9,7 @@ from typing import Dict, Any
from PIL import Image from PIL import Image
import numpy as np import numpy as np
from tinygrad import Device, GlobalCounters, dtypes, Tensor, TinyJit from tinygrad import Device, GlobalCounters, dtypes, Tensor, TinyJit
from tinygrad.helpers import Timing, Context, getenv, fetch, colored, tqdm, flatten from tinygrad.helpers import Timing, Context, getenv, fetch, colored, tqdm, flatten, profile_marker
from tinygrad.nn import Conv2d, GroupNorm from tinygrad.nn import Conv2d, GroupNorm
from tinygrad.nn.state import torch_load, load_state_dict, get_state_dict from tinygrad.nn.state import torch_load, load_state_dict, get_state_dict
from extra.models.clip import Closed, Tokenizer, FrozenOpenClipEmbedder from extra.models.clip import Closed, Tokenizer, FrozenOpenClipEmbedder
@@ -266,13 +266,16 @@ if __name__ == "__main__":
parser.add_argument('--fakeweights', action='store_true', help="Skip loading checkpoints and use fake weights") parser.add_argument('--fakeweights', action='store_true', help="Skip loading checkpoints and use fake weights")
args = parser.parse_args() args = parser.parse_args()
profile_marker("create model")
model = StableDiffusion() model = StableDiffusion()
# load in weights profile_marker("load in weights")
with WallTimeEvent(BenchEvent.LOAD_WEIGHTS): with WallTimeEvent(BenchEvent.LOAD_WEIGHTS):
if not args.fakeweights: if not args.fakeweights:
model_bin = fetch('https://huggingface.co/CompVis/stable-diffusion-v-1-4-original/resolve/main/sd-v1-4.ckpt', 'sd-v1-4.ckpt') model_bin = fetch('https://huggingface.co/CompVis/stable-diffusion-v-1-4-original/resolve/main/sd-v1-4.ckpt', 'sd-v1-4.ckpt')
load_state_dict(model, torch_load(model_bin)['state_dict'], verbose=False, strict=False, realize=False) state_dict = torch_load(model_bin)['state_dict']
profile_marker("state dict loaded")
load_state_dict(model, state_dict, verbose=False, strict=False, realize=False)
if args.fp16: if args.fp16:
for k,v in get_state_dict(model).items(): for k,v in get_state_dict(model).items():
@@ -281,12 +284,13 @@ if __name__ == "__main__":
Tensor.realize(*get_state_dict(model).values()) Tensor.realize(*get_state_dict(model).values())
# run through CLIP to get context profile_marker("run clip (conditional)")
tokenizer = Tokenizer.ClipTokenizer() tokenizer = Tokenizer.ClipTokenizer()
prompt = Tensor([tokenizer.encode(args.prompt)]) prompt = Tensor([tokenizer.encode(args.prompt)])
context = model.cond_stage_model.transformer.text_model(prompt).realize() context = model.cond_stage_model.transformer.text_model(prompt).realize()
print("got CLIP context", context.shape) print("got CLIP context", context.shape)
profile_marker("run clip (unconditional)")
prompt = Tensor([tokenizer.encode("")]) prompt = Tensor([tokenizer.encode("")])
unconditional_context = model.cond_stage_model.transformer.text_model(prompt).realize() unconditional_context = model.cond_stage_model.transformer.text_model(prompt).realize()
print("got unconditional CLIP context", unconditional_context.shape) print("got unconditional CLIP context", unconditional_context.shape)
@@ -310,6 +314,7 @@ if __name__ == "__main__":
step_times = [] step_times = []
with Context(BEAM=getenv("LATEBEAM")): with Context(BEAM=getenv("LATEBEAM")):
for index, timestep in (t:=tqdm(list(enumerate(timesteps))[::-1])): for index, timestep in (t:=tqdm(list(enumerate(timesteps))[::-1])):
profile_marker(f"step {len(timesteps)-index-1}")
GlobalCounters.reset() GlobalCounters.reset()
st = time.perf_counter_ns() st = time.perf_counter_ns()
t.set_description("%3d %3d" % (index, timestep)) t.set_description("%3d %3d" % (index, timestep))
@@ -319,24 +324,26 @@ if __name__ == "__main__":
latent = run(model, unconditional_context, context, latent, Tensor([timestep]), alphas[tid], alphas_prev[tid], Tensor([args.guidance])) latent = run(model, unconditional_context, context, latent, Tensor([timestep]), alphas[tid], alphas_prev[tid], Tensor([args.guidance]))
if args.timing: Device[Device.DEFAULT].synchronize() if args.timing: Device[Device.DEFAULT].synchronize()
step_times.append((time.perf_counter_ns() - st)*1e-6) step_times.append((time.perf_counter_ns() - st)*1e-6)
# done with diffusion model
del run del run
del model.model
if (assert_time:=getenv("ASSERT_MIN_STEP_TIME")): if (assert_time:=getenv("ASSERT_MIN_STEP_TIME")):
min_time = min(step_times) min_time = min(step_times)
assert min_time < assert_time, f"Speed regression, expected min step time of < {assert_time} ms but took: {min_time} ms" assert min_time < assert_time, f"Speed regression, expected min step time of < {assert_time} ms but took: {min_time} ms"
# upsample latent space to image with autoencoder profile_marker("run decoder") # upsample latent space to image with autoencoder
x = model.decode(latent) x = model.decode(latent).realize()
print(x.shape) print(x.shape)
# save image profile_marker("save image")
im = Image.fromarray(x.numpy()) im = Image.fromarray(x.numpy())
print(f"saving {args.out}") print(f"saving {args.out}")
im.save(args.out) im.save(args.out)
# Open image. # Open image.
if not args.noshow: im.show() if not args.noshow: im.show()
# validation!
if args.prompt == default_prompt and args.steps == 6 and args.seed == 0 and args.guidance == 7.5: if args.prompt == default_prompt and args.steps == 6 and args.seed == 0 and args.guidance == 7.5:
profile_marker("validate")
ref_image = Tensor(np.array(Image.open(Path(__file__).parent / "stable_diffusion_seed0.png"))) ref_image = Tensor(np.array(Image.open(Path(__file__).parent / "stable_diffusion_seed0.png")))
distance = (((x.cast(dtypes.float) - ref_image.cast(dtypes.float)) / ref_image.max())**2).mean().item() distance = (((x.cast(dtypes.float) - ref_image.cast(dtypes.float)) / ref_image.max())**2).mean().item()
assert distance < 3e-3, colored(f"validation failed with {distance=}", "red") # higher distance with WINO assert distance < 3e-3, colored(f"validation failed with {distance=}", "red") # higher distance with WINO
+9 -6
View File
@@ -64,14 +64,17 @@ nvcmds = {getattr(nv_gpu, x):(x, getattr(nv_gpu, "struct_"+x+"_PARAMS", getattr(
x.startswith("NV") and x[6:].startswith("_CTRL_") and isinstance(getattr(nv_gpu, x), int)} x.startswith("NV") and x[6:].startswith("_CTRL_") and isinstance(getattr(nv_gpu, x), int)}
def get_classes(): def get_classes():
hdrpy = (pathlib.Path(__file__).parent.parent.parent / "tinygrad/runtime/autogen/nv_570.py").read_text() res = {}
clss = re.search(r'NV01_ROOT.*?NV_SEMAPHORE_SURFACE = \(0x000000da\) # macro', hdrpy, re.DOTALL).group() known_classes = {"NV01_DEVICE_0", "NV01_ROOT", "NV1_MEMORY_SYSTEM", "NV01_MEMORY_VIRTUAL", "NV1_MEMORY_USER", "NV50_MEMORY_VIRTUAL", "NV_FERMI_VASPACE_A",
pattern = r'([0-9a-zA-Z_]*) = +\((0x[0-9a-fA-F]+)\)' "NV20_SUBDEVICE_0"}
matches = re.findall(pattern, clss, re.MULTILINE) for nm,val in nv_gpu.__dict__.items():
return {int(num, base=16):name for name, num in matches} if not isinstance(val, int): continue
if 0x3000 < val < 0xffff: res[val] = nm
if nm in known_classes: res[val] = nm
return res
nvclasses = get_classes() nvclasses = get_classes()
nvuvms = {getattr(nv_gpu, x):x for x in dir(nv_gpu) if x.startswith("UVM_") and nv_gpu.__dict__.get(x+"_PARAMS")} nvuvms = {getattr(nv_gpu, x):x for x in dir(nv_gpu) if x.startswith("UVM_") and nv_gpu.__dict__.get(x+"_PARAMS")}
nvqcmds = {int(getattr(nv_gpu, x)):x for x in dir(nv_gpu) if x[:7] in {"NVC6C0_", "NVC56F_", "NVC6B5_"} and isinstance(getattr(nv_gpu, x), int)} nvqcmds = {int(getattr(nv_gpu, x)):x for x in dir(nv_gpu) if x[:7] in {"NVC9B0_", "NVC6C0_", "NVC56F_", "NVC6B5_"} and isinstance(getattr(nv_gpu, x), int)}
global_ioctl_id = 0 global_ioctl_id = 0
gpus_user_modes = [] gpus_user_modes = []
+4 -3
View File
@@ -8,10 +8,10 @@ from sz import NONCORE_DIRS
# llama 3 tokenizer # llama 3 tokenizer
tokenizer = Tokenizer(fetch("https://huggingface.co/bofenghuang/Meta-Llama-3-8B/resolve/main/original/tokenizer.model").as_posix()) tokenizer = Tokenizer(fetch("https://huggingface.co/bofenghuang/Meta-Llama-3-8B/resolve/main/original/tokenizer.model").as_posix())
def read_code(base_path): def read_code(base_path, full=False):
ret = [] ret = []
for path, _, files in os.walk(os.path.join(base_path, "tinygrad")): for path, _, files in os.walk(os.path.join(base_path, "tinygrad")):
if not getenv("CORE") and any(path.split("./")[1].startswith(x) for x in NONCORE_DIRS): continue if not full and any(path.split("./")[1].startswith(x) for x in NONCORE_DIRS): continue
for name in files: for name in files:
if not name.endswith(".py"): continue if not name.endswith(".py"): continue
if 'tinygrad/runtime/autogen' in path.replace('\\', '/'): continue if 'tinygrad/runtime/autogen' in path.replace('\\', '/'): continue
@@ -23,9 +23,10 @@ def read_code(base_path):
if __name__ == "__main__": if __name__ == "__main__":
parser = argparse.ArgumentParser(description="Analyze and optionally save tinygrad code.") parser = argparse.ArgumentParser(description="Analyze and optionally save tinygrad code.")
parser.add_argument("--output", help="Output file to write the combined code to.") parser.add_argument("--output", help="Output file to write the combined code to.")
parser.add_argument("--full", action="store_true", help="All directories")
args = parser.parse_args() args = parser.parse_args()
ret = read_code(".") ret = read_code(".", args.full)
table = [] table = []
for name,code in ret: for name,code in ret:
+37 -17
View File
@@ -15,13 +15,6 @@ from tinygrad.device import Device, ProfileDeviceEvent
from extra.sqtt.attempt_sqtt_parse import parse_sqtt_print_packets from extra.sqtt.attempt_sqtt_parse import parse_sqtt_print_packets
# TODO: should really check for AM driver / USB
if not OSX:
def set_power(x): system(f"sudo /opt/rocm/bin/amd-smi set -l {x}")
@atexit.register
def reset_power(): set_power("auto")
set_power("stable_std")
dev = Device["AMD"] dev = Device["AMD"]
@contextlib.contextmanager @contextlib.contextmanager
@@ -32,9 +25,9 @@ def save_sqtt():
yield sqtt yield sqtt
events = dev.profile_events+[ProfileDeviceEvent("AMD", props=dev.device_props())] events = dev.profile_events+[ProfileDeviceEvent("AMD", props=dev.device_props())]
rctx = decode(events) #rctx = decode(events)
assert len(rctx.inst_execs) > 0, "empty sqtt output" #assert len(rctx.inst_execs) > 0, "empty sqtt output"
sqtt.update(rctx.inst_execs) #sqtt.update(rctx.inst_execs)
for e in events: for e in events:
if isinstance(e, ProfileSQTTEvent): if isinstance(e, ProfileSQTTEvent):
@@ -48,7 +41,6 @@ template = """.text
.type matmul,@function .type matmul,@function
matmul: matmul:
INSTRUCTION INSTRUCTION
s_endpgm
.rodata .rodata
.p2align 6 .p2align 6
@@ -71,7 +63,7 @@ amdhsa.kernels:
.private_segment_fixed_size: 0 .private_segment_fixed_size: 0
.wavefront_size: 32 .wavefront_size: 32
.sgpr_count: 8 .sgpr_count: 8
.vgpr_count: 32 .vgpr_count: 8
.max_flat_workgroup_size: 1024 .max_flat_workgroup_size: 1024
.kernarg_segment_align: 8 .kernarg_segment_align: 8
.kernarg_segment_size: 8 .kernarg_segment_size: 8
@@ -86,21 +78,49 @@ amdhsa.kernels:
.end_amdgpu_metadata .end_amdgpu_metadata
""" """
def run_asm(src): def run_asm(src, num_workgroups=1, num_waves=1):
NUM_WORKGROUPS = 1
WAVE_SIZE = 32 WAVE_SIZE = 32
NUM_WAVES = 1
t = Tensor.empty(0x1000).realize() t = Tensor.empty(0x1000).realize()
buf = t.uop.buffer.ensure_allocated() buf = t.uop.buffer.ensure_allocated()
lib = dev.compiler.compile(template.replace("INSTRUCTION", '\n'.join(src))) lib = dev.compiler.compile(template.replace("INSTRUCTION", '\n'.join(src)))
dev.compiler.disassemble(lib) dev.compiler.disassemble(lib)
fxn = AMDProgram(dev, "matmul", lib) fxn = AMDProgram(dev, "matmul", lib)
fxn(buf._buf, global_size=(NUM_WORKGROUPS,1,1), local_size=(WAVE_SIZE*NUM_WAVES,1,1), wait=True) fxn(buf._buf, global_size=(num_workgroups,1,1), local_size=(WAVE_SIZE*num_waves,1,1), wait=True)
if __name__ == "__main__": if __name__ == "__main__":
with save_sqtt() as sqtt:
run_asm([
"s_nop 100",
"s_nop 100",
"s_load_b64 s[0:1], s[0:1], null",
"s_waitcnt lgkmcnt(0)",
"s_nop 100",
"s_nop 100",
"s_add_i32 s2, s2, 10",
"s_add_i32 s2, s2, 10",
"s_nop 100",
"s_nop 100",
"v_mov_b32_e32 v0, 0",
"v_mov_b32_e32 v0, 0",
"s_nop 100",
"s_nop 100",
"v_dual_fmac_f32 v2, v48, v24 :: v_dual_fmac_f32 v9, v37, v51",
"v_dual_fmac_f32 v2, v48, v24 :: v_dual_fmac_f32 v9, v37, v51",
"s_nop 100",
"s_nop 100",
"global_load_b128 v[2:5], v0, s[0:1]",
"global_load_b128 v[2:5], v0, s[0:1]",
"s_nop 100",
"s_nop 100",
"s_sendmsg sendmsg(MSG_DEALLOC_VGPRS)",
"s_endpgm",
], num_workgroups=1, num_waves=1)
exit(0)
with save_sqtt() as sqtt: with save_sqtt() as sqtt:
#(Tensor.empty(16,16) @ Tensor.empty(16,16)).elu().realize() #(Tensor.empty(16,16) @ Tensor.empty(16,16)).elu().realize()
Tensor.empty(1).elu().realize() #Tensor.empty(1, 64).sum(axis=1).realize()
Tensor.empty(1).log2().realize()
exit(0) exit(0)
with save_sqtt() as sqtt: with save_sqtt() as sqtt:
+402 -397
View File
@@ -1,66 +1,169 @@
import pickle import pickle, sys
from tinygrad.helpers import getenv from tinygrad.helpers import getenv, Timing, colored
from extra.sqtt.roc import decode, ProfileSQTTEvent from extra.sqtt.roc import decode, ProfileSQTTEvent
# do these enums match fields in the packets?
#from tinygrad.runtime.support.amd import import_soc
#soc = import_soc([11])
#perf_sel = {getattr(soc, k):k for k in dir(soc) if k.startswith("SQ_PERF_")}
# Instruction packets (one per ISA op) # Instruction packets (one per ISA op)
# NOTE: these are bad guesses and may be wrong! feel free to update if you know better # NOTE: these are bad guesses and may be wrong! feel free to update if you know better
# some names were taken from SQ_TT_TOKEN_MASK_TOKEN_EXCLUDE_SHIFT # some names were taken from SQ_TT_TOKEN_MASK_TOKEN_EXCLUDE_SHIFT
# we see 18 opcodes
# opcodes(18): 1 2 3 4 5 6 8 9 F 10 11 12 14 15 16 17 18 19
# if you exclude everything, you are left with 6
# opcodes( 6): 10 11 14 15 16 17
# sometimes we see a lot of B, but not repeatable
# not seen
# 7 A C
# NOTE: INST runs before EXEC
OPCODE_COLORS = {
# dispatches are BLACK
0x1: "BLACK",
0x18: "BLACK",
# execs are yellow
0x2: "yellow",
0x3: "yellow",
0x4: "YELLOW",
0x5: "YELLOW",
# waves are blue
0x8: "blue",
0x9: "blue",
0x6: "cyan",
0xb: "cyan",
}
OPCODE_NAMES = { OPCODE_NAMES = {
# gated by SQ_TT_TOKEN_EXCLUDE_VALUINST_SHIFT (but others must be enabled for it to show)
0x01: "VALUINST",
# gated by SQ_TT_TOKEN_EXCLUDE_VMEMEXEC_SHIFT # gated by SQ_TT_TOKEN_EXCLUDE_VMEMEXEC_SHIFT
0x02: "VMEMEXEC", 0x02: "VMEMEXEC",
# gated by SQ_TT_TOKEN_EXCLUDE_ALUEXEC_SHIFT # gated by SQ_TT_TOKEN_EXCLUDE_ALUEXEC_SHIFT
0x03: "ALUEXEC", 0x03: "ALUEXEC",
# gated by SQ_TT_TOKEN_EXCLUDE_VALUINST_SHIFT (but others must be enabled for it to show) # gated by SQ_TT_TOKEN_EXCLUDE_IMMEDIATE_SHIFT
0x01: "VALUINST", 0x04: "IMMEDIATE",
0x05: "IMMEDIATE_MASK",
# gated by SQ_TT_TOKEN_EXCLUDE_WAVERDY_SHIFT # gated by SQ_TT_TOKEN_EXCLUDE_WAVERDY_SHIFT
0x06: "WAVERDY", 0x06: "WAVERDY",
# gated by SQ_TT_TOKEN_EXCLUDE_WAVESTARTEND_SHIFT # gated by SQ_TT_TOKEN_EXCLUDE_WAVESTARTEND_SHIFT
0x08: "WAVEEND", 0x08: "WAVEEND",
0x09: "WAVESTART", 0x09: "WAVESTART",
# gated by SQ_TT_TOKEN_EXCLUDE_IMMEDIATE_SHIFT # gated by SQ_TT_TOKEN_EXCLUDE_WAVEALLOC_SHIFT
0x04: "IMMEDIATE_4", 0x0B: "WAVEALLOC", # FFF00
0x05: "IMMEDIATE_5",
# some gated by SQ_TT_TOKEN_EXCLUDE_REG_SHIFT, some always there # gated by NOT SQ_TT_TOKEN_EXCLUDE_PERF_SHIFT
0x14: "REG", 0x0D: "PERF",
# gated by SQ_TT_TOKEN_EXCLUDE_EVENT_SHIFT # gated by SQ_TT_TOKEN_EXCLUDE_EVENT_SHIFT
0x12: "EVENT", 0x12: "EVENT",
0x13: "EVENT_BIG", # FFFFF800
# some gated by SQ_TT_TOKEN_EXCLUDE_REG_SHIFT, some always there. something is broken with the timing on this
0x14: "REG",
# gated by SQ_TT_TOKEN_EXCLUDE_INST_SHIFT # gated by SQ_TT_TOKEN_EXCLUDE_INST_SHIFT
0x18: "INST", 0x18: "INST",
# gated by SQ_TT_TOKEN_EXCLUDE_UTILCTR_SHIFT # gated by SQ_TT_TOKEN_EXCLUDE_UTILCTR_SHIFT
0x19: "UTILCTR", 0x19: "UTILCTR",
# ------------------------------------------------------------------------ # this is the first (8 byte) packet in the bitstream
# 0x070x0F: pure timestamp-ish deltas 0x17: "LAYOUT_HEADER", # layout/mode/group + selectors A/B (reversed)
# ------------------------------------------------------------------------
0x07: "TS_DELTA_S8_W3", # shift=8, width=3 (small delta) # pure time (no extra bits)
0x0F: "TS_DELTA_SHORT",
0x10: "NOP",
0x11: "TS_WAVE_STATE", # almost pure time, has a small flag
# not a good name, but seen and understood mostly
0x15: "SNAPSHOT", # small delta + 50-ish bits of snapshot
0x16: "TS_DELTA_OR_MARK", # 36-bit long delta or 36-bit marker
# packets we haven't seen / rarely see 0x0b
0x07: "TS_DELTA_S8_W3_7", # shift=8, width=3 (small delta)
0x0A: "TS_DELTA_S5_W2_A", # shift=5, width=2 0x0A: "TS_DELTA_S5_W2_A", # shift=5, width=2
0x0B: "TS_DELTA_S5_W3_A", # shift=5, width=3
0x0C: "TS_DELTA_S5_W3_B", # shift=5, width=3 (different consumer) 0x0C: "TS_DELTA_S5_W3_B", # shift=5, width=3 (different consumer)
0x0D: "TS_DELTA_S5_W3_C", # shift=5, width=3
0x0E: "TS_DELTA_S7_W2", # shift=7, width=2
0x0F: "TS_DELTA_SHORT_PLUS4", # short delta; ROCm adds +4 before accumulate
# ------------------------------------------------------------------------
# 0x100x19: timestamps, layout headers, events, perf
# ------------------------------------------------------------------------
0x10: "PSEUDO_NEED_MORE_BITS", # not a real packet; decoder refill hint
0x11: "TS_WAVE_STATE_SAMPLE", # wave stall/termination sample (byte at +10)
0x13: "EVT_SMALL_GENERIC", # same structural family as 0x08/0x12/0x19
0x15: "PERFCOUNTER_SNAPSHOT", # small delta + 50-ish bits of snapshot
0x16: "TS_DELTA36_OR_MARK", # 36-bit long delta or 36-bit marker
0x17: "LAYOUT_MODE_HEADER", # layout/mode/group + selectors A/B
} }
# SALU = 0x0 / s_mov_b32
# SMEM = 0x1 / s_load_b*
# JUMP = 0x3 / s_cbranch_scc0
# NEXT = 0x4 / s_cbranch_execz
# MESSAGE = 0x9 / s_sendmsg
# VALU = 0xb / v_(exp,log)_f32_e32
# VALU = 0xd / v_lshlrev_b64
# VALU = 0xe / v_mad_u64_u32
# VMEM = 0x21 / global_load_b32
# VMEM = 0x22 / global_load_b32
# VMEM = 0x24 / global_store_b32
# VMEM = 0x25 / global_store_b64
# VMEM = 0x27 / global_store
# VMEM = 0x28 / global_store_b64
# LDS = 0x29 / ds_load_b128
# LDS = 0x2b / ds_store_b32
# LDS = 0x2e / ds_store_b128
# ???? = 0x5a / hidden global_load instruction
# ???? = 0x5b / hidden global_load instruction
# ???? = 0x5c / hidden global_store instruction
# VALU = 0x73 / v_cmpx_eq_u32_e32 (not normal VALUINST)
OPNAME = {
0x0: "SALU",
0x1: "SMEM",
0x3: "JUMP",
0x4: "NEXT",
0x9: "MESSAGE",
0xb: "VALU",
0xd: "VALU",
0xe: "VALU",
0x10: "__END",
0x21: "VMEM_LOAD",
0x22: "VMEM_LOAD",
0x24: "VMEM_STORE",
0x25: "VMEM_STORE",
0x26: "VMEM_STORE",
0x27: "VMEM_STORE",
0x28: "VMEM_STORE",
0x29: "LDS_LOAD",
0x2b: "LDS_STORE",
0x2e: "LDS_STORE",
0x50: "__SIMD_LDS_LOAD",
0x51: "__SIMD_LDS_LOAD",
0x54: "__SIMD_LDS_STORE",
0x5a: "__SIMD_VMEM_LOAD",
0x5b: "__SIMD_VMEM_LOAD",
0x5c: "__SIMD_VMEM_STORE",
0x5d: "__SIMD_VMEM_STORE",
0x5e: "__SIMD_VMEM_STORE",
0x5f: "__SIMD_VMEM_STORE",
0x72: "SALU_OR",
0x73: "VALU_CMPX",
}
ALUSRC = {
1: "SALU",
2: "VALU",
3: "VALU_ALT",
}
MEMSRC = {
0: "LDS",
1: "__LDS",
2: "VMEM",
3: "__VMEM",
}
# these tables are from rocprof trace decoder # these tables are from rocprof trace decoder
# rocprof_trace_decoder_parse_data-0x11c6a0 # rocprof_trace_decoder_parse_data-0x11c6a0
# parse_sqtt_180 = b *rocprof_trace_decoder_parse_data-0x11c6a0+0x110040 # parse_sqtt_180 = b *rocprof_trace_decoder_parse_data-0x11c6a0+0x110040
# ---------- 1. local_138: 256-byte state->token table ---------- # ---------- 1. local_138: 256-byte state->opcode table ----------
STATE_TO_TOKEN: bytes = bytes([ STATE_TO_OPCODE: bytes = bytes([
0x10, 0x16, 0x18, 0x01, 0x05, 0x0b, 0x0c, 0x00, 0x0f, 0x14, 0x18, 0x01, 0x09, 0x04, 0x03, 0x02, 0x10, 0x16, 0x18, 0x01, 0x05, 0x0b, 0x0c, 0x00, 0x0f, 0x14, 0x18, 0x01, 0x09, 0x04, 0x03, 0x02,
0x10, 0x17, 0x18, 0x01, 0x06, 0x08, 0x0d, 0x00, 0x0f, 0x14, 0x18, 0x01, 0x0a, 0x04, 0x03, 0x02, 0x10, 0x17, 0x18, 0x01, 0x06, 0x08, 0x0d, 0x00, 0x0f, 0x14, 0x18, 0x01, 0x0a, 0x04, 0x03, 0x02,
0x10, 0x07, 0x18, 0x01, 0x05, 0x0b, 0x0c, 0x00, 0x0f, 0x14, 0x18, 0x01, 0x09, 0x04, 0x03, 0x02, 0x10, 0x07, 0x18, 0x01, 0x05, 0x0b, 0x0c, 0x00, 0x0f, 0x14, 0x18, 0x01, 0x09, 0x04, 0x03, 0x02,
@@ -79,17 +182,47 @@ STATE_TO_TOKEN: bytes = bytes([
0x10, 0x15, 0x18, 0x01, 0x06, 0x08, 0x0d, 0x00, 0x0f, 0x14, 0x18, 0x01, 0x0a, 0x04, 0x03, 0x02, 0x10, 0x15, 0x18, 0x01, 0x06, 0x08, 0x0d, 0x00, 0x0f, 0x14, 0x18, 0x01, 0x0a, 0x04, 0x03, 0x02,
]) ])
# opcode mask (the bits used to determine the opcode, worked out by looking at the repeats in STATE_TO_OPCODE)
opcode_mask = {
0x10: 0b1111,
0x16: 0b1111111,
0x17: 0b1111111,
0x07: 0b1111111,
0x19: 0b1111111,
0x11: 0b1111111,
0x12: 0b11111111,
0x13: 0b11111111,
0x15: 0b1111111,
0x18: 0b111,
0x1: 0b111,
0x5: 0b11111,
0x6: 0b11111,
0xb: 0b11111,
0x8: 0b11111,
0xc: 0b11111,
0xd: 0b11111,
0xf: 0b1111,
0x14: 0b1111,
0x9: 0b11111,
0xa: 0b11111,
0x4: 0b1111,
0x3: 0b1111,
0x2: 0b1111,
}
# ---------- 2. DAT_0012e280: nibble budget per opcode&0x1F ---------- # ---------- 2. DAT_0012e280: nibble budget per opcode&0x1F ----------
NIBBLE_BUDGET = [ NIBBLE_BUDGET = [
0x08, 0x0C, 0x08, 0x08, 0x0C, 0x18, 0x18, 0x40, 0x08, 0x0C, 0x08, 0x08, 0x0C, 0x18, 0x18, 0x40, 0x14, 0x20, 0x30, 0x14, 0x34, 0x1C, 0x30, 0x08,
0x14, 0x20, 0x30, 0x14, 0x34, 0x1C, 0x30, 0x08, 0x04, 0x18, 0x18, 0x20, 0x40, 0x40, 0x30, 0x40, 0x14, 0x30, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00,
0x04, 0x18, 0x18, 0x20, 0x40, 0x40, 0x30, 0x40,
0x14, 0x30, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00,
] ]
assert len(NIBBLE_BUDGET) == 32
# ---------- 3. delta_map from your hash nodes ---------- # ---------- 3. delta_map from your hash nodes ----------
@@ -108,7 +241,8 @@ DELTA_MAP_DEFAULT = {
0x0B: (5, 3), # shift=5, end=8 0x0B: (5, 3), # shift=5, end=8
0x0C: (5, 3), # shift=5, end=8 0x0C: (5, 3), # shift=5, end=8
0x0D: (5, 3), # shift=5, end=8 0x0D: (5, 3), # shift=5, end=8
0x0E: (7, 2), # shift=7, end=9 # NOTE: 0x0e can never be decoded, it's not in the STATE_TO_OPCODE table
#0x0E: (7, 2), # shift=7, end=9
0x0F: (4, 4), # shift=4, end=8 0x0F: (4, 4), # shift=4, end=8
0x10: (0, 0), # shift=0, end=0 (no delta) 0x10: (0, 0), # shift=0, end=0 (no delta)
0x11: (7, 9), # shift=7, end=16 0x11: (7, 9), # shift=7, end=16
@@ -124,307 +258,203 @@ DELTA_MAP_DEFAULT = {
# ---------- 4. One-line-per-packet parser ---------- # ---------- 4. One-line-per-packet parser ----------
def decode_packet_fields(opcode: int, reg: int, delta: int) -> str: def reg_mask(opcode):
nb_bits = NIBBLE_BUDGET[opcode & 0x1F]
shift, width = DELTA_MAP_DEFAULT[opcode]
delta_mask = ((1 << width) - 1) << shift
assert delta_mask & opcode_mask[opcode] == 0, "masks shouldn't overlap"
return ((1 << nb_bits) - 1) & ~(delta_mask | opcode_mask[opcode])
def decode_packet_fields(opcode: int, reg: int) -> str:
""" """
Decode packet payloads conservatively, using: Decode packet payloads conservatively, using:
- NIBBLE_BUDGET[opcode & 0x1F] to mask reg down to true width. - NIBBLE_BUDGET[opcode & 0x1F] to mask reg down to true width.
- DELTA_MAP_DEFAULT[opcode] to expose the "primary" field (often delta). - DELTA_MAP_DEFAULT[opcode] to expose the "primary" field (often delta).
- Per-opcode layouts derived from rocprof's decompiled consumers. - Per-opcode layouts derived from rocprof's decompiled consumers.
""" """
# --- 0. Restrict to real packet bits --------------------------------- # --- 0. Restrict to real packet bits not used in delta ---------------------------------
nb_bits = NIBBLE_BUDGET[opcode & 0x1F] pkt = reg & reg_mask(opcode)
if nb_bits <= 0 or nb_bits >= 64:
pkt = reg & ((1 << 64) - 1)
else:
pkt = reg & ((1 << nb_bits) - 1)
fields: list[str] = [] fields: list[str] = []
shift, width = DELTA_MAP_DEFAULT.get(opcode, (0, 0)) match opcode:
if width: case 0x01: # VALUINST
field_mask = (1 << width) - 1 # 6 bit field
shaped_field = (pkt >> shift) & field_mask flag = (pkt >> 6) & 1
else: wave = pkt >> 7
field_mask = 0 fields.append(f"wave={wave:x}")
shaped_field = 0 if flag: fields.append("flag")
case 0x02: # VMEMEXEC
# 2 bit field (pipe is a guess)
src = pkt>>6
fields.append(f"src={src} [{MEMSRC.get(src, '')}]")
case 0x03: # ALUEXEC
# 2 bit field
src = pkt>>6
fields.append(f"src={src} [{ALUSRC.get(src, '')}]")
case 0x04: # IMMEDIATE_4
# 5 bit field (actually 4)
wave = pkt >> 7
fields.append(f"wave={wave:x}")
case 0x05: # IMMEDIATE_5
# 16 bit field
# 1 bit per wave
fields.append(f"mask={pkt>>8:016b}")
case 0x6:
# wave ready FFFF00
# 16 bit field
# 1 bit per wave
fields.append(f"mask={pkt>>8:016b}")
case 0x0d:
# 20 bit field
fields.append(f"arg = {pkt>>8:X}")
case 0x12:
fields.append(f"event = {pkt>>11:X}")
case 0x15:
fields.append(f"snap = {pkt>>10:X}")
case 0x19:
# wave end
fields.append(f"ctr = {pkt>>9:X}")
case 0xf:
extracted_delta = (reg >> 4) & 0xF
fields.append(f"strange_delta=0x{extracted_delta:x}")
case 0x11:
# DELTA_MAP_DEFAULT: shift=7, width=9 -> small delta.
# FF0000 is the mask
coarse = pkt >> 16
fields.append(f"coarse=0x{coarse:02x}")
# From decomp:
# - when layout<3 and coarse&1, it sets a "has interesting wave" flag
# - when coarse&8, it marks all live waves as "terminated"
if coarse & 0x01:
fields.append("flag_wave_interest=1")
if coarse & 0x08:
fields.append("flag_terminate_all=1")
case 0x8:
# wave end, this is 20 bits (FFF00)
flag7 = (pkt >> 8) & 1
simd = (pkt >> 9) & 3
cu = ((pkt >> 11) & 0x7) | (flag7 << 3)
wave = (pkt >> 15) & 0x1f
fields.append(f"wave={wave:x}")
fields.append(f"simd={simd}")
fields.append(f"cu={cu}")
case 0x9:
# From case 9 (WAVESTART) in multiple consumers:
# flag7 = (w >> 7) & 1 (low bit of uVar41)
# cls2 = (w >> 8) & 3 (class / group)
# slot4 = (w >> 10) & 0xf (slot / group index)
# idx_lo = (w >> 0xd) & 0x1f (low index, layout<4 path)
# idx_hi = (w >> 0xf) & 0x1f (high index, layout>=4 path)
# id7 = (w >> 0x19) & 0x7f (7-bit id)
flag7 = (pkt >> 7) & 1
simd = (pkt >> 8) & 3
cu = ((pkt >> 10) & 0x7) | (flag7 << 3)
wave = (pkt >> 13) & 0x1F
id7 = (pkt >> 17)
fields.append(f"wave={wave:x}")
fields.append(f"simd={simd}")
fields.append(f"cu={cu}")
fields.append(f"id7=0x{id7:x}")
case 0x18:
# FFF88 is the mask
# From case 0x18:
# low3 = w & 7
# grp3 = (w >> 3) or (w >> 4) & 7 (layout-dependent)
# flags = bits 6 (B6) and 7 (B7)
# hi8 = (w >> 0xc) & 0xff (layout 4 path)
# hi7 = (w >> 0xd) & 0x7f (other layouts)
# idx5 = (w >> 7) or (w >> 8) & 0x1f, used as wave index
flag1 = (pkt >> 3) & 1
flag2 = (pkt >> 7) & 1
wave = (pkt >> 8) & 0x1F
op = (pkt >> 13)
fields.append(f"wave={wave:x}")
fields.append(f"op=0x{op:02x} [{OPNAME.get(op, '')}]")
if flag1: fields.append("flag1")
if flag2: fields.append("flag2")
case 0x14:
subop = (pkt >> 16) & 0xFFFF # (short)(w >> 0x10)
val32 = (pkt >> 32) & 0xFFFFFFFF # (uint)(w >> 0x20)
slot = (pkt >> 7) & 0x7 # index in local_168[...] tables
hi_byte = (pkt >> 8) & 0xFF # determines config vs marker
# ===================================================================== fields.append(f"subop=0x{subop:04x}")
# 1. Timestamp-centric opcodes (actually drive 'time') fields.append(f"slot={slot}")
# ===================================================================== fields.append(f"val32=0x{val32:08x}")
if opcode == 0x0F: # TS_DELTA_SHORT_PLUS4 if hi_byte & 0x80:
# In the caller, delta already has +4 applied. # Config flavour: writes config words into per-slot state arrays.
raw_delta = shaped_field fields.append("kind=config")
fields.append(f"raw_delta={raw_delta}") if subop == 0x000C:
fields.append(f"ts_short_plus4={delta}") fields.append("slot=lo")
return ", ".join(fields) elif subop == 0x000D:
fields.append("slot=hi")
else:
# COR marker: subop 0xC342, payload "COR\0" → start of a COR region.
if subop == 0xC342:
fields.append("kind=cor_stream")
if val32 == 0x434F5200:
fields.append("cor_magic='COR\\0'")
case 0x16:
# Bits:
# bit8 -> 0x100
# bit9 -> 0x200
# bits 12..47 -> 36-bit field used as delta or marker
bit8 = bool(pkt & 0x100)
bit9 = bool(pkt & 0x200)
if not bit9:
mode = "delta"
elif not bit8:
mode = "marker"
else:
mode = "other"
# need to use reg here
val36 = (reg >> 12) & ((1 << 36) - 1)
fields.append(f"mode={mode}")
if mode != "delta":
fields.append(f"val36=0x{val36:x}")
case 0x17:
# From decomp (two sites with identical logic):
# layout = (w >> 7) & 0x3f
# mode = (w >> 0xd) & 3
# group = (w >> 0xf) & 7
# sel_a = (w >> 0x1c) & 0xf
# sel_b = (w >> 0x21) & 7
# flag4 = (w >> 0x3b) & 1 (only meaningful when layout == 4)
layout = (pkt >> 7) & 0x3F
simd = (pkt >> 13) & 0x3 # you can change this by changing traced simd
group = (pkt >> 15) & 0x7
sel_a = (pkt >> 0x1C) & 0xF
sel_b = (pkt >> 0x21) & 0x7
flag4 = (pkt >> 0x3B) & 0x1
if opcode == 0x11: # TS_WAVE_STATE_SAMPLE fields.append(f"layout={layout}")
# DELTA_MAP_DEFAULT: shift=7, width=9 -> small delta. fields.append(f"group={group}")
raw_delta = shaped_field fields.append(f"simd={simd}")
coarse = (pkt >> (shift + width)) & 0xFF # matches byte at +10 in C fields.append(f"sel_a={sel_a}")
fields.append(f"raw_delta={raw_delta}") fields.append(f"sel_b={sel_b}")
if coarse: if layout == 4:
fields.append(f"coarse_state=0x{coarse:02x}") fields.append(f"layout4_flag={flag4}")
# From decomp: case _:
# - when layout<3 and coarse&1, it sets a "has interesting wave" flag fields.append(f"{pkt:X} & {reg_mask(opcode):X}")
# - when coarse&8, it marks all live waves as "terminated" return ",".join(fields)
if coarse & 0x01:
fields.append("flag_wave_interest=1")
if coarse & 0x08:
fields.append("flag_terminate_all=1")
return ", ".join(fields)
if opcode == 0x16: # TS_DELTA36_OR_MARK FILTER_LEVEL = getenv("FILTER", 1)
# Bits:
# bit8 -> 0x100
# bit9 -> 0x200
# bits 12..47 -> 36-bit field used as delta or marker
bit8 = bool(pkt & 0x100)
bit9 = bool(pkt & 0x200)
if not bit9:
mode = "delta"
elif not bit8:
mode = "marker"
else:
mode = "other"
val36 = (pkt >> 12) & ((1 << 36) - 1)
fields.append(f"mode={mode}")
if mode != "delta":
fields.append(f"val36=0x{val36:x}")
return ", ".join(fields)
# For 0x07, 0x0A0x0E, we know they drive time (via DELTA_MAP_DEFAULT), DEFAULT_FILTER: tuple[int, ...] = tuple()
# but we don't see any other fields used in the decomp. # NOP + pure time + "sample"
if opcode in (0x07, 0x0A, 0x0B, 0x0C, 0x0D, 0x0E): if FILTER_LEVEL >= 0: DEFAULT_FILTER += (0x10, 0xf, 0x11)
if width: # reg + event + sample + marker
raw_delta = shaped_field # TODO: events are probably good
leftover = pkt & ~(field_mask << shift) if FILTER_LEVEL >= 1: DEFAULT_FILTER += (0x14, 0x12, 0x16)
fields.append(f"raw_delta={raw_delta}") # instruction runs + valuinst
if leftover: if FILTER_LEVEL >= 2: DEFAULT_FILTER += (0x01, 0x02, 0x03)
fields.append(f"payload=0x{leftover:x}") # instructions dispatch (inst, immed)
return ", ".join(fields) if FILTER_LEVEL >= 3: DEFAULT_FILTER += (0x4, 0x5, 0x18)
# waves
if FILTER_LEVEL >= 4: DEFAULT_FILTER += (0x6, 0x8, 0x9)
# ===================================================================== def parse_sqtt_print_packets(data: bytes, filter=DEFAULT_FILTER, verbose=True) -> None:
# 2. Small "meta + tiny delta" packets (0x010x06)
# =====================================================================
if opcode == 0x01: # META_ID12_TS_SMALL
id12 = pkt & 0xFFF
fields.append(f"id12=0x{id12:03x}")
if width:
fields.append(f"field_s{shift}_w{width}={shaped_field}")
return ", ".join(fields)
if opcode == 0x02: # META_FLAG8_TS_SMALL
flag8 = pkt & 0xFF
fields.append(f"flag8=0x{flag8:02x}")
if width:
fields.append(f"field_s{shift}_w{width}={shaped_field}")
return ", ".join(fields)
if opcode == 0x03: # META_SUBEVENT8_TS_SMALL
sub8 = pkt & 0xFF
fields.append(f"subevent8=0x{sub8:02x}")
if width:
fields.append(f"field_s{shift}_w{width}={shaped_field}")
return ", ".join(fields)
if opcode == 0x04: # META_BASE_INDEX12_TS
idx12 = pkt & 0xFFF
fields.append(f"base_index12=0x{idx12:03x}")
if width:
fields.append(f"field_s{shift}_w{width}={shaped_field}")
return ", ".join(fields)
if opcode in (0x05, 0x06): # META_DESC24_TS_A/B
desc24 = pkt & 0xFFFFFF
fields.append(f"desc24=0x{desc24:06x}")
if width:
fields.append(f"field_s{shift}_w{width}={shaped_field}")
return ", ".join(fields)
# =====================================================================
# 3. Opcode 0x14: exec/config record (+ COR marker)
# =====================================================================
if opcode == 0x14: # INST_EXEC_OR_CFG
subop = (pkt >> 16) & 0xFFFF # (short)(w >> 0x10)
val32 = (pkt >> 32) & 0xFFFFFFFF # (uint)(w >> 0x20)
slot = (pkt >> 7) & 0x7 # index in local_168[...] tables
hi_byte = (pkt >> 8) & 0xFF # determines config vs marker
fields.append(f"subop=0x{subop:04x}")
fields.append(f"slot={slot}")
fields.append(f"val32=0x{val32:08x}")
if hi_byte & 0x80:
# Config flavour: writes config words into per-slot state arrays.
fields.append("kind=config")
if subop == 0x000C:
fields.append("cfg_target=local_168[slot].lo")
elif subop == 0x000D:
fields.append("cfg_target=local_168[slot].hi")
else:
# COR marker: subop 0xC342, payload "COR\0" → start of a COR region.
if subop == 0xC342:
fields.append("kind=cor_stream")
if val32 == 0x434F5200:
fields.append("cor_magic='COR\\0'")
return ", ".join(fields)
# =====================================================================
# 4. Opcode 0x17: layout / mode header
# =====================================================================
if opcode == 0x17: # LAYOUT_MODE_HEADER
# From decomp (two sites with identical logic):
# layout = (w >> 7) & 0x3f
# mode = (w >> 0xd) & 3
# group = (w >> 0xf) & 7
# sel_a = (w >> 0x1c) & 0xf
# sel_b = (w >> 0x21) & 7
# flag4 = (w >> 0x3b) & 1 (only meaningful when layout == 4)
layout = (pkt >> 7) & 0x3F
mode = (pkt >> 13) & 0x3
group = (pkt >> 15) & 0x7
sel_a = (pkt >> 0x1C) & 0xF
sel_b = (pkt >> 0x21) & 0x7
flag4 = (pkt >> 0x3B) & 0x1
fields.append(f"layout={layout}")
fields.append(f"group={group}")
fields.append(f"mode={mode}")
fields.append(f"sel_a={sel_a}")
fields.append(f"sel_b={sel_b}")
if layout == 4:
fields.append(f"layout4_flag={flag4}")
return ", ".join(fields)
# =====================================================================
# 5. Opcode 0x09: state / route config record
# =====================================================================
if opcode == 0x09: # PERF_ROUTE_CONFIG
# From case 9 in multiple consumers:
# flag7 = (w >> 7) & 1 (low bit of uVar41)
# cls2 = (w >> 8) & 3 (class / group)
# slot4 = (w >> 10) & 0xf (slot / group index)
# idx_lo = (w >> 0xd) & 0x1f (low index, layout<4 path)
# idx_hi = (w >> 0xf) & 0x1f (high index, layout>=4 path)
# id7 = (w >> 0x19) & 0x7f (7-bit id)
flag7 = (pkt >> 7) & 0x1
cls2 = (pkt >> 8) & 0x3
slot4 = (pkt >> 10) & 0xF
idx_lo = (pkt >> 13) & 0x1F
idx_hi = (pkt >> 15) & 0x1F
id7 = (pkt >> 0x19) & 0x7F
fields.append(f"flag7={flag7}")
fields.append(f"cls2={cls2}")
fields.append(f"slot4=0x{slot4:x}")
fields.append(f"idx_lo5=0x{idx_lo:x}")
fields.append(f"idx_hi5=0x{idx_hi:x}")
fields.append(f"id7=0x{id7:x}")
return ", ".join(fields)
# =====================================================================
# 6. Opcode 0x18: perf/event selector (FUN_0010aba0)
# =====================================================================
if opcode == 0x18: # PERF_EVENT_SELECT
# From case 0x18:
# low3 = w & 7
# grp3 = (w >> 3) or (w >> 4) & 7 (layout-dependent)
# flags = bits 6 (B6) and 7 (B7)
# hi8 = (w >> 0xc) & 0xff (layout 4 path)
# hi7 = (w >> 0xd) & 0x7f (other layouts)
# idx5 = (w >> 7) or (w >> 8) & 0x1f, used as wave index
low3 = pkt & 0x7
grp3_a = (pkt >> 3) & 0x7
grp3_b = (pkt >> 4) & 0x7
flag_b6 = (pkt >> 6) & 0x1
flag_b7 = (pkt >> 7) & 0x1
idx5_a = (pkt >> 7) & 0x1F
idx5_b = (pkt >> 8) & 0x1F
hi8 = (pkt >> 12) & 0xFF
hi7 = (pkt >> 13) & 0x7F
fields.append(f"low3=0x{low3:x}")
fields.append(f"grp3_a=0x{grp3_a:x}")
fields.append(f"grp3_b=0x{grp3_b:x}")
fields.append(f"flag_b6={flag_b6}")
fields.append(f"flag_b7={flag_b7}")
fields.append(f"idx5_a=0x{idx5_a:x}")
fields.append(f"idx5_b=0x{idx5_b:x}")
fields.append(f"hi8=0x{hi8:02x}")
fields.append(f"hi7=0x{hi7:02x}")
return ", ".join(fields)
# =====================================================================
# 7. Opcode 0x15: perfcounter snapshot
# =====================================================================
if opcode == 0x15: # PERFCOUNTER_SNAPSHOT
# NIBBLE_BUDGET gives full 64 bits here.
# DELTA_MAP_DEFAULT: shift=7, width=3 → tiny delta field.
raw_delta = shaped_field if width else 0
# low bits below the delta field
snap_low = pkt & ((1 << shift) - 1) if shift else 0
# everything above delta field
snap_hi = pkt >> (shift + width) if width else (pkt >> shift)
fields.append(f"raw_delta={raw_delta}")
fields.append(f"snap_low_s{shift}=0x{snap_low:x}")
fields.append(f"snap_hi=0x{snap_hi:x}")
return ", ".join(fields)
# =====================================================================
# 8. Small event-ish packets (0x08 / 0x12 / 0x13 / 0x19)
# =====================================================================
if opcode in (0x08, 0x12, 0x13, 0x19):
# These are all "small event / metric" style tokens. The exact semantics
# depend on layout (0x17) and accumulated state (local_500 etc), so we
# expose:
# - low 8 bits as kind byte
# - rest as opaque payload.
kind = pkt & 0xFF
payload = pkt >> 8
fields.append(f"kind_byte=0x{kind:02x}")
if payload:
fields.append(f"payload=0x{payload:x}")
return ", ".join(fields)
# =====================================================================
# 9. Pseudo opcode 0x10: never a "real" packet
# =====================================================================
if opcode == 0x10: # PSEUDO_NEED_MORE_BITS
# The main loop never prints these; they're just a control token.
return ""
# =====================================================================
# 10. Generic fallback: expose the DELTA_MAP_DEFAULT field + leftover
# =====================================================================
if width:
fields.append(f"field_s{shift}_w{width}={shaped_field}")
leftover = pkt & ~(field_mask << shift)
if leftover:
fields.append(f"payload=0x{leftover:x}")
return ", ".join(fields)
# 0xb is time something
# 0xd is time something
# 0xf is small time advance
# 0x11 is time advance
# 0x16 is big time advance + markers
# 0x14 is REG
DEFAULT_FILTER = (0xb, 0xd, 0xf, 0x11, 0x16, 0x14) if getenv("FILTER", 1) else None
def parse_sqtt_print_packets(data: bytes, max_tokens: int = 100000, filter=DEFAULT_FILTER) -> None:
""" """
Minimal debug: print ONE LINE per decoded token (packet). Minimal debug: print ONE LINE per decoded token (packet).
@@ -433,111 +463,86 @@ def parse_sqtt_print_packets(data: bytes, max_tokens: int = 100000, filter=DEFAU
""" """
n = len(data) n = len(data)
time = 0 time = 0
last_printed_time = 0
reg = 0 # shift register reg = 0 # shift register
offset = 0 # bit offset, in steps of 4 (one nibble) offset = 0 # bit offset, in steps of 4 (one nibble)
nib_budget = 0x40 nib_budget = 0x40
flags = 0 flags = 0
token_index = 0 token_index = 0
opcodes_seen = set()
while (offset >> 3) < n and token_index < max_tokens: while (offset >> 3) < n:
# Remember where we started refilling for this step (bit offset),
# but the *logical* start of the current packet is last_real_offset.
refill_start = offset
# 1) Fill register with nibbles according to nib_budget # 1) Fill register with nibbles according to nib_budget
if nib_budget != 0: if nib_budget != 0:
target = refill_start + 4 + ((nib_budget - 1) & ~3) target = offset + 4 + ((nib_budget - 1) & ~3)
cur = refill_start while offset != target and (offset >> 3) < n:
while cur != target and (cur >> 3) < n: byte = data[offset >> 3]
byte_index = cur >> 3 nib = (byte >> (offset & 4)) & 0xF
byte = data[byte_index]
shift = 4 if (cur & 4) else 0 # low then high nibble
nib = (byte >> shift) & 0xF
reg = ((reg >> 4) | (nib << 60)) & ((1 << 64) - 1) reg = ((reg >> 4) | (nib << 60)) & ((1 << 64) - 1)
cur += 4 offset += 4
offset = cur
# 2) Decode token from low 8 bits # 2) Decode token from low 8 bits
state = reg & 0xFF opcode = STATE_TO_OPCODE[reg & 0xFF]
opcode = STATE_TO_TOKEN[state] opcodes_seen.add(opcode)
# 3) Handle pseudo-token 0x10: need more bits, don't print. Looks like a NOP. # 4) Set next nibble budget based on opcode
if opcode == 0x10: nib_budget = NIBBLE_BUDGET[opcode & 0x1F]
# "need more bits" pseudo-token: adjust nibble budget and continue
nib_budget = 4
if (offset >> 3) >= n:
break
# Do NOT count this as a real packet; do not update last_real_offset.
continue
# 4) Set next nibble budget # 5) Get delta
nb_index = opcode & 0x1F shift, width = DELTA_MAP_DEFAULT[opcode]
nib_budget = NIBBLE_BUDGET[nb_index] delta = (reg >> shift) & ((1 << width) - 1)
time_before = time
note = "" # 6) Update time and handle special opcodes 0xF/0x16
# 5) Special opcode 0x16 (timestamp / marker)
if opcode == 0x16: if opcode == 0x16:
two_bits = (reg >> 8) & 0x3 two_bits = (reg >> 8) & 0x3
if two_bits == 1: if two_bits == 1:
flags |= 0x01 flags |= 0x01
# Common 36-bit field at bits [12..47] # Common 36-bit field at bits [12..47]
if (reg & 0x200) == 0: if (reg & 0x200) == 0:
# delta mode: add 36-bit delta to time # delta mode: add 36-bit delta to time
delta = (reg >> 12) & ((1 << 36) - 1) pass
time += delta elif (reg & 0x100) == 0:
else:
# marker / other modes: no time advance # marker / other modes: no time advance
if (reg & 0x100) == 0: # real marker: bit9=1, bit8=0, non-zero payload
# real marker: bit9=1, bit8=0, non-zero payload # "other" 0x16 variants, ignored for timing
# "other" 0x16 variants, ignored for timing delta = 0
delta = 0
else:
# 6) Generic opcode (including 0x0F)
shift, width = DELTA_MAP_DEFAULT[opcode]
mask = (1 << width) - 1
delta = (reg >> shift) & mask
# TODO: add more opcode parsers here that add notes to other opcodes
if opcode == 0x0F:
delta_with_fix = delta + 4
time += delta_with_fix
delta = delta_with_fix
else: else:
time += delta raise RuntimeError("unknown 0x16 delta")
elif opcode == 0x0F:
# opcode 0x0F has an offset of 4 to the delta
# update: it's actually computed to be 8 to match WAVESTART
delta = delta + 8
# Append extra decoded fields into the note string # Append extra decoded fields into the note string
note = decode_packet_fields(opcode, reg, delta) note = decode_packet_fields(opcode, reg)
if filter is None or opcode not in filter:
my_reg = reg
my_reg &= (1 << nib_budget) - 1
print(
f"{token_index:4d} "
f"off={offset//4:5d} "
f"op=0x{opcode:02x} "
f"{OPCODE_NAMES[opcode]:24s} "
f" time={time_before:8d}+{delta:8d} "
f"{my_reg:16X} "
f"{note}"
)
# this delta happens before the instruction
time += delta
token_index += 1 token_index += 1
if verbose and (filter is None or opcode not in filter):
print(f"{time:8d} +{time-last_printed_time:8d} : "+colored(f"{OPCODE_NAMES[opcode]:18s} ", OPCODE_COLORS.get(opcode, "white"))+f"{note}")
last_printed_time = time
# Optional summary at the end # Optional summary at the end
print(f"# done: tokens={token_index}, final_time={time}, flags=0x{flags:02x}") print(f"# done: tokens={token_index:_}, final_time={time}, flags=0x{flags:02x}")
if verbose:
print(f"opcodes({len(opcodes_seen):2d}):",
' '.join([colored(f"{op:2X}", "WHITE" if op in opcodes_seen else "BLACK") for op in sorted(opcode_mask)]))
def parse(fn:str): def parse(fn:str):
dat = pickle.load(open(fn, "rb")) with Timing(f"unpickle {fn}: "): dat = pickle.load(open(fn, "rb"))
ctx = decode(dat) if getenv("ROCM", 0):
with Timing(f"decode {fn}: "): ctx = decode(dat)
dat_sqtt = [x for x in dat if isinstance(x, ProfileSQTTEvent)] dat_sqtt = [x for x in dat if isinstance(x, ProfileSQTTEvent)]
print(f"got {len(dat_sqtt)} SQTT events in {fn}") print(f"got {len(dat_sqtt)} SQTT events in {fn}")
return dat_sqtt return dat_sqtt
if __name__ == "__main__": if __name__ == "__main__":
#dat_sqtt = parse("extra/sqtt/examples/profile_empty_run_0.pkl") fn = "extra/sqtt/examples/profile_gemm_run_0.pkl"
#dat_sqtt = parse("extra/sqtt/examples/profile_plus_run_0.pkl") dat_sqtt = parse(sys.argv[1] if len(sys.argv) > 1 else fn)
dat_sqtt = parse("extra/sqtt/examples/profile_gemm_run_0.pkl") for i,dat in enumerate(dat_sqtt):
blob_0 = dat_sqtt[0].blob with Timing(f"decode pkt {i} with len {len(dat.blob):_}: "):
parse_sqtt_print_packets(blob_0[8:]) parse_sqtt_print_packets(dat.blob, verbose=getenv("V", 1))
+40 -19
View File
@@ -1,4 +1,5 @@
import ctypes, pathlib, argparse, pickle, re, functools, dataclasses, itertools, threading import ctypes, pathlib, argparse, pickle, re, functools, dataclasses, itertools, threading
from typing import Generator
from tinygrad.helpers import temp, unwrap, DEBUG from tinygrad.helpers import temp, unwrap, DEBUG
from tinygrad.device import ProfileEvent, ProfileDeviceEvent, ProfileProgramEvent from tinygrad.device import ProfileEvent, ProfileDeviceEvent, ProfileProgramEvent
from tinygrad.runtime.ops_amd import ProfileSQTTEvent, ProfilePMCEvent from tinygrad.runtime.ops_amd import ProfileSQTTEvent, ProfilePMCEvent
@@ -31,31 +32,50 @@ def llvm_disasm(arch:str, lib:bytes) -> dict[int, tuple[str, int]]:
@dataclasses.dataclass(frozen=True) @dataclasses.dataclass(frozen=True)
class InstExec: class InstExec:
typ:str typ:str
inst:str pc:int
stall:int stall:int
dur:int dur:int
time:int time:int
@dataclasses.dataclass(frozen=True) @dataclasses.dataclass(frozen=True)
class WaveExec: class WaveSlot:
wave_id:int wave_id:int
cu:int cu:int
simd:int simd:int
se:int se:int
@property
def simd_loc(self) -> str: return f"SE:{self.se} CU:{self.cu} SIMD:{self.simd}"
@property
def wave_loc(self) -> str: return f"{self.simd_loc} WAVE:{self.wave_id}"
@dataclasses.dataclass(frozen=True)
class WaveExec(WaveSlot):
begin_time:int begin_time:int
end_time:int end_time:int
insts:list[InstExec] insts:bytearray
def unpack_insts(self) -> Generator[InstExec, None, None]:
sz = ctypes.sizeof(struct:=rocprof.rocprofiler_thread_trace_decoder_inst_t)
insts_array = (struct*(len(self.insts)//sz)).from_buffer(self.insts)
for inst in insts_array:
inst_typ = rocprof.enum_rocprofiler_thread_trace_decoder_inst_category_t.get(inst.category)
yield InstExec(inst_typ, inst.pc.address, inst.stall, inst.duration, inst.time)
@dataclasses.dataclass(frozen=True)
class OccEvent(WaveSlot):
time:int
start:int
class _ROCParseCtx: class _ROCParseCtx:
def __init__(self, dev_evs:dict[str, ProfileDeviceEvent], sqtt_evs:list[ProfileSQTTEvent], prog_evs:list[ProfileProgramEvent]): def __init__(self, dev_evs:dict[str, ProfileDeviceEvent], sqtt_evs:list[ProfileSQTTEvent], prog_evs:list[ProfileProgramEvent]):
self.dev_evs, self.sqtt_evs, self.prog_evs = dev_evs, iter(sqtt_evs), prog_evs self.dev_evs, self.sqtt_evs, self.prog_evs = dev_evs, iter(sqtt_evs), prog_evs
self.disasms:dict[tuple[str, int], tuple[str, int]] = {} self.disasms:dict[str, dict[int, tuple[str, int]]] = {}
self.inst_execs:dict[str, list[WaveExec]] = {} self.inst_execs:dict[str, list[WaveExec]] = {}
self.occ_events:dict[str, list[OccEvent]] = {}
for prog in prog_evs: for prog in prog_evs:
arch = "gfx%d%x%x" % ((trgt:=unwrap(dev_evs[prog.device].props)['gfx_target_version']) // 10000, (trgt // 100) % 100, trgt % 100) arch = "gfx%d%x%x" % ((trgt:=unwrap(dev_evs[prog.device].props)['gfx_target_version']) // 10000, (trgt // 100) % 100, trgt % 100)
for addr, info in llvm_disasm(arch, unwrap(prog.lib)).items(): base = unwrap(prog.base)
self.disasms[(prog.name, unwrap(prog.base) + addr)] = info self.disasms[prog.name] = asm = {base+addr:info for addr,info in llvm_disasm(arch, unwrap(prog.lib)).items()}
def next_sqtt(self): def next_sqtt(self):
x = next(self.sqtt_evs, None) x = next(self.sqtt_evs, None)
@@ -65,22 +85,19 @@ class _ROCParseCtx:
return self.active_blob return self.active_blob
def on_occupancy_ev(self, ev:rocprof.rocprofiler_thread_trace_decoder_occupancy_t): def on_occupancy_ev(self, ev:rocprof.rocprofiler_thread_trace_decoder_occupancy_t):
if DEBUG >= 5: print("OCC", ev.time, self.active_se, ev.cu, ev.simd, ev.wave_id, ev.start) if DEBUG >= 5: print(f"OCC {ev.time=} {self.active_se=} {ev.cu=} {ev.simd=} {ev.wave_id=} {ev.start=}")
self.occ_events.setdefault(unwrap(self.active_kern), []).append(OccEvent(ev.wave_id, ev.cu, ev.simd, unwrap(self.active_se), ev.time, ev.start))
def on_wave_ev(self, ev:rocprof.rocprofiler_thread_trace_decoder_wave_t): def on_wave_ev(self, ev:rocprof.rocprofiler_thread_trace_decoder_wave_t):
if DEBUG >= 5: print("WAVE", ev.wave_id, self.active_se, ev.cu, ev.simd, ev.contexts, ev.begin_time, ev.end_time) if DEBUG >= 5: print(f"WAVE {ev.wave_id=} {self.active_se=} {ev.cu=} {ev.simd=} {ev.contexts=} {ev.begin_time=} {ev.end_time=}")
# Skip wave events without instruction timings, occupancy events give the start and duration.
if ev.instructions_size == 0: return
inst_execs:list[InstExec] = [] insts_blob = bytearray(sz:=ev.instructions_size * ctypes.sizeof(rocprof.rocprofiler_thread_trace_decoder_inst_t))
for j in range(ev.instructions_size): ctypes.memmove((ctypes.c_char * sz).from_buffer(insts_blob), ev.instructions_array, sz)
inst_ev = ev.instructions_array[j]
inst_typ = rocprof.enum_rocprofiler_thread_trace_decoder_inst_category_t.get(inst_ev.category)
inst_disasm = self.disasms[(unwrap(self.active_kern), unwrap(inst_ev.pc.address))][0]
inst_execs.append(InstExec(inst_typ, inst_disasm, inst_ev.stall, inst_ev.duration, inst_ev.time))
if DEBUG >= 8: print(inst_execs[-1])
if ev.instructions_size > 0: self.inst_execs.setdefault(unwrap(self.active_kern), []).append(WaveExec(ev.wave_id, ev.cu, ev.simd, unwrap(self.active_se), ev.begin_time,
self.inst_execs.setdefault(unwrap(self.active_kern), []).append(WaveExec(ev.wave_id, ev.cu, ev.simd, unwrap(self.active_se), ev.begin_time, ev.end_time, insts_blob))
ev.end_time, inst_execs))
def decode(profile:list[ProfileEvent]) -> _ROCParseCtx: def decode(profile:list[ProfileEvent]) -> _ROCParseCtx:
dev_events:dict[str, ProfileDeviceEvent] = {} dev_events:dict[str, ProfileDeviceEvent] = {}
@@ -107,13 +124,17 @@ def decode(profile:list[ProfileEvent]) -> _ROCParseCtx:
for ev in (rocprof.rocprofiler_thread_trace_decoder_occupancy_t * n).from_address(events_ptr): ROCParseCtx.on_occupancy_ev(ev) for ev in (rocprof.rocprofiler_thread_trace_decoder_occupancy_t * n).from_address(events_ptr): ROCParseCtx.on_occupancy_ev(ev)
case rocprof.ROCPROFILER_THREAD_TRACE_DECODER_RECORD_WAVE: case rocprof.ROCPROFILER_THREAD_TRACE_DECODER_RECORD_WAVE:
for ev in (rocprof.rocprofiler_thread_trace_decoder_wave_t * n).from_address(events_ptr): ROCParseCtx.on_wave_ev(ev) for ev in (rocprof.rocprofiler_thread_trace_decoder_wave_t * n).from_address(events_ptr): ROCParseCtx.on_wave_ev(ev)
case rocprof.ROCPROFILER_THREAD_TRACE_DECODER_RECORD_REALTIME:
if DEBUG >= 5:
pairs = [(ev.shader_clock, ev.realtime_clock) for ev in (rocprof.rocprofiler_thread_trace_decoder_realtime_t * n).from_address(events_ptr)]
print(f"REALTIME {pairs}")
case _: case _:
if DEBUG >= 5: print(rocprof.enum_rocprofiler_thread_trace_decoder_record_type_t.get(record_type), events_ptr, n) if DEBUG >= 5: print(rocprof.enum_rocprofiler_thread_trace_decoder_record_type_t.get(record_type), events_ptr, n)
return rocprof.ROCPROFILER_THREAD_TRACE_DECODER_STATUS_SUCCESS return rocprof.ROCPROFILER_THREAD_TRACE_DECODER_STATUS_SUCCESS
@rocprof.rocprof_trace_decoder_isa_callback_t @rocprof.rocprof_trace_decoder_isa_callback_t
def isa_cb(instr_ptr, mem_size_ptr, size_ptr, pc, _): def isa_cb(instr_ptr, mem_size_ptr, size_ptr, pc, _):
instr, mem_size_ptr[0] = ROCParseCtx.disasms[(unwrap(ROCParseCtx.active_kern), pc.address)] instr, mem_size_ptr[0] = ROCParseCtx.disasms[unwrap(ROCParseCtx.active_kern)][pc.address]
# this is the number of bytes to next instruction, set to 0 for end_pgm # this is the number of bytes to next instruction, set to 0 for end_pgm
if instr == "s_endpgm": mem_size_ptr[0] = 0 if instr == "s_endpgm": mem_size_ptr[0] = 0
+12
View File
@@ -0,0 +1,12 @@
#!/bin/bash
AMD=1 AMD_LLVM=1 python -m pytest -n=1 test/test_ops.py test/test_dtype.py test/test_dtype_alu.py test/test_linearizer.py test/test_randomness.py test/test_jit.py test/test_graph.py test/test_multitensor.py --durations=20
AMD=1 AMD_LLVM=0 python -m pytest -n=1 test/test_ops.py test/test_dtype.py test/test_dtype_alu.py test/test_linearizer.py test/test_randomness.py test/test_jit.py test/test_graph.py test/test_multitensor.py --durations=20
CNT=1 AMD_LLVM=0 DEBUG=2 FP8E4M3=1 HALF=0 BFLOAT16=0 SHOULD_USE_TC=1 python extra/gemm/simple_matmul.py
CNT=1 AMD_LLVM=0 DEBUG=2 FP8E4M3=0 HALF=1 BFLOAT16=0 SHOULD_USE_TC=1 python extra/gemm/simple_matmul.py
CNT=1 AMD_LLVM=0 DEBUG=2 FP8E4M3=0 HALF=0 BFLOAT16=1 SHOULD_USE_TC=1 python extra/gemm/simple_matmul.py
CNT=1 AMD_LLVM=1 DEBUG=2 FP8E4M3=0 HALF=1 BFLOAT16=0 SHOULD_USE_TC=1 python extra/gemm/simple_matmul.py
CNT=1 AMD_LLVM=1 DEBUG=2 FP8E4M3=0 HALF=0 BFLOAT16=1 SHOULD_USE_TC=1 python extra/gemm/simple_matmul.py
CNT=1 AMD_LLVM=1 DEBUG=2 FP8E4M3=1 HALF=0 BFLOAT16=0 SHOULD_USE_TC=1 python extra/gemm/simple_matmul.py
+9 -9
View File
@@ -56,13 +56,13 @@ class Group:
self.ker.push_store(dst_store, dst) self.ker.push_store(dst_store, dst)
return dst.after(dst_store).reshape(dst.shape) return dst.after(dst_store).reshape(dst.shape)
def mma_AB(self, c:UOp|RT, a:UOp|RT, b:UOp|RT, after=True): def mma_AB(self, c:UOp|RT, a:UOp|RT, b:UOp|RT):
c, a, b = cast(UOp, c), cast(UOp, a), cast(UOp, b) c, a, b = cast(UOp, c), cast(UOp, a), cast(UOp, b)
assert self.warps == 1 assert self.warps == 1
for height in self.ker.range(c.shape[-3], track=False): for height in self.ker.range(c.shape[-3], track=False):
for width in self.ker.range(c.shape[-2], track=False): for width in self.ker.range(c.shape[-2], track=False):
for inner in self.ker.range(a.shape[-2], AxisType.REDUCE, track=False): for inner in self.ker.range(a.shape[-2], axis_type=AxisType.REDUCE, track=False):
wmma_arg = ("WMMA_8_16_16_bfloat16_float", (8, 16, 16), dtypes.bfloat16, dtypes.float, "CUDA", 32, (((4, 2), (3, 2), (8, 2)), ((4, 2), (3, 2)), ((4, 2), (3, 2))), ()) wmma_arg = ("WMMA_8_16_16_bfloat16_float", (8, 16, 16), dtypes.bfloat16, dtypes.float, "CUDA", 32, (((4, 2), (3, 2), (8, 2)), ((4, 2), (3, 2)), ((4, 2), (3, 2))), ())
a_in = UOp.vectorize(*[a[height, inner, i] for i in range(8)]) a_in = UOp.vectorize(*[a[height, inner, i] for i in range(8)])
@@ -77,15 +77,15 @@ class Group:
c_store = UOp.group(*c_i).end(height, width, inner) c_store = UOp.group(*c_i).end(height, width, inner)
self.ker.push_store(c_store, c) self.ker.push_store(c_store, c)
return c.after(c_store).reshape(c.shape) if after else c_store return c.after(c_store).reshape(c.shape)
def mma_ABt(self, c:UOp|RT, a:UOp|RT, b:UOp|RT, after=True): def mma_ABt(self, c:UOp|RT, a:UOp|RT, b:UOp|RT):
c, a, b = cast(UOp, c), cast(UOp, a), cast(UOp, b) c, a, b = cast(UOp, c), cast(UOp, a), cast(UOp, b)
assert self.warps == 1 assert self.warps == 1
for height in self.ker.range(c.shape[-3], track=False): for height in self.ker.range(c.shape[-3], track=False):
for width in self.ker.range(c.shape[-2], track=False): for width in self.ker.range(c.shape[-2], track=False):
for inner in self.ker.range(a.shape[-2], AxisType.REDUCE, track=False): for inner in self.ker.range(a.shape[-2], axis_type=AxisType.REDUCE, track=False):
wmma_arg = ("WMMA_8_16_16_bfloat16_float", (8, 16, 16), dtypes.bfloat16, dtypes.float, "CUDA", 32, (((4, 2), (3, 2), (8, 2)), ((4, 2), (3, 2)), ((4, 2), (3, 2))), ()) wmma_arg = ("WMMA_8_16_16_bfloat16_float", (8, 16, 16), dtypes.bfloat16, dtypes.float, "CUDA", 32, (((4, 2), (3, 2), (8, 2)), ((4, 2), (3, 2)), ((4, 2), (3, 2))), ())
a_in = UOp.vectorize(*[a[height, inner, i] for i in range(8)]) a_in = UOp.vectorize(*[a[height, inner, i] for i in range(8)])
@@ -100,7 +100,7 @@ class Group:
c_store = UOp.group(*c_i).end(height, width, inner) c_store = UOp.group(*c_i).end(height, width, inner)
self.ker.push_store(c_store, c) self.ker.push_store(c_store, c)
return c.after(c_store).reshape(c.shape) if after else c_store return c.after(c_store).reshape(c.shape)
map_rid = 400 map_rid = 400
def map(self, a:ALL_TILES, op:Callable[[UOp], UOp]|Callable[[UOp, tuple], UOp]): def map(self, a:ALL_TILES, op:Callable[[UOp], UOp]|Callable[[UOp, tuple], UOp]):
@@ -135,8 +135,8 @@ class Group:
red_reg = red_reg.after(reg_store).reshape(red_reg.shape) red_reg = red_reg.after(reg_store).reshape(red_reg.shape)
for outer in self.ker.range(2, track=False): for outer in self.ker.range(2, track=False):
for width in self.ker.range(src.shape[-2], AxisType.REDUCE, track=False): for width in self.ker.range(src.shape[-2], axis_type=AxisType.REDUCE, track=False):
for inner in self.ker.range(4, AxisType.REDUCE, track=False): for inner in self.ker.range(4, axis_type=AxisType.REDUCE, track=False):
elem_index = inner + 2 * (inner // 2) + outer * 2 elem_index = inner + 2 * (inner // 2) + outer * 2
reg_store = red_reg[outer].store(op(red_reg[outer], src[height, width, elem_index])).end(inner, width, outer) reg_store = red_reg[outer].store(op(red_reg[outer], src[height, width, elem_index])).end(inner, width, outer)
red_reg = red_reg.after(reg_store).reshape(red_reg.shape) red_reg = red_reg.after(reg_store).reshape(red_reg.shape)
@@ -148,7 +148,7 @@ class Group:
# reduce from shared memory # reduce from shared memory
for outer in self.ker.range(2, track=False): for outer in self.ker.range(2, track=False):
for inner in self.ker.range(3, AxisType.REDUCE, track=False): for inner in self.ker.range(3, axis_type=AxisType.REDUCE, track=False):
offset = (self.laneid // 4) * 4 + ((self.laneid + inner + 1) % 4) offset = (self.laneid // 4) * 4 + ((self.laneid + inner + 1) % 4)
reg_store = red_reg[outer].store(op(red_reg[outer], red_local[offset, outer])).end(inner, outer) reg_store = red_reg[outer].store(op(red_reg[outer], red_local[offset, outer])).end(inner, outer)
red_reg = red_reg.after(reg_store).reshape(red_reg.shape) red_reg = red_reg.after(reg_store).reshape(red_reg.shape)
+12 -5
View File
@@ -6,13 +6,15 @@ from extra.thunder.tiny.tk.tiles import GL, ST, RT, RV
class _tk_range: class _tk_range:
user_rid = 0 user_rid = 0
def __init__(self, end:int, axis_type:AxisType): self.end, self.axis_type, self.done = end, axis_type, False def __init__(self, start:int, end:int, step:int, axis_type:AxisType):
self.start, self.end, self.step = start, end, step
self.axis_type, self.done = axis_type, False
def __iter__(self): return self def __iter__(self): return self
def __next__(self): def __next__(self):
if not self.done: if not self.done:
self.done = True self.done = True
_tk_range.user_rid += 1 _tk_range.user_rid += 1
self._rng = UOp.range(self.end, _tk_range.user_rid-1, axis_type=self.axis_type) self._rng = UOp.range(self.end // self.step, _tk_range.user_rid-1, axis_type=self.axis_type) * self.step + self.start
return self._rng return self._rng
raise StopIteration raise StopIteration
@@ -43,8 +45,9 @@ class Kernel(AbstractContextManager):
@property @property
def warpgroup(self): return self.group(4) def warpgroup(self): return self.group(4)
def range(self, end:int, axis_type:AxisType=AxisType.LOOP, track:bool=True): def range(self, start:int, end:int=0, step:int=1, axis_type:AxisType=AxisType.LOOP, track:bool=True):
rng = _tk_range(end, axis_type) if end == 0: start, end = 0, start
rng = _tk_range(start, end, step, axis_type)
if track: self.range_stack.append(rng) if track: self.range_stack.append(rng)
return rng return rng
@@ -80,7 +83,11 @@ class Kernel(AbstractContextManager):
rngs = [] rngs = []
while self.range_stack: rngs.append(self.range_stack.pop(0)._rng) while self.range_stack: rngs.append(self.range_stack.pop(0)._rng)
return self.store_stack.pop()[0]._uop.end(*rngs).sink(arg=KernelInfo(opts_to_apply=())).simplify() last_store = self.store_stack.pop()[0]
if hasattr(last_store, '_uop'): uop = last_store._uop
else: uop = last_store
return uop.end(*rngs).sink(arg=KernelInfo(opts_to_apply=())).simplify()
def endrange(self): def endrange(self):
last_store = self.store_stack.pop() last_store = self.store_stack.pop()
+23 -11
View File
@@ -11,9 +11,9 @@ def unwrap(x):
if isinstance(x, dict): return {k: unwrap(v) for k,v in x.items()} if isinstance(x, dict): return {k: unwrap(v) for k,v in x.items()}
return x return x
def wrap(x, ker, cls): def wrap(x, s):
if isinstance(x, UOp): return cls(x, ker) if isinstance(x, UOp): return s.ruop(x)
if isinstance(x, (list, tuple)): return type(x)(wrap(y, ker, cls) for y in x) if isinstance(x, (list, tuple)): return type(x)(wrap(y, s) for y in x)
return x return x
def autowrap(source_cls, blacklist=None): def autowrap(source_cls, blacklist=None):
@@ -31,10 +31,10 @@ def autowrap(source_cls, blacklist=None):
if callable(val): if callable(val):
@functools.wraps(val) @functools.wraps(val)
def proxy(*args, **kwargs): def proxy(*args, **kwargs):
return wrap(val(*unwrap(args), **unwrap(kwargs)), self.ker, cls) return wrap(val(*unwrap(args), **unwrap(kwargs)), self)
return proxy return proxy
if name in UOp.__slots__: return val if name in UOp.__slots__: return val
return wrap(val, self.ker, cls) return wrap(val, self)
cls.__getattr__ = __getattr__ cls.__getattr__ = __getattr__
for name in dir(source_cls): for name in dir(source_cls):
@@ -46,9 +46,9 @@ def autowrap(source_cls, blacklist=None):
else: else:
original = getattr(source_cls, name) original = getattr(source_cls, name)
if callable(original): if callable(original):
def make_proxy(op_name, func): def make_proxy(_, func):
def proxy(self, *args, **kwargs): def proxy(self, *args, **kwargs):
return wrap(func(self._uop, *unwrap(args), **unwrap(kwargs)), self.ker, cls) return wrap(func(self._uop, *unwrap(args), **unwrap(kwargs)), self)
return proxy return proxy
setattr(cls, name, make_proxy(name, original)) setattr(cls, name, make_proxy(name, original))
@@ -69,7 +69,7 @@ class TileMathMixin(MathMixin):
if isinstance(self, RT) and isinstance(src[0], RV): uop = self.ker.warp.map(self._uop, lambda x, idx: UOp.alu(x, op, inner_op(src[0]._uop[idx[0], 0, (idx[2]%4)//2]))) if isinstance(self, RT) and isinstance(src[0], RV): uop = self.ker.warp.map(self._uop, lambda x, idx: UOp.alu(x, op, inner_op(src[0]._uop[idx[0], 0, (idx[2]%4)//2])))
else: uop = self.ker.warp.map(self._uop, lambda x, idx: UOp.alu(x, op, inner_op(src[0]._uop[*idx]))) else: uop = self.ker.warp.map(self._uop, lambda x, idx: UOp.alu(x, op, inner_op(src[0]._uop[*idx])))
else: raise NotImplementedError else: raise NotImplementedError
return type(self)(uop, self.ker) return self.ruop(uop)
def const_like(self, b): return b def const_like(self, b): return b
# override ops that do compute on the src uop # override ops that do compute on the src uop
@@ -83,6 +83,9 @@ class GL:
def __init__(self, uop, ker): def __init__(self, uop, ker):
self._uop, self.ker = uop, ker self._uop, self.ker = uop, ker
def ruop(self, uop):
return GL(uop, self.ker)
@classmethod @classmethod
def create(cls, shape, dtype, ker): def create(cls, shape, dtype, ker):
uop = ker.alloc(shape, dtype, AddrSpace.GLOBAL) uop = ker.alloc(shape, dtype, AddrSpace.GLOBAL)
@@ -93,6 +96,9 @@ class ST:
def __init__(self, uop, ker): def __init__(self, uop, ker):
self._uop, self.ker = uop, ker self._uop, self.ker = uop, ker
def ruop(self, uop):
return ST(uop, self.ker)
@classmethod @classmethod
def create(cls, shape, dtype, ker): def create(cls, shape, dtype, ker):
uop = ker.alloc(shape, dtype, AddrSpace.LOCAL) uop = ker.alloc(shape, dtype, AddrSpace.LOCAL)
@@ -107,6 +113,9 @@ class RT(TileMathMixin):
def __init__(self, uop, ker): def __init__(self, uop, ker):
self._uop, self.ker = uop, ker self._uop, self.ker = uop, ker
def ruop(self, uop):
return RT(uop, self.ker)
@classmethod @classmethod
def create(cls, shape, dtype, ker): def create(cls, shape, dtype, ker):
assert len(shape) == 2 assert len(shape) == 2
@@ -121,8 +130,11 @@ class RT(TileMathMixin):
@autowrap(UOp) @autowrap(UOp)
class RV(TileMathMixin): class RV(TileMathMixin):
def __init__(self, uop, ker): def __init__(self, uop, layout, ker):
self._uop, self.ker = uop, ker self._uop, self.layout, self.ker = uop, layout, ker
def ruop(self, uop):
return RV(uop, self.layout, self.ker)
@classmethod @classmethod
def create(cls, length, dtype, layout, ker): def create(cls, length, dtype, layout, ker):
@@ -138,6 +150,6 @@ class RV(TileMathMixin):
case _: raise NotImplementedError(f"rv layout {layout} not implemented") case _: raise NotImplementedError(f"rv layout {layout} not implemented")
uop = ker.alloc((outer_dim, inner_dim, 2), dtype, AddrSpace.REG) uop = ker.alloc((outer_dim, inner_dim, 2), dtype, AddrSpace.REG)
return RV(uop, ker) return RV(uop, layout, ker)
ALL_TILES = UOp | GL | ST | RT | RV ALL_TILES = UOp | GL | ST | RT | RV
+1 -1
View File
@@ -8,4 +8,4 @@ if __name__ == "__main__":
parser.add_argument("--dest", type=str, required=True, help="destination path to save the file") parser.add_argument("--dest", type=str, required=True, help="destination path to save the file")
args = parser.parse_args() args = parser.parse_args()
Tensor(bytes.fromhex(args.hash), device="CPU").load(args.len).to(f"disk:{args.dest}").realize() Tensor(bytes.fromhex(args.hash), device="CPU").fs_load(args.len).to(f"disk:{args.dest}").realize()
+7 -5
View File
@@ -1,4 +1,4 @@
import json, multiprocessing import json, multiprocessing, functools
from pathlib import Path from pathlib import Path
from tinygrad.tensor import Tensor from tinygrad.tensor import Tensor
@@ -14,23 +14,25 @@ def fetch_file(item):
path.parent.mkdir(parents=True, exist_ok=True) path.parent.mkdir(parents=True, exist_ok=True)
try: try:
pt = Tensor(bytes.fromhex(h), device="CPU").load(size).to(f"disk:{path.as_posix()}").realize() pt = Tensor(bytes.fromhex(h), device="CPU").fs_load(size).to(f"disk:{path.as_posix()}").realize()
except Exception as e: except Exception as e:
print(f"error fetching {path}, {h}, {size}: {e}") print(f"error fetching {path}, {h}, {size}: {e}")
raise raise
pt.uop.buffer.deallocate() pt.uop.buffer.deallocate()
def fetch_mapping(): def fetch_mapping(h, l):
mapping_tensor = Tensor(bytes.fromhex("d734f5e3be9f1e9d863bfaa4fc6c1ef2")).load(175866113).realize() mapping_tensor = Tensor(bytes.fromhex(h)).fs_load(l).realize()
mapping = mapping_tensor.data().tobytes().decode() mapping = mapping_tensor.data().tobytes().decode()
mapping = json.loads(mapping) mapping = json.loads(mapping)
mapped_files = mapping.items() mapped_files = mapping.items()
return list(mapped_files) return list(mapped_files)
if __name__ == "__main__": if __name__ == "__main__":
h, l = getenv("HASH", "d734f5e3be9f1e9d863bfaa4fc6c1ef2"), getenv("LENGTH", 175866113)
with multiprocessing.Pool(processes=1) as pool: with multiprocessing.Pool(processes=1) as pool:
mapped_files = pool.apply(fetch_mapping) mapped_files = pool.apply(functools.partial(fetch_mapping, h, l))
print(f"fetched mapping for {len(mapped_files)} files") print(f"fetched mapping for {len(mapped_files)} files")
+2 -2
View File
@@ -8,7 +8,7 @@ raid_root = Path("/raid")
def upload_file(path: Path): def upload_file(path: Path):
pt = Tensor(path).realize() pt = Tensor(path).realize()
h = pt.store().realize() h = pt.fs_store().realize()
pt.uop.realized.deallocate() pt.uop.realized.deallocate()
return h.data().hex(), path, pt.nbytes() return h.data().hex(), path, pt.nbytes()
@@ -26,6 +26,6 @@ if __name__ == "__main__":
mapping = json.dumps(mapping).encode() mapping = json.dumps(mapping).encode()
mapping_tensor = Tensor(mapping, device="CPU") mapping_tensor = Tensor(mapping, device="CPU")
h = mapping_tensor.store().realize() h = mapping_tensor.fs_store().realize()
print(f"final hash: {h.data().hex()}, size: {len(mapping)}") print(f"final hash: {h.data().hex()}, size: {len(mapping)}")
+258 -161
View File
@@ -4,10 +4,10 @@
# A006 Lambda argument `input` is shadowing a Python builtin # A006 Lambda argument `input` is shadowing a Python builtin
from tinygrad import Tensor, dtypes, Device from tinygrad import Tensor, dtypes, Device
from tinygrad.uop.ops import Ops from tinygrad.uop.ops import Ops
from tinygrad.helpers import getenv, prod from tinygrad.helpers import getenv, prod, strides_for_shape, argfix
import torch.lib import torch.lib
TORCH_DEBUG = getenv("TORCH_DEBUG") TORCH_DEBUG = getenv("TORCH_DEBUG")
import torch, pathlib, math, operator, functools, inspect import torch, pathlib, math, operator, functools, weakref
torch.autograd.grad_mode.set_multithreading_enabled(False) torch.autograd.grad_mode.set_multithreading_enabled(False)
from tinygrad.dtype import _from_torch_dtype, _to_torch_dtype from tinygrad.dtype import _from_torch_dtype, _to_torch_dtype
@@ -18,7 +18,17 @@ def _to_torch_device(device: str): return torch.device("tiny", int(device.partit
import torch.utils.cpp_extension import torch.utils.cpp_extension
mod = torch.utils.cpp_extension.load(name="custom_device_extension", sources=[str(pathlib.Path(__file__).parent / "wrapped_tensor.cpp")]) mod = torch.utils.cpp_extension.load(name="custom_device_extension", sources=[str(pathlib.Path(__file__).parent / "wrapped_tensor.cpp")])
def wrap(x:Tensor) -> torch.Tensor: return mod.wrap(x, _to_torch_dtype(x.dtype), _to_torch_device(x.device).index) def calculate_storage_offset(x: Tensor) -> int:
offset = 0
for u in x.uop.toposort():
if u.op == Ops.SHRINK:
u_strides = strides_for_shape(u.src[0].shape)
for i, (start, _) in enumerate(u.marg): offset += start * u_strides[i]
return offset
def wrap(x: Tensor) -> torch.Tensor:
x._strides = strides_for_shape(x.shape) # always recalculate
if (not hasattr(x, '_storage_offset')) or (not x.uop.is_realized): x._storage_offset = calculate_storage_offset(x)
return mod.wrap(x, _to_torch_dtype(x.dtype), _to_torch_device(x.device).index)
def unwrap(x:torch.Tensor) -> Tensor: def unwrap(x:torch.Tensor) -> Tensor:
assert isinstance(x, torch.Tensor), f"x isn't {type(x)}" assert isinstance(x, torch.Tensor), f"x isn't {type(x)}"
return mod.unwrap(x) return mod.unwrap(x)
@@ -35,17 +45,20 @@ torch.utils.generate_methods_for_privateuse1_backend()
aten = torch.ops.aten aten = torch.ops.aten
# track view relationships for in place operations # track view relationships for in place operations
def is_view(tensor: Tensor): return hasattr(tensor, "_view_base")
def canonical_base(view: Tensor): return getattr(view, "_view_base", view) def canonical_base(view: Tensor): return getattr(view, "_view_base", view)
def derived_views(base: Tensor): return [t for tref in getattr(base, "_views", set()) if (t:=tref()) is not None] def derived_views(base: Tensor): return [t for tref in getattr(base, "_views", set()) if (t:=tref()) is not None]
def unwrap_args(args, kwargs):
return [unwrap(x) if isinstance(x, torch.Tensor) else x for x in args], {k:unwrap(v) if isinstance(v, torch.Tensor) else v for k,v in kwargs.items()}
def wrap_view_op(fn): def wrap_view_op(fn):
def _wrap(*args,**kwargs): @functools.wraps(fn)
args = [unwrap(x) if isinstance(x, torch.Tensor) else x for x in args] def _wrap(*args, **kwargs):
kwargs = {k:unwrap(v) if isinstance(v, torch.Tensor) else v for k,v in kwargs.items()} args, kwargs = unwrap_args(args, kwargs)
ret = fn(*args,**kwargs) ret = fn(*args, **kwargs)
ret._view_base = base = canonical_base(args[0]) base = canonical_base(args[0])
if not hasattr(base, "_views"): base._views = set() ret._view_base = base
base._views = getattr(base, "_views", set())
base._views.add(weakref.ref(ret)) base._views.add(weakref.ref(ret))
ret._view_ops = _get_view_ops(args[0]) + [(fn, args[1:], kwargs)]
return wrap(ret) return wrap(ret)
return _wrap return _wrap
@@ -58,48 +71,83 @@ view_ops = {
"aten.transpose.int": Tensor.transpose, "aten.transpose.int": Tensor.transpose,
"aten.squeeze.dim": Tensor.squeeze, "aten.squeeze.dim": Tensor.squeeze,
"aten.unsqueeze": Tensor.unsqueeze, "aten.unsqueeze": Tensor.unsqueeze,
"aten.detach": Tensor.detach,
"aten.select.int": lambda self, dim, idx: self[(slice(None),) * (dim%self.ndim) + (idx,)], "aten.select.int": lambda self, dim, idx: self[(slice(None),) * (dim%self.ndim) + (idx,)],
} "aten.permute": Tensor.permute,
"aten.alias": lambda self: self,
}
# torch 2.10 handles this natively
if tuple(map(int, torch.__version__.split('.')[:2])) < (2, 10): view_ops.update({"aten.detach": Tensor.detach})
for k,v in view_ops.items(): torch.library.impl(k.replace("aten.", "aten::"), "privateuseone")(wrap_view_op(v)) for k,v in view_ops.items(): torch.library.impl(k.replace("aten.", "aten::"), "privateuseone")(wrap_view_op(v))
# in place operations with views def _get_view_ops(view): return getattr(view, "_view_ops", [])
def realize_with_views(self: Tensor, views: Tensor):
if not self.uop.st.contiguous: self.replace(self.contiguous()) def _apply_view_ops(target, ops):
self.replace(self.clone().realize()) for fn, args, kwargs in ops: target = fn(target, *args, **kwargs)
for v in views: return target
if v.uop.base.op is Ops.BUFFER_VIEW: continue # skip subbuffer, we just use the real buffer view
ret = self # similar to https://github.com/pytorch/pytorch/blob/main/aten/src/ATen/InferSize.h
st = ShapeTracker(self.uop.st.views + v.uop.st.views) # TODO: is this right? def _reshape_target_shape(shape:tuple[int, ...], args) -> tuple[int, ...]|None:
for mo in cached_to_movement_ops(self.shape, st): ret = apply_mop(ret, mo) if not (req := argfix(*args)): return None
v.replace(ret) new_shape, infer_idx = [], -1
def maybe_realize_storage(self: Tensor) -> bool: for i, s in enumerate(req):
if realize:=is_view(self): realize_with_views((base:=canonical_base(self)), derived_views(base)) if s is None: s = shape[i] if i < len(shape) else None
return realize if not isinstance(s, int): return None
def inplace_fn(outvars: str|list[str]): if s == -1:
if type(outvars) is str: outvars = [outvars] if infer_idx != -1: return None
def decorator(fn): infer_idx = len(new_shape)
sig = inspect.signature(fn) new_shape.append(s)
def wrapper(*args, **kwargs): total = prod(shape)
bound = sig.bind(*args, **kwargs) if infer_idx != -1:
outs = [kwargs.get(v, bound.arguments.get(v)) for v in outvars] known = prod(x for x in new_shape if x != -1)
outs = [unwrap(o) if isinstance(o, torch.Tensor) else o for o in outs] if known == 0:
realize = any(maybe_realize_storage(o) for o in outs) if total != 0: return None
ret = fn(*args, **kwargs) new_shape[infer_idx] = 0
if realize: Tensor.realize(*(o for o in outs)) else: new_shape[infer_idx] = total // known
return ret return tuple(new_shape) if prod(new_shape) == total else None
return wrapper
return decorator # TODO: can we get rid of this? only for test_flatten_reshape_add
def _try_simple_reshape_view_write(base: Tensor, view: Tensor, val: Tensor) -> bool:
if not (ops := _get_view_ops(view)): return False
shapes = [base.shape]
for fn, args, _ in ops:
if fn is Tensor.reshape:
if not (next_shape := _reshape_target_shape(shapes[-1], args)): return False
shapes.append(next_shape)
if shapes[-1] != view.shape: return False
for s in reversed(shapes[:-1]): val = val.reshape(s)
base.assign(val)
return True
def _view_write(base: Tensor, view: Tensor, value: Tensor) -> None:
val = value if value.dtype == base.dtype else value.cast(base.dtype)
if view.shape == base.shape: return base.assign(val)
if _try_simple_reshape_view_write(base, view, val): return
idx_base = Tensor.arange(base.numel(), device=base.device, dtype=dtypes.int32).reshape(base.shape)
idx_view = _apply_view_ops(idx_base, _get_view_ops(view)).reshape(-1)
flat_base = base.reshape(base.numel()).contiguous()
flat_base[idx_view] = val.reshape(-1)
base.assign(flat_base.reshape(base.shape))
def _apply_inplace(target: Tensor, value: Tensor) -> None:
val = value if value.dtype == target.dtype else value.cast(target.dtype)
base = canonical_base(target)
views = derived_views(base)
if not views: return target.assign(val)
view_ops_map = {v: _get_view_ops(v) for v in views}
if target is base or target.uop is base.uop: base.assign(val)
else: _view_write(base, target, val)
for v in views: v.replace(_apply_view_ops(base, view_ops_map[v]))
# *** bad functions on CPU *** # *** bad functions on CPU ***
@torch.library.impl("aten::_index_put_impl_", "privateuseone") @torch.library.impl("aten::_index_put_impl_", "privateuseone")
@inplace_fn("self")
def _index_put_impl_(self, indices, values, accumulate=False, unsafe=False): def _index_put_impl_(self, indices, values, accumulate=False, unsafe=False):
# TODO: move to tinygrad # TODO: move to tinygrad
ret = aten._index_put_impl_(self.cpu(), [x.cpu() if isinstance(x, torch.Tensor) else None for x in indices], values.cpu(), accumulate, unsafe).to(self.device) ret = aten._index_put_impl_(self.cpu(), [x.cpu() if isinstance(x, torch.Tensor) else None for x in indices], values.cpu(), accumulate, unsafe).to(self.device)
return wrap(unwrap(self).assign(unwrap(ret))) unwrap(self).assign(unwrap(ret))
return self
@torch.library.impl("aten::index_put", "privateuseone") @torch.library.impl("aten::index_put", "privateuseone")
def index_put(self, indices, values, accumulate=False): def index_put(self, indices, values, accumulate=False):
@@ -150,43 +198,23 @@ for i in [
def index_tensor(x, y): def index_tensor(x, y):
return wrap(unwrap(x)[[unwrap(_y.to(x.device)) if _y is not None else slice(None) for _y in y]]) return wrap(unwrap(x)[[unwrap(_y.to(x.device)) if _y is not None else slice(None) for _y in y]])
@torch.library.impl("aten::zero_", "privateuseone")
@inplace_fn("x")
def zero_(x):
if TORCH_DEBUG: print(f"zero_ {x.shape}")
tt = unwrap(x)
tt.assign(tt.zeros_like())
@torch.library.impl("aten::fill_.Scalar", "privateuseone")
@inplace_fn("x")
def fill_scalar(x, y):
if TORCH_DEBUG: print(f"fill_.Scalar {x.shape} {y}")
tt = unwrap(x)
tt.assign(tt.full_like(y))
@torch.library.impl("aten::_local_scalar_dense", "privateuseone") @torch.library.impl("aten::_local_scalar_dense", "privateuseone")
def _local_scalar_dense(tensor): return unwrap(tensor).item() def _local_scalar_dense(tensor): return unwrap(tensor).item()
@functools.cache
def cached_to_movement_ops(shape, st) -> list:
mops = to_movement_ops(st)
if mops[0] == (MovementOps.RESHAPE, shape): mops = mops[1:]
return mops
from tinygrad.shape.shapetracker import ShapeTracker, View
from extra.to_movement_ops import to_movement_ops, apply_mop, MovementOps
@wrap_view_op @wrap_view_op
def _as_strided(tensor:Tensor, size, stride, storage_offset=None): def _as_strided(tensor:Tensor, size, stride, storage_offset=0):
# multiple as_strided do not compound base = getattr(tensor, "_as_strided_base", canonical_base(tensor)).flatten()
base = canonical_base(tensor) if prod(size) == 1: return base[storage_offset].reshape(size)
# TODO: this is heavyweight indices = Tensor.zeros(size, dtype=dtypes.int32, device=base.device) + storage_offset
st = ShapeTracker(base.uop.st.views + (View.create(tuple(size), tuple(stride), storage_offset),)) for dim, (sz, st) in enumerate(zip(size, stride)):
ret = base if st != 0:
if TORCH_DEBUG >= 1: print("**** as_strided", tensor.shape, size, stride, st) dim_range = Tensor.arange(sz, device=base.device, dtype=dtypes.int32) * st
if prod(size) == 1: return ret.flatten()[storage_offset].reshape(size) shape_for_broadcast = [1] * dim + [sz] + [1] * (len(size) - dim - 1)
for mo in cached_to_movement_ops(tuple(base.shape), st): ret = apply_mop(ret, mo) indices = indices + dim_range.reshape(shape_for_broadcast)
return ret result = base[indices.flatten()].reshape(size)
result._as_strided_base = base
return result
@torch.library.impl("aten::as_strided", "privateuseone") @torch.library.impl("aten::as_strided", "privateuseone")
def as_strided(tensor:torch.Tensor, size, stride, storage_offset=None): def as_strided(tensor:torch.Tensor, size, stride, storage_offset=None):
@@ -245,15 +273,14 @@ def convolution_overrideable(input, weight, bias, stride, padding, dilation, tra
if TORCH_DEBUG >= 1: if TORCH_DEBUG >= 1:
print(f"convolution {input.shape=} {weight.shape=} {stride=} {padding=} {dilation=} {transposed=} {output_padding=} {groups=}") print(f"convolution {input.shape=} {weight.shape=} {stride=} {padding=} {dilation=} {transposed=} {output_padding=} {groups=}")
input, weight, bias = unwrap(input), unwrap(weight), unwrap(bias) if bias is not None else None input, weight, bias = unwrap(input), unwrap(weight), unwrap(bias) if bias is not None else None
# TODO: fix test_biased_conv2d fails without realize() if not transposed: return wrap(input.conv2d(weight, bias, groups=groups, stride=stride, dilation=dilation, padding=padding))
if not transposed: return wrap(input.conv2d(weight, bias, groups=groups, stride=stride, dilation=dilation, padding=padding).realize()) return wrap(input.conv_transpose2d(weight, bias, groups=groups, stride=stride, dilation=dilation, padding=padding, output_padding=output_padding))
return wrap(input.conv_transpose2d(weight, bias, groups=groups, stride=stride, dilation=dilation, padding=padding, output_padding=output_padding).realize())
@torch.library.impl("aten::convolution_backward_overrideable", "privateuseone") @torch.library.impl("aten::convolution_backward_overrideable", "privateuseone")
def convolution_backward_overrideable(grad_out, input, weight, stride, padding, dilation, transposed, output_padding, groups, output_mask): def convolution_backward_overrideable(grad_out, input, weight, stride, padding, dilation, transposed, output_padding, groups, output_mask):
if TORCH_DEBUG >= 1: if TORCH_DEBUG >= 1:
print(f"convolution_backward {input.shape=} {weight.shape=} {stride=} {padding=} {dilation=} {transposed=} {output_padding=} {groups=}") print(f"convolution_backward {input.shape=} {weight.shape=} {stride=} {padding=} {dilation=} {transposed=} {output_padding=} {groups=}")
grad_out, input, weight, bias = unwrap(grad_out), unwrap(input), unwrap(weight), Tensor.zeros(weight.shape[0], device=_from_torch_device(weight.device)) grad_out, input, weight, bias = unwrap(grad_out).detach(), unwrap(input).detach(), unwrap(weight).detach(), Tensor.zeros(weight.shape[0], device=_from_torch_device(weight.device))
if not transposed: out = Tensor.conv2d(input, weight, bias, groups=groups, stride=stride, dilation=dilation, padding=padding) if not transposed: out = Tensor.conv2d(input, weight, bias, groups=groups, stride=stride, dilation=dilation, padding=padding)
else: else:
bias = Tensor.zeros(weight.shape[1] * groups) bias = Tensor.zeros(weight.shape[1] * groups)
@@ -315,55 +342,57 @@ for i,pre in enumerate(["", "bi", "tri"]):
torch.library.impl(f"aten::_upsample_nearest_exact{i+1}d", "privateuseone")(functools.partial(upsample, mode="nearest-exact")) torch.library.impl(f"aten::_upsample_nearest_exact{i+1}d", "privateuseone")(functools.partial(upsample, mode="nearest-exact"))
@torch.library.impl("aten::scatter_add.out", "privateuseone") @torch.library.impl("aten::scatter_add.out", "privateuseone")
@inplace_fn("out")
def scatter_add(self, dim, index, src, out): def scatter_add(self, dim, index, src, out):
self, index, src, out = unwrap(self), unwrap(index), unwrap(src), unwrap(out) self, index, src, out_unwrapped = unwrap(self), unwrap(index), unwrap(src), unwrap(out)
if self.shape == (): return wrap(out.assign(src)) if self.shape == (): _apply_inplace(out_unwrapped, src)
return wrap(out.assign(Tensor.scatter_reduce(self, dim, index, src, reduce='sum'))) else: _apply_inplace(out_unwrapped, Tensor.scatter_reduce(self, dim, index, src, reduce='sum'))
return out
@torch.library.impl("aten::_copy_from", "privateuseone") def _copy_between_devices(src, dest, cast_dtype, to_device, non_blocking=False):
def _copy_from(src: torch.Tensor, dest, non_blocking=False):
realize = dest.is_tiny and maybe_realize_storage(unwrap(dest))
cast_dtype = _from_torch_dtype(dest.dtype)
if src.is_tiny and dest.is_tiny: if src.is_tiny and dest.is_tiny:
to_device = _from_torch_device(dest.device) src_t, dest_t = unwrap(src), unwrap(dest)
src,dest = unwrap(src),unwrap(dest) if dest_t.uop.is_contiguous() or dest_t.uop.is_realized: src_t = src_t.contiguous()
# TODO we need to properly match dest shape and strides, not blindly assign _apply_inplace(dest_t, src_t.cast(cast_dtype).to(to_device))
if dest.uop.st.contiguous or dest.uop.is_realized: src = src.contiguous() # this only solves some cases
dest.assign(src.cast(cast_dtype).to(to_device))
if realize: Tensor.realize(dest)
elif src.is_tiny and dest.is_cpu: elif src.is_tiny and dest.is_cpu:
# TODO: is there a better way?
dest.resize_(src.numel()).resize_(src.shape) dest.resize_(src.numel()).resize_(src.shape)
dest.copy_(torch.from_numpy(unwrap(src).cast(cast_dtype).numpy())) dest.copy_(torch.from_numpy(unwrap(src).cast(cast_dtype).numpy()))
elif src.is_cpu and dest.is_tiny: elif src.is_cpu and dest.is_tiny:
to_device = _from_torch_device(dest.device)
# TODO we need to properly match dest shape and strides, not blindly assign
unwrap(dest).assign(Tensor(src.numpy()).cast(cast_dtype).to(to_device)) unwrap(dest).assign(Tensor(src.numpy()).cast(cast_dtype).to(to_device))
if realize: Tensor.realize(unwrap(dest))
else: else:
raise NotImplementedError(f"can't copy from {src.device} -> {dest.device}") raise NotImplementedError(f"can't copy from {src.device} -> {dest.device}")
@torch.library.impl("aten::_copy_from", "privateuseone")
def _copy_from(src: torch.Tensor, dest, non_blocking=False):
cast_dtype = _from_torch_dtype(dest.dtype)
to_device = _from_torch_device(dest.device)
_copy_between_devices(src, dest, cast_dtype, to_device, non_blocking)
return dest
@torch.library.impl("aten::copy_", "privateuseone")
def copy_(self, src, non_blocking=False):
cast_dtype = _from_torch_dtype(self.dtype)
to_device = _from_torch_device(self.device)
_copy_between_devices(src, self, cast_dtype, to_device, non_blocking)
return self
@torch.library.impl("aten::cat.out", "privateuseone") @torch.library.impl("aten::cat.out", "privateuseone")
@inplace_fn("out")
def cat_out(tensors, dim=0, out=None): def cat_out(tensors, dim=0, out=None):
unwrap(out).assign(Tensor.cat(*[unwrap(x) for x in tensors], dim=dim)) _apply_inplace(unwrap(out), Tensor.cat(*[unwrap(x) for x in tensors], dim=dim))
return out
@torch.library.impl("aten::topk.values", "privateuseone") @torch.library.impl("aten::topk.values", "privateuseone")
@inplace_fn(["values", "indices"])
def topk_values(input, k, dim=None, largest=True, sorted=True, values=None, indices=None): def topk_values(input, k, dim=None, largest=True, sorted=True, values=None, indices=None):
out_values, out_indices = unwrap(input).topk(k, dim if dim is not None else -1, largest, sorted) out_values, out_indices = unwrap(input).topk(k, dim if dim is not None else -1, largest, sorted)
unwrap(values).assign(out_values) _apply_inplace(unwrap(values), out_values)
unwrap(indices).assign(out_indices.cast(dtypes.int64)) _apply_inplace(unwrap(indices), out_indices.cast(dtypes.int64))
return wrap(out_values), wrap(out_indices) return values, indices
@torch.library.impl("aten::sort.values_stable", "privateuseone") @torch.library.impl("aten::sort.values_stable", "privateuseone")
@inplace_fn(["values", "indices"])
def sort_values(input, dim=-1, descending=False, stable=True, values=None, indices=None): def sort_values(input, dim=-1, descending=False, stable=True, values=None, indices=None):
out_values, out_indices = unwrap(input).sort(dim, descending) out_values, out_indices = unwrap(input).sort(dim, descending)
unwrap(values).assign(out_values) _apply_inplace(unwrap(values), out_values)
unwrap(indices).assign(out_indices.cast(dtypes.int64)) _apply_inplace(unwrap(indices), out_indices.cast(dtypes.int64))
return wrap(out_values), wrap(out_indices) return values, indices
@torch.library.impl("aten::_linalg_svd", "privateuseone") @torch.library.impl("aten::_linalg_svd", "privateuseone")
def _linalg_svd(self, full_matrices=False): def _linalg_svd(self, full_matrices=False):
@@ -373,7 +402,6 @@ def _linalg_svd(self, full_matrices=False):
# register some decompositions # register some decompositions
from torch._decomp import get_decompositions from torch._decomp import get_decompositions
decomps = [ decomps = [
aten.native_batch_norm, aten.native_batch_norm_backward,
aten.native_layer_norm_backward, aten.native_layer_norm_backward,
aten.linalg_cross, aten.linalg_cross,
aten.addmm, aten.addmm,
@@ -510,7 +538,6 @@ tiny_backend_out = {**{f"aten.{x}.out":getattr(Tensor,x) for x in simple_tensor_
# we add the "out" here # we add the "out" here
def wrap_out(f): def wrap_out(f):
@inplace_fn("out")
def _wrap_out(*args, **kwargs): def _wrap_out(*args, **kwargs):
out = kwargs.pop('out') out = kwargs.pop('out')
assigned = f(*args, **kwargs) assigned = f(*args, **kwargs)
@@ -518,22 +545,33 @@ def wrap_out(f):
assert out.shape == assigned.shape, f"shape mismatch: {assigned.shape} -> {out.shape}" assert out.shape == assigned.shape, f"shape mismatch: {assigned.shape} -> {out.shape}"
assert out.device == assigned.device, f"device mismatch: {assigned.device} -> {out.device}" assert out.device == assigned.device, f"device mismatch: {assigned.device} -> {out.device}"
assert out.dtype == assigned.dtype, f"dtype mismatch: {assigned.dtype} -> {out.dtype}" assert out.dtype == assigned.dtype, f"dtype mismatch: {assigned.dtype} -> {out.dtype}"
if out.uop.is_realized: assigned = assigned.contiguous() # TODO: how does this map to torch's semantics
return out.assign(assigned) return out.assign(assigned)
return _wrap_out return _wrap_out
def _inplace_op(t, new_value):
if not hasattr(t, "_view_base") and not getattr(canonical_base(t), "_views", set()): t.replace(new_value)
else: _apply_inplace(t, new_value)
return t
tiny_backend = {**{k:wrap_out(v) for k,v in tiny_backend_out.items()}, **{ tiny_backend = {**{k:wrap_out(v) for k,v in tiny_backend_out.items()}, **{
"aten.remainder.Scalar_Tensor": lambda x,y: x%y, "aten.remainder.Scalar_Tensor": lambda x,y: x%y,
"aten.floor_divide": lambda x,y: x//y, "aten.floor_divide": lambda x,y: x//y,
"aten.floor_divide_.Tensor": inplace_fn("x")(lambda x,y: x.assign(x//y)), "aten.floor_divide_.Tensor": lambda x,y: x//y,
# TODO: use tinygrad methods, but they require x to be unsigned # TODO: use tinygrad methods, but they require x to be unsigned
"aten.__lshift__.Scalar": lambda x,y: x*(2**y), "aten.__lshift__.Scalar": lambda x,y: x*(2**y),
"aten.__ilshift__.Scalar": inplace_fn("x")(lambda x,y: x.assign(x*(2**y))), "aten.__ilshift__.Scalar": lambda x,y: x*(2**y),
"aten.__rshift__.Scalar": lambda x,y: x//(2**y), "aten.__rshift__.Scalar": lambda x,y: x//(2**y),
"aten.__irshift__.Scalar": inplace_fn("x")(lambda x,y: x.assign(x//(2**y))), "aten.__irshift__.Scalar": lambda x,y: x//(2**y),
# inplace ops using replace for fusion
"aten.zero_": lambda x: x.zeros_like(),
"aten.fill_.Scalar": lambda x, y: x.full_like(y),
"aten.add_.Tensor": lambda self, other, alpha=1.0: self + other * alpha,
"aten.add_.Scalar": lambda self, other, alpha=1.0: self + other * alpha,
"aten.mul_.Tensor": lambda self, other: self * other,
"aten.mul_.Scalar": lambda self, other: self * other,
# relu doesn't have an out form? # relu doesn't have an out form?
"aten.relu": Tensor.relu, "aten.relu": Tensor.relu,
"aten.relu_": inplace_fn("x")(lambda x: x.assign(x.relu())), "aten.relu_": lambda x: x.relu(),
"aten.mean": Tensor.mean, "aten.mean": Tensor.mean,
"aten.mean.dim": Tensor.mean, "aten.mean.dim": Tensor.mean,
"aten.min": Tensor.min, "aten.min": Tensor.min,
@@ -554,19 +592,17 @@ tiny_backend = {**{k:wrap_out(v) for k,v in tiny_backend_out.items()}, **{
"aten.repeat": lambda x,*repeats: Tensor.repeat(x,*repeats).contiguous(), # not a view "aten.repeat": lambda x,*repeats: Tensor.repeat(x,*repeats).contiguous(), # not a view
"aten._softmax": lambda self,dim,half_to_float: self.softmax(dim), "aten._softmax": lambda self,dim,half_to_float: self.softmax(dim),
"aten._log_softmax": lambda self,dim,half_to_float: self.log_softmax(dim), "aten._log_softmax": lambda self,dim,half_to_float: self.log_softmax(dim),
"aten.random_": inplace_fn("self")(lambda self: "aten.random_": lambda self: Tensor.randint(*self.shape, low=dtypes.min(self.dtype), high=dtypes.max(self.dtype), device=self.device, dtype=self.dtype),
self.assign(Tensor.randint(*self.shape, low=dtypes.min(self.dtype), high=dtypes.max(self.dtype), device=self.device, dtype=self.dtype))), "aten.random_.from": lambda self, from_, to: Tensor.randint(*self.shape, low=from_, high=to, device=self.device, dtype=self.dtype),
"aten.random_.from": inplace_fn("self")(lambda self, from_, to: "aten.uniform_": lambda self, low=0, high=1: Tensor.uniform(*self.shape, low=low, high=high, dtype=self.dtype),
self.assign(Tensor.randint(*self.shape, low=from_, high=to, device=self.device, dtype=self.dtype))), "aten.normal_": lambda self, mean=0, std=1: Tensor.normal(*self.shape, mean=mean, std=std, dtype=self.dtype),
"aten.uniform_": inplace_fn("self")(lambda self, low=0, high=1: self.assign(Tensor.uniform(*self.shape, low=low, high=high, dtype=self.dtype))),
"aten.normal_": inplace_fn("self")(lambda self, mean=0, std=1: self.assign(Tensor.normal(*self.shape, mean=mean, std=std, dtype=self.dtype))),
# these don't work in out form, they have size 0 # these don't work in out form, they have size 0
"aten.abs": Tensor.abs, "aten.abs": Tensor.abs,
"aten.logical_not": Tensor.logical_not, "aten.logical_not": Tensor.logical_not,
"aten.logical_or_": inplace_fn("x")(lambda x, y: x.assign(x | y)), "aten.logical_or_": lambda x, y: x | y,
"aten.multinomial": Tensor.multinomial, "aten.multinomial": Tensor.multinomial,
"aten.masked_fill_.Scalar": inplace_fn("self")(lambda self, mask, value: self.assign(self.masked_fill(mask, value))), "aten.masked_fill_.Scalar": lambda self, mask, value: self.masked_fill(mask, value),
"aten.masked_fill_.Tensor": inplace_fn("self")(lambda self, mask, value: self.assign(self.masked_fill(mask, value))), "aten.masked_fill_.Tensor": lambda self, mask, value: self.masked_fill(mask, value),
"aten.masked_fill.Scalar": Tensor.masked_fill, "aten.masked_fill.Scalar": Tensor.masked_fill,
"aten.masked_fill.Tensor": Tensor.masked_fill, "aten.masked_fill.Tensor": Tensor.masked_fill,
"aten.masked_select": Tensor.masked_select, "aten.masked_select": Tensor.masked_select,
@@ -580,7 +616,7 @@ tiny_backend = {**{k:wrap_out(v) for k,v in tiny_backend_out.items()}, **{
"aten.asinh": Tensor.asinh, "aten.asinh": Tensor.asinh,
"aten.mul": Tensor.mul, "aten.mul": Tensor.mul,
"aten.atanh": Tensor.atanh, "aten.atanh": Tensor.atanh,
"aten.fill_.Tensor": Tensor.full, # TODO: looks wrong "aten.fill_.Tensor": lambda self, value: Tensor.full(self.shape, value.reshape(()).item(), device=self.device, dtype=self.dtype),
"aten.flip": Tensor.flip, "aten.flip": Tensor.flip,
"aten.scatter_reduce.two": Tensor.scatter_reduce, "aten.scatter_reduce.two": Tensor.scatter_reduce,
"aten.squeeze_.dim": lambda self, dim: self.replace(self.squeeze(dim), allow_shape_mismatch=True), # TODO: inplace view op, here? "aten.squeeze_.dim": lambda self, dim: self.replace(self.squeeze(dim), allow_shape_mismatch=True), # TODO: inplace view op, here?
@@ -601,20 +637,51 @@ tiny_backend = {**{k:wrap_out(v) for k,v in tiny_backend_out.items()}, **{
"aten.unfold": Tensor.unfold, "aten.unfold": Tensor.unfold,
}} }}
# operations that need inplace treatment (use _inplace_op instead of wrap_fxn) AKA return original tensor
inplace_ops = {
"aten.zero_",
"aten.fill_.Scalar",
"aten.fill_.Tensor",
"aten.add_.Tensor",
"aten.add_.Scalar",
"aten.mul_.Tensor",
"aten.mul_.Scalar",
"aten.floor_divide_.Tensor",
"aten.__ilshift__.Scalar",
"aten.__irshift__.Scalar",
"aten.relu_",
"aten.random_",
"aten.random_.from",
"aten.uniform_",
"aten.normal_",
"aten.logical_or_",
"aten.masked_fill_.Scalar",
"aten.masked_fill_.Tensor",
}
def wrap_fxn(k,f): def wrap_fxn(k,f):
def nf(*args, **kwargs): def nf(*args, **kwargs):
if TORCH_DEBUG: if TORCH_DEBUG:
print(k, len(args), [x.shape if isinstance(x, torch.Tensor) else x for x in args], print(k, len(args), [x.shape if isinstance(x, torch.Tensor) else x for x in args],
{k:v.shape if isinstance(v, torch.Tensor) else v for k,v in kwargs.items()}) {k:v.shape if isinstance(v, torch.Tensor) else v for k,v in kwargs.items()})
args = [unwrap(x) if isinstance(x, torch.Tensor) else x for x in args] args, kwargs = unwrap_args(args, kwargs)
kwargs = {k:unwrap(v) if isinstance(v, torch.Tensor) else v for k,v in kwargs.items()}
out = f(*args, **kwargs) out = f(*args, **kwargs)
if isinstance(out, Tensor): return wrap(out) if isinstance(out, Tensor): return wrap(out)
elif isinstance(out, tuple): return tuple(wrap(x) for x in out) elif isinstance(out, tuple): return tuple(wrap(x) for x in out)
else: raise RuntimeError(f"unknown output type {type(out)}") else: raise RuntimeError(f"unknown output type {type(out)}")
return nf return nf
for k,v in tiny_backend.items(): torch.library.impl(k.replace("aten.", "aten::"), "privateuseone")(wrap_fxn(k,v)) def wrap_inplace(k,f):
def nf(*args, **kwargs):
orig = args[0]
args, kwargs = unwrap_args(args, kwargs)
_inplace_op(args[0], f(*args, **kwargs))
return orig
return nf
for k,v in tiny_backend.items():
wrapper = wrap_inplace if k in inplace_ops else wrap_fxn
torch.library.impl(k.replace("aten.", "aten::"), "privateuseone")(wrapper(k,v))
@torch.library.impl("aten::equal", "privateuseone") @torch.library.impl("aten::equal", "privateuseone")
def equal(x: torch.Tensor, y: torch.Tensor): return (x==y).all().item() def equal(x: torch.Tensor, y: torch.Tensor): return (x==y).all().item()
@@ -628,42 +695,72 @@ if TORCH_DEBUG:
return func(*args, **(kwargs or {})) return func(*args, **(kwargs or {}))
(_dispatch_log:=DispatchLog()).__enter__() # NOTE: must be kept alive (_dispatch_log:=DispatchLog()).__enter__() # NOTE: must be kept alive
# NOTE: patch torch optimizer step to avoid continously growing the computation graph # this implementation is needed to allow the batchnorm kernels to fuse in e.g. mnist training
import weakref # aten::native_batch_norm does more than Tensor.batchnorm
_torch_modules_with_buffers: weakref.WeakSet[torch.nn.Module] = weakref.WeakSet() @torch.library.impl("aten::native_batch_norm", "privateuseone")
def register_torch_buffer(mod, _name, _buffer): _torch_modules_with_buffers.add(mod) def native_batch_norm(input, weight, bias, running_mean, running_var, training, momentum, eps):
def get_real_tinygrad_buffers(): input_t, weight_t, bias_t = unwrap(input), unwrap(weight) if weight is not None else None, unwrap(bias) if bias is not None else None
res = set() running_mean_t, running_var_t = unwrap(running_mean) if running_mean is not None else None, unwrap(running_var) if running_var is not None else None
for mod in _torch_modules_with_buffers: if training:
for _,b in mod.named_buffers(recurse=False): batch_var, batch_mean = input_t.var_mean(axis=tuple(x for x in range(input_t.ndim) if x != 1), correction=0)
if b is not None and b.is_tiny: batch_invstd = batch_var.add(eps).rsqrt()
res.add(unwrap(b)) out = input_t.batchnorm(weight_t, bias_t, batch_mean, batch_invstd)
return res if running_mean_t is not None and running_var_t is not None:
torch.nn.modules.module.register_module_buffer_registration_hook(register_torch_buffer) numel_ratio = input_t.numel() / (input_t.numel() - input_t.shape[1])
running_mean_t.assign((1 - momentum) * running_mean_t + momentum * batch_mean.detach())
running_var_t.assign((1 - momentum) * running_var_t + momentum * numel_ratio * batch_var.detach())
return wrap(out), wrap(batch_mean), wrap(batch_invstd)
else:
out = input_t.batchnorm(weight_t, bias_t, running_mean_t, running_var_t.add(eps).rsqrt())
return wrap(out), wrap(running_mean_t), wrap(running_var_t.add(eps).rsqrt())
from torch.nn.modules import Module @torch.library.impl("aten::native_batch_norm_backward", "privateuseone")
def param_hook(_grad): def native_batch_norm_backward(grad_out, input, weight, running_mean, running_var, save_mean, save_invstd, train, eps, output_mask):
if _grad is not None and _grad.is_tiny: Tensor.realize(unwrap(_grad)) grad_out_t, input_t = unwrap(grad_out), unwrap(input)
def module_hook(module:Module, _name, _submodule): weight_t = unwrap(weight) if weight is not None else None
for param in _submodule.parameters(recurse=False): save_mean_t = unwrap(save_mean)
if param.requires_grad: param.register_hook(param_hook) save_invstd_t = unwrap(save_invstd)
torch.nn.modules.module.register_module_module_registration_hook(module_hook) out = input_t.batchnorm(weight_t, None, save_mean_t, save_invstd_t)
targets = [t for t, m in zip([input_t, weight_t], output_mask[:2]) if t is not None and m]
if targets:
grads = out.gradient(*targets, gradient=grad_out_t)
grad_input = grads.pop(0) if output_mask[0] else None
grad_weight = grads.pop(0) if output_mask[1] and weight_t is not None else None
else:
grad_input, grad_weight = None, None
grad_bias = grad_out_t.sum(axis=tuple(x for x in range(grad_out_t.ndim) if x != 1)) if output_mask[2] else None
return (wrap(grad_input) if grad_input is not None else None,
wrap(grad_weight) if grad_weight is not None else None,
wrap(grad_bias) if grad_bias is not None else None)
def realize_optimizer_step(optimizer: torch.optim.Optimizer, *args, **kwargs): # _pad_circular is not CompositeImplicitAutograd (unlike reflect/replicate pad)
tinygrad_tensors = [] # we need torch.autograd.Function with explicit AutogradPrivateUse1 registration
for param_group in optimizer.param_groups: class _PadCircular(torch.autograd.Function):
for param in param_group["params"]: @staticmethod
if param is None: continue def forward(ctx, input, padding):
tinygrad_tensors.append(param.data) ctx.save_for_backward(input)
for state_dict in optimizer.state.values(): ctx.padding = padding
for _, value in state_dict.items(): return pad_forward(input, padding, mode="circular")
if torch.is_tensor(value): tinygrad_tensors.append(value) @staticmethod
real_tinygrad_tensors = [unwrap(x) for x in tinygrad_tensors if x.is_tiny] def backward(ctx, grad_output):
real_tinygrad_tensors += get_real_tinygrad_buffers() input, = ctx.saved_tensors
if len(real_tinygrad_tensors): Tensor.realize(*real_tinygrad_tensors) return pad_backward(grad_output, input, ctx.padding, mode="circular"), None
_optimizer_init = torch.optim.Optimizer.__init__ @torch.library.impl("aten::_pad_circular", "privateuseone")
def _optimizer_patched_init(self, *args, **kwargs): def _pad_circular(self, padding): return _PadCircular.apply(self, padding)
_optimizer_init(self, *args, **kwargs)
self.register_step_post_hook(realize_optimizer_step) @torch.library.impl("aten::_pad_circular", "AutogradPrivateUse1")
torch.optim.Optimizer.__init__ = _optimizer_patched_init def _pad_circular_autograd(self, padding): return _PadCircular.apply(self, padding)
# only needed for test_diag_backward_gradient_values
# was going through torch before, but now we are using tinygrad directly and tracking views
# Tensor.diagonal does not support all cases tests in the tests
@torch.library.impl("aten::diagonal", "privateuseone")
@wrap_view_op
def diagonal(self, offset=0, dim1=0, dim2=1):
if offset != 0: raise NotImplementedError(f"diagonal with {offset=} not implemented")
dim1, dim2 = dim1 % self.ndim, dim2 % self.ndim
if dim1 != self.ndim - 2 or dim2 != self.ndim - 1: raise NotImplementedError(f"diagonal with {dim1=}, {dim2=} not implemented, only last two dims supported")
batch_shape, m, n = self.shape[:-2], self.shape[-2], self.shape[-1]
diag_len = min(m, n)
return self.reshape(*batch_shape, m*n).pad(tuple((0,0) for _ in batch_shape) + ((0, diag_len),)).reshape(*batch_shape, diag_len, n+1)[..., :, 0]
+10 -2
View File
@@ -1,12 +1,13 @@
from PIL import Image from PIL import Image
from tinygrad.helpers import getenv from tinygrad.helpers import getenv, GlobalCounters
import torch, torchvision, pathlib import torch, torchvision, pathlib, warnings
import torchvision.transforms as transforms import torchvision.transforms as transforms
import extra.torch_backend.backend import extra.torch_backend.backend
device = "tiny" device = "tiny"
torch.set_default_device(device) torch.set_default_device(device)
if __name__ == "__main__": if __name__ == "__main__":
GlobalCounters.reset()
img = Image.open(pathlib.Path(__file__).parent.parent.parent / "test/models/efficientnet/Chicken.jpg").convert('RGB') img = Image.open(pathlib.Path(__file__).parent.parent.parent / "test/models/efficientnet/Chicken.jpg").convert('RGB')
transform = transforms.Compose([ transform = transforms.Compose([
transforms.Resize(256), transforms.CenterCrop(224), transforms.ToTensor(), transforms.Resize(256), transforms.CenterCrop(224), transforms.ToTensor(),
@@ -19,3 +20,10 @@ if __name__ == "__main__":
out = model(img).detach().cpu().numpy() out = model(img).detach().cpu().numpy()
print("output:", out.shape, out.argmax()) print("output:", out.shape, out.argmax())
assert out.argmax() == 7 # cock assert out.argmax() == 7 # cock
kernel_count = GlobalCounters.kernel_count
assert kernel_count > 0, "No kernels, test failed"
expected_kernels = 228
expectation = f"ResNet18 kernels are {kernel_count} vs {expected_kernels} expected."
if kernel_count < expected_kernels: warnings.warn(f"{expectation} Expectation can be lowered.", UserWarning)
assert kernel_count <= expected_kernels, f"{expectation}"
+669 -3
View File
@@ -2,7 +2,7 @@
import unittest import unittest
import torch import torch
import numpy as np import numpy as np
from tinygrad.helpers import getenv, Context, GlobalCounters from tinygrad.helpers import getenv, GlobalCounters
if getenv("TINY_BACKEND2"): if getenv("TINY_BACKEND2"):
import extra.torch_backend.backend2 import extra.torch_backend.backend2
device = "cpu" device = "cpu"
@@ -25,7 +25,7 @@ class TestTorchBackend(unittest.TestCase):
a = torch.ones(4, device=device) a = torch.ones(4, device=device)
np.testing.assert_equal(a.cpu().numpy(), [1,1,1,1]) np.testing.assert_equal(a.cpu().numpy(), [1,1,1,1])
def test_numpy_ones(self): def test_numpy_ones_int32(self):
a = torch.ones(4, dtype=torch.int32, device=device) a = torch.ones(4, dtype=torch.int32, device=device)
assert a.dtype == torch.int32 assert a.dtype == torch.int32
np.testing.assert_equal(a.cpu().numpy(), [1,1,1,1]) np.testing.assert_equal(a.cpu().numpy(), [1,1,1,1])
@@ -219,7 +219,6 @@ class TestTorchBackend(unittest.TestCase):
a = torch.ones(4, device=device) a = torch.ones(4, device=device)
print(str(a)) print(str(a))
@unittest.skip("failed")
def test_floor_div(self): def test_floor_div(self):
a = torch.tensor([10., 7., 5.], device=device) a = torch.tensor([10., 7., 5.], device=device)
b = torch.tensor([3., 2., 2.], device=device) b = torch.tensor([3., 2., 2.], device=device)
@@ -248,5 +247,672 @@ class TestTorchBackend(unittest.TestCase):
def test_diagonal_rectangular(self): self._test_diagonal(4, 5, 6) def test_diagonal_rectangular(self): self._test_diagonal(4, 5, 6)
def test_diagonal_4d(self): self._test_diagonal(2, 3, 4, 5) def test_diagonal_4d(self): self._test_diagonal(2, 3, 4, 5)
def test_pad_circular_simple(self):
a = torch.arange(4, dtype=torch.float32, device=device).reshape(1,1,2,2)
padded = torch.nn.functional.pad(a, (1,1,1,1), mode="circular")
expected = np.array([[[[3.,2.,3.,2.], [1.,0.,1.,0.], [3.,2.,3.,2.], [1.,0.,1.,0.]]]], dtype=np.float32)
np.testing.assert_allclose(padded.cpu().numpy(), expected)
def test_pad_circular_backward(self):
a = torch.arange(4, dtype=torch.float32, device=device).reshape(1,1,2,2).requires_grad_(True)
padded = torch.nn.functional.pad(a, (1,1,1,1), mode="circular")
loss = padded.sum()
loss.backward()
expected_grad = np.array([[[[4., 4.], [4., 4.]]]], dtype=np.float32)
np.testing.assert_allclose(a.grad.cpu().numpy(), expected_grad)
def test_matmul_backward(self):
x = torch.randn(3, 4, device=device, dtype=torch.float32, requires_grad=True)
y = torch.randn(4, 5, device=device, dtype=torch.float32, requires_grad=True)
z = (x @ y).sum()
z.backward()
assert x.grad is not None
assert y.grad is not None
assert x.grad.shape == x.shape
assert y.grad.shape == y.shape
def test_matmul_broadcast_backward(self):
x = torch.randn(2, 3, 4, device=device, dtype=torch.float32, requires_grad=True)
y = torch.randn(4, 5, device=device, dtype=torch.float32, requires_grad=True)
z = (x @ y).sum()
z.backward()
assert x.grad is not None
assert y.grad is not None
assert x.grad.shape == x.shape
assert y.grad.shape == y.shape
def test_diag_vector_to_matrix(self):
vec = torch.tensor([1., 2., 3., 4., 5.], dtype=torch.float32, device=device)
mat = torch.diag(vec)
expected = np.diag([1., 2., 3., 4., 5.])
np.testing.assert_allclose(mat.cpu().numpy(), expected, rtol=1e-5)
assert mat.shape == (5, 5)
def test_diagonal_matrix_to_vector(self):
mat = torch.tensor([[1., 2., 3.], [4., 5., 6.], [7., 8., 9.]], dtype=torch.float32, device=device)
vec = torch.linalg.diagonal(mat)
expected = np.array([1., 5., 9.])
np.testing.assert_allclose(vec.cpu().numpy(), expected, rtol=1e-5)
assert vec.shape == (3,)
def test_permute_2(self):
a = torch.randn(2, 3, 4, dtype=torch.float32, device=device)
b = a.permute(2, 0, 1)
assert b.shape == (4, 2, 3)
np.testing.assert_equal(b.cpu().numpy(), a.cpu().numpy().transpose(2, 0, 1))
def test_batchnorm_unsqueeze(self):
bn = torch.nn.BatchNorm2d(4).to(device)
x = torch.randn(8, 4, 3, 3, device=device)
out = bn(x)
self.assertEqual(out.shape, x.shape)
def test_slice_inplace_zero(self):
a = torch.ones((3, 3), device=device)
b = a[1:, 1:]
b.zero_()
expected = np.array([[1., 1., 1.],
[1., 0., 0.],
[1., 0., 0.]])
np.testing.assert_equal(a.cpu().numpy(), expected)
def test_slice_inplace_fill(self):
a = torch.ones((3, 3), device=device)
b = a[1:, 1:]
b.fill_(5.0)
expected = np.array([[1., 1., 1.],
[1., 5., 5.],
[1., 5., 5.]])
np.testing.assert_equal(a.cpu().numpy(), expected)
def test_fill_tensor_value(self):
a = torch.zeros((2, 2), dtype=torch.float32, device=device)
value = torch.tensor(3, dtype=torch.int64, device=device)
a.fill_(value)
expected = np.full((2, 2), 3, dtype=np.float32)
np.testing.assert_equal(a.cpu().numpy(), expected)
def test_slice_inplace_mul(self):
a = torch.ones((3, 3), device=device)
b = a[1:, 1:]
b *= 2
expected = np.array([[1., 1., 1.],
[1., 2., 2.],
[1., 2., 2.]])
np.testing.assert_equal(a.cpu().numpy(), expected)
def test_permute_slice_zero(self):
a = torch.ones((3, 3), device=device)
b = a[1:, 1:].permute(1, 0)
b.zero_()
expected = np.array([[1., 1., 1.],
[1., 0., 0.],
[1., 0., 0.]])
np.testing.assert_equal(a.cpu().numpy(), expected)
def test_permute_slice_mul(self):
a = torch.ones((3, 3), device=device)
b = a[1:, 1:].permute(1, 0)
b *= 2
expected = np.array([[1., 1., 1.],
[1., 2., 2.],
[1., 2., 2.]])
np.testing.assert_equal(a.cpu().numpy(), expected)
def test_simple_slice_setitem(self):
a = torch.tensor([10, 20, 30], device=device)
a[1] = 99
np.testing.assert_equal(a.cpu().numpy(), [10, 99, 30])
def test_2d_slice_setitem(self):
a = torch.zeros((3, 3), device=device)
a[1, 2] = 99
self.assertEqual(a[1, 2].item(), 99)
self.assertEqual(a.sum().item(), 99)
def test_view_copy(self):
a = torch.tensor([10, 20, 30], device=device)
view = a[1]
view.copy_(torch.tensor(88, device=device))
np.testing.assert_equal(a.cpu().numpy(), [10, 88, 30])
def test_diag_2d_input(self):
a = torch.tensor([[1, 2, 3], [4, 5, 6], [7, 8, 9]], device=device)
d = torch.diag(a)
np.testing.assert_equal(d.cpu().numpy(), [1, 5, 9])
def test_diag_1d_input(self):
a = torch.tensor([1, 2, 3], device=device)
d = torch.diag(a)
expected = [[1, 0, 0], [0, 2, 0], [0, 0, 3]]
np.testing.assert_equal(d.cpu().numpy(), expected)
def test_permute_view_tracking(self):
a = torch.ones((2, 3, 4), device=device)
b = a.permute(2, 0, 1)
self.assertEqual(b.shape, (4, 2, 3))
def test_detach_view_creation(self):
a = torch.tensor([1.0, 2.0, 3.0], device=device)
b = a.detach()
np.testing.assert_equal(b.cpu().numpy(), [1.0, 2.0, 3.0])
def test_view_zero_inplace(self):
a = torch.ones((4, 4), device=device)
view = a[1:3, 1:3]
view.zero_()
self.assertEqual(view.sum().item(), 0)
def test_view_fill_inplace(self):
a = torch.zeros((4, 4), device=device)
view = a[1:3, 1:3]
view.fill_(5)
self.assertEqual(view.sum().item(), 20)
def test_permute_contiguous(self):
a = torch.tensor([[1, 2], [3, 4]], device=device)
b = a.permute(1, 0)
c = b.contiguous()
expected = [[1, 3], [2, 4]]
np.testing.assert_equal(c.cpu().numpy(), expected)
def test_diag_2d_extract_diagonal(self):
a = torch.tensor([[1, 2], [3, 4]], device=device)
result = torch.diag(a)
np.testing.assert_equal(result.cpu().numpy(), [1, 4])
def test_slice_inplace_multiply_offset_preservation(self):
a = torch.tensor([1, 2, 3], device=device)
a[1:] *= 2
np.testing.assert_equal(a.cpu().numpy(), [1, 4, 6])
def test_slice_inplace_mul_pattern(self):
a = torch.tensor([1, 2, 3, 4], device=device)
a[:2] *= 3
a[2:] *= 2
np.testing.assert_equal(a.cpu().numpy(), [3, 6, 6, 8])
def test_chained_slice_column(self):
a = torch.arange(16, dtype=torch.float32, device=device).reshape(4, 4)
torch_res = a[:, 1:2][:, 0:1].cpu().numpy()
cpu_res = torch.arange(16, dtype=torch.float32).reshape(4, 4)[:, 1:2][:, 0:1].numpy()
np.testing.assert_equal(torch_res, cpu_res)
def test_slice_with_step(self):
a = torch.arange(20, dtype=torch.float32, device=device)
torch_res = a[::2][1:4].cpu().numpy()
cpu_res = torch.arange(20, dtype=torch.float32)[::2][1:4].numpy()
np.testing.assert_equal(torch_res, cpu_res)
def test_slice_negative_dim(self):
a = torch.arange(13, dtype=torch.int32, device=device).repeat(8, 1)
torch_chunks = a.chunk(3, -1)
cpu_chunks = torch.arange(13, dtype=torch.int32).repeat(8, 1).chunk(3, -1)
assert len(torch_chunks) == len(cpu_chunks)
for i in range(len(torch_chunks)):
np.testing.assert_equal(torch_chunks[i].cpu().numpy(), cpu_chunks[i].numpy())
def test_dot_vector_matrix(self):
a = torch.arange(65, dtype=torch.float32, device=device)
b = torch.arange(65*45, dtype=torch.float32, device=device).reshape(65, 45)
torch_res = a.matmul(b).reshape(-1).cpu().numpy()
cpu_res = torch.arange(65, dtype=torch.float32).matmul(torch.arange(65*45, dtype=torch.float32).reshape(65, 45)).numpy()
np.testing.assert_equal(torch_res, cpu_res)
def test_alias_passthrough(self):
a = torch.randn(3, 3, device=device)
alias_view = torch.ops.aten.alias(a)
alias_view += 1
np.testing.assert_equal(a.cpu().numpy(), alias_view.cpu().numpy())
def test_split_simple_vector(self):
a = torch.arange(10, dtype=torch.float32, device=device)
torch_chunks = a.split([1,4,5])
cpu_chunks = torch.arange(10, dtype=torch.float32).split([1,4,5])
for tc, cc in zip(torch_chunks, cpu_chunks):
np.testing.assert_equal(tc.cpu().numpy(), cc.cpu().numpy())
def test_split_matches_torch(self):
a = torch.arange(10, dtype=torch.float32, device=device)
torch_chunks = a.split([1,4,5])
tiny_chunks = [chunk.cpu().numpy() for chunk in torch_chunks]
cpu_chunks = [torch.arange(10, dtype=torch.float32).split([1,4,5])[i].numpy() for i in range(3)]
for tr, cr in zip(tiny_chunks, cpu_chunks): np.testing.assert_equal(tr, cr)
def test_sum_matches_torch(self):
a = torch.arange(6, dtype=torch.float32, device=device).reshape(2,3)
torch_res = a.sum().cpu().numpy()
cpu_res = torch.arange(6, dtype=torch.float32).reshape(2,3).sum().numpy()
np.testing.assert_equal(torch_res, cpu_res)
def test_view_matches_torch(self):
a = torch.arange(6, dtype=torch.float32, device=device)
torch_res = a.view(2, 3).cpu().numpy()
cpu_res = torch.arange(6, dtype=torch.float32).view(2, 3).numpy()
np.testing.assert_equal(torch_res, cpu_res)
def test_view_zero_with_indices(self):
a = torch.tensor([1, 2, 3, 4], device=device)
a[1:3].zero_()
np.testing.assert_equal(a.cpu().numpy(), [1, 0, 0, 4])
def test_view_fill_with_indices(self):
a = torch.tensor([1, 2, 3, 4], device=device)
a[::2].fill_(9)
np.testing.assert_equal(a.cpu().numpy(), [9, 2, 9, 4])
def test_nested_slice_inplace_ops(self):
a = torch.tensor([1, 2, 3, 4, 5, 6], device=device)
a[:3] += 10
a[3:] *= 2
np.testing.assert_equal(a.cpu().numpy(), [11, 12, 13, 8, 10, 12])
def test_diag_1d(self):
a = torch.tensor([1, 2, 3], device=device)
result = torch.diag(a)
expected = [[1, 0, 0], [0, 2, 0], [0, 0, 3]]
np.testing.assert_equal(result.cpu().numpy(), expected)
def test_diag_backward(self):
a = torch.randn(5, dtype=torch.float32, device=device, requires_grad=True)
b = torch.diag(a)
b.sum().backward()
assert a.grad is not None
def test_diagonal(self):
a = torch.tensor([[1., 2., 3.], [4., 5., 6.], [7., 8., 9.]], dtype=torch.float32, device=device, requires_grad=True)
b = torch.diagonal(a)
expected = torch.tensor([1., 5., 9.], dtype=torch.float32)
self.assertEqual(b.shape, (3,))
np.testing.assert_allclose(b.detach().cpu().numpy(), expected.numpy(), rtol=1e-5)
def test_diagonal_backward(self):
a = torch.randn(5, 5, dtype=torch.float32, device=device, requires_grad=True)
b = torch.diagonal(a)
b.sum().backward()
assert a.grad is not None
def test_expand_backward(self):
a = torch.randn(4, 3, 1, 6, dtype=torch.float32, device=device, requires_grad=True)
b = a.expand(4, 3, 2, 6)
b.sum().backward()
assert a.grad is not None
def test_einsum_backward(self):
a = torch.randn(10, 10, dtype=torch.float32, device=device, requires_grad=True)
b = torch.einsum('ij->ji', a)
b.sum().backward()
assert a.grad is not None
def test_diag_backward_gradient_values(self):
a = torch.tensor([1.0, 2.0, 3.0], dtype=torch.float32, device=device, requires_grad=True)
b = torch.diag(a)
loss = b.sum()
loss.backward()
expected_grad = torch.ones(3, dtype=torch.float32)
np.testing.assert_allclose(a.grad.cpu().numpy(), expected_grad.numpy(), rtol=1e-5)
def test_diag_backward_gradient_values_2d_to_1d(self):
a = torch.tensor([[1.0, 2.0, 3.0],
[4.0, 5.0, 6.0],
[7.0, 8.0, 9.0]], dtype=torch.float32, device=device, requires_grad=True)
b = torch.diagonal(a)
loss = b.sum()
loss.backward()
expected_grad = torch.tensor([[1.0, 0.0, 0.0],
[0.0, 1.0, 0.0],
[0.0, 0.0, 1.0]], dtype=torch.float32)
np.testing.assert_allclose(a.grad.cpu().numpy(), expected_grad.numpy(), rtol=1e-5)
def test_expand_backward_gradient_values(self):
a = torch.tensor([[1.0], [2.0], [3.0]], dtype=torch.float32, device=device, requires_grad=True)
b = a.expand(3, 4)
loss = b.sum()
loss.backward()
expected_grad = torch.tensor([[4.0], [4.0], [4.0]], dtype=torch.float32)
np.testing.assert_allclose(a.grad.cpu().numpy(), expected_grad.numpy(), rtol=1e-5)
def test_expand_backward_with_leading_dims(self):
a = torch.tensor([[1.0, 2.0]], dtype=torch.float32, device=device, requires_grad=True)
b = a.expand(3, 1, 2)
loss = b.sum()
loss.backward()
expected_grad = torch.tensor([[3.0, 3.0]], dtype=torch.float32)
np.testing.assert_allclose(a.grad.cpu().numpy(), expected_grad.numpy(), rtol=1e-5)
def test_diag_2d_to_1d_backward(self):
a = torch.tensor([[1.0, 2.0], [3.0, 4.0]], dtype=torch.float32, device=device, requires_grad=True)
b = torch.diag(a)
loss = b.sum()
loss.backward()
expected_grad = torch.tensor([[1.0, 0.0], [0.0, 1.0]], dtype=torch.float32)
np.testing.assert_allclose(a.grad.cpu().numpy(), expected_grad.numpy(), rtol=1e-5)
def test_expand_complex_backward(self):
a = torch.tensor([[[1.0, 2.0]]], dtype=torch.float32, device=device, requires_grad=True)
b = a.expand(2, 3, 2)
loss = b.sum()
loss.backward()
expected_grad = torch.tensor([[[6.0, 6.0]]], dtype=torch.float32)
np.testing.assert_allclose(a.grad.cpu().numpy(), expected_grad.numpy(), rtol=1e-5)
def test_diag_backward_with_scaling(self):
a = torch.tensor([1.0, 2.0, 3.0], dtype=torch.float32, device=device, requires_grad=True)
b = torch.diag(a)
loss = (b * torch.tensor([[2.0, 0.0, 0.0],
[0.0, 3.0, 0.0],
[0.0, 0.0, 4.0]], device=device)).sum()
loss.backward()
expected_grad = torch.tensor([2.0, 3.0, 4.0], dtype=torch.float32)
np.testing.assert_allclose(a.grad.cpu().numpy(), expected_grad.numpy(), rtol=1e-5)
def test_repeat_basic(self):
a = torch.tensor([1, 2, 3], dtype=torch.float32, device=device)
b = a.repeat(2, 1)
expected = torch.tensor([[1, 2, 3], [1, 2, 3]], dtype=torch.float32)
np.testing.assert_equal(b.cpu().numpy(), expected.numpy())
def test_repeat_multidim(self):
a = torch.arange(6, dtype=torch.float32, device=device).reshape(2, 3)
b = a.repeat(2, 3)
expected = torch.arange(6, dtype=torch.float32).reshape(2, 3).repeat(2, 3)
np.testing.assert_equal(b.cpu().numpy(), expected.numpy())
def test_repeat_backward(self):
a = torch.tensor([[1.0, 2.0]], dtype=torch.float32, device=device, requires_grad=True)
b = a.repeat(3, 2)
loss = b.sum()
loss.backward()
expected_grad = torch.tensor([[6.0, 6.0]], dtype=torch.float32)
np.testing.assert_allclose(a.grad.cpu().numpy(), expected_grad.numpy(), rtol=1e-5)
def test_cumsum_1d(self):
a = torch.tensor([1, 2, 3, 4], dtype=torch.float32, device=device)
b = torch.cumsum(a, dim=0)
expected = torch.tensor([1, 3, 6, 10], dtype=torch.float32)
np.testing.assert_equal(b.cpu().numpy(), expected.numpy())
def test_cumsum_2d(self):
a = torch.arange(12, dtype=torch.float32, device=device).reshape(3, 4)
b = torch.cumsum(a, dim=0)
expected = torch.arange(12, dtype=torch.float32).reshape(3, 4).cumsum(dim=0)
np.testing.assert_equal(b.cpu().numpy(), expected.numpy())
c = torch.cumsum(a, dim=1)
expected = torch.arange(12, dtype=torch.float32).reshape(3, 4).cumsum(dim=1)
np.testing.assert_equal(c.cpu().numpy(), expected.numpy())
def test_cumsum_backward(self):
a = torch.tensor([1.0, 2.0, 3.0, 4.0], dtype=torch.float32, device=device, requires_grad=True)
b = torch.cumsum(a, dim=0)
loss = b.sum()
loss.backward()
expected_grad = torch.tensor([4.0, 3.0, 2.0, 1.0], dtype=torch.float32)
np.testing.assert_allclose(a.grad.cpu().numpy(), expected_grad.numpy(), rtol=1e-5)
def test_constant_pad_nd_1d(self):
a = torch.tensor([1, 2, 3], dtype=torch.float32, device=device)
b = torch.nn.functional.pad(a, (1, 2), mode='constant', value=0)
expected = torch.tensor([0, 1, 2, 3, 0, 0], dtype=torch.float32)
np.testing.assert_equal(b.cpu().numpy(), expected.numpy())
def test_constant_pad_nd_2d(self):
a = torch.arange(6, dtype=torch.float32, device=device).reshape(2, 3)
b = torch.nn.functional.pad(a, (1, 1, 1, 1), mode='constant', value=0)
expected = torch.nn.functional.pad(torch.arange(6, dtype=torch.float32).reshape(2, 3), (1, 1, 1, 1), mode='constant', value=0)
np.testing.assert_equal(b.cpu().numpy(), expected.numpy())
def test_constant_pad_nd_2d_backward(self):
a = torch.tensor([[1.0, 2.0], [3.0, 4.0]], dtype=torch.float32, device=device, requires_grad=True)
b = torch.nn.functional.pad(a, (1, 1, 1, 1), mode='constant', value=0)
loss = b.sum()
loss.backward()
expected_grad = torch.ones((2, 2), dtype=torch.float32)
np.testing.assert_allclose(a.grad.cpu().numpy(), expected_grad.numpy(), rtol=1e-5)
def test_negative_strides_cumsum_backward(self):
a = torch.randn(5, device=device, requires_grad=True)
b = torch.cumsum(a, dim=0)
b.sum().backward()
grad = a.grad.cpu().numpy()
self.assertEqual(len(grad), 5)
def test_cumsum_fix_gradient_values(self):
a = torch.tensor([1.0, 2.0, 3.0, 4.0], dtype=torch.float32, device=device, requires_grad=True)
b = torch.cumsum(a, dim=0)
loss = b.sum()
loss.backward()
expected = np.array([4.0, 3.0, 2.0, 1.0])
np.testing.assert_allclose(a.grad.cpu().numpy(), expected, rtol=1e-5)
def test_diag_1d_to_2d(self):
a = torch.tensor([1.0, 2.0, 3.0], dtype=torch.float32, device=device, requires_grad=True)
b = torch.diag(a)
expected = [[1, 0, 0], [0, 2, 0], [0, 0, 3]]
np.testing.assert_equal(b.detach().cpu().numpy(), expected)
def test_diag_2d_to_1d(self):
c = torch.tensor([[1, 2, 3], [4, 5, 6], [7, 8, 9]], dtype=torch.float32, device=device)
d = torch.diag(c)
np.testing.assert_equal(d.cpu().numpy(), [1, 5, 9])
def test_biased_conv2d(self):
# Test case for two sequential conv2d with same weights/bias and ReLU in between, this is as special case from test_ops.py
torch.manual_seed(0)
C = 8
x_cpu = torch.randn(1, C, 5, 5, requires_grad=True)
w_cpu = torch.randn(C, C, 1, 1, requires_grad=True)
b_cpu = torch.randn(C, requires_grad=True)
x_tiny = x_cpu.detach().to(device).requires_grad_(True)
w_tiny = w_cpu.detach().to(device).requires_grad_(True)
b_tiny = b_cpu.detach().to(device).requires_grad_(True)
out_cpu = torch.nn.functional.conv2d(torch.nn.functional.conv2d(x_cpu, w_cpu, b_cpu).relu(), w_cpu, b_cpu)
out_tiny = torch.nn.functional.conv2d(torch.nn.functional.conv2d(x_tiny, w_tiny, b_tiny).relu(), w_tiny, b_tiny)
grad_out = torch.randn_like(out_cpu)
out_cpu.backward(grad_out)
out_tiny.backward(grad_out.to(device))
np.testing.assert_allclose(x_tiny.grad.cpu().numpy(), x_cpu.grad.numpy(), atol=1e-4, rtol=1e-3)
np.testing.assert_allclose(w_tiny.grad.cpu().numpy(), w_cpu.grad.numpy(), atol=1e-4, rtol=1e-3)
np.testing.assert_allclose(b_tiny.grad.cpu().numpy(), b_cpu.grad.numpy(), atol=1e-4, rtol=1e-3)
from tinygrad import Tensor
class TestBackendHelpers(unittest.TestCase):
def test_calculate_storage_offset_no_shrink(self):
t = Tensor.ones(3, 4)
assert extra.torch_backend.backend.calculate_storage_offset(t) == 0
def test_calculate_storage_offset_with_shrink(self):
t = Tensor.ones(10, 10)[2:5, 3:7]
# strides for (10, 10) are [10, 1]
# offset = 2*10 + 3*1 = 23
assert extra.torch_backend.backend.calculate_storage_offset(t) == 23
def test_calculate_storage_offset_multiple_shrinks(self):
t = Tensor.ones(5, 6, 7)[1:3, 2:4, 3:5]
# strides for (5, 6, 7) are [42, 7, 1]
# offset = 1*42 + 2*7 + 3*1 = 42 + 14 + 3 = 59
assert extra.torch_backend.backend.calculate_storage_offset(t) == 59
def test_calculate_storage_offset_with_reshape(self):
t = Tensor.ones(10, 10)
orig_offset = extra.torch_backend.backend.calculate_storage_offset(t)
assert orig_offset == 0
t = t.reshape(100)
assert extra.torch_backend.backend.calculate_storage_offset(t) == orig_offset
def test_slice_values_match_torch(self):
torch_cpu = torch.arange(100, dtype=torch.float32).reshape(10, 10)
torch_tiny = torch_cpu.to(device)
sliced_cpu = torch_cpu[2:5, 3:7]
sliced_tiny = torch_tiny[2:5, 3:7]
np.testing.assert_equal(sliced_tiny.cpu().numpy(), sliced_cpu.numpy())
def test_slice_values_match_torch_3d(self):
torch_cpu_3d = torch.arange(210, dtype=torch.float32).reshape(5, 6, 7)
torch_tiny_3d = torch_cpu_3d.to(device)
sliced_cpu_3d = torch_cpu_3d[1:3, 2:4, 3:5]
sliced_tiny_3d = torch_tiny_3d[1:3, 2:4, 3:5]
np.testing.assert_equal(sliced_tiny_3d.cpu().numpy(), sliced_cpu_3d.numpy())
def test_topk_out(self):
a = torch.tensor([1, 3, 2, 4], device=device)
values = torch.empty(2, device=device)
indices = torch.empty(2, dtype=torch.int64, device=device)
ret_values, ret_indices = torch.topk(a, k=2, out=(values, indices))
np.testing.assert_equal(values.cpu().numpy(), [4, 3])
np.testing.assert_equal(indices.cpu().numpy(), [3, 1])
assert ret_values is values
assert ret_indices is indices
def test_sort_out(self):
a = torch.tensor([3, 1, 4, 2], device=device)
values = torch.empty(4, device=device)
indices = torch.empty(4, dtype=torch.int64, device=device)
ret_values, ret_indices = torch.sort(a, out=(values, indices))
np.testing.assert_equal(values.cpu().numpy(), [1, 2, 3, 4])
np.testing.assert_equal(indices.cpu().numpy(), [1, 3, 0, 2])
assert ret_values is values
assert ret_indices is indices
def test_cat_out(self):
a = torch.tensor([1, 2], device=device)
b = torch.tensor([3, 4], device=device)
out = torch.empty(4, device=device)
ret = torch.cat([a, b], out=out)
np.testing.assert_equal(out.cpu().numpy(), [1, 2, 3, 4])
assert ret is out
def test_scatter_add_out(self):
src = torch.tensor([[1, 2, 3], [4, 5, 6]], device=device, dtype=torch.float32)
index = torch.tensor([[0, 1, 2], [0, 1, 2]], device=device)
input = torch.zeros(3, 3, device=device, dtype=torch.float32)
out = torch.zeros(3, 3, device=device, dtype=torch.float32)
ret = torch.scatter_add(input, 0, index, src, out=out)
expected = torch.tensor([[5, 0, 0], [0, 7, 0], [0, 0, 9]], dtype=torch.float32)
np.testing.assert_allclose(out.cpu().numpy(), expected.cpu().numpy())
assert ret is out
def test_floor_divide_inplace_identity(self):
x = torch.tensor([10, 20, 30, 40], dtype=torch.int32, device=device)
y = torch.tensor([2, 4, 5, 8], dtype=torch.int32, device=device)
ret = x.floor_divide_(y)
assert ret is x
np.testing.assert_equal(x.cpu().numpy(), [5, 5, 6, 5])
def test_lshift_inplace_identity(self):
x = torch.tensor([1, 2, 3, 4], dtype=torch.int32, device=device)
ret = x.__ilshift__(2)
assert ret is x
np.testing.assert_equal(x.cpu().numpy(), [4, 8, 12, 16])
def test_rshift_inplace_identity(self):
x = torch.tensor([16, 32, 48, 64], dtype=torch.int32, device=device)
ret = x.__irshift__(2)
assert ret is x
np.testing.assert_equal(x.cpu().numpy(), [4, 8, 12, 16])
def test_relu_inplace_identity(self):
x = torch.tensor([-1.0, 2.0, -3.0, 4.0], device=device)
ret = x.relu_()
assert ret is x
np.testing.assert_equal(x.cpu().numpy(), [0.0, 2.0, 0.0, 4.0])
def test_random_inplace_identity(self):
x = torch.zeros(10, dtype=torch.int32, device=device)
ret = x.random_()
assert ret is x
assert x.shape == (10,)
def test_random_from_inplace_identity(self):
x = torch.zeros(10, dtype=torch.int32, device=device)
ret = x.random_(5, 10)
assert ret is x
# values should be in range [5, 10)
assert torch.all(x >= 5).item() and torch.all(x < 10).item()
def test_uniform_inplace_identity(self):
x = torch.zeros(10, device=device)
ret = x.uniform_(0.0, 1.0)
assert ret is x
# values should be in range [0, 1)
assert torch.all(x >= 0.0).item() and torch.all(x < 1.0).item()
def test_normal_inplace_identity(self):
x = torch.zeros(100, device=device)
ret = x.normal_(0.0, 1.0)
assert ret is x
# just check that values changed from zeros
assert not torch.all(x == 0.0).item()
def test_logical_or_inplace_identity(self):
x = torch.tensor([True, False, True, False], device=device)
y = torch.tensor([False, False, True, True], device=device)
ret = x.logical_or_(y)
assert ret is x
np.testing.assert_equal(x.cpu().numpy(), [True, False, True, True])
def test_masked_fill_scalar_inplace_identity(self):
x = torch.tensor([1.0, 2.0, 3.0, 4.0], device=device)
mask = torch.tensor([True, False, True, False], device=device)
ret = x.masked_fill_(mask, 0.0)
assert ret is x
np.testing.assert_equal(x.cpu().numpy(), [0.0, 2.0, 0.0, 4.0])
def test_masked_fill_tensor_inplace_identity(self):
x = torch.tensor([1.0, 2.0, 3.0, 4.0], device=device)
mask = torch.tensor([True, False, True, False], device=device)
value = torch.tensor(99.0, device=device)
ret = x.masked_fill_(mask, value)
assert ret is x
np.testing.assert_equal(x.cpu().numpy(), [99.0, 2.0, 99.0, 4.0])
def test_zero_inplace_identity(self):
x = torch.tensor([1.0, 2.0, 3.0, 4.0], device=device)
ret = x.zero_()
assert ret is x
np.testing.assert_equal(x.cpu().numpy(), [0.0, 0.0, 0.0, 0.0])
def test_fill_scalar_inplace_identity(self):
x = torch.tensor([1.0, 2.0, 3.0, 4.0], device=device)
ret = x.fill_(5.0)
assert ret is x
np.testing.assert_equal(x.cpu().numpy(), [5.0, 5.0, 5.0, 5.0])
def test_fill_tensor_inplace_identity(self):
x = torch.tensor([1.0, 2.0, 3.0, 4.0], device=device)
value = torch.tensor(7.0, device=device)
ret = x.fill_(value)
assert ret is x
np.testing.assert_equal(x.cpu().numpy(), [7.0, 7.0, 7.0, 7.0])
def test_add_tensor_inplace_identity(self):
x = torch.tensor([1.0, 2.0, 3.0, 4.0], device=device)
y = torch.tensor([10.0, 20.0, 30.0, 40.0], device=device)
ret = x.add_(y)
assert ret is x
np.testing.assert_equal(x.cpu().numpy(), [11.0, 22.0, 33.0, 44.0])
def test_add_scalar_inplace_identity(self):
x = torch.tensor([1.0, 2.0, 3.0, 4.0], device=device)
ret = x.add_(10.0)
assert ret is x
np.testing.assert_equal(x.cpu().numpy(), [11.0, 12.0, 13.0, 14.0])
def test_mul_tensor_inplace_identity(self):
x = torch.tensor([1.0, 2.0, 3.0, 4.0], device=device)
y = torch.tensor([2.0, 3.0, 4.0, 5.0], device=device)
ret = x.mul_(y)
assert ret is x
np.testing.assert_equal(x.cpu().numpy(), [2.0, 6.0, 12.0, 20.0])
def test_mul_scalar_inplace_identity(self):
x = torch.tensor([1.0, 2.0, 3.0, 4.0], device=device)
ret = x.mul_(2.0)
assert ret is x
np.testing.assert_equal(x.cpu().numpy(), [2.0, 4.0, 6.0, 8.0])
if __name__ == "__main__": if __name__ == "__main__":
unittest.main() unittest.main()
+144
View File
@@ -0,0 +1,144 @@
# simple tests
import unittest
import torch
import warnings
from tinygrad.helpers import getenv, GlobalCounters
if getenv("TINY_BACKEND2"):
import extra.torch_backend.backend2
device = "cpu"
else:
import extra.torch_backend.backend
device = "tiny"
class TestKernelFusionRegression(unittest.TestCase):
def _realize(self, t): _ = t.detach().cpu().numpy()
def _check_kernel_count(self, fn, expected_kernels):
torch.manual_seed(42)
GlobalCounters.reset()
fn().detach().cpu().numpy()
expectation = f"{GlobalCounters.kernel_count} vs {expected_kernels} expected."
if GlobalCounters.kernel_count < expected_kernels: warnings.warn(f"{expectation} Expectation can be lowered.", UserWarning)
self.assertLessEqual(GlobalCounters.kernel_count, expected_kernels, f"{expectation}")
def test_elementwise_fusion(self):
def fn():
x = torch.randn(128, 128, device=device)
return (x + 1.0) * 2.0 - 0.5
self._check_kernel_count(fn, 6)
def test_relu_fusion(self):
def fn():
x = torch.randn(1, 3, 32, 32, device=device)
conv = torch.nn.Conv2d(3, 16, 3, padding=1).to(device)
with torch.no_grad():
return torch.nn.functional.relu(conv(x))
self._check_kernel_count(fn, 8)
def test_batchnorm_fusion(self):
def fn():
x = torch.randn(2, 3, 16, 16, device=device)
conv = torch.nn.Conv2d(3, 8, 3, padding=1).to(device)
bn = torch.nn.BatchNorm2d(8).to(device)
bn.eval()
with torch.no_grad():
return torch.nn.functional.relu(bn(conv(x)))
self._check_kernel_count(fn, 16)
def test_reduce_fusion(self):
def fn():
x = torch.randn(64, 64, device=device)
return (x * 2.0).sum()
self._check_kernel_count(fn, 7)
def test_matmul_elementwise_fusion(self):
def fn():
x = torch.randn(32, 32, device=device)
w = torch.randn(32, 32, device=device)
return torch.nn.functional.relu(x @ w + 1.0)
self._check_kernel_count(fn, 6)
def test_pooling_fusion(self):
def fn():
x = torch.randn(1, 8, 16, 16, device=device)
return torch.nn.functional.max_pool2d(x * 2.0, 2)
self._check_kernel_count(fn, 5)
def test_residual_add_relu_fusion(self):
def fn():
x = torch.randn(1, 8, 16, 16, device=device)
identity = torch.randn(1, 8, 16, 16, device=device)
out = x + identity
return torch.nn.functional.relu(out)
self._check_kernel_count(fn, 6)
def test_inplace_add_relu_fusion(self):
def fn():
x = torch.randn(1, 16, 32, 32, device=device)
y = torch.randn(1, 16, 32, 32, device=device)
x += y
return torch.nn.functional.relu(x)
self._check_kernel_count(fn, 6)
def test_conv_bn_add_relu_fusion(self):
def fn():
x = torch.randn(1, 8, 16, 16, device=device)
identity = torch.randn(1, 8, 16, 16, device=device)
conv = torch.nn.Conv2d(8, 8, 3, padding=1, bias=False).to(device)
bn = torch.nn.BatchNorm2d(8).to(device)
bn.eval()
with torch.no_grad():
out = bn(conv(x))
out += identity
return torch.nn.functional.relu(out)
self._check_kernel_count(fn, 16)
def test_multiple_inplace_ops_fusion(self):
def fn():
x = torch.randn(64, 64, device=device)
x += 1.0
x *= 2.0
return torch.nn.functional.relu(x)
self._check_kernel_count(fn, 4)
def test_view_inplace_no_fusion_break(self):
def fn():
x = torch.randn(4, 64, device=device)
view = x[1:3]
view += 1.0
return x.sum()
self._check_kernel_count(fn, 8)
def test_batchnorm_running_stats_update(self):
def fn():
x = torch.randn(2, 8, 8, 8, device=device)
bn = torch.nn.BatchNorm2d(8).to(device)
bn.train()
with torch.no_grad():
return bn(x)
self._check_kernel_count(fn, 10)
# this is a minimal extra/other_mnist/beautiful_mnist_torch.py to cover fusion for training with optimizer
def test_mnist_training_fusion(self):
def fn():
model = torch.nn.Sequential(
torch.nn.Conv2d(1, 8, 3, padding=1),
torch.nn.ReLU(),
torch.nn.MaxPool2d(2),
torch.nn.Flatten(),
torch.nn.Linear(8*14*14, 10)
).to(device)
optimizer = torch.optim.Adam(model.parameters(), 1e-3)
x = torch.randn(32, 1, 28, 28, device=device)
labels = torch.randint(0, 10, (32,), device=device)
out = model(x)
loss = torch.nn.functional.cross_entropy(out, labels)
optimizer.zero_grad()
loss.backward()
optimizer.step()
return loss
self._check_kernel_count(fn, 33)
if __name__ == "__main__":
unittest.main()
+2 -9
View File
@@ -113,16 +113,9 @@ int register_hook() {
int temp_register_hook = register_hook(); int temp_register_hook = register_hook();
at::Tensor wrap_tensor(py::object &py_obj, c10::ScalarType dtype, c10::DeviceIndex device_index) { at::Tensor wrap_tensor(py::object &py_obj, c10::ScalarType dtype, c10::DeviceIndex device_index) {
// TODO: we have to get the dtype and the shape from the tinygrad Tensor
std::vector<int64_t> sizes = py_obj.attr("shape").cast<std::vector<int64_t>>(); std::vector<int64_t> sizes = py_obj.attr("shape").cast<std::vector<int64_t>>();
std::vector<int64_t> strides = py_obj.attr("_strides").cast<std::vector<int64_t>>();
py::list views = py_obj.attr("uop").attr("st").attr("views"); int64_t storage_offset = py_obj.attr("_storage_offset").cast<int64_t>();
std::vector<int64_t> strides = views[views.size() - 1].attr("strides").cast<std::vector<int64_t>>();
int64_t storage_offset = 0;
for (auto& v: views) {
storage_offset += v.attr("offset").cast<int64_t>(); // TODO: is this correct?
}
return at::detail::make_tensor<at::TinyOpaqueTensorImpl<std::shared_ptr<c10::SafePyObject>>>( return at::detail::make_tensor<at::TinyOpaqueTensorImpl<std::shared_ptr<c10::SafePyObject>>>(
at::DispatchKeySet(at::DispatchKey::PrivateUse1), at::DispatchKeySet(at::DispatchKey::PrivateUse1),
c10::scalarTypeToTypeMeta(dtype), c10::scalarTypeToTypeMeta(dtype),
+10 -4
View File
@@ -4,6 +4,8 @@ import token
import tokenize import tokenize
import itertools import itertools
from tabulate import tabulate from tabulate import tabulate
from tinygrad.uop import Ops
from tinygrad.helpers import ContextVar
TOKEN_WHITELIST = [token.OP, token.NAME, token.NUMBER, token.STRING] TOKEN_WHITELIST = [token.OP, token.NAME, token.NUMBER, token.STRING]
@@ -79,11 +81,15 @@ if __name__ == "__main__":
print(tabulate([headers] + sorted(table, key=lambda x: -x[1]), headers="firstrow", floatfmt=".1f")+"\n") print(tabulate([headers] + sorted(table, key=lambda x: -x[1]), headers="firstrow", floatfmt=".1f")+"\n")
groups = sorted([('/'.join(x[0].rsplit("/", 1)[0].split("/")[0:2]), x[1], x[2]) for x in table]) groups = sorted([('/'.join(x[0].rsplit("/", 1)[0].split("/")[0:2]), x[1], x[2]) for x in table])
dir_sizes = {} dir_sizes = {}
for dir_name, group in itertools.groupby(groups, key=lambda x:x[0]): for dir_name, _group in itertools.groupby(groups, key=lambda x:x[0]):
group = list(_group)
dir_sizes[dir_name] = sum([x[1] for x in group]) dir_sizes[dir_name] = sum([x[1] for x in group])
print(f"{dir_name:30s} : {dir_sizes[dir_name]:6d}") print(f"{dir_name:30s} : {dir_sizes[dir_name]:6d} in {len(group):2d} files")
print(f"\n core line count: {sum([v for k,v in dir_sizes.items() if k not in NONCORE_DIRS])}") print()
print(f" ops: {len(Ops)}")
print(f" flags: {len(ContextVar._cache)}")
print(f" core lines: {sum([v for k,v in dir_sizes.items() if k not in NONCORE_DIRS])}")
total_lines = sum([x[1] for x in table]) total_lines = sum([x[1] for x in table])
print(f"total line count: {total_lines}") print(f"total lines: {total_lines}")
max_line_count = int(os.getenv("MAX_LINE_COUNT", "-1")) max_line_count = int(os.getenv("MAX_LINE_COUNT", "-1"))
assert max_line_count == -1 or total_lines <= max_line_count, f"OVER {max_line_count} LINES" assert max_line_count == -1 or total_lines <= max_line_count, f"OVER {max_line_count} LINES"
+34
View File
@@ -0,0 +1,34 @@
# benchmark speed of pyrender for all created UOps saved with TRACK_MATCH_STATS=2
import functools, pickle
from tinygrad.uop.ops import UOp, Ops
from tinygrad.helpers import tqdm, temp, time_to_str, cpu_profile
BENCHMARK_OPS = {Ops.INDEX, Ops.BUFFERIZE}
@functools.cache
def create_uop(a:int) -> UOp:
op, dtype, src, arg, *rest = trace.uop_fields[a]
return UOp(op, dtype, tuple(create_uop(s) for s in src), arg, *rest)
if __name__ == "__main__":
# load rewrite trace
with open(temp("rewrites.pkl", append_user=True), "rb") as f:
trace = pickle.load(f)
# benchmark
result:list[tuple[str, int]] = []
try:
for steps in tqdm(trace.rewrites):
for r in steps:
for _,yn,_,__ in r.matches:
y = create_uop(yn)
if y.op in BENCHMARK_OPS:
with cpu_profile("pyrender") as e:
try: ren = y.render()
except Exception: ren = "PYRENDER_ERR"
result.append((ren, float(e.en-e.st)/1e6))
finally:
N = 10
print(f"Slowst {N} renders from {len(result)} samples:")
for ren,tm in sorted(result, key=lambda x:x[1], reverse=True)[:N]:
print(f"{time_to_str(tm).strip():<10s} {ren}")
+2
View File
@@ -2,6 +2,7 @@ import gc
from tinygrad import Tensor, UOp, Device, nn from tinygrad import Tensor, UOp, Device, nn
from tinygrad.engine.realize import method_cache, get_program from tinygrad.engine.realize import method_cache, get_program
from tinygrad.schedule.indexing import apply_movement_op from tinygrad.schedule.indexing import apply_movement_op
from tinygrad.uop.divandmod import fold_divmod_general
from test.test_tiny import TestTiny from test.test_tiny import TestTiny
def uops_allocated(): return sum([isinstance(x, UOp) for x in gc.get_objects()]) def uops_allocated(): return sum([isinstance(x, UOp) for x in gc.get_objects()])
@@ -69,6 +70,7 @@ if __name__ == "__main__":
# these caches will keep uops alive # these caches will keep uops alive
method_cache.clear() method_cache.clear()
apply_movement_op.cache_clear() apply_movement_op.cache_clear()
fold_divmod_general.cache_clear()
Tensor._device_seeds.clear() Tensor._device_seeds.clear()
Tensor._device_rng_counters.clear() Tensor._device_rng_counters.clear()
+66
View File
@@ -0,0 +1,66 @@
import random
import z3
from tinygrad.uop.ops import UOp, Ops
from tinygrad.uop.validate import uops_to_z3
from tinygrad.helpers import DEBUG, Context, colored
seed = random.randint(0, 100)
print(f"Seed: {seed}")
random.seed(seed)
def get_random_term(ranges, factors):
# 10% chance of nesting
if random.randint(0,9) == 0: return get_random_expr(ranges, factors)
return random.choice(ranges)*random.choice(factors)*random.choice([1, 1, 1, -1])
def get_random_expr(ranges, factors):
num_terms = random.randint(2,4)
x = UOp.sum(*[get_random_term(ranges, factors) for _ in range(num_terms)])
return x.alu(random.choice([Ops.IDIV, Ops.MOD]), x.ufix(random.choice(factors)*random.choice([1, 1, 1, -1])))
if __name__ == "__main__":
skipped = 0
for i in range(700):
if i % 100 == 0:
print(f"Running test {i}")
upper_bounds = [*list(range(1, 4)), 16, 33, 53, 64, 256]
variable_names = ["i", "j", "k"]
variables = [UOp.variable(s, 1, random.choice(upper_bounds)) for s in variable_names]
factors = variables+upper_bounds
# add some products
for _ in range(2): factors.append(random.choice(variables)*random.choice(variables))
# add some adds
for _ in range(2): factors.append(random.choice(variables)+random.choice(factors))
num_ranges = 4
ranges = [UOp.range(random.choice(factors), i) for i in range(num_ranges)]
variable_names += [f"r{i}" for i in range(num_ranges)]
expr = get_random_expr(ranges, factors)
with Context(CORRECT_DIVMOD_FOLDING=1):
simplified_expr = expr.simplify()
if DEBUG>=1:
print(expr.render(simplify=False), " --> ", simplified_expr.render(simplify=False))
solver = z3.Solver()
solver.set(timeout=3000) # some expressions take very long verify, but its very unlikely they actually return sat
z3_expr, z3_simplified_expr, *z3_vars = uops_to_z3(solver, expr, simplified_expr, *variables, *ranges)
check = solver.check(z3_simplified_expr != z3_expr)
if check == z3.unknown and DEBUG>=1:
skipped += 1
print("skipped z3 verification due to timeout")
elif check == z3.sat:
print(colored("simplify INCORRECT!", "red"))
print(solver.model())
var_vals = {s:solver.model()[z] for s,z in zip(variable_names, z3_vars)}
print("reproduce with:")
print("var_vals = ", var_vals)
print("globals = var_vals|{'cdiv':cdiv,'cmod':cmod}")
print("expr = ast.simplify()")
print("assert eval(ast.render(pm=renderer_infer, simplify=False),globals) == eval(expr.render(pm=renderer_infer, simplify=False),globals)")
print()
assert False
if DEBUG >= 2: print(f"validated {expr.render()}")
print(f"Skipped {skipped} expressions due to timeout")
+3 -1
View File
@@ -36,7 +36,9 @@ def trunc_log(x):
logging.info("\n".join(lines)) logging.info("\n".join(lines))
# user config # user config
SKIP_PROCESS_REPLAY = (k:="[skip_process_replay]") in os.getenv("COMMIT_MESSAGE", "") or k in os.getenv("PR_TITLE", "") # NOTE: process replay is slow so it's now disabled by default. add [pr] to enable it
#SKIP_PROCESS_REPLAY = (k:="[skip_process_replay]") in os.getenv("COMMIT_MESSAGE", "") or k in os.getenv("PR_TITLE", "")
SKIP_PROCESS_REPLAY = not ASSERT_DIFF and not ((k:="[p]") in os.getenv("COMMIT_MESSAGE", "") or k in os.getenv("PR_TITLE", ""))
if REF == "master": SKIP_PROCESS_REPLAY = True if REF == "master": SKIP_PROCESS_REPLAY = True
class ProcessReplayWarning(Warning): pass class ProcessReplayWarning(Warning): pass
+61 -3
View File
@@ -1,6 +1,7 @@
import unittest import unittest
import pathlib import pathlib
from examples.whisper import init_whisper, load_file_waveform, transcribe_file, transcribe_waveform from examples.whisper import init_whisper, load_file_waveform, transcribe_file, transcribe_waveform
import examples.mlperf.metrics as metrics
from tinygrad.helpers import CI, fetch, CPU_LLVM from tinygrad.helpers import CI, fetch, CPU_LLVM
from tinygrad import Device, dtypes from tinygrad import Device, dtypes
from tinygrad.device import is_dtype_supported from tinygrad.device import is_dtype_supported
@@ -14,7 +15,39 @@ TEST_FILE_2 = str(pathlib.Path(__file__).parent / "whisper/test2.wav")
TRANSCRIPTION_2 = "a slightly longer audio file so that we can test batch transcriptions of varying length." TRANSCRIPTION_2 = "a slightly longer audio file so that we can test batch transcriptions of varying length."
# TODO this file will possibly not survive long. find another 1-2 minute sound file online to transcribe # TODO this file will possibly not survive long. find another 1-2 minute sound file online to transcribe
TEST_FILE_3_URL = 'https://homepage.ntu.edu.tw/~karchung/miniconversations/mc45.mp3' TEST_FILE_3_URL = 'https://homepage.ntu.edu.tw/~karchung/miniconversations/mc45.mp3'
TRANSCRIPTION_3 = "Just lie back and relax. Is the level of pressure about right? Yes, it's fine, and I'd like conditioner please. Sure. I'm going to start the second lathering now. Would you like some Q-tips? How'd you like it cut? I'd like my bangs and the back trimmed, and I'd like the rest thinned out a bit and layered. Where would you like the part? On the left, right about here. Here, have a look. What do you think? It's fine. Here's a thousand anti-dollars. It's 30-ant extra for the rants. Here's your change and receipt. Thank you, and please come again. So how do you like it? It could have been worse, but you'll notice that I didn't ask her for her card. Hmm, yeah. Maybe you can try that place over there next time." # noqa: E501 TRANSCRIPTION_3 = """Just lie back and relax.
Is the level of pressure about right?
Yes, it's fine. And I'd like conditioner, please.
Sure. I'm going to start the second lathering now.
Would you like some Q-tips?
How'd you like it cut?
I'd like my bangs and the back trimmed,
and I'd like the rest thinned out a bit and layered.
Where would you like the part?
On the left, right about here.
Here, have a look. What do you think?
It's fine. Here's thousand NT dollars.
It's 30 NT extra for the rinse. Here's your change and receipt.
Thank you, and please come again!
So, how do you like it?
It could have been worse. But you'll notice that I didn't ask her for her card.
Hmm, yeah.
Mm, maybe you can try that place over there next time."""
TRANSCRIPTION_3_ALT = "Just lie back and relax. Is the level of pressure about right? Yes, it's fine. And I'd like conditioner please. Sure. I'm going to start the second lathering now. Would you like some Q-tips? How'd you like it cut? I'd like my bangs on the back trimmed, and I'd like the rest to stand out a bit and layered. Where would you like the part? On the left, right about here. Here. Have a look. What do you think? It's fine. Here's a thousand and eighty dollars. It's thirty and t extra for the rants. Here's your change and receipt. Thank you, and please come again. So how do you like it? It could have been worse, but you'll notice that I didn't ask her for her card. Hmm, yeah. Maybe you can try that place over there next time." #noqa: E501
# NOTE: same as TRANSCRIPTION_3 but with minor changes that should only amount to ~0.079 WER difference (see test_wer_same)
# 'and' --> 'on'
# 'thinned' --> 'to stand'
# 'nt' --> 'and eighty'
# '30 nt' --> 'thirty and t'
# 'rinse' --> 'rants'
# 'mm' --> ''
def wer_helper(result: str, reference: str)->float:
result = metrics.normalize_string(result)
reference = metrics.normalize_string(reference)
wer, _, _ = metrics.word_error_rate([result], [reference])
return wer
@unittest.skipIf(Device.DEFAULT in ["CPU"], "slow") @unittest.skipIf(Device.DEFAULT in ["CPU"], "slow")
@unittest.skipUnless(is_dtype_supported(dtypes.float16), "need float16 support") @unittest.skipUnless(is_dtype_supported(dtypes.float16), "need float16 support")
@@ -30,6 +63,15 @@ class TestWhisper(unittest.TestCase):
del cls.model del cls.model
del cls.enc del cls.enc
def assertWER(self, actual: str, expected: str, threshold: float):
__tracebackhide__ = True # Hide traceback for py.test
wer = wer_helper(actual, expected)
if wer > threshold:
err = f"WER={wer:.3f} > {threshold}"
raise AssertionError(
err
)
def test_transcribe_file1(self): def test_transcribe_file1(self):
self.assertEqual(transcribe_file(self.model, self.enc, TEST_FILE_1), TRANSCRIPTION_1) self.assertEqual(transcribe_file(self.model, self.enc, TEST_FILE_1), TRANSCRIPTION_1)
@@ -56,7 +98,7 @@ class TestWhisper(unittest.TestCase):
def test_transcribe_long(self): def test_transcribe_long(self):
waveform = [load_file_waveform(fetch(TEST_FILE_3_URL))] waveform = [load_file_waveform(fetch(TEST_FILE_3_URL))]
transcription = transcribe_waveform(self.model, self.enc, waveform) transcription = transcribe_waveform(self.model, self.enc, waveform)
self.assertEqual(TRANSCRIPTION_3, transcription) self.assertWER(transcription, TRANSCRIPTION_3, 0.085)
@unittest.skipIf(CI or (Device.DEFAULT == "CPU" and CPU_LLVM), "too long for CI") @unittest.skipIf(CI or (Device.DEFAULT == "CPU" and CPU_LLVM), "too long for CI")
def test_transcribe_long_no_batch(self): def test_transcribe_long_no_batch(self):
@@ -64,8 +106,24 @@ class TestWhisper(unittest.TestCase):
trancriptions = transcribe_waveform(self.model, self.enc, waveforms) trancriptions = transcribe_waveform(self.model, self.enc, waveforms)
self.assertEqual(2, len(trancriptions)) self.assertEqual(2, len(trancriptions))
self.assertEqual(TRANSCRIPTION_3, trancriptions[0]) self.assertWER(trancriptions[0], TRANSCRIPTION_3, 0.085)
self.assertEqual(TRANSCRIPTION_1, trancriptions[1]) self.assertEqual(TRANSCRIPTION_1, trancriptions[1])
def test_wer_same(self):
reference = TRANSCRIPTION_3
self.assertWER(TRANSCRIPTION_3_ALT, reference, 0.079)
def test_wer_different(self):
reference = TRANSCRIPTION_3
self.assertWER("[no speech]", reference, 1.0)
def test_wer_different_2(self):
reference = TRANSCRIPTION_3
self.assertWER("", reference, 1.0)
def test_wer_different_3(self):
reference = TRANSCRIPTION_3
self.assertWER(reference[:len(reference)//2], reference, 0.524)
if __name__ == '__main__': if __name__ == '__main__':
unittest.main() unittest.main()
+10
View File
@@ -765,6 +765,16 @@ class TestMultiTensor(unittest.TestCase):
with self.assertRaises(RuntimeError): with self.assertRaises(RuntimeError):
Tensor.rand_like(t, device=(d3, d4)) Tensor.rand_like(t, device=(d3, d4))
def test_full_like_on_shard(self, axis=None):
t = Tensor.empty((16, 16)).shard(devices_2, axis=axis)
t2 = Tensor.full_like(t, 1.0)
self.assertEqual(t.shape, t2.shape)
self.assertEqual(t.device, t2.device)
self.assertEqual(t.dtype, t2.dtype)
self.assertEqual(t.uop.axis, t2.uop.axis)
t2.realize()
def test_full_like_on_shard_axis(self): self.test_full_like_on_shard(0)
def test_dropout_on_shard(self): def test_dropout_on_shard(self):
with Tensor.train(): with Tensor.train():
X = Tensor.ones(256).to(devices_2) X = Tensor.ones(256).to(devices_2)
+3
View File
@@ -2699,6 +2699,9 @@ class TestOps(unittest.TestCase):
a = Tensor(3.14) a = Tensor(3.14)
np.testing.assert_allclose(Tensor.stack(a, a).numpy(), Tensor([3.14, 3.14]).numpy()) np.testing.assert_allclose(Tensor.stack(a, a).numpy(), Tensor([3.14, 3.14]).numpy())
def test_stack_max(self):
helper_test_op(None, lambda x, y: torch.stack((x, y)).max(axis=0)[0], lambda x, y: Tensor.stack(x, y).max(axis=0), vals=[[1.], [2.]])
def test_repeat(self): def test_repeat(self):
x = Tensor.randn(4, 6, 3) x = Tensor.randn(4, 6, 3)
base_repeats = [2, 4, 3] base_repeats = [2, 4, 3]
+1 -1
View File
@@ -20,7 +20,7 @@ class TestPickle(unittest.TestCase):
self.assertEqual(pm2.rewrite(sink).key, tt.key) self.assertEqual(pm2.rewrite(sink).key, tt.key)
def test_pickle_main_pattern_matcher(self): def test_pickle_main_pattern_matcher(self):
from tinygrad.codegen.late.devectorizer import sym from tinygrad.uop.symbolic import sym
ssym = pickle.dumps(sym) ssym = pickle.dumps(sym)
dsym = pickle.loads(ssym) dsym = pickle.loads(ssym)
self.assertEqual(dsym.patterns[0][0].location, sym.patterns[0][0].location) self.assertEqual(dsym.patterns[0][0].location, sym.patterns[0][0].location)
+3 -3
View File
@@ -35,9 +35,9 @@ class TestTiny(unittest.TestCase):
out = Tensor.cat(Tensor.ones(8).contiguous(), Tensor.zeros(8).contiguous()) out = Tensor.cat(Tensor.ones(8).contiguous(), Tensor.zeros(8).contiguous())
self.assertListEqual(out.tolist(), [1]*8+[0]*8) self.assertListEqual(out.tolist(), [1]*8+[0]*8)
def test_sum(self): def test_sum(self, N=getenv("SUM_N", 256)):
out = Tensor.ones(256).contiguous().sum() out = Tensor.ones(N).contiguous().sum()
self.assertEqual(out.item(), 256) self.assertEqual(out.item(), N)
def test_gemm(self, N=getenv("GEMM_N", 64), out_dtype=dtypes.float): def test_gemm(self, N=getenv("GEMM_N", 64), out_dtype=dtypes.float):
a = Tensor.ones(N,N).contiguous() a = Tensor.ones(N,N).contiguous()
+1 -1
View File
@@ -517,7 +517,7 @@ class TestUOpStr(unittest.TestCase):
class TestUPatHelpers(unittest.TestCase): class TestUPatHelpers(unittest.TestCase):
def test_location(self): def test_location(self):
self.assertEqual(sym.patterns[-1][0].location[0].replace("\\", "/").split("/")[-1], "math.py") self.assertEqual(sym.patterns[-1][0].location[0].replace("\\", "/").split("/")[-1], "symbolic.py")
self.assertEqual(shared_spec.patterns[0][0].location[0].replace("\\", "/").split("/")[-1], "spec.py") self.assertEqual(shared_spec.patterns[0][0].location[0].replace("\\", "/").split("/")[-1], "spec.py")
test_upat = UPat(Ops.CONST, dtypes.bool) test_upat = UPat(Ops.CONST, dtypes.bool)
self.assertEqual(test_upat.location[0].split("/")[-1], __file__.replace("\\", "/").split("/")[-1]) self.assertEqual(test_upat.location[0].split("/")[-1], __file__.replace("\\", "/").split("/")[-1])
+4 -4
View File
@@ -5,16 +5,16 @@ from tinygrad.runtime.support.c import Struct
class TestAutogen(unittest.TestCase): class TestAutogen(unittest.TestCase):
def test_packed_struct_sizeof(self): def test_packed_struct_sizeof(self):
layout = [('a', ctypes.c_char), ('b', ctypes.c_int, 5), ('c', ctypes.c_char)] layout = [('a', ctypes.c_char), ('b', ctypes.c_int, 5), ('c', ctypes.c_char)]
class X(ctypes.Structure): _fields_, _layout_ = layout, 'gcc-sysv'
class Y(ctypes.Structure): _fields_, _pack_, _layout_ = layout, 1, 'ms' class Y(ctypes.Structure): _fields_, _pack_, _layout_ = layout, 1, 'ms'
class Z(Struct): _packed_, _fields_ = True, layout class Z(Struct): pass
self.assertNotEqual(ctypes.sizeof(X), 4) # ctypes bug! gcc-13.3.0 says this should have size 4 Z._packed_, Z._fields_ = True, layout
self.assertEqual(ctypes.sizeof(Y), 6) self.assertEqual(ctypes.sizeof(Y), 6)
self.assertEqual(ctypes.sizeof(Z), 3) self.assertEqual(ctypes.sizeof(Z), 3)
layout = [('a', ctypes.c_int, 31), ('b', ctypes.c_int, 31), ('c', ctypes.c_int, 1), ('d', ctypes.c_int, 1)] layout = [('a', ctypes.c_int, 31), ('b', ctypes.c_int, 31), ('c', ctypes.c_int, 1), ('d', ctypes.c_int, 1)]
class Foo(ctypes.Structure): _fields_, _layout_ = layout, 'gcc-sysv' class Foo(ctypes.Structure): _fields_, _layout_ = layout, 'gcc-sysv'
class Bar(ctypes.Structure): _fields_, _pack_, _layout_ = layout, 1, 'ms' class Bar(ctypes.Structure): _fields_, _pack_, _layout_ = layout, 1, 'ms'
class Baz(Struct): _fields_, _packed_ = layout, True class Baz(Struct): pass
Baz._packed_, Baz._fields_ = True, layout
self.assertEqual(ctypes.sizeof(Foo), 12) self.assertEqual(ctypes.sizeof(Foo), 12)
self.assertEqual(ctypes.sizeof(Bar), 12) self.assertEqual(ctypes.sizeof(Bar), 12)
self.assertEqual(ctypes.sizeof(Baz), 8) self.assertEqual(ctypes.sizeof(Baz), 8)
+9
View File
@@ -1,4 +1,5 @@
import unittest, time import unittest, time
from tinygrad.helpers import Profiling
from tinygrad.uop.ops import UOp from tinygrad.uop.ops import UOp
from tinygrad.dtype import dtypes from tinygrad.dtype import dtypes
@@ -38,6 +39,14 @@ class TestMicrobenchmarks(unittest.TestCase):
a = UOp.const(dtypes.int, 2) a = UOp.const(dtypes.int, 2)
for _ in range(N): (a+a).simplify() for _ in range(N): (a+a).simplify()
class TestMicroprofile(unittest.TestCase):
def test_uop_simplify_complex(self):
x = UOp.variable("x", 0, 10)
y = UOp.variable("y", 0, 10)
expr = (x*2)+5+(x*4)+(y*2)+y
with Profiling():
for _ in range(1000): expr.simplify()
if __name__ == '__main__': if __name__ == '__main__':
unittest.main() unittest.main()
+22
View File
@@ -430,5 +430,27 @@ class TestImageSimplification(unittest.TestCase):
load = get_load_image_uop((128, 768, 4), valid, (alu0, alu1)) load = get_load_image_uop((128, 768, 4), valid, (alu0, alu1))
self.check(load, None, "((((idx1*24)+r3)+(r5*3))+-3)", "(((idx2*2)+r4)+-1)") self.check(load, None, "((((idx1*24)+r3)+(r5*3))+-3)", "(((idx2*2)+r4)+-1)")
def test_simplify7(self):
# DEBUG=2 ALLOWED_KERNEL_COUNT=123 ALLOWED_READ_IMAGE=1397 ALLOWED_GATED_READ_IMAGE=94 FLOAT16=1 CL=1 IMAGE=2 python examples/openpilot/compile3.py https://gitlab.com/commaai/openpilot-lfs.git/gitlab-lfs/objects/cf6376aa9a090f0da26c280ef69eabf9bbdd51d1faac9ed392919c3db69be916 # noqa: E501
# kernel 143
gidx0 = Special("gidx0", 32)
lidx0 = Special("lidx0", 16)
lidx1 = Special("lidx1", 8)
r0 = Range(0, 7)
# buf.render()='UOp(Ops.DEFINE_GLOBAL, dtypes.imageh((32, 1024, 4)), arg=1, src=())'
alu0 = ((gidx0*2+(lidx0*128+r0*64+lidx1*8+-183)%64*64+(lidx0*128+r0*64+lidx1*8+-183)//64%32*4096+1)//4%1024)
alu1 = ((gidx0*2+(lidx0*128+r0*64+lidx1*8+-183)%64*64+(lidx0*128+r0*64+lidx1*8+-183)//64%32*4096+1)//4096)
valid = ((lidx1<7)&((((lidx0*2+r0)<3)!=1)&((lidx0*2+r0)<35)))
load = get_load_image_uop((32, 1024, 4), valid, (alu0, alu1))
self.check(load, None, "(lidx1*128+gidx0//2+144)", "(lidx0*2+r0+-3)")
# TODO: this is the same idx as above, but simplifying idx too early makes it hard to drop the valid
alu0 = ((gidx0*2+lidx1*512+(lidx0*8192+r0*4096)+-11711)//4%1024)
alu1 = (lidx0*2+r0+-3)
valid = ((lidx1<7)&((((lidx0*2+r0)<3)!=1)&((lidx0*2+r0)<35)))
load = get_load_image_uop((32, 1024, 4), valid, (alu0, alu1))
self.check(load, "(lidx1<7)", "((gidx0*2+lidx1*512+(lidx0*8192+r0*4096)+-11711)//4%1024)", "(lidx0*2+r0+-3)")
if __name__ == '__main__': if __name__ == '__main__':
unittest.main() unittest.main()
+35
View File
@@ -159,3 +159,38 @@ class TestFuzzFailure(unittest.TestCase):
num = expr.simplify().substitute({v1:v1_val, v2:v2_val, v3:v3_val}).ssimplify() num = expr.simplify().substitute({v1:v1_val, v2:v2_val, v3:v3_val}).ssimplify()
rn = expr.substitute({v1:v1_val, v2:v2_val, v3:v3_val}).ssimplify() rn = expr.substitute({v1:v1_val, v2:v2_val, v3:v3_val}).ssimplify()
self.assertEqual(num, rn) self.assertEqual(num, rn)
def test_fuzz_failure11(self):
v1=Variable("v1", 0, 16)
v2=Variable("v2", 0, 128)
v3=Variable("v3", 0, 5)
expr = UOp(Ops.MOD, dtypes.index, arg=None, src=(
UOp(Ops.ADD, dtypes.index, arg=None, src=(
UOp(Ops.MOD, dtypes.index, arg=None, src=(
UOp(Ops.ADD, dtypes.index, arg=None, src=(
UOp(Ops.MAX, dtypes.index, arg=None, src=(
UOp(Ops.MUL, dtypes.index, arg=None, src=(
x5:=UOp(Ops.DEFINE_VAR, dtypes.index, arg=('v2', 0, 128), src=()),
UOp(Ops.CONST, dtypes.index, arg=0, src=()),)),
UOp(Ops.CONST, dtypes.index, arg=8, src=()),)),
UOp(Ops.MUL, dtypes.index, arg=None, src=(
x5,
UOp(Ops.CONST, dtypes.index, arg=-2, src=()),)),)),
x10:=UOp(Ops.CONST, dtypes.index, arg=5, src=()),)),
UOp(Ops.ADD, dtypes.index, arg=None, src=(
UOp(Ops.ADD, dtypes.index, arg=None, src=(
UOp(Ops.IDIV, dtypes.index, arg=None, src=(
x14:=UOp(Ops.DEFINE_VAR, dtypes.index, arg=('v1', 0, 16), src=()),
UOp(Ops.CONST, dtypes.index, arg=6, src=()),)),
UOp(Ops.CONST, dtypes.index, arg=4, src=()),)),
UOp(Ops.ADD, dtypes.index, arg=None, src=(
x14,
UOp(Ops.CONST, dtypes.index, arg=1, src=()),)),)),)),
x10,))
v1_val, v2_val, v3_val = UOp.const(dtypes.int, 0), UOp.const(dtypes.int, 7),UOp.const(dtypes.int, 0)
num = expr.simplify().substitute({v1:v1_val, v2:v2_val, v3:v3_val}).ssimplify()
rn = expr.substitute({v1:v1_val, v2:v2_val, v3:v3_val}).ssimplify()
self.assertEqual(num, rn)
if __name__ == '__main__':
unittest.main()
+4 -4
View File
@@ -3,19 +3,19 @@ from tinygrad import Tensor
class TestLoadStore(unittest.TestCase): class TestLoadStore(unittest.TestCase):
def test_load_shape(self): def test_load_shape(self):
t = Tensor(bytes(16)).load(1024).kernelize() t = Tensor(bytes(16)).fs_load(1024).kernelize()
assert t.shape == (1024,), t.shape assert t.shape == (1024,), t.shape
def test_store_shape(self): def test_store_shape(self):
t = Tensor.zeros(1024).store().kernelize() t = Tensor.zeros(1024).fs_store().kernelize()
assert t.shape == (16,), t.shape assert t.shape == (16,), t.shape
def test_load_large_shape(self): def test_load_large_shape(self):
t = Tensor(bytes(16)).load(10_000_000).kernelize() t = Tensor(bytes(16)).fs_load(10_000_000).kernelize()
assert t.shape == (10_000_000,), t.shape assert t.shape == (10_000_000,), t.shape
def test_store_large_shape(self): def test_store_large_shape(self):
t = Tensor.zeros(10_000_000).store().kernelize() t = Tensor.zeros(10_000_000).fs_store().kernelize()
assert t.shape == (16,), t.shape assert t.shape == (16,), t.shape
if __name__ == "__main__": if __name__ == "__main__":
+1
View File
@@ -128,6 +128,7 @@ class TestProgressBar(unittest.TestCase):
self._compare_bars(tinytqdm_output, tqdm_output) self._compare_bars(tinytqdm_output, tqdm_output)
if n > 5: break if n > 5: break
@unittest.skip("this is flaky")
@patch('sys.stderr', new_callable=StringIO) @patch('sys.stderr', new_callable=StringIO)
@patch('shutil.get_terminal_size') @patch('shutil.get_terminal_size')
def test_set_description(self, mock_terminal_size, mock_stderr): def test_set_description(self, mock_terminal_size, mock_stderr):
+5 -1
View File
@@ -15,7 +15,7 @@ def check_uop_against_string(self, v:UOp, s:str):
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.index: s_eval = UOp.const(dtypes.index, s_eval)
elif isinstance(s_eval, (bool, int, float)): s_eval = UOp.const(dtypes.from_py(s_eval), 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") s_eval = graph_rewrite(s_eval, commutative, name="cannonicalize eval")
self.assertIs(s_eval, v, f"eval did not match simplified: {s_eval} != {v} for {s}") 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 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 uconst(val): return UOp.const(dtypes.index, val)
@@ -679,6 +679,10 @@ class TestSymbolic(unittest.TestCase):
b = Variable("b", 0, 3) b = Variable("b", 0, 3)
c = Variable("c", 0, 3) c = Variable("c", 0, 3)
d = Variable("d", -3, 3) d = Variable("d", -3, 3)
self.helper_test_variable((a<2), 0, 1, "(a<2)")
self.helper_test_variable((a<=2), 0, 1, "((2<a)!=True)")
self.helper_test_variable((a>1), 0, 1, "(1<a)")
self.helper_test_variable((a>=1), 0, 1, "((a<1)!=True)")
self.helper_test_variable((a<1).ne(True), 0, 1, "((a<1)!=True)") self.helper_test_variable((a<1).ne(True), 0, 1, "((a<1)!=True)")
self.helper_test_variable((a+b<1).ne(True), 0, 1, "(((a+b)<1)!=True)") self.helper_test_variable((a+b<1).ne(True), 0, 1, "(((a+b)<1)!=True)")
self.helper_test_variable((a*3+b*4<1).ne(True), 0, 1, "(((a+b)<1)!=True)") self.helper_test_variable((a*3+b*4<1).ne(True), 0, 1, "(((a+b)<1)!=True)")
+3 -4
View File
@@ -4,7 +4,7 @@ from collections import defaultdict
from dataclasses import dataclass from dataclasses import dataclass
from tinygrad.dtype import dtypes, ImageDType, DType, AddrSpace, Invalid, PtrDType from tinygrad.dtype import dtypes, ImageDType, DType, AddrSpace, Invalid, PtrDType
from tinygrad.uop.ops import UOp, Ops, UPat, PatternMatcher, graph_rewrite, GroupOp, identity_element from tinygrad.uop.ops import UOp, Ops, UPat, PatternMatcher, graph_rewrite, GroupOp, identity_element
from tinygrad.uop.symbolic import uop_given_valid, parse_valid, sym, symbolic, invalid_gate from tinygrad.uop.symbolic import uop_given_valid, parse_valid, symbolic, invalid_gate
from tinygrad.helpers import getenv, flatten, AMX, prod from tinygrad.helpers import getenv, flatten, AMX, prod
from tinygrad.renderer import Renderer from tinygrad.renderer import Renderer
@@ -26,7 +26,6 @@ def simplify_valid_load(buf:UOp, start_idx:UOp, valid:UOp) -> UOp|None:
# for X0 + X1 + ... >= 1, check if it's out of bound when Xi = 0 for all i # for X0 + X1 + ... >= 1, check if it's out of bound when Xi = 0 for all i
if not is_upper_bound and c == 1 and all(u.op in GroupOp.Irreducible and u.vmin == 0 for u in X.split_uop(Ops.ADD)): if not is_upper_bound and c == 1 and all(u.op in GroupOp.Irreducible and u.vmin == 0 for u in X.split_uop(Ops.ADD)):
testidx = functools.reduce(lambda nowidx,u: nowidx.substitute({u:u.const_like(0)}), X.split_uop(Ops.ADD), idx) testidx = functools.reduce(lambda nowidx,u: nowidx.substitute({u:u.const_like(0)}), X.split_uop(Ops.ADD), idx)
testidx = testidx.simplify()
if testidx.gep(0).vmax < 0 or testidx.gep(1).vmax < 0: if testidx.gep(0).vmax < 0 or testidx.gep(1).vmax < 0:
drop_stmt.append(stmt) drop_stmt.append(stmt)
continue continue
@@ -36,7 +35,7 @@ def simplify_valid_load(buf:UOp, start_idx:UOp, valid:UOp) -> UOp|None:
test_value = c + 1 if is_upper_bound else c - 1 test_value = c + 1 if is_upper_bound else c - 1
for i,b in zip(idx.src, (buf.dtype.shape[1], buf.dtype.shape[0])): for i,b in zip(idx.src, (buf.dtype.shape[1], buf.dtype.shape[0])):
if i.is_increasing(): if i.is_increasing():
rw = i.substitute({X:X.const_like(test_value)}).simplify() rw = i.substitute({X:X.const_like(test_value)})
if rw.vmin >= b or rw.vmax < 0: if rw.vmin >= b or rw.vmax < 0:
drop_stmt.append(stmt) drop_stmt.append(stmt)
break break
@@ -314,7 +313,7 @@ pm_reduce = PatternMatcher([
# tensor core built in accumulate # tensor core built in accumulate
(UPat(Ops.WMMA, name="wmma") + UPat.var("add"), (UPat(Ops.WMMA, name="wmma") + UPat.var("add"),
lambda add, wmma: UOp(wmma.op, wmma.dtype, (wmma.src[0], wmma.src[1], wmma.src[2]+add), wmma.arg)), lambda add, wmma: UOp(wmma.op, wmma.dtype, (wmma.src[0], wmma.src[1], wmma.src[2]+add), wmma.arg)),
])+sym ])
# add loads # add loads
+5 -6
View File
@@ -18,6 +18,7 @@ class Scheduler:
self.ast, self.ren = ast, ren self.ast, self.ren = ast, ren
self.dont_use_locals = self.ast.arg.dont_use_locals if self.ast.arg is not None else False self.dont_use_locals = self.ast.arg.dont_use_locals if self.ast.arg is not None else False
self.applied_opts = list(self.ast.arg.applied_opts) if self.ast.arg is not None else [] self.applied_opts = list(self.ast.arg.applied_opts) if self.ast.arg is not None else []
self.opt_range = itertools.count(start=max([x.arg[0] for x in self.rngs], default=0)+1)
@property @property
def rngs(self): def rngs(self):
@@ -29,8 +30,6 @@ class Scheduler:
def full_shape(self): return [ssimplify(x.src[0]) for x in self.rngs] def full_shape(self): return [ssimplify(x.src[0]) for x in self.rngs]
@property @property
def axis_types(self): return [x.arg[-1] for x in self.rngs] def axis_types(self): return [x.arg[-1] for x in self.rngs]
@property
def maxarg(self): return max([x.arg[0] for x in self.rngs], default=0)
# strings like ['g0', 'g1', 'l0', 'l1', 'l2', 'l3', 'l4', 'l5', 'R0', 'r0', 'r1', 'r2', 'u0', 'u1', 'u2'] # strings like ['g0', 'g1', 'l0', 'l1', 'l2', 'l3', 'l4', 'l5', 'R0', 'r0', 'r1', 'r2', 'u0', 'u1', 'u2']
def shape_str(self) -> list[str]: def shape_str(self) -> list[str]:
@@ -95,7 +94,7 @@ class Scheduler:
def shift_to(self, rng:UOp, amount:int, new_type:AxisType, top:bool=False, input_new_rng=None): def shift_to(self, rng:UOp, amount:int, new_type:AxisType, top:bool=False, input_new_rng=None):
if (old_sz:=rng.src[0].divides(amount)) is None: if (old_sz:=rng.src[0].divides(amount)) is None:
raise KernelOptError(f"{amount} can't divide {rng.src[0]} in {self.colored_shape()}") raise KernelOptError(f"{amount} can't divide {rng.src[0]} in {self.colored_shape()}")
new_rng = UOp.range(amount, self.maxarg+1, new_type) if input_new_rng is None else input_new_rng new_rng = UOp.range(amount, next(self.opt_range), new_type) if input_new_rng is None else input_new_rng
replaced_rng = rng.replace(src=(UOp.const(dtypes.int, old_sz),)) replaced_rng = rng.replace(src=(UOp.const(dtypes.int, old_sz),))
sub_axis = (new_rng * old_sz + replaced_rng) if top else (replaced_rng * amount + new_rng) sub_axis = (new_rng * old_sz + replaced_rng) if top else (replaced_rng * amount + new_rng)
self.ast = self.ast.substitute({rng:sub_axis}, name=f"shift {rng.arg[:-1]} {amount} {str(new_type).split('.')[1].lower()}") self.ast = self.ast.substitute({rng:sub_axis}, name=f"shift {rng.arg[:-1]} {amount} {str(new_type).split('.')[1].lower()}")
@@ -231,9 +230,9 @@ class Scheduler:
for tc in tensor_cores: for tc in tensor_cores:
if tc.dtype_in == in0.dtype.scalar() and tc.dtype_in == in1.dtype.scalar() and tc.dtype_out == reduceop.dtype.scalar(): if tc.dtype_in == in0.dtype.scalar() and tc.dtype_in == in1.dtype.scalar() and tc.dtype_out == reduceop.dtype.scalar():
# tensor cores have three ranges. X, Y, and REDUCE # tensor cores have three ranges. X, Y, and REDUCE
in0_ranges = sorted([u for u in in0.ranges if u not in in1.ranges], key=lambda x: -x.arg[0]) in0_ranges = sorted([u for u in in0.ranges if u not in in1.ranges], key=lambda x: x.arg[0], reverse=True)
in1_ranges = sorted([u for u in in1.ranges if u not in in0.ranges], key=lambda x: -x.arg[0]) in1_ranges = sorted([u for u in in1.ranges if u not in in0.ranges], key=lambda x: x.arg[0], reverse=True)
red_ranges = sorted(reduceop.src[1:], key=lambda x: -x.arg[0]) red_ranges = sorted(reduceop.src[1:], key=lambda x: x.arg[0], reverse=True)
if DEBUG >= 3: if DEBUG >= 3:
print(f"TC({axis}): {[(x.arg[0],x.vmax+1) for x in in0_ranges]}", print(f"TC({axis}): {[(x.arg[0],x.vmax+1) for x in in0_ranges]}",
f"{[(x.arg[0],x.vmax+1) for x in in1_ranges]} {[(x.arg[0],x.vmax+1) for x in red_ranges]}") f"{[(x.arg[0],x.vmax+1) for x in in1_ranges]} {[(x.arg[0],x.vmax+1) for x in red_ranges]}")
+1 -1
View File
@@ -142,7 +142,7 @@ pm_reduce_simplify = pm_reduce_unparented + PatternMatcher([
# remove REDUCE on load, comes from indexing a tensor with another tensor # remove REDUCE on load, comes from indexing a tensor with another tensor
def no_load(u:UOp) -> bool: return not any(x.op is Ops.INDEX for x in u.backward_slice_with_self) def no_load(u:UOp) -> bool: return not any(x.op is Ops.INDEX for x in u.backward_slice_with_self)
pm_load_collapse = PatternMatcher([ pm_load_collapse = PatternMatcher([
(UPat(Ops.REDUCE, src=(UPat.var("u"), UPat()), name="red"), reduce_load_collapse), (UPat(Ops.REDUCE, arg=Ops.ADD, src=(UPat.var("u"), UPat()), name="red"), reduce_load_collapse),
# we want to make sure we dont do math on a loaded index since that can cause overflow, this undoes the rule in pm_reduce_load_collapse # we want to make sure we dont do math on a loaded index since that can cause overflow, this undoes the rule in pm_reduce_load_collapse
((UPat.var("x", dtypes.index)+UPat.var("y"))<UPat.var("c"), lambda x,y,c: x < c-y if no_load(y) and no_load(c) and not no_load(x) else None), ((UPat.var("x", dtypes.index)+UPat.var("y"))<UPat.var("c"), lambda x,y,c: x < c-y if no_load(y) and no_load(c) and not no_load(x) else None),
]) ])
+2 -2
View File
@@ -5,7 +5,7 @@ from typing import Any, Generic, TypeVar, Iterator, Sequence, cast, Generator
import importlib, inspect, functools, pathlib, os, platform, contextlib, sys, re, atexit, pickle, decimal import importlib, inspect, functools, pathlib, os, platform, contextlib, sys, re, atexit, pickle, decimal
from tinygrad.helpers import CI, OSX, LRU, getenv, diskcache_get, diskcache_put, DEBUG, GlobalCounters, flat_mv, PROFILE, temp, colored, CPU_LLVM from tinygrad.helpers import CI, OSX, LRU, getenv, diskcache_get, diskcache_put, DEBUG, GlobalCounters, flat_mv, PROFILE, temp, colored, CPU_LLVM
from tinygrad.helpers import Context, CCACHE, ALLOW_DEVICE_USAGE, MAX_BUFFER_SIZE, cpu_events, ProfileEvent, ProfilePointEvent, dedup from tinygrad.helpers import Context, CCACHE, ALLOW_DEVICE_USAGE, MAX_BUFFER_SIZE, cpu_events, ProfileEvent, ProfilePointEvent, dedup
from tinygrad.helpers import unwrap_class_type, suppress_finalizing, AMD_LLVM, select_first_inited, VIZ from tinygrad.helpers import unwrap_class_type, suppress_finalizing, select_first_inited, VIZ
from tinygrad.dtype import DType, ImageDType, PtrDType, dtypes, _to_np_dtype from tinygrad.dtype import DType, ImageDType, PtrDType, dtypes, _to_np_dtype
from tinygrad.renderer import Renderer from tinygrad.renderer import Renderer
@@ -329,7 +329,7 @@ def is_dtype_supported(dtype:DType, device:str|None=None) -> bool:
return device in {"AMD", "PYTHON", "NULL"} return device in {"AMD", "PYTHON", "NULL"}
if dtype in dtypes.fp8s: if dtype in dtypes.fp8s:
if device in {"CUDA", "NV"}: return not CI and not getenv(f"{device}_PTX") and not getenv("NV_NAK") if device in {"CUDA", "NV"}: return not CI and not getenv(f"{device}_PTX") and not getenv("NV_NAK")
if device == "AMD": return not CI and not AMD_LLVM and getattr(Device["AMD"], "target") in {(9,4,2), (9,5,0)} if device == "AMD": return not CI and getattr(Device["AMD"], "target") in {(9,4,2), (9,5,0)}
return device in {"PYTHON", "NULL"} return device in {"PYTHON", "NULL"}
if device == "WEBGPU": return dtype in [dtypes.bool, dtypes.char, dtypes.uchar, dtypes.short, if device == "WEBGPU": return dtype in [dtypes.bool, dtypes.char, dtypes.uchar, dtypes.short,
dtypes.ushort, dtypes.float, dtypes.int32, dtypes.uint32, dtypes.half] dtypes.ushort, dtypes.float, dtypes.int32, dtypes.uint32, dtypes.half]
+10 -4
View File
@@ -147,8 +147,10 @@ def temp(x:str, append_user:bool=False) -> str:
class Context(contextlib.ContextDecorator): class Context(contextlib.ContextDecorator):
def __init__(self, **kwargs): self.kwargs = kwargs def __init__(self, **kwargs): self.kwargs = kwargs
def __enter__(self): def __enter__(self):
self.old_context:dict[str, int] = {k:v.value for k,v in ContextVar._cache.items()} self.old_context:dict[str, int] = {}
for k,v in self.kwargs.items(): ContextVar._cache[k].value = v for k,v in self.kwargs.items():
self.old_context[k] = ContextVar._cache[k].value
ContextVar._cache[k].value = v
def __exit__(self, *args): def __exit__(self, *args):
for k,v in self.old_context.items(): ContextVar._cache[k].value = v for k,v in self.old_context.items(): ContextVar._cache[k].value = v
@@ -279,7 +281,7 @@ class ProfilePointEvent(ProfileEvent):
cpu_events:list[ProfileEvent] = [] cpu_events:list[ProfileEvent] = []
@contextlib.contextmanager @contextlib.contextmanager
def cpu_profile(name:str|TracingKey, device="CPU", is_copy=False, display=True) -> Generator[ProfileRangeEvent, None, None]: def cpu_profile(name:str|TracingKey, device="TINY", is_copy=False, display=True) -> Generator[ProfileRangeEvent, None, None]:
res = ProfileRangeEvent(device, name, perf_counter_us(), is_copy=is_copy) res = ProfileRangeEvent(device, name, perf_counter_us(), is_copy=is_copy)
try: yield res try: yield res
finally: finally:
@@ -382,7 +384,11 @@ def fetch(url:str, name:pathlib.Path|str|None=None, subdir:str|None=None, gunzip
# *** Exec helpers # *** Exec helpers
def system(cmd, **kwargs): return subprocess.check_output(cmd.split(), **kwargs).decode().strip() def system(cmd:str, **kwargs) -> str:
st = time.perf_counter()
ret = subprocess.check_output(cmd.split(), **kwargs).decode().strip()
if DEBUG >= 1: print(f"system: '{cmd}' returned {len(ret)} bytes in {(time.perf_counter() - st)*1e3:.2f} ms")
return ret
def cpu_objdump(lib, objdump_tool='objdump'): def cpu_objdump(lib, objdump_tool='objdump'):
with tempfile.NamedTemporaryFile(delete=True) as f: with tempfile.NamedTemporaryFile(delete=True) as f:
+3 -3
View File
@@ -2,7 +2,7 @@
import itertools import itertools
from tinygrad.helpers import dedup, flatten, getenv, unwrap, FUSE_OPTIM from tinygrad.helpers import dedup, flatten, getenv, unwrap, FUSE_OPTIM
from tinygrad.tensor import Tensor from tinygrad.tensor import Tensor
from tinygrad.dtype import dtypes, least_upper_dtype from tinygrad.dtype import dtypes, least_upper_dtype, to_dtype
class Optimizer: class Optimizer:
""" """
@@ -24,9 +24,9 @@ class Optimizer:
if self.fused: self.pos_params = list(itertools.accumulate(self.params, lambda x,y: x+y.numel(), initial=0)) if self.fused: self.pos_params = list(itertools.accumulate(self.params, lambda x,y: x+y.numel(), initial=0))
def _new_optim_param(self) -> list[Tensor]: def _new_optim_param(self) -> list[Tensor]:
param_dtype = getenv("OPTIM_DTYPE", "float32") param_dtype = to_dtype(getenv("OPTIM_DTYPE", "float32"))
if self.fused: return [Tensor.zeros(self.pos_params[-1], dtype=param_dtype, device=self.device, requires_grad=False).contiguous()] if self.fused: return [Tensor.zeros(self.pos_params[-1], dtype=param_dtype, device=self.device, requires_grad=False).contiguous()]
return [Tensor.zeros(*t.shape, dtype=param_dtype, device=t.device, requires_grad=False).contiguous() for t in self.params] return [Tensor.zeros_like(t, dtype=param_dtype, requires_grad=False).contiguous() for t in self.params]
def zero_grad(self): def zero_grad(self):
""" """
+11 -10
View File
@@ -22,10 +22,10 @@ base_rewrite = PatternMatcher([
(UPat(Ops.CAST, name="x"), lambda ctx,x: (UPat(Ops.CAST, name="x"), lambda ctx,x:
f"__builtin_convertvector({ctx[x.src[0]]}, {ctx.render_dtype(x.dtype)})" if x.dtype.count > 1 and not isinstance(x.dtype, PtrDType) else None), f"__builtin_convertvector({ctx[x.src[0]]}, {ctx.render_dtype(x.dtype)})" if x.dtype.count > 1 and not isinstance(x.dtype, PtrDType) else None),
(UPat(Ops.CAST, name="x"), lambda ctx,x: f"({ctx.render_cast(x.dtype, ctx[x.src[0]])})"), (UPat(Ops.CAST, name="x"), lambda ctx,x: f"({ctx.render_cast(x.dtype, ctx[x.src[0]])})"),
(UPat(Ops.BITCAST, name="x"), lambda ctx,x: f"(*(({ctx.buffer_prefix}{ctx.render_dtype(x.dtype)}*)&{ctx[x.src[0]]}))"), (UPat(Ops.BITCAST, name="x"), lambda ctx,x:
f"__builtin_bit_cast({ctx.render_dtype(x.dtype)}, ({ctx.render_dtype(x.src[0].dtype)})({ctx[x.src[0]]}))"),
(UPat(Ops.DEFINE_LOCAL, name="x"), lambda ctx,x: f"{ctx.smem_align}{ctx.smem_prefix}{ctx.render_dtype(x.dtype.base)} {ctx[x]}[{x.dtype.size}];"), (UPat(Ops.DEFINE_LOCAL, name="x"), lambda ctx,x: f"{ctx.smem_align}{ctx.smem_prefix}{ctx.render_dtype(x.dtype.base)} {ctx[x]}[{x.dtype.size}];"),
(UPat(Ops.BARRIER), lambda ctx: ctx.barrier), (UPat(Ops.BARRIER), lambda ctx: ctx.barrier),
(UPat(Ops.PRECAST, name="x"), lambda ctx,x: ctx[x.src[0]]),
(UPat(Ops.SPECIAL, name="x"), lambda ctx,x: f"{ctx.code_for_workitem[x.arg[0]](x.arg[-1])}; /* {(x.src[0]).render()} */"), (UPat(Ops.SPECIAL, name="x"), lambda ctx,x: f"{ctx.code_for_workitem[x.arg[0]](x.arg[-1])}; /* {(x.src[0]).render()} */"),
# const # const
(UPat(Ops.CONST, arg=math.inf, name="x"), lambda ctx, x: f"({ctx.render_cast(x.dtype, ctx.infinity)})"), (UPat(Ops.CONST, arg=math.inf, name="x"), lambda ctx, x: f"({ctx.render_cast(x.dtype, ctx.infinity)})"),
@@ -60,9 +60,6 @@ base_rewrite = PatternMatcher([
]) ])
extra_pm = PatternMatcher([ extra_pm = PatternMatcher([
# insert a PRECAST before BITCAST to force it to be rendered. not needed on all backends?
(UPat(Ops.BITCAST, name="x"), lambda x: UOp(Ops.BITCAST, x.dtype, (UOp(Ops.PRECAST, x.src[0].dtype, x.src),))
if x.src[0].op not in {Ops.PRECAST, Ops.LOAD, Ops.CUSTOM} else None),
# devectorize any bools # devectorize any bools
(UPat((*GroupOp.ALU, Ops.CAST, Ops.BITCAST, Ops.INDEX), dtype=dtypes.bool, name="alu"), no_vectorized_alu), (UPat((*GroupOp.ALU, Ops.CAST, Ops.BITCAST, Ops.INDEX), dtype=dtypes.bool, name="alu"), no_vectorized_alu),
# CAST (from bool) can't be vectorized # CAST (from bool) can't be vectorized
@@ -181,7 +178,7 @@ class CStyleLanguage(Renderer):
elif u.op is Ops.RANGE: r[u] = f"{axis_letters[u.arg[-1]]}idx"+range_str(u) elif u.op is Ops.RANGE: r[u] = f"{axis_letters[u.arg[-1]]}idx"+range_str(u)
else: else:
prefix = {Ops.WMMA: "wmma", Ops.DEFINE_LOCAL: "temp", Ops.CONST: "const", prefix = {Ops.WMMA: "wmma", Ops.DEFINE_LOCAL: "temp", Ops.CONST: "const",
Ops.CAST: "cast", Ops.BITCAST: "cast", Ops.GEP: "gep", Ops.VECTORIZE: "cast", Ops.PRECAST: "precast", Ops.CAST: "cast", Ops.BITCAST: "cast", Ops.GEP: "gep", Ops.VECTORIZE: "cast",
Ops.INDEX: "bidx", Ops.DEFINE_REG: "acc", Ops.LOAD: "val"}.get(u.op, "alu") Ops.INDEX: "bidx", Ops.DEFINE_REG: "acc", Ops.LOAD: "val"}.get(u.op, "alu")
r[u] = f"{prefix}{c[prefix]}" r[u] = f"{prefix}{c[prefix]}"
@@ -278,7 +275,7 @@ class OpenCLRenderer(CStyleLanguage):
dtypes.bfloat16: "ushort" } dtypes.bfloat16: "ushort" }
string_rewrite = PatternMatcher([ string_rewrite = PatternMatcher([
(UPat(Ops.BITCAST, name="x"), lambda ctx,x: f"as_{ctx.render_dtype(x.dtype)}({ctx[x.src[0]]})"), (UPat(Ops.BITCAST, name="x"), lambda ctx,x: f"as_{ctx.render_dtype(x.dtype)}(({ctx.render_dtype(x.src[0].dtype)})({ctx[x.src[0]]}))"),
# load/store image (OpenCL) # load/store image (OpenCL)
(UPat(Ops.LOAD, dtype=dtypes.float.vec(4), src=(UPat.var('buf').index(UPat.var('idx', dtypes.int.vec(2)), UPat.var("gate")), UPat.var("var"))), (UPat(Ops.LOAD, dtype=dtypes.float.vec(4), src=(UPat.var('buf').index(UPat.var('idx', dtypes.int.vec(2)), UPat.var("gate")), UPat.var("var"))),
lambda ctx,buf,idx,var,gate: f"({ctx[gate]}?read_imagef({ctx[buf]}, smp, {ctx[idx]}):{ctx[var]})"), lambda ctx,buf,idx,var,gate: f"({ctx[gate]}?read_imagef({ctx[buf]}, smp, {ctx[idx]}):{ctx[var]})"),
@@ -338,7 +335,7 @@ class MetalRenderer(CStyleLanguage):
]) + extra_pm ]) + extra_pm
string_rewrite = PatternMatcher([ string_rewrite = PatternMatcher([
(UPat(Ops.BITCAST, name="x"), lambda ctx,x: f"as_type<{ctx.render_dtype(x.dtype)}>({ctx[x.src[0]]})"), (UPat(Ops.BITCAST, name="x"), lambda ctx,x: f"as_type<{ctx.render_dtype(x.dtype)}>(({ctx.render_dtype(x.src[0].dtype)})({ctx[x.src[0]]}))"),
]) + base_rewrite ]) + base_rewrite
def render_kernel(self, function_name, kernel, bufs, uops, prefix=None): def render_kernel(self, function_name, kernel, bufs, uops, prefix=None):
@@ -385,6 +382,10 @@ class CUDARenderer(CStyleLanguage):
extra_matcher = create_non_native_float_pats(dtypes.fp8s, casting=False) + PatternMatcher([ extra_matcher = create_non_native_float_pats(dtypes.fp8s, casting=False) + PatternMatcher([
(UPat(Ops.CAST, dtypes.fp8s, UPat.var("x", dtypes.fp8s), name='y'), lambda x,y: x.cast(dtypes.float).cast(y.dtype) if x.dtype!=y.dtype else None), (UPat(Ops.CAST, dtypes.fp8s, UPat.var("x", dtypes.fp8s), name='y'), lambda x,y: x.cast(dtypes.float).cast(y.dtype) if x.dtype!=y.dtype else None),
]) + extra_pm ]) + extra_pm
string_rewrite = PatternMatcher([
(UPat(Ops.BITCAST, name="x"), lambda ctx,x: f"tg_bitcast<{ctx.render_dtype(x.dtype)}>(({ctx.render_dtype(x.src[0].dtype)})({ctx[x.src[0]]}))"),
]) + base_rewrite
def render_vector_prefix(self, dt:DType) -> str: def render_vector_prefix(self, dt:DType) -> str:
vec, scal = self.render_dtype(dt), self.render_dtype(dt.scalar()), vec, scal = self.render_dtype(dt), self.render_dtype(dt.scalar()),
elems, header = ', '.join(_nms[:dt.count]), ', '.join([f"{scal} {x}" for x in _nms[:dt.count]]) elems, header = ', '.join(_nms[:dt.count]), ', '.join([f"{scal} {x}" for x in _nms[:dt.count]])
@@ -392,8 +393,8 @@ class CUDARenderer(CStyleLanguage):
def render_kernel(self, function_name, kernel, bufs, uops, prefix=None): def render_kernel(self, function_name, kernel, bufs, uops, prefix=None):
# TODO: why is dtypes.bfloat16.name == "__bf16"? would be easier not override dtypes.name # TODO: why is dtypes.bfloat16.name == "__bf16"? would be easier not override dtypes.name
prefix = ["#define INFINITY (__int_as_float(0x7f800000))","#define NAN (__int_as_float(0x7fffffff))"] prefix = ["#define INFINITY (__int_as_float(0x7f800000))", "#define NAN (__int_as_float(0x7fffffff))",
"template <class T, class F> __device__ __forceinline__ T tg_bitcast(F v) { union U { F f; T t; }; U u; u.f = v; return u.t; }"]
used_dtypes = uops_to_dtypes(uops) used_dtypes = uops_to_dtypes(uops)
if any(dt.scalar() in dtypes.fp8s for dt in used_dtypes): prefix.append("#include <cuda_fp8.h>") if any(dt.scalar() in dtypes.fp8s for dt in used_dtypes): prefix.append("#include <cuda_fp8.h>")
if any(dt.scalar() == dtypes.half for dt in used_dtypes): prefix.append("#include <cuda_fp16.h>") if any(dt.scalar() == dtypes.half for dt in used_dtypes): prefix.append("#include <cuda_fp16.h>")
+33 -21
View File
@@ -2,21 +2,22 @@ from typing import cast
import math, struct, sys import math, struct, sys
from tinygrad.codegen.opt import tc from tinygrad.codegen.opt import tc
from tinygrad.renderer import Renderer from tinygrad.renderer import Renderer
from tinygrad.renderer.cstyle import AMDRenderer from tinygrad.renderer.cstyle import AMDRenderer, create_non_native_float_pats
from tinygrad.uop.decompositions import xexp2, xlog2 from tinygrad.uop.decompositions import xexp2, xlog2
from tinygrad.uop.ops import UOp, PatternMatcher, UPat, Ops, GroupOp, range_str from tinygrad.uop.ops import UOp, PatternMatcher, UPat, Ops, GroupOp, range_str
from tinygrad.dtype import dtypes, DType, PtrDType, truncate from tinygrad.dtype import dtypes, float_to_fp8, DType, PtrDType, truncate
from tinygrad.helpers import prod, AMX from tinygrad.helpers import prod, AMX
def ldt(dt:DType): def ldt(dt:DType):
if dt.vcount > 1: return f"<{dt.vcount} x {ldt(dt.scalar())}>" if dt.vcount > 1: return f"<{dt.vcount} x {ldt(dt.scalar())}>"
if isinstance(dt, PtrDType): return ldt(dt.base) + "*" if isinstance(dt, PtrDType): return ldt(dt.base) + "*"
return {dtypes.void: "void", dtypes.bool: "i1", dtypes.int8: "i8", dtypes.int16: "i16", dtypes.int32: "i32", dtypes.int64: "i64", return {dtypes.void: "void", dtypes.bool: "i1", dtypes.int8: "i8", dtypes.int16: "i16", dtypes.int32: "i32", dtypes.int64: "i64",
dtypes.uint8: "i8", dtypes.uint16: "i16", dtypes.uint32: "i32", dtypes.uint64: "i64", dtypes.uint8: "i8", dtypes.uint16: "i16", dtypes.uint32: "i32", dtypes.uint64: "i64", dtypes.fp8e4m3: "i8", dtypes.fp8e5m2: "i8",
dtypes.float16: "half", dtypes.bfloat16: "bfloat", dtypes.float32: "float", dtypes.float64: "double"}[dt] dtypes.float16: "half", dtypes.bfloat16: "bfloat", dtypes.float32: "float", dtypes.float64: "double"}[dt]
def lconst(x, dtype:DType): def lconst(x, dtype:DType):
if dtype in dtypes.floats: if dtype in dtypes.floats:
if dtype in dtypes.fp8s: return float_to_fp8(x, dtype)
if math.isinf(x) or math.isnan(x): return "0x%02X%02X%02X%02X%02X%02X%02X%02X" % tuple(struct.pack("d",x)[::-1]) if math.isinf(x) or math.isnan(x): return "0x%02X%02X%02X%02X%02X%02X%02X%02X" % tuple(struct.pack("d",x)[::-1])
return truncate[dtype](x) return truncate[dtype](x)
return int(x) return int(x)
@@ -47,13 +48,14 @@ def render_wmma_amx(ctx, wmma: UOp) -> str:
f' {ctx[wmma]} = load {ldt(wmma.dtype)}, ptr {ctx[wmma]}_amx2, align {wmma.dtype.itemsize}']) f' {ctx[wmma]} = load {ldt(wmma.dtype)}, ptr {ctx[wmma]}_amx2, align {wmma.dtype.itemsize}'])
def render_wmma_amd(ctx, wmma: UOp, cdna=False) -> str: def render_wmma_amd(ctx, wmma: UOp, cdna=False) -> str:
dt_map = {dtypes.half: "f16", dtypes.float: "f32", dtypes.ushort: "bf16.1k" if cdna else "bf16", dtypes.bfloat16: "bf16.1k" if cdna else "bf16"} dt_map = {dtypes.half: "f16", dtypes.float: "f32", dtypes.ushort: "bf16.1k" if cdna else "bf16", dtypes.bfloat16: "bf16.1k" if cdna else "bf16",
dtypes.fp8e4m3: ".fp8.fp8", dtypes.fp8e5m2: ".bf8.bf8"}
# https://github.com/llvm/llvm-project/blob/main/clang/test/CodeGenOpenCL/builtins-amdgcn-mfma.cl # https://github.com/llvm/llvm-project/blob/main/clang/test/CodeGenOpenCL/builtins-amdgcn-mfma.cl
N,M,K = wmma.arg[1] N,M,K = wmma.arg[1]
if cdna: if cdna:
if K == 32: dt_map.update({dtypes.half: ".f16", dtypes.bfloat16: ".bf16"}) if K == 32: dt_map.update({dtypes.half: ".f16", dtypes.bfloat16: ".bf16"})
return f" {ctx[wmma]} = call {ldt(wmma.dtype)} @llvm.amdgcn.mfma.{dt_map[wmma.src[-1].dtype.scalar()]}" + \ return f" {ctx[wmma]} = call {ldt(wmma.dtype)} @llvm.amdgcn.mfma.{dt_map[wmma.src[-1].dtype.scalar()]}" + \
f".{N}x{M}x{K}{dt_map[wmma.src[0].dtype.scalar()]}(" + ", ".join([f"{ldt(w.dtype)} {ctx[w]}" for w in wmma.src]) + ", i32 0, i32 0, i32 0)" f".{N}x{M}x{K}{dt_map[wmma.arg[2]]}(" + ", ".join([f"{ldt(w.dtype)} {ctx[w]}" for w in wmma.src]) + ", i32 0, i32 0, i32 0)"
# https://github.com/llvm/llvm-project/blob/main/llvm/test/CodeGen/AMDGPU/GlobalISel/llvm.amdgcn.wmma_32.ll # https://github.com/llvm/llvm-project/blob/main/llvm/test/CodeGen/AMDGPU/GlobalISel/llvm.amdgcn.wmma_32.ll
# example: %wmma0 = call <8 x float> @llvm.amdgcn.wmma.f32.16x16x16.f16(<16 x half> %v99,<16 x half> %v100,<8 x float> %v101) # example: %wmma0 = call <8 x float> @llvm.amdgcn.wmma.f32.16x16x16.f16(<16 x half> %v99,<16 x half> %v100,<8 x float> %v101)
return f" {ctx[wmma]} = call {ldt(wmma.dtype)} @llvm.amdgcn.wmma.{dt_map[wmma.src[-1].dtype.scalar()]}.16x16x16." + \ return f" {ctx[wmma]} = call {ldt(wmma.dtype)} @llvm.amdgcn.wmma.{dt_map[wmma.src[-1].dtype.scalar()]}.16x16x16." + \
@@ -136,27 +138,16 @@ class LLVMRenderer(Renderer):
has_local = False has_local = False
global_max: tuple[int, ...] | None = None global_max: tuple[int, ...] | None = None
string_rewrite = base_rewrite + PatternMatcher([(UPat(Ops.WMMA, name="wmma"), render_wmma_amx)]) string_rewrite = base_rewrite + PatternMatcher([(UPat(Ops.WMMA, name="wmma"), render_wmma_amx)])
code_for_op = {Ops.FDIV: lambda: None} code_for_op = {Ops.FDIV: lambda: None, Ops.CMPLT: lambda: None}
if AMX: tensor_cores = tc.amx if AMX: tensor_cores = tc.amx
extra_matcher = PatternMatcher([ extra_matcher = create_non_native_float_pats((dtypes.bfloat16,))
# rewrite MAX to CMPLT + WHERE
(UPat(Ops.MAX, name="m"), lambda m: (m.src[0] < m.src[1]).where(m.src[1], m.src[0])),
# copied from cstyle.py, upcast to float32 all the ops that don't support bfloat16
(UPat((Ops.SQRT, Ops.EXP2, Ops.LOG2, Ops.SIN), dtype=dtypes.bfloat16, name="x"),
lambda x: (UOp(x.op, dtypes.float, tuple(vv.cast(dtypes.float) for vv in x.src), x.arg).cast(dtypes.bfloat16))),
# copied from cstyle.py, add float intermediate casting
(UPat(Ops.CAST, name="x", src=UPat.var("y", dtypes.bfloat16)),lambda x,y: y.cast(dtypes.float).cast(x.dtype) if x.dtype!=dtypes.float else None),
(UPat(Ops.CAST, dtypes.bfloat16, UPat.var("x")),lambda x: x.cast(dtypes.float).cast(dtypes.bfloat16) if x.dtype!=dtypes.float else None),
])
def render(self, uops: list[UOp]) -> str: return "\n".join((k:=self._render_kernel(uops))[0] + (k[1], self._render_footer(uops))) def render(self, uops: list[UOp]) -> str: return "\n".join((k:=self._render_kernel(uops))[0] + (k[1], self._render_footer(uops)))
def _render_footer(self, uops: list[UOp]) -> str: return 'attributes #0 = { alwaysinline nounwind "no-builtins" "no-trapping-math"="true" }' def _render_footer(self, uops: list[UOp]) -> str: return 'attributes #0 = { alwaysinline nounwind "no-builtins" "no-trapping-math"="true" }'
def _render_fn(self, name:str, args:list[tuple[str,DType]], kernel:list[str], prefix:list[str]|None=None) -> str: def _render_fn(self, name:str, args:list[tuple[str,DType]], kernel:list[str], prefix:list[str]|None=None) -> str:
# NOTE: CPUAllocator promises 0x20 alignment # NOTE: CPUAllocator promises 0x20 alignment
sargs = ", ".join([f"{ldt(dt)}{' noalias align 32' if isinstance(dt, PtrDType) else ''} {name}" for name,dt in args]) sargs = ", ".join([f"{ldt(dt)}{' noalias align 32' if isinstance(dt, PtrDType) else ''} {name}" for name,dt in args])
sprefix = "".join([f" {x}" for x in (prefix or []) + [self.abi] if x is not None]) return "\n".join((prefix or []) + [f"define{' ' + self.abi if self.abi else ''} void @{name}({sargs}) #0", "{"] + kernel + [" ret void\n}"])
return "\n".join([f"define{sprefix} void @{name}({sargs}) #0", "{"] + kernel + [" ret void\n}"])
def _render_kernel(self, uops: list[UOp], prefix:list[str]|None=None) -> tuple[tuple[str, ...], str]: def _render_kernel(self, uops: list[UOp], prefix:list[str]|None=None) -> tuple[tuple[str, ...], str]:
r: dict[UOp, str] = {} r: dict[UOp, str] = {}
args: list[tuple[str, DType]] = [] args: list[tuple[str, DType]] = []
@@ -226,8 +217,13 @@ class AMDLLVMRenderer(LLVMRenderer):
(UPat(tuple(llvm_intrinsics), name="x"), (UPat(tuple(llvm_intrinsics), name="x"),
lambda ctx, x: f" {ctx[x]} = call {ldt(x.dtype)} @llvm.{llvm_intrinsics[x.op]}.{ldt(x.dtype.scalar())}({ldt(x.src[0].dtype)} {ctx[x.src[0]]})"), lambda ctx, x: f" {ctx[x]} = call {ldt(x.dtype)} @llvm.{llvm_intrinsics[x.op]}.{ldt(x.dtype.scalar())}({ldt(x.src[0].dtype)} {ctx[x.src[0]]})"),
(UPat(Ops.BARRIER), lambda ctx: barrier), (UPat(Ops.BARRIER), lambda ctx: barrier),
(UPat(Ops.CAST, dtypes.fp8s, (UPat.var("y", dtypes.float),), name="x",), lambda ctx,x,y:
f" {ctx[x]} = call i8 @f32_to_fp8({ldt(x.src[0].dtype)} {ctx[x.src[0]]}, i1 {'1' if x.dtype == dtypes.fp8e5m2 else '0'})"),
(UPat(Ops.CAST, dtypes.float, (UPat.var("y", dtypes.fp8s),), name="x",), lambda ctx,x,y:
f" {ctx[x.src[0]]}_i32 = zext i8 {ctx[x.src[0]]} to i32\n"
f" {ctx[x]} = call float @llvm.amdgcn.cvt.f32.{'bf8' if y.dtype == dtypes.fp8e5m2 else 'fp8'}(i32 {ctx[x.src[0]]}_i32, i32 0)"),
]) + base_rewrite ]) + base_rewrite
extra_matcher = LLVMRenderer.extra_matcher + PatternMatcher([ extra_matcher = LLVMRenderer.extra_matcher + create_non_native_float_pats(dtypes.fp8s) + PatternMatcher([
(UPat(Ops.CAST, name="x", dtype=dtypes.half.vec(16), src=UPat.var("y", dtypes.half.vec(8))), (UPat(Ops.CAST, name="x", dtype=dtypes.half.vec(16), src=UPat.var("y", dtypes.half.vec(8))),
lambda x, y: UOp(Ops.VECTORIZE, dtypes.half.vec(16), tuple(y.gep(i // 2) if i % 2 == 0 else UOp.const(dtypes.half, 0.0) for i in range(16)))), lambda x, y: UOp(Ops.VECTORIZE, dtypes.half.vec(16), tuple(y.gep(i // 2) if i % 2 == 0 else UOp.const(dtypes.half, 0.0) for i in range(16)))),
(UPat(Ops.CAST, name="x", dtype=dtypes.half.vec(8), src=UPat.var("y", dtypes.half.vec(16))), (UPat(Ops.CAST, name="x", dtype=dtypes.half.vec(8), src=UPat.var("y", dtypes.half.vec(16))),
@@ -236,6 +232,19 @@ class AMDLLVMRenderer(LLVMRenderer):
(UPat(Ops.LOG2, dtype=dtypes.double, src=(UPat.var("d"),)), xlog2), (UPat(Ops.LOG2, dtype=dtypes.double, src=(UPat.var("d"),)), xlog2),
(UPat(Ops.EXP2, dtype=dtypes.double, src=(UPat.var("d"),)), xexp2), (UPat(Ops.EXP2, dtype=dtypes.double, src=(UPat.var("d"),)), xexp2),
]) ])
def render(self, uops: list[UOp]) -> str:
prefix = ["""define i8 @f32_to_fp8(float %val, i1 %is_bf8) {
entry: %ival = bitcast float %val to i32\n %exp = and i32 %ival, 2139095040\n %is_special = icmp eq i32 %exp, 2139095040
br i1 %is_special, label %select_clip, label %clip
clip: br i1 %is_bf8, label %bf8_clip, label %fp8_clip
bf8_clip: %clamped_bf8 = call float @llvm.amdgcn.fmed3.f32(float %val, float 57344.0, float -57344.0)\n br label %select_clip
fp8_clip: %clamped_fp8 = call float @llvm.amdgcn.fmed3.f32(float %val, float 448.0, float -448.0) \n br label %select_clip
select_clip: %phi_val = phi float [%val, %entry], [%clamped_bf8, %bf8_clip], [%clamped_fp8, %fp8_clip]\n br i1 %is_bf8, label %do_bf8, label %do_fp8
do_bf8: %packed_bf8 = call i32 @llvm.amdgcn.cvt.pk.bf8.f32(float %phi_val, float %phi_val, i32 0, i1 false)\n br label %exit
do_fp8: %packed_fp8 = call i32 @llvm.amdgcn.cvt.pk.fp8.f32(float %phi_val, float %phi_val, i32 0, i1 false)\n br label %exit
exit: %packed = phi i32 [%packed_bf8, %do_bf8], [%packed_fp8, %do_fp8]\n %trunc = trunc i32 %packed to i8\n ret i8 %trunc
}""".replace(": ", ":\n ")] if any(u.dtype in dtypes.fp8s for u in uops) else []
return "\n".join((k:=self._render_kernel(uops, prefix))[0] + (k[1], self._render_footer(uops)))
def _render_footer(self, uops: list[UOp]) -> str: def _render_footer(self, uops: list[UOp]) -> str:
# TODO: this is copied from cstyle # TODO: this is copied from cstyle
local_dims = [u.src[0] for u in uops if u.op is Ops.SPECIAL and u.arg[0] == "l"] local_dims = [u.src[0] for u in uops if u.op is Ops.SPECIAL and u.arg[0] == "l"]
@@ -252,7 +261,10 @@ class AMDLLVMRenderer(LLVMRenderer):
self.extra_matcher += PatternMatcher([ self.extra_matcher += PatternMatcher([
(UPat(Ops.WMMA, name="x", dtype=dtypes.float.vec(4)), (UPat(Ops.WMMA, name="x", dtype=dtypes.float.vec(4)),
lambda x: UOp(Ops.WMMA, dtypes.float.vec(4), (x.src[0].bitcast(dtypes.uint16.vec(4)), x.src[1].bitcast(dtypes.uint16.vec(4)), lambda x: UOp(Ops.WMMA, dtypes.float.vec(4), (x.src[0].bitcast(dtypes.uint16.vec(4)), x.src[1].bitcast(dtypes.uint16.vec(4)),
x.src[2]), (*x.arg,)) if x.src[0].dtype == dtypes.bfloat16.vec(4) else None) x.src[2]), (*x.arg,)) if x.src[0].dtype == dtypes.bfloat16.vec(4) else None),
(UPat(Ops.WMMA, name="x", dtype=dtypes.float.vec(4)),
lambda x: UOp(Ops.WMMA, dtypes.float.vec(4), (x.src[0].bitcast(dtypes.uint64), x.src[1].bitcast(dtypes.uint64),
x.src[2]), (*x.arg,)) if x.src[0].dtype in (dtypes.fp8e4m3.vec(8), dtypes.fp8e5m2.vec(8)) else None),
]) ])
if self.arch.split(":")[0] == "gfx1100": if self.arch.split(":")[0] == "gfx1100":
self.extra_matcher += PatternMatcher([ self.extra_matcher += PatternMatcher([
+28 -12
View File
@@ -1,6 +1,6 @@
from __future__ import annotations from __future__ import annotations
from typing import cast, ClassVar from typing import cast, ClassVar
import os, ctypes, struct, hashlib, functools, importlib, mmap, errno, array, contextlib, sys, weakref, itertools, collections import os, ctypes, struct, hashlib, functools, importlib, mmap, errno, array, contextlib, sys, weakref, itertools, collections, atexit
assert sys.platform != 'win32' assert sys.platform != 'win32'
from dataclasses import dataclass from dataclasses import dataclass
from tinygrad.runtime.support.hcq import HCQCompiled, HCQAllocator, HCQBuffer, HWQueue, CLikeArgsState, HCQSignal, HCQProgram, FileIOInterface from tinygrad.runtime.support.hcq import HCQCompiled, HCQAllocator, HCQBuffer, HWQueue, CLikeArgsState, HCQSignal, HCQProgram, FileIOInterface
@@ -20,7 +20,8 @@ from tinygrad.runtime.support.amd import AMDReg, AMDIP, import_module, import_so
from tinygrad.runtime.support.system import System, PCIIfaceBase, PCIAllocationMeta, PCIDevice, USBPCIDevice, MAP_FIXED, MAP_NORESERVE from tinygrad.runtime.support.system import System, PCIIfaceBase, PCIAllocationMeta, PCIDevice, USBPCIDevice, MAP_FIXED, MAP_NORESERVE
if getenv("IOCTL"): import extra.hip_gpu_driver.hip_ioctl # noqa: F401 # pylint: disable=unused-import if getenv("IOCTL"): import extra.hip_gpu_driver.hip_ioctl # noqa: F401 # pylint: disable=unused-import
SQTT, SQTT_ITRACE_SE_MASK, PMC = ContextVar("SQTT", VIZ.value>=2), ContextVar("SQTT_ITRACE_SE_MASK", 0b11), ContextVar("PMC", 0) SQTT, SQTT_ITRACE_SE_MASK, SQTT_LIMIT_SE = ContextVar("SQTT", VIZ.value>=2), ContextVar("SQTT_ITRACE_SE_MASK", 0b11), ContextVar("SQTT_LIMIT_SE", 0)
PMC = ContextVar("PMC", 0)
EVENT_INDEX_PARTIAL_FLUSH = 4 # based on a comment in nvd.h EVENT_INDEX_PARTIAL_FLUSH = 4 # based on a comment in nvd.h
WAIT_REG_MEM_FUNCTION_EQ = 3 # == WAIT_REG_MEM_FUNCTION_EQ = 3 # ==
WAIT_REG_MEM_FUNCTION_NEQ = 4 # != WAIT_REG_MEM_FUNCTION_NEQ = 4 # !=
@@ -193,11 +194,18 @@ class AMDComputeQueue(HWQueue):
bind_point=(__BIND_POINT_COMPUTE:=1), api_pso_hash=data64_le(prg.libhash[0]))) bind_point=(__BIND_POINT_COMPUTE:=1), api_pso_hash=data64_le(prg.libhash[0])))
self.sqtt_userdata(sqtt.struct_rgp_sqtt_marker_event(has_thread_dims=1, cmd_id=next(prg.dev.sqtt_next_cmd_id)), *global_size) self.sqtt_userdata(sqtt.struct_rgp_sqtt_marker_event(has_thread_dims=1, cmd_id=next(prg.dev.sqtt_next_cmd_id)), *global_size)
se_cap = max(prod([x if isinstance(x, int) else 1 for x in global_size]) // 4, 1) // 32 if SQTT_LIMIT_SE:
for xcc in range(self.dev.xccs): # Calculate number of CUs per SE to enable based on blocks count. 4 is maximum simd per CU, but on rdna we can trace only 1.
with self.pred_exec(xcc_mask=1 << xcc): cu_per_se = prod([x if isinstance(x, int) else 1 for x in global_size]) // (((self.dev.max_cu_id + 1) // self.dev.se_cnt) * 4)
for i in range(8 if prg.dev.target >= (11,0,0) else 4): for xcc in range(self.dev.xccs):
self.wreg(getattr(self.gc, f'regCOMPUTE_STATIC_THREAD_MGMT_SE{i}'), min(0xffffffff, (1 << (se_cap + (1 if i == 0 else 0))) - 1)) with self.pred_exec(xcc_mask=1 << xcc):
for i in range(8 if prg.dev.target >= (11,0,0) else 4):
if SQTT_LIMIT_SE > 1: mask = 1 if SQTT_ITRACE_SE_MASK.value & (1 << i) else 0 # only run unmasked shader engines
else:
sa_mask = (1 << (self.dev.iface.props['cu_per_simd_array'] // 2)) - 1
cu_mask = (1 << (cu_per_se + (1 if i == 0 else 0))) - 1
mask = lo32((cu_mask & sa_mask) | (cu_mask & (sa_mask << 16)) << 16)
self.wreg(getattr(self.gc, f'regCOMPUTE_STATIC_THREAD_MGMT_SE{i}'), mask)
def sqtt_userdata(self, data, *extra_dwords): def sqtt_userdata(self, data, *extra_dwords):
data_ints = [x[0] for x in struct.iter_unpack('<I', bytes(data))] + list(extra_dwords) data_ints = [x[0] for x in struct.iter_unpack('<I', bytes(data))] + list(extra_dwords)
@@ -776,8 +784,16 @@ class KFDIface:
raise RuntimeError("\n".join(report)) raise RuntimeError("\n".join(report))
def is_in_profile_mode(self): def require_profile_mode(self, can_set_mode=True):
return self.dev.target[0] == 9 or FileIOInterface(f'{self.dev_sysfs_path}/power_dpm_force_performance_level').read()[:16] == 'profile_standard' if self.dev.target[0] == 9: return
fn = f'{self.dev_sysfs_path}/power_dpm_force_performance_level'
if (perflevel:=FileIOInterface(fn).read().strip()) != 'profile_standard':
if can_set_mode:
atexit.register(lambda: os.system(f"echo '{perflevel}' | sudo tee {fn} > /dev/null"))
os.system(f"echo 'profile_standard' | sudo tee {fn} > /dev/null")
self.require_profile_mode(can_set_mode=False)
else:
raise RuntimeError("PMC/SQTT requires stable power state: run `amd-smi set -l stable_std` for KFD iface")
class PCIIface(PCIIfaceBase): class PCIIface(PCIIfaceBase):
gpus:ClassVar[list[str]] = [] gpus:ClassVar[list[str]] = []
@@ -788,7 +804,7 @@ class PCIIface(PCIIfaceBase):
self._setup_adev(self.pci_dev) self._setup_adev(self.pci_dev)
self.pci_dev.write_config(pci.PCI_COMMAND, self.pci_dev.read_config(pci.PCI_COMMAND, 2) | pci.PCI_COMMAND_MASTER, 2) self.pci_dev.write_config(pci.PCI_COMMAND, self.pci_dev.read_config(pci.PCI_COMMAND, 2) | pci.PCI_COMMAND_MASTER, 2)
def is_in_profile_mode(self): return True def require_profile_mode(self): return True
def _setup_adev(self, pci_dev:PCIDevice, dma_regions:list[tuple[int, MMIOInterface]]|None=None): def _setup_adev(self, pci_dev:PCIDevice, dma_regions:list[tuple[int, MMIOInterface]]|None=None):
self.dev_impl:AMDev = AMDev(pci_dev, dma_regions) self.dev_impl:AMDev = AMDev(pci_dev, dma_regions)
@@ -925,7 +941,7 @@ class AMDDevice(HCQCompiled):
self.pmc_enabled = PROFILE and PMC > 0 self.pmc_enabled = PROFILE and PMC > 0
if self.pmc_enabled: if self.pmc_enabled:
if self.target[0] not in {9, 11, 12}: raise RuntimeError(f'PMC are not supported on gc:{self.target}') if self.target[0] not in {9, 11, 12}: raise RuntimeError(f'PMC are not supported on gc:{self.target}')
if not self.iface.is_in_profile_mode(): raise RuntimeError("PMC requires stable power state: run `amd-smi set -l stable_std` for KFD iface") self.iface.require_profile_mode()
self.pmc_sched:list[PMCSample] = [] self.pmc_sched:list[PMCSample] = []
self.pmc_counters = import_pmc(self.target) self.pmc_counters = import_pmc(self.target)
@@ -943,7 +959,7 @@ class AMDDevice(HCQCompiled):
self.sqtt_enabled = PROFILE and SQTT > 0 self.sqtt_enabled = PROFILE and SQTT > 0
if self.sqtt_enabled: if self.sqtt_enabled:
if self.target[0] not in {9, 11, 12}: raise RuntimeError(f'SQ Thread Tracing is not supported on gc:{self.target}') if self.target[0] not in {9, 11, 12}: raise RuntimeError(f'SQ Thread Tracing is not supported on gc:{self.target}')
if not self.iface.is_in_profile_mode(): raise RuntimeError("SQTT requires stable power state: run `amd-smi set -l stable_std` for KFD iface") self.iface.require_profile_mode()
SQTT_BUFFER_SIZE = getenv("SQTT_BUFFER_SIZE", 256) # in mb, per shader engine SQTT_BUFFER_SIZE = getenv("SQTT_BUFFER_SIZE", 256) # in mb, per shader engine
self.sqtt_buffers = [self.allocator.alloc(SQTT_BUFFER_SIZE << 20, BufferSpec(nolru=True, uncached=True)) for _ in range(self.se_cnt)] self.sqtt_buffers = [self.allocator.alloc(SQTT_BUFFER_SIZE << 20, BufferSpec(nolru=True, uncached=True)) for _ in range(self.se_cnt)]
+3 -2
View File
@@ -360,7 +360,7 @@ class NVKIface:
self.dma_class:int = next(c for c in [nv_gpu.BLACKWELL_DMA_COPY_B, nv_gpu.AMPERE_DMA_COPY_B] if c in self.nvclasses) self.dma_class:int = next(c for c in [nv_gpu.BLACKWELL_DMA_COPY_B, nv_gpu.AMPERE_DMA_COPY_B] if c in self.nvclasses)
usermode = self.rm_alloc(self.dev.subdevice, self.usermode_class) usermode = self.rm_alloc(self.dev.subdevice, self.usermode_class)
return usermode, MMIOInterface(self._gpu_map_to_cpu(usermode, mmio_sz:=0x10000, flags=2), mmio_sz, fmt='I') return usermode, MMIOInterface(self._gpu_map_to_cpu(usermode, mmio_sz:=0x10000), mmio_sz, fmt='I')
def setup_vm(self, vaspace): def setup_vm(self, vaspace):
self.rm_control(self.dev.subdevice, nv_gpu.NV2080_CTRL_CMD_GPU_GET_GID_INFO, raw_uuid:=nv_gpu.NV2080_CTRL_GPU_GET_GID_INFO_PARAMS( self.rm_control(self.dev.subdevice, nv_gpu.NV2080_CTRL_CMD_GPU_GET_GID_INFO, raw_uuid:=nv_gpu.NV2080_CTRL_GPU_GET_GID_INFO_PARAMS(
@@ -514,7 +514,8 @@ class NVDevice(HCQCompiled[HCQSignal]):
channel_params = nv_gpu.NV_CHANNEL_GROUP_ALLOCATION_PARAMETERS(engineType=nv_gpu.NV2080_ENGINE_TYPE_GRAPHICS) channel_params = nv_gpu.NV_CHANNEL_GROUP_ALLOCATION_PARAMETERS(engineType=nv_gpu.NV2080_ENGINE_TYPE_GRAPHICS)
channel_group = self.iface.rm_alloc(self.nvdevice, nv_gpu.KEPLER_CHANNEL_GROUP_A, channel_params) channel_group = self.iface.rm_alloc(self.nvdevice, nv_gpu.KEPLER_CHANNEL_GROUP_A, channel_params)
gpfifo_area = self.iface.alloc(0x200000, contiguous=True, cpu_access=True, force_devmem=True, map_flags=0x10d0000) gpfifo_area = self.iface.alloc(0x200000, contiguous=True, cpu_access=True, force_devmem=True,
map_flags=(nv_gpu.NVOS33_FLAGS_CACHING_TYPE_WRITECOMBINED<<23))
ctxshare_params = nv_gpu.NV_CTXSHARE_ALLOCATION_PARAMETERS(hVASpace=vaspace, flags=nv_gpu.NV_CTXSHARE_ALLOCATION_FLAGS_SUBCONTEXT_ASYNC) ctxshare_params = nv_gpu.NV_CTXSHARE_ALLOCATION_PARAMETERS(hVASpace=vaspace, flags=nv_gpu.NV_CTXSHARE_ALLOCATION_FLAGS_SUBCONTEXT_ASYNC)
ctxshare = self.iface.rm_alloc(channel_group, nv_gpu.FERMI_CONTEXT_SHARE_A, ctxshare_params) ctxshare = self.iface.rm_alloc(channel_group, nv_gpu.FERMI_CONTEXT_SHARE_A, ctxshare_params)
+3 -3
View File
@@ -1,4 +1,4 @@
import socket, json, asyncio, threading import socket, json, asyncio, threading, math
from contextlib import asynccontextmanager from contextlib import asynccontextmanager
from tinygrad.device import Compiled, Allocator from tinygrad.device import Compiled, Allocator
from tinygrad.helpers import DEBUG, getenv from tinygrad.helpers import DEBUG, getenv
@@ -92,9 +92,9 @@ class TinyFSAllocator(Allocator[TinyFSDevice]):
if dest.device.op == "LOAD": if dest.device.op == "LOAD":
locs = self.dev.sfile.readline() locs = self.dev.sfile.readline()
dest.copyout_queue = json.loads(locs) dest.copyout_queue = json.loads(locs)
dest.hash_buf[:] = src.tobytes() dest.hash_buf = src.tobytes()
elif dest.device.op == "STORE": elif dest.device.op == "STORE":
expected_hashes = dest.size // Tensor.CHUNK_SIZE expected_hashes = math.ceil(dest.size / Tensor.CHUNK_SIZE)
dest.hash_buf = bytearray(expected_hashes * 16) dest.hash_buf = bytearray(expected_hashes * 16)
self.dev.sfile.readinto(dest.hash_buf) self.dev.sfile.readinto(dest.hash_buf)
+3 -1
View File
@@ -65,7 +65,9 @@ def import_ip_offsets(ip): return type("IPOFF", (object,), import_header(f"inclu
def import_pmc(ip) -> dict[str, tuple[str, int]]: def import_pmc(ip) -> dict[str, tuple[str, int]]:
res:dict[str, tuple[str, int]] = {} res:dict[str, tuple[str, int]] = {}
arch = f"gfx{ip[0]}{ip[1]:x}{ip[2]:x}"
# NOTE: precise arch for mi300+, generic for others, since rocm headers lack some archs
arch = f"gfx{ip[0]}{ip[1]:x}{ip[2]:x}" if ip[0] == 9 else f"gfx{ip[0]}"
for sec in header_download("rocprofiler-compute/src/rocprof_compute_soc/profile_configs/counter_defs.yaml", url=ROCM_URL).split('- name: ')[1:]: for sec in header_download("rocprofiler-compute/src/rocprof_compute_soc/profile_configs/counter_defs.yaml", url=ROCM_URL).split('- name: ')[1:]:
for arch_spec in sec.split('- architectures:')[1:]: for arch_spec in sec.split('- architectures:')[1:]:
+1 -1
View File
@@ -241,7 +241,7 @@ def gen(dll, files, args=[], prolog=[], rules=[], epilog=[], recsym=False, use_e
it = iter(toks[1:]) it = iter(toks[1:])
_args = [nm(t) for t in itertools.takewhile(lambda t:nm(t)!=')', it) if clang.clang_getTokenKind(t) == clang.CXToken_Identifier] _args = [nm(t) for t in itertools.takewhile(lambda t:nm(t)!=')', it) if clang.clang_getTokenKind(t) == clang.CXToken_Identifier]
if len(body:=list(it)) == 0: continue if len(body:=list(it)) == 0: continue
macros += [f"{nm(c)} = lambda {','.join(_args)}: {readext(f, loc(body[0]), clang.clang_getRangeEnd(extent(toks[-1])))}"] macros += [f"{nm(c)} = lambda{' ' * bool(_args)}{','.join(_args)}: {readext(f,loc(body[0]),clang.clang_getRangeEnd(extent(toks[-1])))}"]
else: macros += [f"{nm(c)} = {readext(f, loc(toks[1]), clang.clang_getRangeEnd(extent(toks[-1])))}"] else: macros += [f"{nm(c)} = {readext(f, loc(toks[1]), clang.clang_getRangeEnd(extent(toks[-1])))}"]
case clang.CXCursor_VarDecl if clang.clang_getCursorLinkage(c) == clang.CXLinkage_Internal: case clang.CXCursor_VarDecl if clang.clang_getCursorLinkage(c) == clang.CXLinkage_Internal:
ty = clang.clang_getCursorType(c) ty = clang.clang_getCursorType(c)
+1 -1
View File
@@ -103,7 +103,7 @@ class HIPCCCompiler(Compiler):
subprocess.run(["hipcc", "-c", "-emit-llvm", "--cuda-device-only", "-O3", "-mcumode", subprocess.run(["hipcc", "-c", "-emit-llvm", "--cuda-device-only", "-O3", "-mcumode",
f"--offload-arch={self.arch}", "-I/opt/rocm/include/hip", "-o", bcf.name, srcf.name] + self.extra_options, check=True) f"--offload-arch={self.arch}", "-I/opt/rocm/include/hip", "-o", bcf.name, srcf.name] + self.extra_options, check=True)
subprocess.run(["hipcc", "-target", "amdgcn-amd-amdhsa", f"-mcpu={self.arch}", subprocess.run(["hipcc", "-target", "amdgcn-amd-amdhsa", f"-mcpu={self.arch}",
"-O3", "-mllvm", "-amdgpu-internalize-symbols", "-c", "-o", libf.name, bcf.name], check=True) "-O3", "-mllvm", "-amdgpu-internalize-symbols", "-c", "-o", libf.name, bcf.name] + self.extra_options, check=True)
return pathlib.Path(libf.name).read_bytes() return pathlib.Path(libf.name).read_bytes()
def disassemble(self, lib:bytes): amdgpu_disassemble(lib) def disassemble(self, lib:bytes): amdgpu_disassemble(lib)
+2 -3
View File
@@ -3,7 +3,7 @@ from typing import cast, Callable, Type, TypeVar, Generic, Any, Sequence
import contextlib, decimal, statistics, time, ctypes, array, os, struct, collections, functools import contextlib, decimal, statistics, time, ctypes, array, os, struct, collections, functools
try: import fcntl # windows misses that try: import fcntl # windows misses that
except ImportError: fcntl = None #type:ignore[assignment] except ImportError: fcntl = None #type:ignore[assignment]
from tinygrad.helpers import PROFILE, getenv, to_mv, ProfileRangeEvent, select_first_inited from tinygrad.helpers import PROFILE, getenv, to_mv, ProfileRangeEvent, select_first_inited, unwrap
from tinygrad.device import BufferSpec, Compiled, LRUAllocator, ProfileDeviceEvent, ProfileProgramEvent, CompilerPairT from tinygrad.device import BufferSpec, Compiled, LRUAllocator, ProfileDeviceEvent, ProfileProgramEvent, CompilerPairT
from tinygrad.uop.ops import sym_infer, sint, UOp from tinygrad.uop.ops import sym_infer, sint, UOp
from tinygrad.runtime.autogen import libc from tinygrad.runtime.autogen import libc
@@ -276,8 +276,7 @@ def hcq_profile(dev:HCQCompiled, enabled, desc, queue_type:Callable[[], HWQueue]
elif enabled and queue_type is not None: elif enabled and queue_type is not None:
queue_type().wait(dev.timeline_signal, dev.timeline_value - 1).timestamp(en).signal(dev.timeline_signal, dev.next_timeline()).submit(dev) queue_type().wait(dev.timeline_signal, dev.timeline_value - 1).timestamp(en).signal(dev.timeline_signal, dev.next_timeline()).submit(dev)
if enabled and PROFILE: if enabled and PROFILE: dev.sig_prof_records.append((unwrap(st), unwrap(en), desc, (queue_type or type(queue)) is dev.hw_copy_queue_t))
dev.sig_prof_records.append((cast(HCQSignal, st), cast(HCQSignal, en), desc, (queue_type or type(queue)) is dev.hw_copy_queue_t))
class HCQArgsState(Generic[ProgramType]): class HCQArgsState(Generic[ProgramType]):
def __init__(self, buf:HCQBuffer, prg:ProgramType, bufs:tuple[HCQBuffer, ...], vals:tuple[sint, ...]=()): def __init__(self, buf:HCQBuffer, prg:ProgramType, bufs:tuple[HCQBuffer, ...], vals:tuple[sint, ...]=()):
+5 -1
View File
@@ -39,6 +39,7 @@ class BufferizeOpts:
# on AddrSpace.LOCAL, device is the id # on AddrSpace.LOCAL, device is the id
device: str|tuple[str, ...]|int|None device: str|tuple[str, ...]|int|None
addrspace: AddrSpace = AddrSpace.GLOBAL addrspace: AddrSpace = AddrSpace.GLOBAL
removable: bool = True
@dataclass @dataclass
class IndexingContext: class IndexingContext:
@@ -68,8 +69,11 @@ def create_bufferize_and_index_based_on_ranges(ctx:IndexingContext, x:UOp):
new_src = s.end(*[r for r in closed_ranges if r.op is Ops.RANGE]) new_src = s.end(*[r for r in closed_ranges if r.op is Ops.RANGE])
del ctx.realize_map[s] del ctx.realize_map[s]
else: else:
# the Bufferize before a COPY is not removable. there should be a better way to do this
removable = x.op is not Ops.COPY and s.op not in ALWAYS_CONTIGUOUS
# None in the device assigns it a number later # None in the device assigns it a number later
opts = BufferizeOpts(device=s.device) if len(ctx.range_map[s][1]) == len(realized_ranges) else BufferizeOpts(None, AddrSpace.LOCAL) opts = BufferizeOpts(device=s.device, removable=removable) if len(ctx.range_map[s][1]) == len(realized_ranges) else \
BufferizeOpts(None, AddrSpace.LOCAL, removable=removable)
new_src = UOp(Ops.BUFFERIZE, s.dtype, src=(new_src,)+closed_ranges, arg=opts, tag=s.tag if opts.addrspace == AddrSpace.GLOBAL else None) new_src = UOp(Ops.BUFFERIZE, s.dtype, src=(new_src,)+closed_ranges, arg=opts, tag=s.tag if opts.addrspace == AddrSpace.GLOBAL else None)
if x in ctx.range_map: new_src = new_src.index(*[r for i,r in enumerate(ctx.range_map[x][0]) if i in realized_ranges]) if x in ctx.range_map: new_src = new_src.index(*[r for i,r in enumerate(ctx.range_map[x][0]) if i in realized_ranges])
new_srcs.append(new_src) new_srcs.append(new_src)
+5 -13
View File
@@ -1,7 +1,7 @@
from dataclasses import dataclass, field from dataclasses import dataclass, field
import itertools import itertools
from tinygrad.dtype import dtypes, PtrDType, ImageDType, AddrSpace from tinygrad.dtype import dtypes, PtrDType, ImageDType, AddrSpace
from tinygrad.uop.ops import PatternMatcher, UPat, Ops, UOp, resolve, GroupOp, _substitute, ssimplify, KernelInfo from tinygrad.uop.ops import PatternMatcher, UPat, Ops, UOp, resolve, GroupOp, _substitute, KernelInfo
from tinygrad.uop.ops import track_rewrites, graph_rewrite, identity_element, sint, AxisType, BottomUpGate, Kernel, _remove_all_tags, range_str from tinygrad.uop.ops import track_rewrites, graph_rewrite, identity_element, sint, AxisType, BottomUpGate, Kernel, _remove_all_tags, range_str
from tinygrad.uop.symbolic import symbolic from tinygrad.uop.symbolic import symbolic
from tinygrad.helpers import argsort, prod, all_same, pluralize, getenv, flatten, dedup, all_int, DEBUG, SPLIT_REDUCEOP, DEBUG_RANGEIFY from tinygrad.helpers import argsort, prod, all_same, pluralize, getenv, flatten, dedup, all_int, DEBUG, SPLIT_REDUCEOP, DEBUG_RANGEIFY
@@ -152,7 +152,7 @@ def remove_bufferize(src:UOp, buf:UOp, idx:UOp):
assert all(x.op in {Ops.RANGE, Ops.CONST} for x in buf.src[1:]) assert all(x.op in {Ops.RANGE, Ops.CONST} for x in buf.src[1:])
# if it's user contiguous, we never remove it # if it's user contiguous, we never remove it
if src.op in ALWAYS_RUN_OPS: return None if src.op in ALWAYS_RUN_OPS or not buf.arg.removable: return None
# we don't want to bufferize threefry, also causes problems because not all platforms support long # we don't want to bufferize threefry, also causes problems because not all platforms support long
if src.op is not Ops.THREEFRY: if src.op is not Ops.THREEFRY:
@@ -177,7 +177,7 @@ def remove_bufferize(src:UOp, buf:UOp, idx:UOp):
accessed_buffers = dedup(accessed_buffers) accessed_buffers = dedup(accessed_buffers)
# if this is generated from multiple buffers, don't remove this buffer # if this is generated from multiple buffers, don't remove this buffer
if len(accessed_buffers) > 2 and not (PCONTIG > 2): return None if len(accessed_buffers) > 3 and not (PCONTIG > 2): return None
# if any reduces access a buffer, don't remove this buffer # if any reduces access a buffer, don't remove this buffer
buffer_in_reduce = False buffer_in_reduce = False
@@ -238,13 +238,7 @@ pm_const_buffer_folding = pm_mops+PatternMatcher([
lambda s: UOp.const(c.dtype, c.arg) if (c:=s.base).op is Ops.CONST else None), lambda s: UOp.const(c.dtype, c.arg) if (c:=s.base).op is Ops.CONST else None),
]) ])
def pre_bufferize(b:UOp, x:UOp, copy:UOp):
nb = b.replace(src=(b.src[0].contiguous(),)+b.src[1:])
return copy.replace(src=(x.replace(src=(nb,)+x.src[1:]), copy.src[1]))
pm_remove_bufferize = PatternMatcher([ pm_remove_bufferize = PatternMatcher([
# hack so remove_bufferize doesnt remove the buffer before a copy
(UPat(Ops.COPY, src=(UPat(GroupOp.All-{Ops.CONTIGUOUS, Ops.COPY}).f(Ops.BUFFERIZE, allow_any_len=True, name="b")
.f(Ops.INDEX, allow_any_len=True, name="x"), UPat()), name="copy"), pre_bufferize),
# remove reindexing with cost function # remove reindexing with cost function
(UPat.var("src").f(Ops.BUFFERIZE, allow_any_len=True, name="buf").f(Ops.INDEX, allow_any_len=True, name="idx"), remove_bufferize), (UPat.var("src").f(Ops.BUFFERIZE, allow_any_len=True, name="buf").f(Ops.INDEX, allow_any_len=True, name="idx"), remove_bufferize),
]) ])
@@ -356,7 +350,7 @@ def flatten_bufferize(x:UOp):
rngs = x.src[1:] rngs = x.src[1:]
ret = ret.forced_reshape(x.shape) ret = ret.forced_reshape(x.shape)
if any(r.op is Ops.RANGE and r.src[0].op is not Ops.CONST for r in rngs): if any(r.op is Ops.RANGE and r.src[0].op is not Ops.CONST for r in rngs):
sym_shape = tuple([ssimplify(r.src[0]) if r.op is not Ops.CONST else 1 for r in rngs]) sym_shape = tuple([r.src[0] if r.op is not Ops.CONST else 1 for r in rngs])
ret = ret.shrink(tuple([(0,x) for x in sym_shape])) ret = ret.shrink(tuple([(0,x) for x in sym_shape]))
return ret.rtag(x.tag) return ret.rtag(x.tag)
pm_flatten_bufferize = PatternMatcher([(UPat(Ops.BUFFERIZE, name="x"), flatten_bufferize)]) pm_flatten_bufferize = PatternMatcher([(UPat(Ops.BUFFERIZE, name="x"), flatten_bufferize)])
@@ -556,9 +550,7 @@ def get_rangeify_map(sink:UOp) -> dict[UOp, UOp]:
# convert movement ops to ranges # convert movement ops to ranges
tsink, rctx = run_rangeify(tsink, DEBUG_RANGEIFY) tsink, rctx = run_rangeify(tsink, DEBUG_RANGEIFY)
tsink = graph_rewrite(tsink, symbolic+pm_reduce_simplify+pm_const_buffer_folding, name="symbolic+reduce_collapse") tsink = graph_rewrite(tsink, symbolic+pm_reduce_simplify+pm_const_buffer_folding+pm_remove_bufferize, name="symbolic+reduce_collapse+debuf")
tsink = graph_rewrite(tsink, pm_remove_bufferize, bottom_up=True, name="remove bufferize with cost function")
tsink = graph_rewrite(tsink, symbolic+pm_reduce_simplify+pm_const_buffer_folding, name="symbolic+reduce_collapse pt 2")
tsink = graph_rewrite(tsink, pm_limit_bufs, ctx=rctx, name="limit buffers") tsink = graph_rewrite(tsink, pm_limit_bufs, ctx=rctx, name="limit buffers")
# rebuild the sink with all the BUFFERIZEs with tags, this is what's ending up in the tensor graph # rebuild the sink with all the BUFFERIZEs with tags, this is what's ending up in the tensor graph
+32 -27
View File
@@ -6,7 +6,7 @@ from typing import Callable, ClassVar, Sequence, cast, get_args, Literal, Suppor
from tinygrad.dtype import DType, DTypeLike, dtypes, ImageDType, ConstType, least_upper_float, least_upper_dtype, sum_acc_dtype, to_dtype, truncate from tinygrad.dtype import DType, DTypeLike, dtypes, ImageDType, ConstType, least_upper_float, least_upper_dtype, sum_acc_dtype, to_dtype, truncate
from tinygrad.dtype import _from_np_dtype, _to_np_dtype from tinygrad.dtype import _from_np_dtype, _to_np_dtype
from tinygrad.helpers import argfix, make_tuple, flatten, prod, all_int, round_up, merge_dicts, argsort, getenv, all_same, fully_flatten from tinygrad.helpers import argfix, make_tuple, flatten, prod, all_int, round_up, merge_dicts, argsort, getenv, all_same, fully_flatten
from tinygrad.helpers import IMAGE, WINO, Metadata, TRACEMETA, ceildiv, fetch, polyN, DEBUG, is_numpy_ndarray, SPEC from tinygrad.helpers import IMAGE, WINO, Metadata, TRACEMETA, ceildiv, fetch, polyN, DEBUG, is_numpy_ndarray, SPEC, TracingKey, cpu_profile
from tinygrad.helpers import suppress_finalizing from tinygrad.helpers import suppress_finalizing
from tinygrad.gradient import compute_gradient from tinygrad.gradient import compute_gradient
from tinygrad.mixin import OpMixin from tinygrad.mixin import OpMixin
@@ -26,18 +26,21 @@ def canonicalize_device(device:str|None) -> str: return Device.canonicalize(devi
# *** all in scope Tensors are here. this gets relevant UOps *** # *** all in scope Tensors are here. this gets relevant UOps ***
all_tensors: dict[weakref.ref[Tensor], None] = {} all_tensors: dict[weakref.ref[Tensor], None] = {}
def _apply_map_to_tensors(applied_map:dict[UOp, UOp], name:str|None=None) -> None: def _apply_map_to_tensors(applied_map:dict[UOp, UOp], name:str) -> None:
scope_tensors = [t for tref in tuple(all_tensors) if (t:=tref()) is not None and with cpu_profile(TracingKey(name), "TINY"):
(t.uop in applied_map or len(applied_map.keys() & t.uop.backward_slice.keys()))] # get tensors in scope
in_scope: dict[UOp, bool] = {}
def visitor(node: UOp) -> bool: return True if node in applied_map else any(in_scope.get(s, False) for s in node.src)
scope_tensors = [t for tref in list(all_tensors) if (t:=tref()) is not None and t.uop.topovisit(visitor, in_scope)]
# get all Tensors and apply the map # get all Tensors and apply the map
sink = UOp.sink(*[t.uop for t in scope_tensors]) sink = UOp.sink(*[t.uop for t in scope_tensors])
new_sink = sink.substitute(applied_map, name=name) new_sink = sink.substitute(applied_map, name=f"substitute {name}")
# set the relevant uop to the realized UOps # set the relevant uop to the realized UOps
for t,s,ns in zip(scope_tensors, sink.src, new_sink.src): for t,s,ns in zip(scope_tensors, sink.src, new_sink.src):
if s is ns: continue if s is ns: continue
t.uop = ns t.uop = ns
# **** Tensor helper functions **** # **** Tensor helper functions ****
@@ -127,7 +130,7 @@ class Tensor(OpMixin):
# create a UOp from the different types of inputs # create a UOp from the different types of inputs
if isinstance(data, UOp): if isinstance(data, UOp):
assert _dtype is None or _dtype==data.dtype, "dtype doesn't match, and casting isn't supported" assert _dtype is None or _dtype==data.dtype, f"dtype doesn't match ({_dtype} vs {data.dtype}), and casting isn't supported"
# if data is dtype.index that means that this is a symbolic int and we need to lower it to something we can make a Tensor out of # if data is dtype.index that means that this is a symbolic int and we need to lower it to something we can make a Tensor out of
if data.dtype==dtypes.index: data = _index_to_concrete_int(data) if data.dtype==dtypes.index: data = _index_to_concrete_int(data)
if data.op is Ops.BIND: # type: ignore # mypy type narrowing is bugged here if data.op is Ops.BIND: # type: ignore # mypy type narrowing is bugged here
@@ -229,7 +232,7 @@ class Tensor(OpMixin):
if SPEC: type_verify(big_sink, tensor_spec) if SPEC: type_verify(big_sink, tensor_spec)
if any(isinstance(x._device, tuple) for x in big_sink.toposort()): if any(isinstance(x._device, tuple) for x in big_sink.toposort()):
_apply_map_to_tensors(get_multi_map(big_sink), "Apply Multi Map") _apply_map_to_tensors(get_multi_map(big_sink), name="Apply Multi Map")
big_sink = UOp.sink(*flatten([x.uop.src if x.uop.op is Ops.MULTI else [x.uop] for x in (self,)+lst])) big_sink = UOp.sink(*flatten([x.uop.src if x.uop.op is Ops.MULTI else [x.uop] for x in (self,)+lst]))
becomes_map = get_rangeify_map(big_sink) becomes_map = get_rangeify_map(big_sink)
@@ -259,8 +262,8 @@ class Tensor(OpMixin):
_apply_map_to_tensors(remove_assign_map, name="Remove After") _apply_map_to_tensors(remove_assign_map, name="Remove After")
# create the schedule # create the schedule
schedule, var_vals = create_schedule_with_vars(sink) with cpu_profile(TracingKey("toposort schedule")): schedule, var_vals = create_schedule_with_vars(sink)
schedule = memory_planner(schedule) with cpu_profile(TracingKey("memory planner")): schedule = memory_planner(schedule)
if (DEBUG >= 1 and len(schedule) > 1) or DEBUG >= 3: print(f"scheduled {len(schedule)} kernels in {(time.perf_counter()-st)*1000:.2f} ms") if (DEBUG >= 1 and len(schedule) > 1) or DEBUG >= 3: print(f"scheduled {len(schedule)} kernels in {(time.perf_counter()-st)*1000:.2f} ms")
return schedule, var_vals return schedule, var_vals
@@ -299,6 +302,7 @@ class Tensor(OpMixin):
assert self.shape == x.shape, f"assign shape mismatch {self.shape} != {x.shape}" assert self.shape == x.shape, f"assign shape mismatch {self.shape} != {x.shape}"
assert self.device == x.device, f"assign device mismatch {self.device} != {x.device}" assert self.device == x.device, f"assign device mismatch {self.device} != {x.device}"
assert self.dtype == x.dtype, f"assign dtype mismatch {self.dtype} != {x.dtype}" assert self.dtype == x.dtype, f"assign dtype mismatch {self.dtype} != {x.dtype}"
assert not isinstance(self.device, tuple) or self.uop.axis == x.uop.axis, f"multi assign axis mismatch {self.uop.axis} != {x.uop.axis}"
return self.replace(self._apply_uop(UOp.assign, x)) return self.replace(self._apply_uop(UOp.assign, x))
def detach(self) -> Tensor: def detach(self) -> Tensor:
@@ -421,7 +425,7 @@ class Tensor(OpMixin):
return self.replace(self.shard(devices, axis)) return self.replace(self.shard(devices, axis))
CHUNK_SIZE = 2**20 CHUNK_SIZE = 2**20
def load(self, size:int) -> Tensor: def fs_load(self, size:int) -> Tensor:
""" """
Load a tensor from storage. Load a tensor from storage.
@@ -449,7 +453,7 @@ class Tensor(OpMixin):
return data[:size] return data[:size]
def store(self) -> Tensor: def fs_store(self) -> Tensor:
""" """
Store a tensor to storage. Store a tensor to storage.
""" """
@@ -738,6 +742,14 @@ class Tensor(OpMixin):
t = (Tensor.arange(n, device=device).unsqueeze(-1) == Tensor.arange(m, device=device)) t = (Tensor.arange(n, device=device).unsqueeze(-1) == Tensor.arange(m, device=device))
return t.cast(dtype or dtypes.default_float).requires_grad_(requires_grad) return t.cast(dtype or dtypes.default_float).requires_grad_(requires_grad)
def _multi_like(self, fxn, *args, **kwargs) -> Tensor:
dtype = kwargs.pop("dtype", self.dtype)
if kwargs.get("device") is not None: raise RuntimeError("cannot specify `device` on `*_like` of a multi device tensor")
if self.uop.axis is None: return fxn(self.shape, *args, dtype=dtype, **kwargs).shard(self.device)
sharded_shape = tuple(s//len(self.device) if a==self.uop.axis else s for a,s in enumerate(self.shape))
stacked = UOp(Ops.MSTACK, dtype=dtype, src=tuple([fxn(sharded_shape, *args, device=d, dtype=dtype, **kwargs).uop for d in self.device]))
return Tensor(UOp.multi(stacked, axis=self.uop.axis), device=self.device, dtype=dtype)
def full_like(self, fill_value:ConstType, **kwargs) -> Tensor: def full_like(self, fill_value:ConstType, **kwargs) -> Tensor:
""" """
Creates a tensor with the same shape as `self`, filled with the given value. Creates a tensor with the same shape as `self`, filled with the given value.
@@ -751,6 +763,7 @@ class Tensor(OpMixin):
print(Tensor.full_like(t, 42).numpy()) print(Tensor.full_like(t, 42).numpy())
``` ```
""" """
if isinstance(self.device, tuple): return self._multi_like(Tensor.full, fill_value, **kwargs)
return Tensor.full(self.shape, fill_value, dtype=kwargs.pop("dtype", self.dtype), device=kwargs.pop("device", self.device), **kwargs) return Tensor.full(self.shape, fill_value, dtype=kwargs.pop("dtype", self.dtype), device=kwargs.pop("device", self.device), **kwargs)
def zeros_like(self, **kwargs) -> Tensor: def zeros_like(self, **kwargs) -> Tensor:
@@ -793,16 +806,8 @@ class Tensor(OpMixin):
print(Tensor.rand_like(t).numpy()) print(Tensor.rand_like(t).numpy())
``` ```
""" """
dtype = kwargs.pop("dtype", self.dtype) if isinstance(self.device, tuple): return self._multi_like(Tensor.rand, **kwargs)
if isinstance(self.device, tuple): return Tensor.rand(*self.shape, device=kwargs.pop("device", self.device), dtype=kwargs.pop("dtype", self.dtype), **kwargs)
if kwargs.get("device") is not None: raise RuntimeError("cannot specify `device` on `rand_like` of a multi device tensor")
if self.uop.axis is None: return Tensor.rand(*self.shape, dtype=dtype, **kwargs).shard(self.device)
contiguous = kwargs.pop("contiguous", True)
sharded_shape = tuple(s//len(self.device) if a==self.uop.axis else s for a,s in enumerate(self.shape))
rands = UOp(Ops.MSTACK, dtype=dtype,
src=tuple([Tensor.rand(sharded_shape, device=d, dtype=dtype, contiguous=contiguous, **kwargs).uop for d in self.device]))
return Tensor(UOp.multi(rands, axis=self.uop.axis), device=self.device, dtype=dtype, **kwargs)
return Tensor.rand(*self.shape, device=kwargs.pop("device", self.device), dtype=dtype, **kwargs)
# ***** rng hlops ***** # ***** rng hlops *****
+12 -18
View File
@@ -13,25 +13,23 @@ class FastEnum(IntEnum):
class Ops(FastEnum): class Ops(FastEnum):
# ** 1 -- defines/special ** # ** 1 -- defines/special **
# TODO: unify these ops into the levels of the memory hierarchy # define GLOBAL/VAR are ptrs to outside the Kernel
DEFINE_GLOBAL = auto(); DEFINE_LOCAL = auto(); DEFINE_REG = auto() DEFINE_GLOBAL = auto(); DEFINE_VAR = auto(); BIND = auto()
# this is for symbolic shapes
DEFINE_VAR = auto(); BIND = auto()
# this is a RANGE for GPU dimensions, similar to symbolic shapes but not exactly # this is a RANGE for GPU dimensions, similar to symbolic shapes but not exactly
SPECIAL = auto() SPECIAL = auto()
# define LOCAL/REG allocate things
DEFINE_LOCAL = auto(); DEFINE_REG = auto()
# ** 2 -- non op uops ** # ** 2 -- non op uops **
# uops that aren't rendered # uops that aren't rendered
NOOP = auto(); SINK = auto(); PRECAST = auto() NOOP = auto(); REWRITE_ERROR = auto()
# AFTER passes src[0] through and promises in the toposort that any consumers of the AFTER run after src[1:] # AFTER passes src[0] through and promises in the toposort that any consumers of the AFTER run after src[1:]
AFTER = auto()
# GROUP is a NOOP that just merges things together # GROUP is a NOOP that just merges things together
GROUP = auto() SINK = auto(); AFTER = auto(); GROUP = auto()
# vector creation / item selection # vector creation / item selection
GEP = auto(); VECTORIZE = auto() GEP = auto(); VECTORIZE = auto()
@@ -76,25 +74,21 @@ class Ops(FastEnum):
# ** 6 -- ops that don't exist in programs ** # ** 6 -- ops that don't exist in programs **
# tensor graph ops # tensor graph ops
UNIQUE = auto(); DEVICE = auto(); KERNEL = auto() UNIQUE = auto(); DEVICE = auto(); KERNEL = auto(); ASSIGN = auto()
ASSIGN = auto()
# buffer ops
BUFFERIZE = auto(); COPY = auto(); BUFFER = auto(); BUFFER_VIEW = auto(); MSELECT = auto(); MSTACK = auto()
# ops that adjust the behavior of the scheduler # ops that adjust the behavior of the scheduler
CONTIGUOUS = auto(); CONTIGUOUS_BACKWARD = auto(); DETACH = auto() CONTIGUOUS = auto(); CONTIGUOUS_BACKWARD = auto(); DETACH = auto()
# movement ops! these only exist in the tensor graph # buffer ops
BUFFERIZE = auto(); COPY = auto(); BUFFER = auto(); BUFFER_VIEW = auto(); MSELECT = auto(); MSTACK = auto()
# the core 6 movement ops! these only exist in the tensor graph
RESHAPE = auto(); PERMUTE = auto(); EXPAND = auto(); PAD = auto(); SHRINK = auto(); FLIP = auto() RESHAPE = auto(); PERMUTE = auto(); EXPAND = auto(); PAD = auto(); SHRINK = auto(); FLIP = auto()
MULTI = auto() # MULTI is really a movement op MULTI = auto() # MULTI is really a movement op
# reduce # reduce
REDUCE_AXIS = auto(); REDUCE = auto(); ALLREDUCE = auto() REDUCE_AXIS = auto(); REDUCE = auto(); ALLREDUCE = auto()
# errors/placeholders
REWRITE_ERROR = auto(); SENTINEL = auto()
# expander ops # expander ops
UNROLL = auto(); CONTRACT = auto(); CAT = auto(); PTRCAT = auto() UNROLL = auto(); CONTRACT = auto(); CAT = auto(); PTRCAT = auto()
+112
View File
@@ -0,0 +1,112 @@
import functools
from tinygrad.uop.ops import PatternMatcher, UPat, Ops, UOp
from tinygrad.dtype import dtypes
from tinygrad.helpers import cdiv, cmod, CORRECT_DIVMOD_FOLDING, unwrap
# NOTE: this cache is only on index UOps and matches the cache in the old ShapeTracker in spirit
@functools.cache
def fold_divmod_general(d: UOp, correct_divmod_folding: bool) -> UOp|None:
x, y = d.src
# cancel_divmod: simple cancel div/mod case when the range of the numerator lies within a single denominator interval
x_min, x_max, y_min, y_max = x.vmin, x.vmax, y.vmin, y.vmax
assert isinstance(x_min, int) and isinstance(x_max, int) and isinstance(y_min, int) and isinstance(y_max, int)
if y_min==y_max==0: raise ZeroDivisionError(f"{'Division' if d.op is Ops.IDIV else 'Mod'} by zero trying to rewrite {x.alu(d.op, y)}")
if y_min*y_max > 0 and (q:=cdiv(x_min,y_min)) == cdiv(x_min,y_max) == cdiv(x_max,y_min) == cdiv(x_max,y_max):
return x - q*y if d.op is Ops.MOD else d.const_like(q)
# split uops for the rest of the processing
x_peeled, const = x.pop_const()
uops_no_const = list(x_peeled.split_uop(Ops.ADD))
# ** Constant Denominator Rules **
# these rules strictly require y to be a scalar constant > 0
if y.op is Ops.CONST and (c := y.arg) > 0:
# remove_nested_mod: remove nested mod in case the inner mod is a multiple of the outer mod, example: (a%4 + b)%2 -> (a+b)%2
if d.op is Ops.MOD and x.vmin >= 0:
new_xs, changed = [], False
for u in uops_no_const:
if u.op is Ops.MOD and u.src[1].divides(c) is not None:
u = u.src[0]
changed = True
new_xs.append(u)
if changed and (new_x:=(UOp.sum(*new_xs) + const)).vmin >= 0: return new_x % y
# Shared decomposition for folding rules
decomp = [(u.divides(f:=u.const_factor()),f) for u in uops_no_const]
terms, factors = zip(*decomp)
# fold_binary_numerator: fold if expression has one non-constant term that takes on two values
if len(terms)==1 and (v:=terms[0]).vmax-v.vmin == 1:
y1 = cmod(factors[0]*v.vmin+const, c) if d.op is Ops.MOD else cdiv(factors[0]*v.vmin+const, c)
y2 = cmod(factors[0]*v.vmax+const, c) if d.op is Ops.MOD else cdiv(factors[0]*v.vmax+const, c)
return (y2-y1)*(v-v.vmin) + y1
# fold_divmod_congruence: fold if a is congruent to an expression whose range is between 0 and c
if not (x.vmin<0 and correct_divmod_folding):
rems = [min((r:=f%c), r-c, key=abs) for f in factors]
if (rem:=sum(r*v for r,v in zip(rems,terms))+const%c).vmin//c==rem.vmax//c:
if d.op is Ops.MOD: return rem - rem.vmin//c*c
return sum((f-r)//c * v for f,r,v in zip(factors,rems,terms)) + (const-const%c+rem.vmin//c*c)//c
# gcd_with_remainder: factor out common gcd from numerator
# Note: this rule uses uops_no_const to exclude the additive constant from the GCD calculation
if x.vmin >= 0:
gcd = UOp.gcd(*uops_no_const, y).simplify()
if gcd.op is Ops.CONST and gcd.arg > 1:
new_x = unwrap(x_peeled.divide_exact(gcd)).simplify() + (const%c)//gcd.arg
if new_x.vmin >= 0:
ret = new_x.alu(d.op, x.ufix(c//gcd.arg))
return ret*gcd + const%gcd.arg if d.op is Ops.MOD else ret+const//c
# nest_div_by_smallest_factor: try and nest the div and see if it allows the numerator to be simplified
if d.op is Ops.IDIV and x.vmin >= 0:
div = min([c] + [abs(f) for u, f in zip(uops_no_const, factors) if u.op not in (Ops.CONST, Ops.VCONST) and abs(f) > 1 and (c%f)==0])
# NOTE: this is recursive!
if div < c and (newxs := fold_divmod_general(x//div, correct_divmod_folding)) is not None and newxs.vmin >= 0:
return newxs // (c // div)
# ** Variable Denominator / Fallback Rules **
# These rules apply to variables OR constants that failed the checks above.
# Reconstruct all uops including const for these checks.
all_uops = uops_no_const + ([x.const_like(const)] if const != 0 else [])
# divide_by_gcd: x//y -> (x//gcd)//(y//gcd)
gcd = UOp.gcd(*all_uops, y).simplify()
if not (gcd.op is Ops.CONST and gcd.arg==1):
ret = unwrap(x.divide_exact(gcd)).alu(d.op, unwrap(y.divide_exact(gcd)))
return ret*gcd if d.op is Ops.MOD else ret
# factor_remainder: (d*x+y)//d -> x+y//d
if y.vmin<0 or x.vmin<0: return None
quo, rem = [], []
for u in all_uops:
if (q:=u.divide_exact(y)) is not None: quo.append(q)
elif d.op is Ops.MOD and y.op is Ops.CONST and (c:=u.const_factor())%y.arg!=c:
rem.append(u.divides(c)*(c%y.arg))
quo.append(u.const_like(0))
else: rem.append(u)
if not quo: return None
new_x = sum(rem)+x.const_like(0)
if new_x.vmin<0: return None
return new_x%y if d.op is Ops.MOD else new_x//y+sum(quo)
div_and_mod_symbolic = PatternMatcher([
# ** 1. Fast Inline Rules **
((UPat.var("x")//UPat.cvar("c") + UPat.cvar("a"))//UPat.cvar("d"), lambda x,c,a,d: (x+a*c)//(c*d)
if c.vmin>0 and d.vmin>0 and ((x.vmin>=0 and a.vmin>=0) or (x.vmax<=0 and a.vmax<=0)) else None), # (x//c+a)//d -> (x+a*c)//(c*d)
(UPat.var("x", dtypes.index) // UPat.var("d"), lambda x,d: -(x//(-d)) if d.vmax < 0 else None),
(UPat.var("x", dtypes.index) // UPat.var("d"), lambda x,d: -((-x)//d) if x.vmax <= 0 else None),
((UPat.var("x", dtypes.index)+UPat.cvar("c", vec=False)).named("n")//UPat.cvar("d", vec=False),
lambda x,c,n,d: ((x+c.arg%d.arg)//d + c.arg//d.arg) if c.arg%d.arg!=c.arg and x.vmin>=0 and n.vmin>=0 and d.arg>0 else None),
((UPat.var("x", dtypes.index)+UPat.cvar("c", vec=False)).named("n")//UPat.cvar("d", vec=False),
lambda x,c,n,d: (-(-(c.arg%d.arg + x - (d.arg-1))//d) + c.arg//d.arg) if x.vmax<=0 and n.vmin>=0 and d.arg>0 else None),
# ** 2. Slow Rules **
(UPat((Ops.IDIV, Ops.MOD), dtypes.index, name="d"), lambda d: fold_divmod_general(d, bool(CORRECT_DIVMOD_FOLDING))),
# NOTE: these have to go at the bottom or TestSymbolicOps.test_var loops
(UPat.var("x", dtypes.index) % UPat.var("d"), lambda x,d: -((-x)%d) if x.vmax <= 0 else None),
(UPat.var("x", dtypes.index) % UPat.var("d"), lambda x,d: (x%(-d)) if d.vmax < 0 else None),
])
+50 -30
View File
@@ -1,5 +1,5 @@
from __future__ import annotations from __future__ import annotations
from typing import Any, Callable, cast, TYPE_CHECKING, Type, Sequence, Iterable from typing import Any, Callable, cast, TYPE_CHECKING, Type, Sequence, Iterable, Final
import sys, time, functools, itertools, math, operator, hashlib, os, types, pickle, pathlib, inspect, weakref, collections import sys, time, functools, itertools, math, operator, hashlib, os, types, pickle, pathlib, inspect, weakref, collections
from dataclasses import dataclass from dataclasses import dataclass
from enum import Enum, auto from enum import Enum, auto
@@ -98,7 +98,6 @@ buffers:weakref.WeakKeyDictionary[UOp, Buffer|MultiBuffer] = weakref.WeakKeyDict
all_metadata:weakref.WeakKeyDictionary[UOp, tuple[Metadata, ...]] = weakref.WeakKeyDictionary() # TODO: should this be here? all_metadata:weakref.WeakKeyDictionary[UOp, tuple[Metadata, ...]] = weakref.WeakKeyDictionary() # TODO: should this be here?
# recursive_property replaces functools.cached_property in recursive UOp functions to prevent RecursionError # recursive_property replaces functools.cached_property in recursive UOp functions to prevent RecursionError
_NOT_FOUND = object()
class recursive_property(property): class recursive_property(property):
def __init__(self, fxn): def __init__(self, fxn):
self.fxn = fxn self.fxn = fxn
@@ -106,10 +105,16 @@ class recursive_property(property):
self.__doc__ = fxn.__doc__ self.__doc__ = fxn.__doc__
def __get__(self, x:UOp|None, owner=None): def __get__(self, x:UOp|None, owner=None):
if x is None: return self if x is None: return self
if (val:=x.__dict__.get(self.nm, _NOT_FOUND)) is _NOT_FOUND: # this is very similar to toposort/topovisit
for s in x.toposort(lambda z: not hasattr(z, self.nm)): stack: list[tuple[UOp, bool]] = [(x, False)]
s.__dict__[self.nm] = val = self.fxn(s) while stack:
return val node, visited = stack.pop()
if self.nm in node.__dict__: continue
if not visited:
stack.append((node, True))
for s in reversed(node.src): stack.append((s, False))
else: node.__dict__[self.nm] = self.fxn(node)
return x.__dict__[self.nm]
# we import this late so we can use resolve/smax in mixins # we import this late so we can use resolve/smax in mixins
from tinygrad.mixin import OpMixin from tinygrad.mixin import OpMixin
@@ -157,17 +162,29 @@ class UOp(OpMixin, metaclass=UOpMetaClass):
def op_in_backward_slice_with_self(self, *ops:Ops): return any(x.op in ops for x in self.backward_slice_with_self) def op_in_backward_slice_with_self(self, *ops:Ops): return any(x.op in ops for x in self.backward_slice_with_self)
def toposort(self, gate:Callable|None=None) -> dict[UOp, None]: def toposort(self, gate:Callable|None=None) -> dict[UOp, None]:
ret: dict[UOp, None] = {} cache: dict[UOp, None] = {}
stack: list[tuple[UOp, bool]] = [(self, False)] # each stack entry is (node, visited_flag) stack: list[tuple[UOp, bool]] = [(self, False)] # each stack entry is (node, visited_flag)
while stack: while stack:
node, visited = stack.pop() node, visited = stack.pop()
if node in ret: continue if node in cache: continue
if not visited: if not visited:
if gate is None or gate(node): if gate is None or gate(node):
stack.append((node, True)) # push node back on stack to process after its srcs stack.append((node, True)) # push node back on stack to process after its srcs
for s in reversed(node.src): stack.append((s, False)) # push srcs on the stack for s in reversed(node.src): stack.append((s, False)) # push srcs on the stack
else: ret[node] = None # second time i'm seeing this node, add it to returned toposort else: cache[node] = None # second time i'm seeing this node, add it to returned toposort
return ret return cache
def topovisit(self, visitor:Callable[[UOp], T], cache:dict[UOp, T]) -> T:
# NOTE: this shares a lot of code with toposort
stack: list[tuple[UOp, bool]] = [(self, False)]
while stack:
node, visited = stack.pop()
if node in cache: continue
if not visited:
stack.append((node, True))
for s in reversed(node.src): stack.append((s, False))
else: cache[node] = visitor(node)
return cache[self]
# returns map of UOps to their consumers in the graph rooted by self # returns map of UOps to their consumers in the graph rooted by self
def get_consumer_map(self) -> dict[UOp, dict[UOp, None]]: return consumer_map_from_toposort(self.toposort()) def get_consumer_map(self) -> dict[UOp, dict[UOp, None]]: return consumer_map_from_toposort(self.toposort())
@@ -200,7 +217,7 @@ class UOp(OpMixin, metaclass=UOpMetaClass):
match self.op: match self.op:
# late ops don't have shape # late ops don't have shape
case Ops.UNIQUE | Ops.DEVICE | Ops.RANGE | Ops.LOAD | Ops.IF | Ops.BARRIER | Ops.CUSTOM | Ops.CUSTOMI | \ case Ops.UNIQUE | Ops.DEVICE | Ops.RANGE | Ops.LOAD | Ops.IF | Ops.BARRIER | Ops.CUSTOM | Ops.CUSTOMI | \
Ops.VECTORIZE | Ops.VCONST | Ops.GEP | Ops.SPECIAL | Ops.UNROLL | Ops.PRECAST | Ops.CONTRACT: Ops.VECTORIZE | Ops.VCONST | Ops.GEP | Ops.SPECIAL | Ops.UNROLL | Ops.CONTRACT:
return None return None
case Ops.INDEX: case Ops.INDEX:
@@ -275,7 +292,7 @@ class UOp(OpMixin, metaclass=UOpMetaClass):
return tuple(1 if i in axis_arg else s for i,s in enumerate(ps)) return tuple(1 if i in axis_arg else s for i,s in enumerate(ps))
# elementwise ops keep the shape the same. all inputs with shape must match # elementwise ops keep the shape the same. all inputs with shape must match
if self.op in (GroupOp.Elementwise-{Ops.BITCAST}).union({Ops.COPY, Ops.ASSIGN, Ops.NOOP, Ops.GROUP, Ops.SINK, Ops.ALLREDUCE, Ops.STORE}): if self.op in GroupOp.ALU.union({Ops.CAST, Ops.COPY, Ops.ASSIGN, Ops.NOOP, Ops.GROUP, Ops.SINK, Ops.ALLREDUCE, Ops.STORE}):
# TODO: remove this hack for 3 op assign # TODO: remove this hack for 3 op assign
input_shapes = [x._shape for x in (self.src[:2] if self.op is Ops.ASSIGN else self.src) if x._shape is not None] input_shapes = [x._shape for x in (self.src[:2] if self.op is Ops.ASSIGN else self.src) if x._shape is not None]
if len(input_shapes) == 0: return None if len(input_shapes) == 0: return None
@@ -321,11 +338,11 @@ class UOp(OpMixin, metaclass=UOpMetaClass):
# *** uop evaluation *** # *** uop evaluation ***
def simplify(self, tracked=False, full_symbolic=True): def simplify(self, tracked=False):
# late import! # late import!
from tinygrad.uop.symbolic import symbolic, commutative from tinygrad.uop.symbolic import symbolic
with Context(TRACK_MATCH_STATS=0 if not tracked else TRACK_MATCH_STATS.value): with Context(TRACK_MATCH_STATS=0 if not tracked else TRACK_MATCH_STATS.value):
return graph_rewrite(self, symbolic if full_symbolic else commutative, name="simplify") return graph_rewrite(self, symbolic, name="simplify")
def ssimplify(self) -> UOp|ConstType: return ret.arg if (ret:=self.simplify()).op is Ops.CONST else ret def ssimplify(self) -> UOp|ConstType: return ret.arg if (ret:=self.simplify()).op is Ops.CONST else ret
def sintify(self) -> sint: return self.arg if self.op is Ops.CONST else self def sintify(self) -> sint: return self.arg if self.op is Ops.CONST else self
def _eval(self, dtype, expected_type:Type[T]) -> T: def _eval(self, dtype, expected_type:Type[T]) -> T:
@@ -866,8 +883,8 @@ def print_uops(uops:list[UOp]):
def get_location() -> tuple[str, int]: def get_location() -> tuple[str, int]:
frm = sys._getframe(1) frm = sys._getframe(1)
# skip over ops.py/mathtraits.py (unless there's nothing but ops.py/mathtraits.py) # skip over ops.py and anything in mixin
while pathlib.Path(frm.f_code.co_filename).name in ("ops.py", "mathtraits.py") and frm.f_back is not None and \ while ((codepath:=pathlib.Path(frm.f_code.co_filename)).name == "ops.py" or codepath.parent.name == "mixin") and frm.f_back is not None and \
not frm.f_back.f_code.co_filename.startswith("<frozen"): not frm.f_back.f_code.co_filename.startswith("<frozen"):
frm = frm.f_back frm = frm.f_back
return frm.f_code.co_filename, frm.f_lineno return frm.f_code.co_filename, frm.f_lineno
@@ -1077,20 +1094,22 @@ def track_rewrites(name:Callable[..., str|TracingKey]|bool=True, replay:bool=Fal
active_rewrites:list[TrackedGraphRewrite] = [] active_rewrites:list[TrackedGraphRewrite] = []
def profile_matches(fxn:Callable): def profile_matches(fxn:Callable):
def wrap(*args, **kwargs): def wrap_profile_matches(*args, **kwargs):
name = str(kwargs.get("name", None) or fxn.__name__) if TRACK_MATCH_STATS >= 2:
assert args and isinstance(args[0], UOp), f"invalid match tracing inputs for {name} with {args}" name = str(kwargs.get("name", None) or fxn.__name__)
if tracking:=(TRACK_MATCH_STATS >= 2): assert args and isinstance(args[0], UOp), f"invalid match tracing inputs for {name} with {args}"
loc = ((frm:=sys._getframe(1)).f_code.co_filename, frm.f_lineno) loc = ((frm:=sys._getframe(1)).f_code.co_filename, frm.f_lineno)
depth = len(active_rewrites) depth = len(active_rewrites)
if not tracked_ctxs: add_trace_group(TracingKey(f"default {fxn.__name__}")) if not tracked_ctxs: add_trace_group(TracingKey(f"default {fxn.__name__}"))
tracked_ctxs[-1].append(ctx:=TrackedGraphRewrite(loc, args[0].trace_num, [], name, depth, kwargs.get("bottom_up", False))) tracked_ctxs[-1].append(ctx:=TrackedGraphRewrite(loc, args[0].trace_num, [], name, depth, kwargs.get("bottom_up", False)))
active_rewrites.append(ctx) active_rewrites.append(ctx)
with cpu_profile(name, "TINY", display=tracking): with cpu_profile(name, "TINY"):
ret = fxn(*args, **kwargs) ret = fxn(*args, **kwargs)
if tracking: active_rewrites.pop() active_rewrites.pop()
return ret return ret
return wrap # without tracking, we just call the function
return fxn(*args, **kwargs)
return wrap_profile_matches
class TrackedPatternMatcher(PatternMatcher): class TrackedPatternMatcher(PatternMatcher):
def rewrite(self, uop:UOp, ctx=None) -> UOp|None: def rewrite(self, uop:UOp, ctx=None) -> UOp|None:
@@ -1151,7 +1170,8 @@ if TRACK_MATCH_STATS or PROFILE:
# *** simple graph rewrite engine *** # *** simple graph rewrite engine ***
with Context(SPEC=0): SENTINEL = UOp(Ops.SENTINEL) # A pure Python sentinel, but *typed* as UOp so it fits all the dict annotations
SENTINEL: Final[UOp] = cast(UOp, object())
class BottomUpGate(Exception): pass class BottomUpGate(Exception): pass
class RewriteContext: class RewriteContext:
def __init__(self, pm, bpm, ctx=None): def __init__(self, pm, bpm, ctx=None):
@@ -1162,12 +1182,12 @@ class RewriteContext:
self.ctx = ctx self.ctx = ctx
self.replace: dict[UOp, UOp] = {} self.replace: dict[UOp, UOp] = {}
def cached_pm_rewrite(self, x:UOp): def cached_pm_rewrite(self, x:UOp) -> UOp|None:
if (ret:=self.pm_cache.get(x,SENTINEL)) is not SENTINEL: return ret if (ret:=self.pm_cache.get(x,SENTINEL)) is not SENTINEL: return ret
ret = self.pm_cache[x] = unwrap(self.pm).rewrite(x, self.ctx) ret = self.pm_cache[x] = unwrap(self.pm).rewrite(x, self.ctx)
return ret return ret
def cached_bpm_rewrite(self, x:UOp): def cached_bpm_rewrite(self, x:UOp) -> UOp|None:
if (ret:=self.bpm_cache.get(x,SENTINEL)) is not SENTINEL: return ret if (ret:=self.bpm_cache.get(x,SENTINEL)) is not SENTINEL: return ret
ret = self.bpm_cache[x] = unwrap(self.bpm).rewrite(x, self.ctx) ret = self.bpm_cache[x] = unwrap(self.bpm).rewrite(x, self.ctx)
return ret return ret
@@ -1350,7 +1370,7 @@ pm_pyrender_extra = PatternMatcher([
(UPat(Ops.REDUCE_AXIS, name="r"), lambda ctx,r: f"{ctx[r.src[0]]}.r({r.arg[0]}, {r.arg[1]})"), (UPat(Ops.REDUCE_AXIS, name="r"), lambda ctx,r: f"{ctx[r.src[0]]}.r({r.arg[0]}, {r.arg[1]})"),
# NOTE: range has srcs sometimes after control flow # NOTE: range has srcs sometimes after control flow
(UPat(Ops.RANGE, src=(UPat(Ops.CONST, name="c"),), allow_any_len=True, name="x"), lambda ctx,x,c: (UPat(Ops.RANGE, src=(UPat(Ops.CONST, name="c"),), allow_any_len=True, name="x"), lambda ctx,x,c:
"UOp.range("+', '.join([str(c.arg)] + [str(y) for y in x.arg])+ "UOp.range("+', '.join([str(c.arg)] + [repr(y) for y in x.arg])+
(f', src={srcs(ctx, x.src[1:])}' if len(x.src) > 1 else '')+(', dtype='+str(x.dtype) if x.dtype is not dtypes.index else '')+")"), (f', src={srcs(ctx, x.src[1:])}' if len(x.src) > 1 else '')+(', dtype='+str(x.dtype) if x.dtype is not dtypes.index else '')+")"),
# TODO: index shouldn't mismatch dtype # TODO: index shouldn't mismatch dtype
(UPat(Ops.INDEX, src=(UPat(), UPat()), allow_any_len=True, name="x"), lambda ctx,x: (UPat(Ops.INDEX, src=(UPat(), UPat()), allow_any_len=True, name="x"), lambda ctx,x:
+2 -5
View File
@@ -17,9 +17,6 @@ from tinygrad.uop.validate import validate_index
shared_spec = PatternMatcher([ shared_spec = PatternMatcher([
(UPat(Ops.SINK, dtypes.void), lambda: True), # NOTE: for testing, we let sinks be anything (UPat(Ops.SINK, dtypes.void), lambda: True), # NOTE: for testing, we let sinks be anything
# SENTINEL should never be anywhere
(UPat(Ops.SENTINEL), lambda: False),
# CONST/DEFINE_VAR are everywhere # CONST/DEFINE_VAR are everywhere
(UPat(Ops.CONST, src=(), name="x"), lambda x: type(x.arg) is type(dtypes.as_const(x.arg, x.dtype))), (UPat(Ops.CONST, src=(), name="x"), lambda x: type(x.arg) is type(dtypes.as_const(x.arg, x.dtype))),
(UPat(Ops.DEFINE_VAR, name="x"), lambda x: isinstance(x.arg[1], int) and isinstance(x.arg[2], int)), (UPat(Ops.DEFINE_VAR, name="x"), lambda x: isinstance(x.arg[1], int) and isinstance(x.arg[2], int)),
@@ -147,8 +144,8 @@ shared_codegen_spec = PatternMatcher([
(UPat().index(UPat()).or_casted().load(), lambda: True), (UPat().index(UPat()).or_casted().load(), lambda: True),
(UPat(Ops.INDEX).or_casted().store(UPat()), lambda: True), (UPat(Ops.INDEX).or_casted().store(UPat()), lambda: True),
# all CUSTOM + PRECAST # CUSTOM (inline and non inline)
(UPat((Ops.CUSTOMI, Ops.CUSTOM, Ops.PRECAST)), lambda: True), (UPat((Ops.CUSTOMI, Ops.CUSTOM)), lambda: True),
# INDEX # INDEX
(UPat(GroupOp.Defines|{Ops.AFTER}, name="buf").index(UPat.var("idx")), validate_index), (UPat(GroupOp.Defines|{Ops.AFTER}, name="buf").index(UPat.var("idx")), validate_index),
+26 -159
View File
@@ -3,8 +3,9 @@ import math, operator, struct, functools
from collections import defaultdict from collections import defaultdict
from tinygrad.uop.ops import Ops, PatternMatcher, UPat, UOp, GroupOp, exec_alu from tinygrad.uop.ops import Ops, PatternMatcher, UPat, UOp, GroupOp, exec_alu
from tinygrad.dtype import ConstType, dtypes, PtrDType, can_safe_cast, Invalid from tinygrad.dtype import ConstType, dtypes, PtrDType, can_safe_cast, Invalid
from tinygrad.helpers import partition, all_same, prod, flatten, get_single_element, cdiv, cmod, CORRECT_DIVMOD_FOLDING, unwrap from tinygrad.helpers import partition, all_same, prod, flatten, get_single_element, unwrap
from tinygrad.uop.decompositions import xpow from tinygrad.uop.decompositions import xpow
from tinygrad.uop.divandmod import div_and_mod_symbolic
# ******** phase 1 of symbolic used to live in ops, it's the most generic folding rules ******** # ******** phase 1 of symbolic used to live in ops, it's the most generic folding rules ********
@@ -24,19 +25,16 @@ def fold_bitcast(root:UOp, c:UOp) -> UOp|None:
invalid_pat = UPat(Ops.CONST, arg=Invalid, name="i") invalid_pat = UPat(Ops.CONST, arg=Invalid, name="i")
invalid_gate = UPat.var("cond").where(UPat.var("x"), invalid_pat) invalid_gate = UPat.var("cond").where(UPat.var("x"), invalid_pat)
# this needs to be before symbolic so that 0*something_that_might_be_invalid doesnt become 0
propagate_invalid = PatternMatcher([ propagate_invalid = PatternMatcher([
# this needs to be before symbolic so that 0*something_that_might_be_invalid doesnt become 0
# propagate invalid, push it past children # propagate invalid, push it past children
(invalid_gate.cast(name="cast"), lambda i,x,cond,cast: x.cast(cast.dtype) if cast.dtype is not dtypes.index else None), (invalid_gate.cast(name="cast"), lambda i,x,cond,cast: x.cast(cast.dtype)),
*((invalid_gate.alu(op, UPat.var("y")).named("alu"), lambda cond,x,y,alu,i: cond.where(x.alu(alu.op,y), i)) *((invalid_gate.alu(op, UPat.var("y")).named("alu"), lambda cond,x,y,alu,i: cond.where(x.alu(alu.op,y), i))
for op in GroupOp.Binary-GroupOp.Comparison), for op in GroupOp.Binary-GroupOp.Comparison),
# TODO: when can this happen? and is it always safe to just drop invalid?
*((invalid_gate.alu(op, UPat.var("y")).named("alu"), lambda cond,x,y,alu,i: x.alu(alu.op,y)) for op in GroupOp.Comparison), *((invalid_gate.alu(op, UPat.var("y")).named("alu"), lambda cond,x,y,alu,i: x.alu(alu.op,y)) for op in GroupOp.Comparison),
# invalid + y -> y same for other ops # invalid + y -> invalid same for other ops
*((invalid_pat.alu(op, UPat(dtype=dtypes.index)).named("alu"), lambda alu,i: i) for op in GroupOp.Binary-GroupOp.Comparison), *((invalid_pat.alu(op, UPat(dtype=dtypes.index)).named("alu"), lambda alu,i: i) for op in GroupOp.Binary-GroupOp.Comparison),
# i < y -> a_bool_value_that_will_never_be_used: we choose a random bool const
*((invalid_pat.alu(op, UPat(dtype=dtypes.index)), lambda i: UOp.const(dtypes.bool, True)) for op in GroupOp.Comparison),
# a.where(b.where(c, d), d) -> (a & b).where(c, d)
(UPat.var("a").where(UPat.var("b").where(UPat.var("c"), UPat.var("d")), UPat.var("d")), lambda a,b,c,d: (a&b).where(c,d)),
]) ])
symbolic_simple = propagate_invalid + PatternMatcher([ symbolic_simple = propagate_invalid + PatternMatcher([
@@ -105,22 +103,17 @@ symbolic_simple = propagate_invalid + PatternMatcher([
# positive const ** x # positive const ** x
(UPat.cvar("c", vec=False).alu(Ops.POW, UPat.var("x")), lambda c,x: c if c.arg == 1 else (x*math.log2(c.arg)).exp2() if c.arg > 0 else None), (UPat.cvar("c", vec=False).alu(Ops.POW, UPat.var("x")), lambda c,x: c if c.arg == 1 else (x*math.log2(c.arg)).exp2() if c.arg > 0 else None),
# rules for threefry # rules for threefry
((UPat.var('x', dtypes.uint64)&0xFFFFFFFF).cast(dtypes.uint32), lambda x: x.cast(dtypes.uint32)&0xFFFFFFFF), # TODO: why is the and needed? ((UPat.var('x', dtypes.uint64)&0xFFFFFFFF).cast(dtypes.uint32), lambda x: x.cast(dtypes.uint32)),
(((UPat.var(None, dtypes.uint64)*(1<<32)) | UPat.var('y', dtypes.uint32).cast(dtypes.uint64)).cast(dtypes.uint32), lambda y: y), (((UPat.var(None, dtypes.uint64)*(1<<32)) | UPat.var('y', dtypes.uint32).cast(dtypes.uint64)).cast(dtypes.uint32), lambda y: y),
(((UPat.var('x', dtypes.uint64)*(1<<32)) | UPat.var(None, dtypes.uint32).cast(dtypes.uint64))//(1<<32), lambda x: x), (((UPat.var('x', dtypes.uint64)*(1<<32)) | UPat.var(None, dtypes.uint32).cast(dtypes.uint64))//(1<<32), lambda x: x),
# hacks for threefry long removal when padded (TODO: genericize)
(UPat.var('x', dtypes.uint32).cast(dtypes.uint64) * UPat.var('y').where(UPat.const(dtypes.uint64, 1<<32), UPat.const(dtypes.uint64, 0)),
lambda x,y: y.where(x, 0).cast(dtypes.uint64) * (1<<32)),
((UPat.var('x', dtypes.uint64)&(UPat.var('y').where(UPat.const(dtypes.uint64, 0xFFFFFFFF), UPat.const(dtypes.uint64, 0)))).cast(dtypes.uint32),
lambda x,y: y.where(x.cast(dtypes.uint32), 0)),
# new decomp rules for threefry
(((UPat.var(None, dtypes.uint64)<<32) | UPat.var('y', dtypes.uint32).cast(dtypes.uint64)).cast(dtypes.uint32), lambda y: y), (((UPat.var(None, dtypes.uint64)<<32) | UPat.var('y', dtypes.uint32).cast(dtypes.uint64)).cast(dtypes.uint32), lambda y: y),
(((UPat.var('x', dtypes.uint64)<<32) | UPat.var(None, dtypes.uint32).cast(dtypes.uint64))>>32, lambda x: x), (((UPat.var('x', dtypes.uint64)<<32) | UPat.var(None, dtypes.uint32).cast(dtypes.uint64))>>32, lambda x: x),
(UPat.var('b').where(UPat.var('x', dtypes.uint32).cast(dtypes.uint64), UPat.const(dtypes.uint64, 0)).cast(dtypes.uint32), lambda b,x: b.where(x,0)),
# ** simple where folding ** # ** simple where folding **
# a conditional with the same results either way is a noop, also fold const conditionals # a conditional with the same results either way is a noop, also fold const conditionals
(UPat.var().where(UPat.var("val"), UPat.var("val")), lambda val: val), (UPat.var().where(UPat.var("val"), UPat.var("val")), lambda val: val),
(UPat.cvar("gate", vec=False).where(UPat.var("c0"), UPat.var("c1")), lambda gate, c0, c1: c0 if gate.arg else c1), (UPat.cvar("gate", vec=False).where(UPat.var("c0"), UPat.var("c1")), lambda gate, c0, c1: c0 if gate.arg else c1),
# a.where(b.where(c, d), d) -> (a & b).where(c, d)
(UPat.var("a").where(UPat.var("b").where(UPat.var("c"), UPat.var("d")), UPat.var("d")), lambda a,b,c,d: (a&b).where(c,d)),
]) ])
# ******** phase 2 builds on phase 1, it includes the old "symbolic", rules that match deeper ******** # ******** phase 2 builds on phase 1, it includes the old "symbolic", rules that match deeper ********
@@ -144,101 +137,6 @@ def canonicalize_simplex(X:UOp) -> UOp|None:
ret.append(u) ret.append(u)
return UOp.sum(*ret) if changed else None return UOp.sum(*ret) if changed else None
def cancel_divmod(d: UOp, x: UOp, y: UOp) -> UOp|None:
# simple cancel div/mod case when the range of the numerator lies within a single denominator interval
x_min, x_max, y_min, y_max = x.vmin, x.vmax, y.vmin, y.vmax
assert isinstance(x_min, int) and isinstance(x_max, int) and isinstance(y_min, int) and isinstance(y_max, int)
if y_min==y_max==0: raise ZeroDivisionError(f"{'Division' if d.op is Ops.IDIV else 'Mod'} by zero trying to rewrite {x.alu(d.op, y)}")
if y_min*y_max > 0 and (q:=cdiv(x_min,y_min)) == cdiv(x_min,y_max) == cdiv(x_max,y_min) == cdiv(x_max,y_max):
return x - q*y if d.op is Ops.MOD else d.const_like(q)
return None
def remove_nested_mod(m: UOp, x: UOp, y: UOp) -> UOp|None:
# remove nested mod in case the inner mod is a multiple of the outer mod
# example: (a%4 + b)%2 -> (a+b)%2
if ((c := y.arg) < 0) or x.vmin<0: return None
new_xs = []
something_changed = False
for u in x.split_uop(Ops.ADD):
if u.op is Ops.MOD:
if u.src[1].divides(c) is not None:
something_changed = True
u = u.src[0]
new_xs.append(u)
new_x: UOp = UOp.sum(*new_xs)
if something_changed and new_x.vmin>=0: return new_x % y
return None
def fold_binary_numerator(d: UOp, x: UOp, y: UOp) -> UOp|None:
# we can fold if the expression has only one non-constant term and this term can only take on two values
if ((c := y.arg) < 0): return None
x,const = x.pop_const()
terms, factors = zip(*[(u.divides(f:=u.const_factor()),f) for u in x.split_uop(Ops.ADD)])
if len(terms)==1 and (v:=terms[0]).vmax-v.vmin == 1:
y1 = cmod(factors[0]*v.vmin+const, c) if d.op is Ops.MOD else cdiv(factors[0]*v.vmin+const, c)
y2 = cmod(factors[0]*v.vmax+const, c) if d.op is Ops.MOD else cdiv(factors[0]*v.vmax+const, c)
return (y2-y1)*(v-v.vmin) + y1
return None
def fold_divmod_congruence(d: UOp, x: UOp, y: UOp) -> UOp|None:
# within a mod we can freely subtract multiples of c, we use this to see if a is congruent to an expression whose vmin/vmax are between 0 and c
if (x.vmin<0 and CORRECT_DIVMOD_FOLDING) or ((c := y.arg) < 0): return None
x,const = x.pop_const()
terms, factors = zip(*[(u.divides(f:=u.const_factor()),f) for u in x.split_uop(Ops.ADD)])
# a//c = (a-a%c)/c, if we can fold a%c, we can fold a//c
rems = [min((r:=f%c), r-c, key=abs) for f in factors]
if (rem:=sum(r*v for r,v in zip(rems,terms))+const%c).vmin//c!=rem.vmax//c: return None
if d.op is Ops.MOD: return rem - rem.vmin//c*c
return sum((f-r)//c * v for f,r,v in zip(factors,rems,terms)) + (const-const%c+rem.vmin//c*c)//c
def divide_by_gcd(d: UOp, x: UOp, y: UOp) -> UOp|None:
# x//y -> (x//gcd)//(y//gcd) or x%y -> gcd*(x//gcd)%(y//gcd)
gcd = UOp.gcd(*x.split_uop(Ops.ADD), y).simplify()
if gcd.op is Ops.CONST and gcd.arg==1: return None
ret = unwrap(x.divide_exact(gcd)).alu(d.op, unwrap(y.divide_exact(gcd)))
return ret*gcd if d.op is Ops.MOD else ret
def gcd_with_remainder(d: UOp, x: UOp, y: UOp):
# (gcd*x+r)//(gcd*d) -> (x+(r%d)//gcd)//d + r//(gcd*d)
# (gcd*x+r)%(gcd*d) -> gcd*(x+(r%d)//gcd)%d + r%gcd
# These only work for floordiv (and the corresponding remainder)! Thats why we check the sign of x,y and new_x
if ((c := y.arg) < 0) or x.vmin<0: return None
x_no_const, const = x.pop_const()
gcd = UOp.gcd(*x_no_const.split_uop(Ops.ADD), y).simplify()
assert gcd.op is Ops.CONST
if gcd.arg==1: return None
new_x = unwrap(x_no_const.divide_exact(gcd)).simplify() + (const%c)//gcd
if new_x.vmin<0: return None
ret = new_x.alu(d.op, x.ufix(c//gcd.arg))
return ret*gcd + const%gcd.arg if d.op is Ops.MOD else ret+const//c
def factor_remainder(d: UOp, x: UOp, y: UOp) -> UOp|None:
# (d*x+y)//d -> x+y//d or (d*x+y)%d
# for mod we go further and take the remainder of all factors to reduce their size
# These only work for floordiv (and the corresponding remainder)! Thats why we check the sign of x,y and new_x
if y.vmin<0 or x.vmin<0: return None
quo, rem = [], []
for u in x.split_uop(Ops.ADD):
if (q:=u.divide_exact(y)) is not None: quo.append(q)
# if this is mod and y is a const, we can make the remainder factor sm
elif d.op is Ops.MOD and y.op is Ops.CONST and (c:=u.const_factor())%y.arg!=c:
rem.append(u.divides(c)*(c%y.arg))
quo.append(u.const_like(0)) # we append this so we can check if something changed
else: rem.append(u)
new_x = sum(rem)+x.const_like(0)
if len(quo)==0 or new_x.vmin<0: return None
return new_x%y if d.op is Ops.MOD else new_x//y+sum(quo)
def nest_div_by_smallest_factor(d: UOp, x: UOp, y: UOp) -> UOp|None:
# we try and nest the div and see if it allows the numerator to be simplified
if ((c := y.arg) < 0): return None
factors = [u.const_factor() for u in x.split_uop(Ops.ADD) if u.op not in (Ops.CONST, Ops.VCONST)]
div = min([y.arg]+[abs(f) for f in factors if abs(f) > 1 and (c%f)==0])
newxs = fold_divmod_congruence(newx:=(x//div), x, y.const_like(div))
if newxs is None: newxs = factor_remainder(newx, x, y.const_like(div))
if div==y.arg or newxs is None or x.vmin<0 or newx.vmin<0: return None
return newxs//(c//div)
def gep_through_wmma(gep:UOp, wmma:UOp): def gep_through_wmma(gep:UOp, wmma:UOp):
out_sz = prod(x[1] for x in wmma.arg[6][-1]) out_sz = prod(x[1] for x in wmma.arg[6][-1])
wmma_idxs = gep.arg[::out_sz] wmma_idxs = gep.arg[::out_sz]
@@ -298,10 +196,10 @@ symbolic = symbolic_simple+commutative+PatternMatcher([
((UPat.var("y") + UPat.var("x")) + UPat.var("x"), lambda y,x: y+x*2), ((UPat.var("y") + UPat.var("x")) + UPat.var("x"), lambda y,x: y+x*2),
((UPat.var("x") / UPat.var("x2")) / UPat.var("x3"), lambda x,x2,x3: x/(x2*x3) if x2 is not x3 else None), # (x/x2)/x3 -> x/(x2*x3) ((UPat.var("x") / UPat.var("x2")) / UPat.var("x3"), lambda x,x2,x3: x/(x2*x3) if x2 is not x3 else None), # (x/x2)/x3 -> x/(x2*x3)
(-1 * (UPat.var("x") + UPat.cvar("c")), lambda x,c: (-x)+(-c)), # -(x+c) -> -x + -c (-1 * (UPat.var("x") + UPat.cvar("c")), lambda x,c: (-x)+(-c)), # -(x+c) -> -x + -c
(UPat.cvar("y") * (UPat.var("x", dtype=dtypes.index) + UPat.cvar("c")), lambda x,y,c: (y*x)+(y*c)), # -(x+c) -> -x + -c (UPat.cvar("y") * (UPat.var("x", dtype=dtypes.index) + UPat.cvar("c")), lambda x,y,c: (y*x)+(y*c)), # y*(x+c) -> y*x + y*c
# ** where folding ** # ** where folding **
(UPat.var("cond", dtype=dtypes.bool).logical_not().where(UPat.var("t"), UPat.var("f")), lambda cond, t, f: cond.where(f,t) (UPat.var("cond", dtype=dtypes.bool).logical_not().where(UPat.var("t"), UPat.var("f")),
if f.arg is not Invalid else None), lambda cond, t, f: cond.where(f,t) if f.arg is not Invalid else None),
# alu of two where with same conds can combine, only do if true branch or false branch is const # alu of two where with same conds can combine, only do if true branch or false branch is const
(UPat(GroupOp.Binary, name="alu", src=(UPat.var("c").where(UPat.var("t"), UPat.var("f")), UPat.var("c").where(UPat.var("tt"), UPat.var("ff")))), \ (UPat(GroupOp.Binary, name="alu", src=(UPat.var("c").where(UPat.var("t"), UPat.var("f")), UPat.var("c").where(UPat.var("tt"), UPat.var("ff")))), \
lambda alu,c,t,tt,f,ff: c.where(t.alu(alu.op, tt), f.alu(alu.op, ff)) if t.op == tt.op == Ops.CONST or f.op == ff.op == Ops.CONST else None), lambda alu,c,t,tt,f,ff: c.where(t.alu(alu.op, tt), f.alu(alu.op, ff)) if t.op == tt.op == Ops.CONST or f.op == ff.op == Ops.CONST else None),
@@ -315,7 +213,6 @@ symbolic = symbolic_simple+commutative+PatternMatcher([
(UPat.maximum(UPat.var("x"), UPat.var("y")), lambda x,y: x if x.vmin >= y.vmax else y if x.vmax <= y.vmin else None), (UPat.maximum(UPat.var("x"), UPat.var("y")), lambda x,y: x if x.vmin >= y.vmax else y if x.vmax <= y.vmin else None),
# TODO: why does this rule break beautiful_mnist? # TODO: why does this rule break beautiful_mnist?
#((UPat.var("x")+UPat.var("z")).maximum(UPat.var("y")+UPat.var("z")), lambda x,y,z: x.maximum(y) + z), #((UPat.var("x")+UPat.var("z")).maximum(UPat.var("y")+UPat.var("z")), lambda x,y,z: x.maximum(y) + z),
#((UPat.var("x")*UPat.cvar("c1")).maximum(UPat.var("x")*UPat.cvar("c2")), max_var_const),
# ** two stage ALU folding ** # ** two stage ALU folding **
*((UPat.var("x").alu(op, UPat.cvar("c1")).alu(op, UPat.cvar("c2")).named("f"), *((UPat.var("x").alu(op, UPat.cvar("c1")).alu(op, UPat.cvar("c2")).named("f"),
lambda f,x,c1,c2: x.alu(f.op,c1.alu(f.op,c2))) for op in GroupOp.Associative), lambda f,x,c1,c2: x.alu(f.op,c1.alu(f.op,c2))) for op in GroupOp.Associative),
@@ -338,34 +235,11 @@ symbolic = symbolic_simple+commutative+PatternMatcher([
# generic lt folding # generic lt folding
(UPat.var("x", dtypes.index)<UPat.cvar("c", vec=False), lambda x,c: lt_folding(x, c.arg) if 0 < c.arg else None), (UPat.var("x", dtypes.index)<UPat.cvar("c", vec=False), lambda x,c: lt_folding(x, c.arg) if 0 < c.arg else None),
(UPat.var("x", dtypes.index)*-1 < UPat.var("y")*-1, lambda x,y: y<x), (UPat.var("x", dtypes.index)*-1 < UPat.var("y")*-1, lambda x,y: y<x),
# canonicalize a simplex with positive coefficients > 0 # canonicalize a simplex with positive coefficients > 0. NOTE: not x < 1 means x > 0
# not x < 1 -> X > 0
((UPat.var("x", dtypes.index)<1).ne(True), lambda x: (newx<1).ne(True) if (newx:=canonicalize_simplex(x)) is not None else None), ((UPat.var("x", dtypes.index)<1).ne(True), lambda x: (newx<1).ne(True) if (newx:=canonicalize_simplex(x)) is not None else None),
# ** div **
# div folding
((UPat.var("x")//UPat.cvar("c") + UPat.cvar("a"))//UPat.cvar("d"), lambda x,c,a,d: (x+a*c)//(c*d)
if c.vmin>0 and d.vmin>0 and ((x.vmin>=0 and a.vmin>=0) or (x.vmax<=0 and a.vmax<=0)) else None), # (x//c+a)//d -> (x+a*c)//(c*d)
# a range mod its own upper bound is just the range # a range mod its own upper bound is just the range
(UPat(Ops.RANGE, src=UPat.var("end"), name="r")%UPat.var("end"), lambda r,end: r), (UPat(Ops.RANGE, src=UPat.var("end"), name="r")%UPat.var("end"), lambda r,end: r),
(UPat(Ops.RANGE, src=UPat.var("end"), name="r")//UPat.var("end"), lambda r,end: r.const_like(0)), (UPat(Ops.RANGE, src=UPat.var("end"), name="r")//UPat.var("end"), lambda r,end: r.const_like(0)),
(UPat((Ops.IDIV, Ops.MOD), dtypes.index, name="d", src=(UPat.var("x"), UPat.var("y"))), cancel_divmod),
(UPat.var("x", dtypes.index) // UPat.var("d"), lambda x,d: -(x//(-d)) if d.vmax < 0 else None),
(UPat((Ops.IDIV, Ops.MOD), dtypes.index, name="d", src=(UPat.var("x"), UPat.cvar("y", vec=False))), fold_binary_numerator),
(UPat((Ops.IDIV, Ops.MOD), dtypes.index, name="d", src=(UPat.var("x"), UPat.cvar("y", vec=False))), fold_divmod_congruence),
(UPat((Ops.IDIV, Ops.MOD), dtypes.index, name="d", src=(UPat.var("x"), UPat.var("y"))), divide_by_gcd),
(UPat((Ops.IDIV, Ops.MOD), dtypes.index, name="d", src=(UPat.var("x"), UPat.cvar("y", vec=False))), gcd_with_remainder),
(UPat(Ops.MOD, dtypes.index, name="m", src=(UPat.var("x"), UPat.cvar("y", vec=False))), remove_nested_mod),
(UPat((Ops.IDIV), dtypes.index, name="d", src=(UPat.var("x"), UPat.cvar("y", vec=False))), nest_div_by_smallest_factor),
(UPat((Ops.IDIV, Ops.MOD), dtypes.index, name="d", src=(UPat.var("x"), UPat.var("y"))), factor_remainder),
(UPat.var("x", dtypes.index) // UPat.var("d"), lambda x,d: -((-x)//d) if x.vmax<=0 else None),
((UPat.var("x", dtypes.index)+UPat.cvar("c", vec=False)).named("n")//UPat.cvar("d", vec=False),
lambda x,c,n,d: ((x+c.arg%d.arg)//d + c.arg//d.arg) if c.arg%d.arg!=c.arg and x.vmin>=0 and n.vmin>=0 and d.arg>0 else None),
((UPat.var("x", dtypes.index)+UPat.cvar("c", vec=False)).named("n")//UPat.cvar("d", vec=False),
lambda x,c,n,d: (-(-(c.arg%d.arg + x - (d.arg-1))//d) + c.arg//d.arg) if x.vmax<=0 and n.vmin>=0 and d.arg>0 else None),
# ** mod **
# mod folding
(UPat.var("x", dtypes.index) % UPat.var("d"), lambda x,d: -((-x)%d) if x.vmax <= 0 else None),
(UPat.var("x", dtypes.index) % UPat.var("d"), lambda x,d: (x%(-d)) if d.vmax < 0 else None),
# cast/long folding # cast/long folding
# if the intermediate cast doesnt narrow we can do it in one cast # if the intermediate cast doesnt narrow we can do it in one cast
(UPat.var('x').cast(name="a").cast(name="b"), lambda x,a,b: x.cast(b.dtype) if can_safe_cast(x.dtype, a.dtype) else None), (UPat.var('x').cast(name="a").cast(name="b"), lambda x,a,b: x.cast(b.dtype) if can_safe_cast(x.dtype, a.dtype) else None),
@@ -382,19 +256,21 @@ symbolic = symbolic_simple+commutative+PatternMatcher([
(UPat(Ops.AFTER, src=(UPat.var("s"),)), lambda s: s), (UPat(Ops.AFTER, src=(UPat.var("s"),)), lambda s: s),
# VECTORIZE/CONST # VECTORIZE/CONST
(UPat(Ops.VECTORIZE, src=UPat(Ops.CONST), name="vec"), lambda vec: UOp.const(vec.dtype, tuple(x.arg for x in vec.src))), (UPat(Ops.VECTORIZE, src=UPat(Ops.CONST), name="vec"), lambda vec: UOp.const(vec.dtype, tuple(x.arg for x in vec.src))),
])+gep_pushing ])+div_and_mod_symbolic+gep_pushing
# ******** we take a small aside to "simplify_valid" to rewrite valids ******** # ******** we take a small aside to "simplify_valid" to rewrite valids ********
def parse_valid(valid:UOp) -> tuple[UOp, bool, int]|None: def parse_valid(v:UOp) -> tuple[UOp, bool, int]|None:
# if it's X <= c, returns X, True, c # if it's X <= c, returns X, True, c
# if it's X >= c, returns X, False, c # if it's X >= c, returns X, False, c
# (X < c).ne(True) -> X >= c if v.op is Ops.CMPNE and v.src[1].op is Ops.CONST and v.src[1].arg == 1 and (s0:=v.src[0]).op is Ops.CMPLT and dtypes.is_int(s0.src[0].dtype):
if valid.op is Ops.CMPNE and valid.src[1].op is Ops.CONST and valid.src[1].arg == 1 and \ # (X < c).ne(True) -> X >= c
(s0:=valid.src[0]).op is Ops.CMPLT and dtypes.is_int(s0.src[0].dtype): return s0.src[0], False, int(s0.src[1].vmin) return s0.src[0], False, int(s0.src[1].vmin)
# X < c -> X <= c-1 if v.op is Ops.CMPLT and dtypes.is_int(v.src[0].dtype):
if valid.op is Ops.CMPLT and dtypes.is_int(valid.src[0].dtype): return valid.src[0], True, int((valid.src[1]).vmax)-1 # X < c -> X <= c-1
return v.src[0], True, int((v.src[1]).vmax)-1
# NOTE: v.src[1].op can be Ops.VCONST
return None return None
def uop_given_valid(valid:UOp, uop:UOp, try_simplex=True) -> UOp: def uop_given_valid(valid:UOp, uop:UOp, try_simplex=True) -> UOp:
@@ -425,7 +301,7 @@ def uop_given_valid(valid:UOp, uop:UOp, try_simplex=True) -> UOp:
# if every branch in candidate gives the same simplified uop, we can rewrite the uop # if every branch in candidate gives the same simplified uop, we can rewrite the uop
newuops = [uop.substitute({X:newX}) for X,newX in candidate] newuops = [uop.substitute({X:newX}) for X,newX in candidate]
if any(u is uop for u in newuops): continue # if any branch doesnt appear in uop, skip if any(u is uop for u in newuops): continue # if any branch doesnt appear in uop, skip
newuops = [u.simplify().substitute({newX:X}).simplify(full_symbolic=False) for (X,newX),u in zip(candidate,newuops)] newuops = [u.simplify().substitute({newX:X}).simplify() for (X,newX),u in zip(candidate,newuops)]
if all_same(newuops): uop = newuops[0] if all_same(newuops): uop = newuops[0]
elif uop.op is Ops.VECTORIZE and len(uop.src) == 2: elif uop.op is Ops.VECTORIZE and len(uop.src) == 2:
if all_same([uops.src[0] for uops in newuops]): uop = uop.replace(src=(newuops[0].src[0], uop.src[1])) if all_same([uops.src[0] for uops in newuops]): uop = uop.replace(src=(newuops[0].src[0], uop.src[1]))
@@ -433,7 +309,7 @@ def uop_given_valid(valid:UOp, uop:UOp, try_simplex=True) -> UOp:
# try all the valids together (but only the whole expressions) # try all the valids together (but only the whole expressions)
if (s_uop:=uop.substitute(sub_dict:=dict(all_candidates))) is not uop: if (s_uop:=uop.substitute(sub_dict:=dict(all_candidates))) is not uop:
uop = s_uop.simplify().substitute({newX:X for X,newX in sub_dict.items()}).simplify(full_symbolic=False) uop = s_uop.simplify().substitute({newX:X for X,newX in sub_dict.items()}).simplify()
return uop return uop
def _valid_priority(v: UOp, valids:list[UOp]): def _valid_priority(v: UOp, valids:list[UOp]):
@@ -466,7 +342,7 @@ def reduce_mul_chain(r:UOp):
def drop_and_clauses(cond:UOp, x:UOp, i:UOp) -> UOp|None: def drop_and_clauses(cond:UOp, x:UOp, i:UOp) -> UOp|None:
if not (dropped_clauses:=[c for c in cond.split_uop(Ops.AND) if not any(r in x.ranges for r in c.ranges)]): return None if not (dropped_clauses:=[c for c in cond.split_uop(Ops.AND) if not any(r in x.ranges for r in c.ranges)]): return None
return UOp.const(dtypes.bool, True).prod(*[c for c in cond.split_uop(Ops.AND) if c not in dropped_clauses]).where(x, i) return UOp.const(dtypes.bool, True).prod(*[c for c in cond.split_uop(Ops.AND) if c not in dropped_clauses]).where(x, i)
pm_drop_and_clauses = PatternMatcher([(UPat.var("cond").where(UPat.var("x", dtype=dtypes.index), invalid_pat), drop_and_clauses)]) pm_drop_and_clauses = PatternMatcher([(invalid_gate, drop_and_clauses)])
def where_on_load(c1, buf, x): def where_on_load(c1, buf, x):
c2 = x.get_valid() c2 = x.get_valid()
@@ -490,24 +366,15 @@ pm_simplify_valid = PatternMatcher([
# simplify valid # simplify valid
(UPat(Ops.AND, name="valid"), simplify_valid), (UPat(Ops.AND, name="valid"), simplify_valid),
# TODO: this regressed openpilot, not having this regressed cifar # TODO: this regressed openpilot, not having this regressed cifar
# (UPat.var("c").where(UPat.var("x", dtype=dtypes.index), invalid_pat), lambda c,x,i: c.where(uop_given_valid(c, x, try_simplex=False), i)), # (invalid_gate, lambda cond,x,i: cond.where(uop_given_valid(cond, x, try_simplex=False), i)),
]) ])
# this is symbolic 2.0 # this is symbolic 2.0
REMOVE_FROM_SINK_LIKE = {Ops.UNROLL, Ops.NOOP, Ops.VECTORIZE, Ops.SINK} REMOVE_FROM_SINK_LIKE = {Ops.UNROLL, Ops.NOOP, Ops.VECTORIZE, Ops.SINK}
sym = symbolic+pm_simplify_valid+PatternMatcher([ sym = symbolic+pm_simplify_valid+PatternMatcher([
# VECTORIZE/GEP
(UPat(Ops.VECTORIZE, src=UPat(Ops.GEP, src=(UPat.var("x"),)), name="vec"), lambda vec,x: x.gep(tuple(y.arg[0] for y in vec.src))),
# reorder ALU/VECTORIZE # reorder ALU/VECTORIZE
(UPat(GroupOp.ALU, src=(UPat(Ops.VECTORIZE, src=UPat(name='x')), UPat(Ops.VECTORIZE, src=UPat(name='y'))), name='alu'), (UPat(GroupOp.ALU, src=(UPat(Ops.VECTORIZE, src=UPat(name='x')), UPat(Ops.VECTORIZE, src=UPat(name='y'))), name='alu'),
lambda x,y,alu: UOp(Ops.VECTORIZE, alu.dtype, (UOp(alu.op, alu.dtype.scalar(), (x,y)),)*alu.dtype.count)), lambda x,y,alu: UOp(Ops.VECTORIZE, alu.dtype, (UOp(alu.op, alu.dtype.scalar(), (x,y)),)*alu.dtype.count)),
# VECTORIZE of a single element is just that element
(UPat(Ops.VECTORIZE, src=(UPat(name='x'),)), lambda x: x),
# VECTORIZE void is GROUP
(UPat(Ops.VECTORIZE, dtype=dtypes.void, name='x'), lambda x: UOp.group(*x.src)),
# tensor core with a 0 input is acc
(UPat(Ops.WMMA, src=(UPat.const(None, 0.0), UPat.var(), UPat.var("acc"))), lambda acc: acc),
(UPat(Ops.WMMA, src=(UPat.var(), UPat.const(None, 0.0), UPat.var("acc"))), lambda acc: acc),
# ** self folding ** # ** self folding **
# x!=0 -> (bool)x # x!=0 -> (bool)x
(UPat.var("x")!=0, lambda x: x.cast(dtypes.bool.vec(x.dtype.count))), (UPat.var("x")!=0, lambda x: x.cast(dtypes.bool.vec(x.dtype.count))),
+1 -1
View File
@@ -25,7 +25,7 @@ try:
# variables # variables
(UPat(Ops.SPECIAL, name="x"), lambda x,ctx: create_bounded(x.arg, 0, ctx[1][x.src[0]]-1, ctx[0])), (UPat(Ops.SPECIAL, name="x"), lambda x,ctx: create_bounded(x.arg, 0, ctx[1][x.src[0]]-1, ctx[0])),
(UPat(Ops.DEFINE_VAR, name="x"), lambda x,ctx: create_bounded(x.arg[0], x.arg[1], x.arg[2], ctx[0])), (UPat(Ops.DEFINE_VAR, name="x"), lambda x,ctx: create_bounded(x.arg[0], x.arg[1], x.arg[2], ctx[0])),
(UPat(Ops.RANGE, name="x"), lambda x,ctx: create_bounded(f"r{x.arg}", 0, ctx[1][x.src[0]]-1, ctx[0])), (UPat(Ops.RANGE, name="x"), lambda x,ctx: create_bounded(x.render(simplify=False), 0, ctx[1][x.src[0]]-1, ctx[0])),
# loads are variables bounded by the min/max of the dtype # loads are variables bounded by the min/max of the dtype
(UPat(Ops.LOAD, dtypes.ints+(dtypes.index,), name="x"), lambda x,ctx: create_bounded(f"load{len(ctx[1])}", x.dtype.min, x.dtype.max, ctx[0])), (UPat(Ops.LOAD, dtypes.ints+(dtypes.index,), name="x"), lambda x,ctx: create_bounded(f"load{len(ctx[1])}", x.dtype.min, x.dtype.max, ctx[0])),
(UPat(Ops.LOAD, dtypes.bool, name="x"), lambda x,ctx: (z3.Bool(f"load{len(ctx[1])}", ctx=ctx[0].ctx), None)), (UPat(Ops.LOAD, dtypes.bool, name="x"), lambda x,ctx: (z3.Bool(f"load{len(ctx[1])}", ctx=ctx[0].ctx), None)),
+1 -3
View File
@@ -56,7 +56,6 @@
} }
ul > ul { ul > ul {
display: none; display: none;
margin-left: 6px;
} }
ul.has-children > p::before { ul.has-children > p::before {
content:"▸ "; content:"▸ ";
@@ -270,12 +269,11 @@
font-size: 10px; font-size: 10px;
} }
#device-list > div { #device-list > div {
min-height: 32px;
width: 134px;
overflow-x: auto; overflow-x: auto;
overflow-y: hidden; overflow-y: hidden;
white-space: nowrap; white-space: nowrap;
display: flex; display: flex;
min-height: 32px;
} }
#device-list > div:hover { #device-list > div:hover {
background-color: rgba(20, 23, 35, 0.3); background-color: rgba(20, 23, 35, 0.3);
+58 -38
View File
@@ -156,9 +156,9 @@ function formatMicroseconds(ts, dur=ts) {
} }
const formatUnit = (d, unit="") => d3.format(".3~s")(d)+unit; const formatUnit = (d, unit="") => d3.format(".3~s")(d)+unit;
const colorScheme = {TINY:["#1b5745", "#354f52", "#354f52", "#1d2e62", "#63b0cd"], const colorScheme = {TINY:new Map([["Schedule","#1b5745"],["get_program","#1d2e62"],["compile","#63b0cd"],["DEFAULT","#354f52"]]),
DEFAULT:["#2b2e39", "#2c2f3a", "#31343f", "#323544", "#2d303a", "#2e313c", "#343746", "#353847", "#3c4050", "#404459", "#444862", "#4a4e65"], DEFAULT:["#2b2e39", "#2c2f3a", "#31343f", "#323544", "#2d303a", "#2e313c", "#343746", "#353847", "#3c4050", "#404459", "#444862", "#4a4e65"],
BUFFER:["#342483", "#3E2E94", "#4938A4", "#5442B4", "#5E4CC2", "#674FCA"], SIMD:["#3600f0"], BUFFER:["#342483", "#3E2E94", "#4938A4", "#5442B4", "#5E4CC2", "#674FCA"], SE:new Map([["OCC", "#101725"], ["INST", "#0A2042"]]),
CATEGORICAL:["#ff8080", "#F4A261", "#C8F9D4", "#8D99AE", "#F4A261", "#ffffa2", "#ffffc0", "#87CEEB"],} CATEGORICAL:["#ff8080", "#F4A261", "#C8F9D4", "#8D99AE", "#F4A261", "#ffffa2", "#ffffc0", "#87CEEB"],}
const cycleColors = (lst, i) => lst[i%lst.length]; const cycleColors = (lst, i) => lst[i%lst.length];
@@ -198,15 +198,17 @@ function focusShape(shape) {
return metadata.replaceChildren(shapeMetadata.get(focusedShape) ?? ""); return metadata.replaceChildren(shapeMetadata.get(focusedShape) ?? "");
} }
async function renderProfiler(path, unit) { const EventTypes = { EXEC:0, BUF:1 };
async function renderProfiler(path, unit, opts) {
displaySelection("#profiler"); displaySelection("#profiler");
metadata.replaceChildren(shapeMetadata.get(focusedShape) ?? ""); metadata.replaceChildren(shapeMetadata.get(focusedShape) ?? "");
// layout once! // layout once!
if (data != null && data.path === path) return updateProgress({ start:false }); if (data != null && data.path === path) return updateProgress({ start:false });
// support non realtime x axis units // support non realtime x axis units
const formatTime = unit === "realtime" ? formatMicroseconds : (s) => `${s} ${unit}`; const formatTime = unit === "realtime" ? formatMicroseconds : (s) => formatUnit(s, " "+unit);
const profiler = d3.select("#profiler").html(""); const profiler = d3.select("#profiler").html("");
const buf = await (await fetch(path)).arrayBuffer(); const buf = cache[path] ?? await fetchValue(path);
const view = new DataView(buf); const view = new DataView(buf);
let offset = 0; let offset = 0;
const u8 = () => { const ret = view.getUint8(offset); offset += 1; return ret; } const u8 = () => { const ret = view.getUint8(offset); offset += 1; return ret; }
@@ -234,30 +236,36 @@ async function renderProfiler(path, unit) {
for (let i=0; i<layoutsLen; i++) { for (let i=0; i<layoutsLen; i++) {
const nameLen = view.getUint8(offset, true); offset += 1; const nameLen = view.getUint8(offset, true); offset += 1;
const k = textDecoder.decode(new Uint8Array(buf, offset, nameLen)); offset += nameLen; const k = textDecoder.decode(new Uint8Array(buf, offset, nameLen)); offset += nameLen;
const div = deviceList.append("div").attr("id", k).text(k).style("padding", padding+"px"); const div = deviceList.append("div").attr("id", k).text(k).style("padding", padding+"px").style("width", opts.width);
const { y:baseY, height:baseHeight } = rect(div.node()); const { y:baseY, height:baseHeight } = rect(div.node());
const colors = colorScheme[k.split(":")[0]] ?? colorScheme.DEFAULT;
const offsetY = baseY-canvasTop+padding/2; const offsetY = baseY-canvasTop+padding/2;
const shapes = [], visible = []; const shapes = [], visible = [];
const EventTypes = {TIMELINE:0, MEMORY:1};
const eventType = u8(), eventsLen = u32(); const eventType = u8(), eventsLen = u32();
if (eventType === EventTypes.TIMELINE) { if (eventType === EventTypes.EXEC) {
const levelHeight = baseHeight-padding; const levelHeight = (baseHeight-padding)*(opts.heightScale ?? 1);
const levels = []; const levels = [];
data.tracks.set(k, { shapes, visible, offsetY, pcolor:"#9ea2ad" }); data.tracks.set(k, { shapes, eventType, visible, offsetY, pcolor:"#9ea2ad" });
let colorKey, ref; let colorKey, ref;
for (let j=0; j<eventsLen; j++) { for (let j=0; j<eventsLen; j++) {
const e = {name:strings[u32()], ref:optional(u32()), key:optional(u32()), st:u32(), dur:f32(), info:strings[u32()] || null}; const e = {name:strings[u32()], ref:optional(u32()), key:optional(u32()), st:u32(), dur:f32(), info:strings[u32()] || null};
// find a free level to put the event // find a free level to put the event
let depth = levels.findIndex(levelEt => e.st >= levelEt); let depth = 0;
const et = e.st+Math.trunc(e.dur); if (opts.levelKey != null) { depth = opts.levelKey(e); levels[depth] = 0; }
if (depth === -1) { else {
depth = levels.length; depth = levels.findIndex(levelEt => e.st >= levelEt);
levels.push(et); const et = e.st+Math.trunc(e.dur);
} else levels[depth] = et; if (depth === -1) {
depth = levels.length;
levels.push(et);
} else levels[depth] = et;
}
if (depth === 0) colorKey = e.name.split(" ")[0]; if (depth === 0) colorKey = e.name.split(" ")[0];
if (!colorMap.has(colorKey)) colorMap.set(colorKey, d3.rgb(cycleColors(colorScheme[k.split(":")[0]] ?? colorScheme.DEFAULT, colorMap.size))); if (!colorMap.has(colorKey)) {
const base = colorMap.get(colorKey), s = Math.min(Math.pow(1/0.7, depth), 240 / Math.max(base.r, base.g, base.b)); const color = colors instanceof Map ? (colors.get(colorKey) || colors.get("DEFAULT")) : cycleColors(colors, colorMap.size);
const fillColor = d3.rgb(base.r*s, base.g*s, base.b*s).toString(); colorMap.set(colorKey, d3.rgb(color));
}
const fillColor = colorMap.get(colorKey).brighter(0.3*depth).toString();
const label = parseColors(e.name).map(({ color, st }) => ({ color, st, width:ctx.measureText(st).width })); const label = parseColors(e.name).map(({ color, st }) => ({ color, st, width:ctx.measureText(st).width }));
let shapeRef = e.ref; let shapeRef = e.ref;
if (shapeRef != null) { ref = {ctx:e.ref, step:0}; shapeRef = ref; } if (shapeRef != null) { ref = {ctx:e.ref, step:0}; shapeRef = ref; }
@@ -277,10 +285,11 @@ async function renderProfiler(path, unit) {
// tiny device events go straight to the rewrite rule // tiny device events go straight to the rewrite rule
const key = k.startsWith("TINY") ? null : `${k}-${j}`; const key = k.startsWith("TINY") ? null : `${k}-${j}`;
if (key != null) shapeMetadata.set(key, html.node()); if (key != null) shapeMetadata.set(key, html.node());
const arg = { tooltipText:colored(e.name).outerHTML+"\n"+formatTime(e.dur)+(e.info != null ? "\n"+e.info : ""), key, ...shapeRef }; const arg = { tooltipText:colored(label).outerHTML+"\n"+formatTime(e.dur)+(e.info != null ? "\n"+e.info : ""), key,
ctx:shapeRef?.ctx, step:shapeRef?.step };
if (e.key != null) shapeMap.set(e.key, arg); if (e.key != null) shapeMap.set(e.key, arg);
// offset y by depth // offset y by depth
shapes.push({x:e.st, y:levelHeight*depth, width:e.dur, height:levelHeight, arg, label, fillColor }); shapes.push({x:e.st, y:levelHeight*depth, width:e.dur, height:levelHeight, arg, label:opts.hideLabels ? null : label, fillColor });
} }
div.style("height", levelHeight*levels.length+padding+"px").style("pointerEvents", "none"); div.style("height", levelHeight*levels.length+padding+"px").style("pointerEvents", "none");
} else { } else {
@@ -366,7 +375,8 @@ async function renderProfiler(path, unit) {
sum.x.push(allX[i], allX[i+1]); sum.x.push(allX[i], allX[i+1]);
const y = maxY.get(allX[i]); sum.y1.push(y, y); sum.y0.push(base0, base0); const y = maxY.get(allX[i]); sum.y1.push(y, y); sum.y0.push(base0, base0);
} }
data.tracks.set(k, { shapes:[sum], visible, offsetY, pcolor:"#c9a8ff", height, peak, scaleFactor:maxheight*4/height, views:[[sum], shapes], valueMap }); data.tracks.set(k, { shapes:[sum], eventType, visible, offsetY, pcolor:"#c9a8ff", height, peak, scaleFactor:maxheight*4/height,
views:[[sum], shapes], valueMap });
div.style("height", height+padding+"px").style("cursor", "pointer").on("click", (e) => { div.style("height", height+padding+"px").style("cursor", "pointer").on("click", (e) => {
const newFocus = e.currentTarget.id === focusedDevice ? null : e.currentTarget.id; const newFocus = e.currentTarget.id === focusedDevice ? null : e.currentTarget.id;
let offset = 0; let offset = 0;
@@ -396,11 +406,11 @@ async function renderProfiler(path, unit) {
xscale.domain(visibleX); xscale.domain(visibleX);
// draw shapes // draw shapes
const paths = []; const paths = [];
for (const [_, { offsetY, shapes, visible, valueMap, pcolor }] of data.tracks) { for (const [_, { shapes, eventType, visible, offsetY, valueMap, pcolor }] of data.tracks) {
visible.length = 0; visible.length = 0;
for (const e of shapes) { for (const e of shapes) {
const p = new Path2D(); const p = new Path2D();
if (e.width == null) { // generic polygon if (eventType === EventTypes.BUF) { // generic polygon
if (e.x[0]>et || e.x.at(-1)<st) continue; if (e.x[0]>et || e.x.at(-1)<st) continue;
const x = e.x.map(xscale); const x = e.x.map(xscale);
p.moveTo(x[0], offsetY+e.y0[0]); p.moveTo(x[0], offsetY+e.y0[0]);
@@ -465,7 +475,7 @@ async function renderProfiler(path, unit) {
drawLine(ctx, [x, x], [0, canvas.clientHeight], { color:m.color }); drawLine(ctx, [x, x], [0, canvas.clientHeight], { color:m.color });
ctx.fillText(m.name, x+2, 1); ctx.fillText(m.name, x+2, 1);
} }
for (const [p, color] of paths) { ctx.lineWidth = 1.4; ctx.strokeStyle = color; ctx.stroke(p); } for (const [p, color] of paths) { ctx.strokeStyle = color; ctx.stroke(p); }
} }
function resize() { function resize() {
@@ -489,16 +499,16 @@ async function renderProfiler(path, unit) {
new ResizeObserver(([e]) => e.contentRect.width > 0 && resize()).observe(profiler.node()); new ResizeObserver(([e]) => e.contentRect.width > 0 && resize()).observe(profiler.node());
function findRectAtPosition(x, y) { function findRectAtPosition(x, y) {
let tid = null; let track = null;
for (const k of data.tracks.keys()) { for (const k of data.tracks.keys()) {
const r = rect(document.getElementById(k)); const r = rect(document.getElementById(k));
if (y >= r.y && y <= r.y+r.height) { tid = k; break; } if (y >= r.y && y <= r.y+r.height) { track = data.tracks.get(k); break; }
} }
if (tid == null) return; if (track == null) return;
const { top, left, width, height } = rect(canvas); const R = rect(canvas);
const X = ((x-left) * (canvas.width/width))/dpr; const X = ((x-R.left) * (canvas.width/R.width))/dpr;
const Y = ((y-top) * (canvas.height/height))/dpr; const Y = ((y-R.top) * (canvas.height/R.height))/dpr;
for (const r of data.tracks.get(tid).visible) { for (const r of track.visible) {
if (Y>=r.y0 && Y<=r.y1 && X>=r.x0 && X<=r.x1) return r.arg; if (Y>=r.y0 && Y<=r.y1 && X>=r.x0 && X<=r.x1) return r.arg;
} }
} }
@@ -591,6 +601,11 @@ hljs.registerLanguage("cpp", (hljs) => ({
contains: [{ begin: '\\b(?:float|half)[0-9]+\\b', className: 'type' }, ...hljs.getLanguage('cpp').contains] contains: [{ begin: '\\b(?:float|half)[0-9]+\\b', className: 'type' }, ...hljs.getLanguage('cpp').contains]
})); }));
async function fetchValue(path) {
const res = await fetch(path);
return (await (res.headers.get("content-type") === "application/json" ? res.json() : res.arrayBuffer()));
}
var ret = []; var ret = [];
var cache = {}; var cache = {};
var ctxs = null; var ctxs = null;
@@ -648,7 +663,7 @@ async function main() {
// ** left sidebar context list // ** left sidebar context list
if (ctxs == null) { if (ctxs == null) {
ctxs = [{ name:"Profiler", steps:[] }]; ctxs = [{ name:"Profiler", steps:[] }];
for (const r of (await (await fetch("/ctxs")).json())) ctxs.push(r); for (const r of await fetchValue("/ctxs")) ctxs.push(r);
const ctxList = document.querySelector(".ctx-list"); const ctxList = document.querySelector(".ctx-list");
for (const [i,{name, steps}] of ctxs.entries()) { for (const [i,{name, steps}] of ctxs.entries()) {
const ul = ctxList.appendChild(document.createElement("ul")); const ul = ctxList.appendChild(document.createElement("ul"));
@@ -663,7 +678,8 @@ async function main() {
while (stack.length && stack.at(-1).depth >= u.depth) stack.pop(); while (stack.length && stack.at(-1).depth >= u.depth) stack.pop();
const list = stack.length > 0 ? stack.at(-1).li : ul; const list = stack.length > 0 ? stack.at(-1).li : ul;
u.li = list.appendChild(document.createElement("ul")); u.li = list.appendChild(document.createElement("ul"));
u.li.id = `step-${i}-${j}`; u.li.id = `step-${i}-${j}`
u.li.style.marginLeft = u.depth > 0 ? "calc(6px + 1ch)" : "6px";
const p = u.li.appendChild(document.createElement("p")); const p = u.li.appendChild(document.createElement("p"));
p.appendChild(colored(`${u.name}`+(u.match_count ? ` - ${u.match_count}` : ''))); p.appendChild(colored(`${u.name}`+(u.match_count ? ` - ${u.match_count}` : '')));
p.onclick = (e) => { p.onclick = (e) => {
@@ -694,15 +710,19 @@ async function main() {
if (url.pathname+url.search !== ckey) e.close(); if (url.pathname+url.search !== ckey) e.close();
else if (e.readyState === EventSource.OPEN) activeSrc = e; else if (e.readyState === EventSource.OPEN) activeSrc = e;
} }
if (ctx.name === "Profiler") return renderProfiler("/get_profile", "realtime"); if (ctx.name === "Profiler") return renderProfiler("/get_profile", "realtime", { width:"132px" });
if (workerUrl == null) await initWorker(); if (workerUrl == null) await initWorker();
if (ckey in cache) { if (ckey in cache) {
ret = cache[ckey]; ret = cache[ckey];
} }
// ** Disassembly view // ** Disassembly view
if (ckey.startsWith("/render")) { if (!ckey.startsWith("/rewrites")) {
if (step.fmt === "timeline") return renderProfiler(ckey, "clk"); // cycles on the x axis if (!(ckey in cache)) cache[ckey] = ret = await fetchValue(ckey);
if (!(ckey in cache)) cache[ckey] = ret = await (await fetch(ckey)).json(); // cycles on the x axis
if (ret instanceof ArrayBuffer) {
opts = {heightScale:0.5, hideLabels:true, levelKey:(e) => parseInt(e.name.split(" ")[1].split(":")[1])};
return renderProfiler(ckey, "clk", opts);
}
displaySelection("#custom"); displaySelection("#custom");
metadata.innerHTML = ""; metadata.innerHTML = "";
const root = d3.create("div").classed("raw-text", true).node(); const root = d3.create("div").classed("raw-text", true).node();
+88 -60
View File
@@ -1,13 +1,13 @@
#!/usr/bin/env python3 #!/usr/bin/env python3
import multiprocessing, pickle, difflib, os, threading, json, time, sys, webbrowser, socket, argparse, socketserver, functools, codecs, io, struct import multiprocessing, pickle, difflib, os, threading, json, time, sys, webbrowser, socket, argparse, socketserver, functools, codecs, io, struct
import subprocess, ctypes, pathlib, traceback import ctypes, pathlib, traceback, itertools
from contextlib import redirect_stdout, redirect_stderr from contextlib import redirect_stdout, redirect_stderr
from decimal import Decimal from decimal import Decimal
from http.server import BaseHTTPRequestHandler from http.server import BaseHTTPRequestHandler
from urllib.parse import parse_qs, urlparse from urllib.parse import parse_qs, urlparse
from typing import Any, TypedDict, TypeVar, Generator, Callable from typing import Any, TypedDict, TypeVar, Generator, Callable
from tinygrad.helpers import colored, getenv, tqdm, unwrap, word_wrap, TRACEMETA, ProfileEvent, ProfileRangeEvent, TracingKey, ProfilePointEvent, temp from tinygrad.helpers import colored, getenv, tqdm, unwrap, word_wrap, TRACEMETA, ProfileEvent, ProfileRangeEvent, TracingKey, ProfilePointEvent, temp
from tinygrad.helpers import printable from tinygrad.helpers import printable, system
from tinygrad.uop.ops import TrackedGraphRewrite, RewriteTrace, UOp, Ops, GroupOp, srender, sint, sym_infer, range_str, pyrender from tinygrad.uop.ops import TrackedGraphRewrite, RewriteTrace, UOp, Ops, GroupOp, srender, sint, sym_infer, range_str, pyrender
from tinygrad.uop.ops import print_uops, range_start, multirange_str from tinygrad.uop.ops import print_uops, range_start, multirange_str
from tinygrad.device import ProfileDeviceEvent, ProfileGraphEvent, ProfileGraphEntry, Device from tinygrad.device import ProfileDeviceEvent, ProfileGraphEvent, ProfileGraphEntry, Device
@@ -25,18 +25,25 @@ uops_colors = {Ops.LOAD: "#ffc0c0", Ops.STORE: "#87CEEB", Ops.CONST: "#e0e0e0",
# VIZ API # VIZ API
# A step is a lightweight descriptor for a trace entry
# Includes a name, metadata and a URL path for fetching the full data
def create_step(name:str, query:tuple[str, int, int], data=None, depth:int=0, **kwargs) -> dict:
return {"name":name, "query":f"{query[0]}?ctx={query[1]}&step={query[2]}", "data":data, "depth":depth, **kwargs}
# ** list all saved rewrites # ** list all saved rewrites
ref_map:dict[Any, int] = {} ref_map:dict[Any, int] = {}
def get_rewrites(t:RewriteTrace) -> list[dict]: def get_rewrites(t:RewriteTrace) -> list[dict]:
ret = [] ret = []
for i,(k,v) in enumerate(zip(t.keys, t.rewrites)): for i,(k,v) in enumerate(zip(t.keys, t.rewrites)):
steps = [{"name":s.name, "loc":s.loc, "match_count":len(s.matches), "code_line":printable(s.loc), "trace":k.tb if j == 0 else None, steps = [create_step(s.name, ("/rewrites", i, j), loc=s.loc, match_count=len(s.matches), code_line=printable(s.loc), trace=k.tb if j==0 else None,
"query":f"/ctxs?ctx={i}&idx={j}", "depth":s.depth} for j,s in enumerate(v)] depth=s.depth) for j,s in enumerate(v)]
if isinstance(k.ret, ProgramSpec): if isinstance(k.ret, ProgramSpec):
steps.append({"name":"View UOp List", "query":f"/render?ctx={i}&fmt=uops", "depth":0}) steps.append(create_step("View UOp List", ("/uops", i, len(steps)), k.ret))
steps.append({"name":"View Program", "query":f"/render?ctx={i}&fmt=src", "depth":0}) steps.append(create_step("View Program", ("/code", i, len(steps)), k.ret))
steps.append({"name":"View Disassembly", "query":f"/render?ctx={i}&fmt=asm", "depth":0}) steps.append(create_step("View Disassembly", ("/asm", i, len(steps)), k.ret))
for key in k.keys: ref_map[key] = i for key in k.keys: ref_map[key] = i
ret.append({"name":k.display_name, "steps":steps}) ret.append({"name":k.display_name, "steps":steps})
return ret return ret
@@ -203,48 +210,51 @@ def mem_layout(dev_events:list[tuple[int, int, float, DevEvent]], start_ts:int,
peaks.append(peak) peaks.append(peak)
return struct.pack("<BIQ", 1, len(events), peak)+b"".join(events) if events else None return struct.pack("<BIQ", 1, len(events), peak)+b"".join(events) if events else None
def err(name:str, msg:str|None=None) -> None:
ctxs.append({"name":"ERR", "steps":[create_step(name, ("render",len(ctxs),0), {"src":msg or traceback.format_exc()})]})
def row_tuple(row:str) -> tuple[int, ...]: return tuple(int(x.split(":")[1]) for x in row.split())
def load_sqtt(profile:list[ProfileEvent]) -> None: def load_sqtt(profile:list[ProfileEvent]) -> None:
from tinygrad.runtime.ops_amd import ProfileSQTTEvent from tinygrad.runtime.ops_amd import ProfileSQTTEvent
if not (sqtt_events:=[e for e in profile if isinstance(e, ProfileSQTTEvent)]): return None if not (sqtt_events:=[e for e in profile if isinstance(e, ProfileSQTTEvent)]): return None
def err(name:str, msg:str|None=None) -> None:
step = {"name":name, "data":{"src":msg or traceback.format_exc()}, "depth":0, "query":f"/render?ctx={len(ctxs)}&step=0&fmt=counters"}
return ctxs.append({"name":"Counters", "steps":[step]})
try: from extra.sqtt.roc import decode try: from extra.sqtt.roc import decode
except Exception: return err("DECODER IMPORT ISSUE") except Exception: return err("DECODER IMPORT ISSUE")
try: rctx = decode(profile) try: rctx = decode(profile)
except Exception: return err("DECODER ERROR") except Exception: return err("DECODER ERROR")
if not rctx.inst_execs: return err("EMPTY SQTT OUTPUT", f"{len(sqtt_events)} SQTT events recorded, none got decoded") if getenv("SQTT_PARSE"):
from extra.sqtt.attempt_sqtt_parse import parse_sqtt_print_packets
for e in sqtt_events: parse_sqtt_print_packets(e.blob)
if not any([rctx.inst_execs, rctx.occ_events]): return err("EMPTY SQTT OUTPUT", f"{len(sqtt_events)} SQTT events recorded, none got decoded")
steps:list[dict] = [] steps:list[dict] = []
units:set[str] = set() for name,disasm in rctx.disasms.items():
for name,waves in rctx.inst_execs.items():
events:list[ProfileEvent] = [] events:list[ProfileEvent] = []
prg = trace.keys[r].ret if (r:=ref_map.get(name)) else None # wave instruction events
steps.append(first:={"name":prg.name if prg is not None else name, "query":f"/render?ctx={len(ctxs)}&step={len(steps)}&fmt=counters", wave_insts:dict[str, dict] = {}
"depth":0, "fmt":"timeline"}) inst_units:dict[str, itertools.count] = {}
for w in rctx.inst_execs.get(name, []):
# Idle: The total time gap between the completion of previous instruction and the beginning of the current instruction. if (u:=w.wave_loc) not in inst_units: inst_units[u] = itertools.count(0)
# The idle time can be caused by: n = next(inst_units[u])
# * Arbiter loss events.append(ProfileRangeEvent(w.simd_loc, f"INST WAVE:{w.wave_id} N:{n}", Decimal(w.begin_time), Decimal(w.end_time)))
# * Source or destination register dependency wave_insts[f"{u} N:{n}"] = {"wave":w, "disasm":disasm, "run_number":n}
# * Instruction cache miss # occupancy events
# Stall: The total number of cycles the hardware pipe couldn't issue an instruction. units:dict[str, itertools.count] = {}
# Duration: Total latency in cycles, defined as "Stall time + Issue time" for gfx9 or "Stall time + Execute time" for gfx10+. wave_start:dict[str, int] = {}
for w in waves: for occ in rctx.occ_events[name]:
units.add(row:=f"SIMD:{w.simd} CU:{w.cu} SE:{w.se}") if (u:=occ.wave_loc) not in units: units[u] = itertools.count(0)
events.append(ProfileRangeEvent(row, wave_name:=f"wave {w.wave_id}", Decimal(w.begin_time), Decimal(w.end_time))) if u in inst_units: continue
rows, prev_instr = [], w.begin_time if occ.start: wave_start[u] = occ.time
for i,e in enumerate(w.insts): else: events.append(ProfileRangeEvent(occ.simd_loc, f"OCC WAVE:{occ.wave_id} N:{next(units[u])}", Decimal(wave_start.pop(u)),Decimal(occ.time)))
rows.append((e.inst, e.time, max(0, e.time-prev_instr), e.dur, e.stall, str(e.typ).split("_")[-1])) if not events: continue
prev_instr = max(prev_instr, e.time + e.dur) # gather and sort all sqtt events for this kernel
summary = [{"label":"Total Cycles", "value":w.end_time-w.begin_time}, {"label":"SIMD", "value":w.simd}, {"label":"CU", "value":w.cu},
{"label":"SE", "value":w.se}]
steps.append({"name":wave_name, "depth":1, "query":f"/render?ctx={len(ctxs)}&step={len(steps)}&fmt=counters",
"data":{"rows":rows, "cols":["Instruction", "Clk", "Idle", "Duration", "Stall", "Type"], "summary":summary}})
events = [ProfilePointEvent(unit, "start", unit, ts=Decimal(0)) for unit in units]+events events = [ProfilePointEvent(unit, "start", unit, ts=Decimal(0)) for unit in units]+events
first["data"] = {"value":get_profile(events), "content_type":"application/octet-stream"} kernel = trace.keys[r].ret if (r:=ref_map.get(name)) else None
steps.append(create_step(kernel.name if kernel is not None else name, ("/counters", len(ctxs), len(steps)),
{"value":get_profile(events, sort_fn=row_tuple), "content_type":"application/octet-stream"}, depth=1))
for k in sorted(wave_insts, key=row_tuple): steps.append(create_step(k, ("/sqtt-insts", len(ctxs), len(steps)), wave_insts[k], depth=2))
ctxs.append({"name":"Counters", "steps":steps}) ctxs.append({"name":"Counters", "steps":steps})
def get_profile(profile:list[ProfileEvent]) -> bytes|None: def get_profile(profile:list[ProfileEvent], sort_fn:Callable[[str], Any]|None=None) -> bytes|None:
# start by getting the time diffs # start by getting the time diffs
for ev in profile: for ev in profile:
if isinstance(ev,ProfileDeviceEvent): device_ts_diffs[ev.device] = (ev.comp_tdiff, ev.copy_tdiff if ev.copy_tdiff is not None else ev.comp_tdiff) if isinstance(ev,ProfileDeviceEvent): device_ts_diffs[ev.device] = (ev.comp_tdiff, ev.copy_tdiff if ev.copy_tdiff is not None else ev.comp_tdiff)
@@ -270,11 +280,11 @@ def get_profile(profile:list[ProfileEvent]) -> bytes|None:
scache:dict[str, int] = {} scache:dict[str, int] = {}
peaks:list[int] = [] peaks:list[int] = []
dtype_size:dict[str, int] = {} dtype_size:dict[str, int] = {}
for k,v in dev_events.items(): for k in sorted(dev_events, key=sort_fn) if sort_fn else dev_events:
v.sort(key=lambda e:e[0]) (v:=dev_events[k]).sort(key=lambda e:e[0])
layout[k] = timeline_layout(v, start_ts, scache) layout[k] = timeline_layout(v, start_ts, scache)
layout[f"{k} Memory"] = mem_layout(v, start_ts, unwrap(end_ts), peaks, dtype_size, scache) layout[f"{k} Memory"] = mem_layout(v, start_ts, unwrap(end_ts), peaks, dtype_size, scache)
groups = sorted(layout.items(), key=lambda x: '' if len(ss:=x[0].split(" ")) == 1 else ss[1]) groups = layout.items() if sort_fn is not None else sorted(layout.items(), key=lambda x: '' if len(ss:=x[0].split(" ")) == 1 else ss[1])
ret = [b"".join([struct.pack("<B", len(k)), k.encode(), v]) for k,v in groups if v is not None] ret = [b"".join([struct.pack("<B", len(k)), k.encode(), v]) for k,v in groups if v is not None]
index = json.dumps({"strings":list(scache), "dtypeSize":dtype_size, "markers":[{"ts":int(e.ts-start_ts), **e.arg} for e in markers]}).encode() index = json.dumps({"strings":list(scache), "dtypeSize":dtype_size, "markers":[{"ts":int(e.ts-start_ts), **e.arg} for e in markers]}).encode()
return struct.pack("<IQII", unwrap(end_ts)-start_ts, max(peaks,default=0), len(index), len(ret))+index+b"".join(ret) return struct.pack("<IQII", unwrap(end_ts)-start_ts, max(peaks,default=0), len(index), len(ret))+index+b"".join(ret)
@@ -282,9 +292,9 @@ def get_profile(profile:list[ProfileEvent]) -> bytes|None:
# ** Assembly analyzers # ** Assembly analyzers
def get_llvm_mca(asm:str, mtriple:str, mcpu:str) -> dict: def get_llvm_mca(asm:str, mtriple:str, mcpu:str) -> dict:
target_args = [f"-mtriple={mtriple}", f"-mcpu={mcpu}"] target_args = f"-mtriple={mtriple} -mcpu={mcpu}"
# disassembly output can include headers / metadata, skip if llvm-mca can't parse those lines # disassembly output can include headers / metadata, skip if llvm-mca can't parse those lines
data = json.loads(subprocess.check_output(["llvm-mca","-skip-unsupported-instructions=parse-failure","--json","-"]+target_args, input=asm.encode())) data = json.loads(system("llvm-mca -skip-unsupported-instructions=parse-failure --json -"+target_args, input=asm.encode()))
cr = data["CodeRegions"][0] cr = data["CodeRegions"][0]
resource_labels = [repr(x)[1:-1] for x in data["TargetInfo"]["Resources"]] resource_labels = [repr(x)[1:-1] for x in data["TargetInfo"]["Resources"]]
rows:list = [[instr] for instr in cr["Instructions"]] rows:list = [[instr] for instr in cr["Instructions"]]
@@ -309,19 +319,36 @@ def get_stdout(f: Callable) -> str:
return buf.getvalue() return buf.getvalue()
def get_render(i:int, j:int, fmt:str) -> dict: def get_render(i:int, j:int, fmt:str) -> dict:
if fmt == "counters": return ctxs[i]["steps"][j]["data"] data = ctxs[i]["steps"][j]["data"]
if not isinstance(prg:=trace.keys[i].ret, ProgramSpec): return {} if fmt == "uops": return {"src":get_stdout(lambda: print_uops(data.uops or [])), "lang":"txt"}
if fmt == "uops": return {"src":get_stdout(lambda: print_uops(prg.uops or [])), "lang":"txt"} if fmt == "code": return {"src":data.src, "lang":"cpp"}
if fmt == "src": return {"src":prg.src, "lang":"cpp"} if fmt == "asm":
compiler = Device[prg.device].compiler compiler = Device[data.device].compiler
disasm_str = get_stdout(lambda: compiler.disassemble(compiler.compile(prg.src))) disasm_str = get_stdout(lambda: compiler.disassemble(compiler.compile(data.src)))
from tinygrad.runtime.support.compiler_cpu import llvm, LLVMCompiler from tinygrad.runtime.support.compiler_cpu import llvm, LLVMCompiler
if isinstance(compiler, LLVMCompiler): if isinstance(compiler, LLVMCompiler):
mtriple = ctypes.string_at(llvm.LLVMGetTargetMachineTriple(tm:=compiler.target_machine)).decode() return get_llvm_mca(disasm_str, ctypes.string_at(llvm.LLVMGetTargetMachineTriple(tm:=compiler.target_machine)).decode(),
mcpu = ctypes.string_at(llvm.LLVMGetTargetMachineCPU(tm)).decode() ctypes.string_at(llvm.LLVMGetTargetMachineCPU(tm)).decode())
ret = get_llvm_mca(disasm_str, mtriple, mcpu) return {"src":disasm_str, "lang":"x86asm"}
else: ret = {"src":disasm_str, "lang":"x86asm"} if fmt == "sqtt-insts":
return ret columns = ["Instruction", "Clk", "Idle", "Duration", "Stall", "Type"]
# Idle: The total time gap between the completion of previous instruction and the beginning of the current instruction.
# The idle time can be caused by:
# * Arbiter loss
# * Source or destination register dependency
# * Instruction cache miss
# Stall: The total number of cycles the hardware pipe couldn't issue an instruction.
# Duration: Total latency in cycles, defined as "Stall time + Issue time" for gfx9 or "Stall time + Execute time" for gfx10+.
prev_instr = (w:=data["wave"]).begin_time
pc_to_inst = data["disasm"]
rows:list[tuple] = []
for e in w.unpack_insts():
rows.append((pc_to_inst[e.pc][0], e.time, max(0, e.time-prev_instr), e.dur, e.stall, str(e.typ).split("_")[-1]))
prev_instr = max(prev_instr, e.time + e.dur)
summary = [{"label":"Total Cycles", "value":w.end_time-w.begin_time}, {"label":"SE", "value":w.se}, {"label":"CU", "value":w.cu},
{"label":"SIMD", "value":w.simd}, {"label":"Wave ID", "value":w.wave_id}, {"label":"Run number", "value":data["run_number"]}]
return {"rows":rows, "cols":columns, "summary":summary}
return data
# ** HTTP server # ** HTTP server
@@ -340,13 +367,14 @@ class Handler(BaseHTTPRequestHandler):
if url.path.endswith(".css"): content_type = "text/css" if url.path.endswith(".css"): content_type = "text/css"
except FileNotFoundError: status_code = 404 except FileNotFoundError: status_code = 404
elif (query:=parse_qs(url.query)): elif (query:=parse_qs(url.query)):
if url.path == "/render": i, j = get_int(query, "ctx"), get_int(query, "step")
render_src = get_render(get_int(query, "ctx"), get_int(query, "step"), query["fmt"][0]) if (fmt:=url.path.lstrip("/")) == "rewrites":
try: return self.stream_json(get_full_rewrite(trace.rewrites[i][j], i))
except (KeyError, IndexError): status_code = 404
else:
render_src = get_render(i, j, fmt)
if "content_type" in render_src: ret, content_type = render_src["value"], render_src["content_type"] if "content_type" in render_src: ret, content_type = render_src["value"], render_src["content_type"]
else: ret, content_type = json.dumps(render_src).encode(), "application/json" else: ret, content_type = json.dumps(render_src).encode(), "application/json"
else:
try: return self.stream_json(get_full_rewrite(trace.rewrites[i:=get_int(query, "ctx")][get_int(query, "idx")], i))
except (KeyError, IndexError): status_code = 404
elif url.path == "/ctxs": elif url.path == "/ctxs":
lst = [{**c, "steps":[{k:v for k, v in s.items() if k != "data"} for s in c["steps"]]} for c in ctxs] lst = [{**c, "steps":[{k:v for k, v in s.items() if k != "data"} for s in c["steps"]]} for c in ctxs]
ret, content_type = json.dumps(lst).encode(), "application/json" ret, content_type = json.dumps(lst).encode(), "application/json"