forked from tinygrad/tinygrad
Compare commits
128
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
82aa943cd4 | ||
|
|
e16782cf9e | ||
|
|
1c47ee729e | ||
|
|
a8f9e69bd9 | ||
|
|
385618d45b | ||
|
|
ffff194e93 | ||
|
|
fba4535289 | ||
|
|
225eb1500f | ||
|
|
1a72ac16a6 | ||
|
|
79055ddb8b | ||
|
|
0c9fbf87e1 | ||
|
|
f2221130bb | ||
|
|
be72b78dcb | ||
|
|
e4fbde5b3b | ||
|
|
1a332afa76 | ||
|
|
a438c277de | ||
|
|
722e7a16ed | ||
|
|
1afa3c0877 | ||
|
|
46cb65e692 | ||
|
|
9c59b3d19e | ||
|
|
a647c9eca6 | ||
|
|
06e39a88a9 | ||
|
|
805de27e07 | ||
|
|
05294bc648 | ||
|
|
5623e765c8 | ||
|
|
331f70aa75 | ||
|
|
583560ab72 | ||
|
|
8e8e53c886 | ||
|
|
e4fead8a86 | ||
|
|
8894a5409d | ||
|
|
6d3385c284 | ||
|
|
b637093be9 | ||
|
|
98e9e73286 | ||
|
|
e7e1935225 | ||
|
|
33773fda87 | ||
|
|
e2cee64050 | ||
|
|
646372490c | ||
|
|
a37f221e44 | ||
|
|
f63ded5817 | ||
|
|
50a443f558 | ||
|
|
9bb17c53ea | ||
|
|
55be95da15 | ||
|
|
cabd4add48 | ||
|
|
13efdf8c31 | ||
|
|
295600dc5a | ||
|
|
a9ed241172 | ||
|
|
c70b06ec19 | ||
|
|
8f0e747b3a | ||
|
|
6372c95094 | ||
|
|
61625a3898 | ||
|
|
acbe6361ab | ||
|
|
ef42334239 | ||
|
|
e8844853ed | ||
|
|
5b823af696 | ||
|
|
df53c62a9f | ||
|
|
d37e1fe065 | ||
|
|
22c08b470c | ||
|
|
567066f51f | ||
|
|
6c5fa349e1 | ||
|
|
d1bb08c5a1 | ||
|
|
e5351699bd | ||
|
|
7c110e1a57 | ||
|
|
888aaab151 | ||
|
|
3e63831b98 | ||
|
|
2ee701a009 | ||
|
|
c80d459d99 | ||
|
|
14eb48b13a | ||
|
|
734bfa07b4 | ||
|
|
f72b1fbca4 | ||
|
|
84f065f2a2 | ||
|
|
44d84228ff | ||
|
|
09f3aae169 | ||
|
|
777cbec5b3 | ||
|
|
7eb0d8e744 | ||
|
|
ba84d415fe | ||
|
|
547304c471 | ||
|
|
4ada51618f | ||
|
|
6b1bae6614 | ||
|
|
3049f3edda | ||
|
|
3af231904e | ||
|
|
faf68c03a8 | ||
|
|
256f81bb02 | ||
|
|
7e0aaadecd | ||
|
|
6be86dde17 | ||
|
|
f9b7586e08 | ||
|
|
263b724143 | ||
|
|
5efa727b83 | ||
|
|
bcdfc109b5 | ||
|
|
006dea4c3e | ||
|
|
f9586b38ba | ||
|
|
7316da3253 | ||
|
|
17aa3379e9 | ||
|
|
4e5a9132e7 | ||
|
|
759557f633 | ||
|
|
3f939f3d3c | ||
|
|
f9851a852f | ||
|
|
fe2876a6d8 | ||
|
|
a23dea202b | ||
|
|
ab9fa964d8 | ||
|
|
be2e24cb25 | ||
|
|
8f1f195b6d | ||
|
|
9a53fcbde4 | ||
|
|
13f10a31dc | ||
|
|
8b26cf2b3d | ||
|
|
bc8e537423 | ||
|
|
af17e07251 | ||
|
|
7a6853fa40 | ||
|
|
82eb63d3ad | ||
|
|
fcd8d0751a | ||
|
|
74b9d33acb | ||
|
|
371c1f2355 | ||
|
|
41a098a82d | ||
|
|
222bb12ddf | ||
|
|
787f0070ed | ||
|
|
ece1415def | ||
|
|
2f0ea29b34 | ||
|
|
bc55bc4849 | ||
|
|
23b90945c3 | ||
|
|
c2075f3613 | ||
|
|
e59313da08 | ||
|
|
6fd7ce3832 | ||
|
|
8002921a04 | ||
|
|
f91e366a17 | ||
|
|
73497af4c0 | ||
|
|
a6360fd94d | ||
|
|
f3692b7406 | ||
|
|
22b8579234 | ||
|
|
58b7e4fab3 |
@@ -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('**/setup.py') }}-${{ env.PYTHON_CACHE_VERSION }}
|
key: venv-${{ runner.os }}-python-${{ steps.setup-python.outputs.python-version }}-${{ inputs.deps }}-${{ inputs.pydeps }}-${{ hashFiles('**/pyproject.toml') }}-${{ env.CACHE_VERSION }}
|
||||||
|
|
||||||
# **** Caching downloads ****
|
# **** Caching downloads ****
|
||||||
|
|
||||||
@@ -70,13 +70,13 @@ runs:
|
|||||||
uses: actions/cache@v4
|
uses: actions/cache@v4
|
||||||
with:
|
with:
|
||||||
path: ~/.cache/tinygrad/downloads/
|
path: ~/.cache/tinygrad/downloads/
|
||||||
key: downloads-cache-${{ inputs.key }}-${{ env.DOWNLOAD_CACHE_VERSION }}
|
key: downloads-cache-${{ inputs.key }}-${{ env.CACHE_VERSION }}
|
||||||
- name: Cache downloads (macOS)
|
- name: Cache downloads (macOS)
|
||||||
if: inputs.key != '' && runner.os == 'macOS'
|
if: inputs.key != '' && runner.os == 'macOS'
|
||||||
uses: actions/cache@v4
|
uses: actions/cache@v4
|
||||||
with:
|
with:
|
||||||
path: ~/Library/Caches/tinygrad/downloads/
|
path: ~/Library/Caches/tinygrad/downloads/
|
||||||
key: osx-downloads-cache-${{ inputs.key }}-${{ env.DOWNLOAD_CACHE_VERSION }}
|
key: osx-downloads-cache-${{ inputs.key }}-${{ env.CACHE_VERSION }}
|
||||||
|
|
||||||
# **** Python deps ****
|
# **** Python deps ****
|
||||||
|
|
||||||
@@ -187,7 +187,7 @@ runs:
|
|||||||
uses: actions/cache@v4
|
uses: actions/cache@v4
|
||||||
with:
|
with:
|
||||||
path: /var/cache/apt/archives/
|
path: /var/cache/apt/archives/
|
||||||
key: ${{ runner.os }}-apt-${{ steps.apt-pkgs.outputs.hash }}-${{ env.APT_CACHE_VERSION }}
|
key: ${{ runner.os }}-apt-${{ steps.apt-pkgs.outputs.hash }}-${{ env.CACHE_VERSION }}
|
||||||
|
|
||||||
- name: Run apt Update + Install
|
- name: Run apt Update + Install
|
||||||
if: runner.os == 'Linux' && (inputs.opencl == 'true' || inputs.amd == 'true' || inputs.cuda == 'true' || inputs.webgpu == 'true' || inputs.llvm == 'true')
|
if: runner.os == 'Linux' && (inputs.opencl == 'true' || inputs.amd == 'true' || inputs.cuda == 'true' || inputs.webgpu == 'true' || inputs.llvm == 'true')
|
||||||
@@ -247,7 +247,7 @@ runs:
|
|||||||
cache-name: cache-gpuocelot-build-1
|
cache-name: cache-gpuocelot-build-1
|
||||||
with:
|
with:
|
||||||
path: ${{ github.workspace }}/gpuocelot/ocelot
|
path: ${{ github.workspace }}/gpuocelot/ocelot
|
||||||
key: ${{ runner.os }}-gpuocelot-b16039dc940dc6bc4ea0a98380495769ff35ed99-rebuild-${{ env.BUILD_CACHE_VERSION }}
|
key: ${{ runner.os }}-gpuocelot-b16039dc940dc6bc4ea0a98380495769ff35ed99-rebuild-${{ env.CACHE_VERSION }}
|
||||||
- name: Clone/compile gpuocelot
|
- name: Clone/compile gpuocelot
|
||||||
if: inputs.ocelot == 'true' && steps.cache-build.outputs.cache-hit != 'true'
|
if: inputs.ocelot == 'true' && steps.cache-build.outputs.cache-hit != 'true'
|
||||||
shell: bash
|
shell: bash
|
||||||
|
|||||||
+122
-43
@@ -1,10 +1,7 @@
|
|||||||
name: Autogen
|
name: Autogen
|
||||||
env:
|
env:
|
||||||
# increment this when downloads substantially change to avoid the internet
|
# increment this when downloads substantially change to avoid the internet
|
||||||
DOWNLOAD_CACHE_VERSION: '12'
|
CACHE_VERSION: '13'
|
||||||
PYTHON_CACHE_VERSION: '4'
|
|
||||||
APT_CACHE_VERSION: '1'
|
|
||||||
BUILD_CACHE_VERSION: '1'
|
|
||||||
CAPTURE_PROCESS_REPLAY: 1
|
CAPTURE_PROCESS_REPLAY: 1
|
||||||
GH_TOKEN: ${{ secrets.GITHUB_TOKEN }}
|
GH_TOKEN: ${{ secrets.GITHUB_TOKEN }}
|
||||||
PYTHONPATH: ${{ github.workspace }}
|
PYTHONPATH: ${{ github.workspace }}
|
||||||
@@ -14,15 +11,15 @@ on:
|
|||||||
branches:
|
branches:
|
||||||
- master
|
- master
|
||||||
pull_request:
|
pull_request:
|
||||||
paths:
|
paths:
|
||||||
- 'tinygrad/runtime/autogen/**/*'
|
- 'tinygrad/runtime/autogen/**/*'
|
||||||
workflow_dispatch:
|
workflow_dispatch:
|
||||||
paths:
|
paths:
|
||||||
- 'tinygrad/runtime/autogen/**/*'
|
- 'tinygrad/runtime/autogen/**/*'
|
||||||
|
|
||||||
jobs:
|
jobs:
|
||||||
autogen:
|
autogen:
|
||||||
name: Autogen
|
name: In-tree Autogen
|
||||||
runs-on: ubuntu-24.04
|
runs-on: ubuntu-24.04
|
||||||
timeout-minutes: 15
|
timeout-minutes: 15
|
||||||
steps:
|
steps:
|
||||||
@@ -34,64 +31,146 @@ jobs:
|
|||||||
opencl: 'true'
|
opencl: 'true'
|
||||||
amd: 'true'
|
amd: 'true'
|
||||||
cuda: 'true'
|
cuda: 'true'
|
||||||
webgpu: 'true'
|
|
||||||
llvm: 'true'
|
llvm: 'true'
|
||||||
|
webgpu: 'true'
|
||||||
|
mesa: 'true'
|
||||||
pydeps: 'pyyaml mako'
|
pydeps: 'pyyaml mako'
|
||||||
- name: Install autogen support packages
|
- name: Install autogen support packages
|
||||||
run: sudo apt-get install -y --no-install-recommends llvm-14-dev libclang-14-dev llvm-20-dev
|
run: sudo apt-get install -y --no-install-recommends libclang-20-dev llvm-20-dev hip-dev libusb-1.0-0-dev
|
||||||
- name: Verify OpenCL autogen
|
- name: Verify OpenCL autogen
|
||||||
run: |
|
run: |
|
||||||
cp tinygrad/runtime/autogen/opencl.py /tmp/opencl.py.bak
|
mv tinygrad/runtime/autogen/opencl.py /tmp/opencl.py.bak
|
||||||
./autogen_stubs.sh opencl
|
python3 -c "from tinygrad.runtime.autogen import opencl"
|
||||||
diff /tmp/opencl.py.bak tinygrad/runtime/autogen/opencl.py
|
diff /tmp/opencl.py.bak tinygrad/runtime/autogen/opencl.py
|
||||||
- name: Verify CUDA autogen
|
- name: Verify CUDA autogen
|
||||||
run: |
|
run: |
|
||||||
cp tinygrad/runtime/autogen/cuda.py /tmp/cuda.py.bak
|
mv tinygrad/runtime/autogen/cuda.py /tmp/cuda.py.bak
|
||||||
cp tinygrad/runtime/autogen/nv_gpu.py /tmp/nv_gpu.py.bak
|
mv tinygrad/runtime/autogen/nvrtc.py /tmp/nvrtc.py.bak
|
||||||
./autogen_stubs.sh cuda
|
mv tinygrad/runtime/autogen/nvjitlink.py /tmp/nvjitlink.py.bak
|
||||||
./autogen_stubs.sh nv
|
mv tinygrad/runtime/autogen/nv_570.py /tmp/nv_570.py.bak
|
||||||
|
mv tinygrad/runtime/autogen/nv.py /tmp/nv.py.bak
|
||||||
|
python3 -c "from tinygrad.runtime.autogen import cuda, nvrtc, nvjitlink, nv_570, nv"
|
||||||
diff /tmp/cuda.py.bak tinygrad/runtime/autogen/cuda.py
|
diff /tmp/cuda.py.bak tinygrad/runtime/autogen/cuda.py
|
||||||
diff /tmp/nv_gpu.py.bak tinygrad/runtime/autogen/nv_gpu.py
|
diff /tmp/nvrtc.py.bak tinygrad/runtime/autogen/nvrtc.py
|
||||||
|
diff /tmp/nvjitlink.py.bak tinygrad/runtime/autogen/nvjitlink.py
|
||||||
|
diff /tmp/nv_570.py.bak tinygrad/runtime/autogen/nv_570.py
|
||||||
|
diff /tmp/nv.py.bak tinygrad/runtime/autogen/nv.py
|
||||||
- name: Verify AMD autogen
|
- name: Verify AMD autogen
|
||||||
run: |
|
run: |
|
||||||
cp tinygrad/runtime/autogen/hsa.py /tmp/hsa.py.bak
|
mv tinygrad/runtime/autogen/comgr.py /tmp/comgr.py.bak
|
||||||
cp tinygrad/runtime/autogen/kfd.py /tmp/kfd.py.bak
|
mv tinygrad/runtime/autogen/hsa.py /tmp/hsa.py.bak
|
||||||
cp tinygrad/runtime/autogen/comgr.py /tmp/comgr.py.bak
|
mv tinygrad/runtime/autogen/hip.py /tmp/hip.py.bak
|
||||||
cp tinygrad/runtime/autogen/amd_gpu.py /tmp/amd_gpu.py.bak
|
mv tinygrad/runtime/autogen/amd_gpu.py /tmp/amd_gpu.py.bak
|
||||||
cp tinygrad/runtime/autogen/sqtt.py /tmp/sqtt.py.bak
|
mv tinygrad/runtime/autogen/sqtt.py /tmp/sqtt.py.bak
|
||||||
./autogen_stubs.sh hsa
|
mv tinygrad/runtime/autogen/rocprof.py /tmp/rocprof.py.bak
|
||||||
./autogen_stubs.sh kfd
|
mv tinygrad/runtime/autogen/am/am.py /tmp/am_am.py.bak
|
||||||
./autogen_stubs.sh comgr
|
mv tinygrad/runtime/autogen/am/pm4_soc15.py /tmp/am_pm4_soc15.py.bak
|
||||||
./autogen_stubs.sh amd
|
mv tinygrad/runtime/autogen/am/pm4_nv.py /tmp/am_pm4_nv.py.bak
|
||||||
./autogen_stubs.sh sqtt
|
mv tinygrad/runtime/autogen/am/sdma_4_0_0.py /tmp/am_sdma_4_0_0.py.bak
|
||||||
diff /tmp/hsa.py.bak tinygrad/runtime/autogen/hsa.py
|
mv tinygrad/runtime/autogen/am/sdma_5_0_0.py /tmp/am_sdma_5_0_0.py.bak
|
||||||
diff /tmp/kfd.py.bak tinygrad/runtime/autogen/kfd.py
|
mv tinygrad/runtime/autogen/am/sdma_6_0_0.py /tmp/am_sdma_6_0_0.py.bak
|
||||||
|
mv tinygrad/runtime/autogen/am/smu_v13_0_0.py /tmp/am_smu_v13_0_0.py.bak
|
||||||
|
mv tinygrad/runtime/autogen/am/smu_v14_0_2.py /tmp/am_smu_v14_0_2.py.bak
|
||||||
|
python3 -c "from tinygrad.runtime.autogen import comgr, hsa, hip, amd_gpu, sqtt, rocprof; from tinygrad.runtime.autogen.am import am, pm4_soc15, pm4_nv, sdma_4_0_0, sdma_5_0_0, sdma_6_0_0, smu_v13_0_0, smu_v14_0_2"
|
||||||
diff /tmp/comgr.py.bak tinygrad/runtime/autogen/comgr.py
|
diff /tmp/comgr.py.bak tinygrad/runtime/autogen/comgr.py
|
||||||
|
diff /tmp/hsa.py.bak tinygrad/runtime/autogen/hsa.py
|
||||||
|
diff /tmp/hip.py.bak tinygrad/runtime/autogen/hip.py
|
||||||
diff /tmp/amd_gpu.py.bak tinygrad/runtime/autogen/amd_gpu.py
|
diff /tmp/amd_gpu.py.bak tinygrad/runtime/autogen/amd_gpu.py
|
||||||
diff /tmp/sqtt.py.bak tinygrad/runtime/autogen/sqtt.py
|
diff /tmp/sqtt.py.bak tinygrad/runtime/autogen/sqtt.py
|
||||||
|
diff /tmp/rocprof.py.bak tinygrad/runtime/autogen/rocprof.py
|
||||||
|
diff /tmp/am_am.py.bak tinygrad/runtime/autogen/am/am.py
|
||||||
|
diff /tmp/am_pm4_soc15.py.bak tinygrad/runtime/autogen/am/pm4_soc15.py
|
||||||
|
diff /tmp/am_pm4_nv.py.bak tinygrad/runtime/autogen/am/pm4_nv.py
|
||||||
|
diff /tmp/am_sdma_4_0_0.py.bak tinygrad/runtime/autogen/am/sdma_4_0_0.py
|
||||||
|
diff /tmp/am_sdma_5_0_0.py.bak tinygrad/runtime/autogen/am/sdma_5_0_0.py
|
||||||
|
diff /tmp/am_sdma_6_0_0.py.bak tinygrad/runtime/autogen/am/sdma_6_0_0.py
|
||||||
|
diff /tmp/am_smu_v13_0_0.py.bak tinygrad/runtime/autogen/am/smu_v13_0_0.py
|
||||||
|
diff /tmp/am_smu_v14_0_2.py.bak tinygrad/runtime/autogen/am/smu_v14_0_2.py
|
||||||
- name: Verify Linux autogen
|
- name: Verify Linux autogen
|
||||||
run: |
|
run: |
|
||||||
cp tinygrad/runtime/autogen/libc.py /tmp/libc.py.bak
|
mv tinygrad/runtime/autogen/libc.py /tmp/libc.py.bak
|
||||||
cp tinygrad/runtime/autogen/io_uring.py /tmp/io_uring.py.bak
|
mv tinygrad/runtime/autogen/kfd.py /tmp/kfd.py.bak
|
||||||
cp tinygrad/runtime/autogen/ib.py /tmp/ib.py.bak
|
mv tinygrad/runtime/autogen/io_uring.py /tmp/io_uring.py.bak
|
||||||
./autogen_stubs.sh libc
|
mv tinygrad/runtime/autogen/ib.py /tmp/ib.py.bak
|
||||||
./autogen_stubs.sh io_uring
|
mv tinygrad/runtime/autogen/pci.py /tmp/pci.py.bak
|
||||||
./autogen_stubs.sh ib
|
mv tinygrad/runtime/autogen/vfio.py /tmp/vfio.py.bak
|
||||||
|
python3 -c "from tinygrad.runtime.autogen import libc, kfd, io_uring, ib, pci, vfio"
|
||||||
diff /tmp/libc.py.bak tinygrad/runtime/autogen/libc.py
|
diff /tmp/libc.py.bak tinygrad/runtime/autogen/libc.py
|
||||||
|
diff /tmp/kfd.py.bak tinygrad/runtime/autogen/kfd.py
|
||||||
diff /tmp/io_uring.py.bak tinygrad/runtime/autogen/io_uring.py
|
diff /tmp/io_uring.py.bak tinygrad/runtime/autogen/io_uring.py
|
||||||
diff /tmp/ib.py.bak tinygrad/runtime/autogen/ib.py
|
diff /tmp/ib.py.bak tinygrad/runtime/autogen/ib.py
|
||||||
- name: Verify WebGPU autogen
|
diff /tmp/pci.py.bak tinygrad/runtime/autogen/pci.py
|
||||||
run: |
|
diff /tmp/vfio.py.bak tinygrad/runtime/autogen/vfio.py
|
||||||
cp tinygrad/runtime/autogen/webgpu.py /tmp/webgpu.py.bak
|
|
||||||
./autogen_stubs.sh webgpu
|
|
||||||
diff /tmp/webgpu.py.bak tinygrad/runtime/autogen/webgpu.py
|
|
||||||
- name: Verify LLVM autogen
|
- name: Verify LLVM autogen
|
||||||
run: |
|
run: |
|
||||||
cp tinygrad/runtime/autogen/llvm.py /tmp/llvm.py.bak
|
mv tinygrad/runtime/autogen/llvm.py /tmp/llvm.py.bak
|
||||||
./autogen_stubs.sh llvm
|
python3 -c "from tinygrad.runtime.autogen import llvm"
|
||||||
diff /tmp/llvm.py.bak tinygrad/runtime/autogen/llvm.py
|
diff /tmp/llvm.py.bak tinygrad/runtime/autogen/llvm.py
|
||||||
|
- name: Verify WebGPU autogen
|
||||||
|
run: |
|
||||||
|
mv tinygrad/runtime/autogen/webgpu.py /tmp/webgpu.py.bak
|
||||||
|
python3 -c "from tinygrad.runtime.autogen import webgpu"
|
||||||
|
diff /tmp/webgpu.py.bak tinygrad/runtime/autogen/webgpu.py
|
||||||
|
- name: Verify Qualcomm autogen
|
||||||
|
run: |
|
||||||
|
mv tinygrad/runtime/autogen/kgsl.py /tmp/kgsl.py.bak
|
||||||
|
mv tinygrad/runtime/autogen/adreno.py /tmp/adreno.py.bak
|
||||||
|
mv tinygrad/runtime/autogen/qcom_dsp.py /tmp/qcom_dsp.py.bak
|
||||||
|
python3 -c "from tinygrad.runtime.autogen import kgsl, adreno, qcom_dsp"
|
||||||
|
diff /tmp/kgsl.py.bak tinygrad/runtime/autogen/kgsl.py
|
||||||
|
diff /tmp/adreno.py.bak tinygrad/runtime/autogen/adreno.py
|
||||||
|
diff /tmp/qcom_dsp.py.bak tinygrad/runtime/autogen/qcom_dsp.py
|
||||||
|
- name: Verify libusb autogen
|
||||||
|
run: |
|
||||||
|
mv tinygrad/runtime/autogen/libusb.py /tmp/libusb.py.bak
|
||||||
|
python3 -c "from tinygrad.runtime.autogen import libusb"
|
||||||
|
diff /tmp/libusb.py.bak tinygrad/runtime/autogen/libusb.py
|
||||||
- name: Verify mesa autogen
|
- name: Verify mesa autogen
|
||||||
run: |
|
run: |
|
||||||
cp tinygrad/runtime/autogen/mesa.py /tmp/mesa.py.bak
|
mv tinygrad/runtime/autogen/mesa.py /tmp/mesa.py.bak
|
||||||
./autogen_stubs.sh mesa
|
python3 -c "from tinygrad.runtime.autogen import mesa"
|
||||||
diff /tmp/mesa.py.bak tinygrad/runtime/autogen/mesa.py
|
diff /tmp/mesa.py.bak tinygrad/runtime/autogen/mesa.py
|
||||||
|
- name: Verify libclang autogen
|
||||||
|
run: |
|
||||||
|
cp tinygrad/runtime/autogen/libclang.py /tmp/libclang.py.bak
|
||||||
|
REGEN=1 python3 -c "from tinygrad.runtime.autogen import libclang"
|
||||||
|
diff /tmp/libclang.py.bak tinygrad/runtime/autogen/libclang.py
|
||||||
|
autogen-mac:
|
||||||
|
name: In-tree Autogen (macos)
|
||||||
|
runs-on: macos-14
|
||||||
|
timeout-minutes: 15
|
||||||
|
steps:
|
||||||
|
- name: Checkout Code
|
||||||
|
uses: actions/checkout@v4
|
||||||
|
- name: Setup Environment
|
||||||
|
uses: ./.github/actions/setup-tinygrad
|
||||||
|
with:
|
||||||
|
llvm: 'true'
|
||||||
|
- name: Verify macos autogen
|
||||||
|
run: |
|
||||||
|
mv tinygrad/runtime/autogen/metal.py /tmp/metal.py.bak
|
||||||
|
LIBCLANG_PATH=/opt/homebrew/opt/llvm@20/lib/libclang.dylib python3 -c "from tinygrad.runtime.autogen import metal"
|
||||||
|
diff /tmp/metal.py.bak tinygrad/runtime/autogen/metal.py
|
||||||
|
autogen-comgr-3:
|
||||||
|
name: In-tree Autogen (comgr 3)
|
||||||
|
runs-on: ubuntu-24.04
|
||||||
|
timeout-minutes: 15
|
||||||
|
steps:
|
||||||
|
- name: Checkout Code
|
||||||
|
uses: actions/checkout@v4
|
||||||
|
- name: Setup Environment
|
||||||
|
uses: ./.github/actions/setup-tinygrad
|
||||||
|
- name: Install autogen support packages
|
||||||
|
run: |
|
||||||
|
wget https://repo.radeon.com/rocm/rocm.gpg.key -O - | gpg --dearmor | sudo tee /etc/apt/keyrings/rocm.gpg > /dev/null
|
||||||
|
sudo tee /etc/apt/sources.list.d/rocm.list <<EOF
|
||||||
|
deb [arch=amd64 signed-by=/etc/apt/keyrings/rocm.gpg] https://repo.radeon.com/rocm/apt/6.4 $(lsb_release -cs) main
|
||||||
|
EOF
|
||||||
|
echo -e 'Package: *\nPin: release o=repo.radeon.com\nPin-Priority: 600' | sudo tee /etc/apt/preferences.d/rocm-pin-600
|
||||||
|
sudo apt -qq update || true
|
||||||
|
sudo apt-get install -y --no-install-recommends libclang-20-dev comgr
|
||||||
|
- name: Verify comgr (3) autogen
|
||||||
|
run: |
|
||||||
|
mv tinygrad/runtime/autogen/comgr_3.py /tmp/comgr_3.py.bak
|
||||||
|
python3 -c "from tinygrad.runtime.autogen import comgr_3"
|
||||||
|
diff /tmp/comgr_3.py.bak tinygrad/runtime/autogen/comgr_3.py
|
||||||
|
|||||||
@@ -54,7 +54,7 @@ jobs:
|
|||||||
- name: Print macOS version
|
- name: Print macOS version
|
||||||
run: sw_vers
|
run: sw_vers
|
||||||
- name: Run Stable Diffusion
|
- name: Run Stable Diffusion
|
||||||
run: BENCHMARK_LOG=stable_diffusion JIT=1 ASSERT_MIN_STEP_TIME=720 python3.11 examples/stable_diffusion.py --fp16 --seed 0 --noshow --timing | tee sd.txt
|
run: BENCHMARK_LOG=stable_diffusion JIT=1 ASSERT_MIN_STEP_TIME=800 python3.11 examples/stable_diffusion.py --fp16 --seed 0 --noshow --timing | tee sd.txt
|
||||||
- name: Run Stable Diffusion without fp16
|
- name: Run Stable Diffusion without fp16
|
||||||
run: BENCHMARK_LOG=stable_diffusion_fp32 JIT=1 ASSERT_MIN_STEP_TIME=800 python3.11 examples/stable_diffusion.py --seed 0 --noshow --timing | tee sd_no_fp16.txt
|
run: BENCHMARK_LOG=stable_diffusion_fp32 JIT=1 ASSERT_MIN_STEP_TIME=800 python3.11 examples/stable_diffusion.py --seed 0 --noshow --timing | tee sd_no_fp16.txt
|
||||||
- name: Run Stable Diffusion v2
|
- name: Run Stable Diffusion v2
|
||||||
@@ -199,7 +199,7 @@ jobs:
|
|||||||
- name: Test speed vs torch
|
- name: Test speed vs torch
|
||||||
run: NV=1 CAPTURE_PROCESS_REPLAY=0 HALF=1 BIG=2 TORCHCUDA=1 python3 test/speed/external_test_speed_v_torch.py | tee torch_speed.txt
|
run: NV=1 CAPTURE_PROCESS_REPLAY=0 HALF=1 BIG=2 TORCHCUDA=1 python3 test/speed/external_test_speed_v_torch.py | tee torch_speed.txt
|
||||||
- name: Test speed vs theoretical
|
- name: Test speed vs theoretical
|
||||||
run: NV=1 IGNORE_BEAM_CACHE=1 DISABLE_COMPILER_CACHE=1 BEAM_DEBUG=1 DEBUG=1 python -m pytest -rA test/external/speed_v_theoretical.py --durations=20
|
run: NV=1 IGNORE_BEAM_CACHE=1 CCACHE=0 BEAM_DEBUG=1 DEBUG=1 python -m pytest -rA test/external/speed_v_theoretical.py --durations=20
|
||||||
- name: Test benchmark allreduce
|
- name: Test benchmark allreduce
|
||||||
run: NV=1 python test/external/external_benchmark_multitensor_allreduce.py
|
run: NV=1 python test/external/external_benchmark_multitensor_allreduce.py
|
||||||
- name: Test tensor cores
|
- name: Test tensor cores
|
||||||
@@ -320,19 +320,20 @@ jobs:
|
|||||||
# run: NV=1 NV_PTX=1 M_START=12 M_STOP=20 M_STEP=1 N_START=6 N_STOP=10 N_STEP=1 K_START=28 K_STOP=36 K_STEP=1 HALF=1 TC_OPT=2 python3 ./extra/gemm/fuzz_matmul.py
|
# run: NV=1 NV_PTX=1 M_START=12 M_STOP=20 M_STEP=1 N_START=6 N_STOP=10 N_STEP=1 K_START=28 K_STOP=36 K_STEP=1 HALF=1 TC_OPT=2 python3 ./extra/gemm/fuzz_matmul.py
|
||||||
- name: Train MNIST
|
- name: Train MNIST
|
||||||
run: time PYTHONPATH=. NV=1 TARGET_EVAL_ACC_PCT=96.0 python3 examples/beautiful_mnist.py | tee beautiful_mnist.txt
|
run: time PYTHONPATH=. NV=1 TARGET_EVAL_ACC_PCT=96.0 python3 examples/beautiful_mnist.py | tee beautiful_mnist.txt
|
||||||
|
# TODO: too slow
|
||||||
- name: Run 10 CIFAR training steps
|
- name: Run 10 CIFAR training steps
|
||||||
run: BENCHMARK_LOG=cifar_10steps ASSERT_MIN_STEP_TIME=270 NV=1 STEPS=10 python3 examples/hlb_cifar10.py | tee train_cifar.txt
|
run: BENCHMARK_LOG=cifar_10steps ASSERT_MIN_STEP_TIME=1300 NV=1 STEPS=10 python3 examples/hlb_cifar10.py | tee train_cifar.txt
|
||||||
- name: Run 10 CIFAR training steps w HALF
|
# - name: Run 10 CIFAR training steps w HALF
|
||||||
run: BENCHMARK_LOG=cifar_10steps_half ASSERT_MIN_STEP_TIME=240 NV=1 STEPS=10 DEFAULT_FLOAT=HALF python3 examples/hlb_cifar10.py | tee train_cifar_half.txt
|
# run: BENCHMARK_LOG=cifar_10steps_half ASSERT_MIN_STEP_TIME=240 NV=1 STEPS=10 DEFAULT_FLOAT=HALF python3 examples/hlb_cifar10.py | tee train_cifar_half.txt
|
||||||
- name: Run 10 CIFAR training steps w BF16
|
# - name: Run 10 CIFAR training steps w BF16
|
||||||
run: BENCHMARK_LOG=cifar_10steps_bf16 ASSERT_MIN_STEP_TIME=270 NV=1 STEPS=10 DEFAULT_FLOAT=BFLOAT16 python3 examples/hlb_cifar10.py | tee train_cifar_bf16.txt
|
# run: BENCHMARK_LOG=cifar_10steps_bf16 ASSERT_MIN_STEP_TIME=270 NV=1 STEPS=10 DEFAULT_FLOAT=BFLOAT16 python3 examples/hlb_cifar10.py | tee train_cifar_bf16.txt
|
||||||
# TODO: too slow
|
# TODO: too slow
|
||||||
# - name: Run 10 CIFAR training steps w winograd
|
# - name: Run 10 CIFAR training steps w winograd
|
||||||
# run: BENCHMARK_LOG=cifar_10steps_half_wino ASSERT_MIN_STEP_TIME=350 NV=1 CAPTURE_PROCESS_REPLAY=0 WINO=1 STEPS=10 DEFAULT_FLOAT=HALF python3 examples/hlb_cifar10.py | tee train_cifar_wino.txt
|
# run: BENCHMARK_LOG=cifar_10steps_half_wino ASSERT_MIN_STEP_TIME=350 NV=1 CAPTURE_PROCESS_REPLAY=0 WINO=1 STEPS=10 DEFAULT_FLOAT=HALF python3 examples/hlb_cifar10.py | tee train_cifar_wino.txt
|
||||||
- name: Run full CIFAR training w 1 GPU
|
# - name: Run full CIFAR training w 1 GPU
|
||||||
run: time BENCHMARK_LOG=cifar NV=1 DEFAULT_FLOAT=HALF STEPS=1000 TARGET_EVAL_ACC_PCT=93.0 python3 examples/hlb_cifar10.py | tee train_cifar_one_gpu.txt
|
# run: time BENCHMARK_LOG=cifar NV=1 DEFAULT_FLOAT=HALF STEPS=1000 TARGET_EVAL_ACC_PCT=93.0 python3 examples/hlb_cifar10.py | tee train_cifar_one_gpu.txt
|
||||||
- name: Run full CIFAR training steps w 6 GPUS
|
# - name: Run full CIFAR training steps w 6 GPUS
|
||||||
run: time BENCHMARK_LOG=cifar_6gpu CAPTURE_PROCESS_REPLAY=0 NV=1 DEFAULT_FLOAT=HALF STEPS=350 BS=1536 GPUS=6 TARGET_EVAL_ACC_PCT=93.0 python3 examples/hlb_cifar10.py | tee train_cifar_six_gpu.txt
|
# run: time BENCHMARK_LOG=cifar_6gpu CAPTURE_PROCESS_REPLAY=0 NV=1 DEFAULT_FLOAT=HALF STEPS=350 BS=1536 GPUS=6 TARGET_EVAL_ACC_PCT=93.0 python3 examples/hlb_cifar10.py | tee train_cifar_six_gpu.txt
|
||||||
- name: Run MLPerf resnet eval on training data
|
- name: Run MLPerf resnet eval on training data
|
||||||
run: time BENCHMARK_LOG=resnet_eval NV=1 MODEL=resnet python3 examples/mlperf/model_eval.py
|
run: time BENCHMARK_LOG=resnet_eval NV=1 MODEL=resnet python3 examples/mlperf/model_eval.py
|
||||||
#- name: Run 10 MLPerf ResNet50 training steps (1 gpu)
|
#- name: Run 10 MLPerf ResNet50 training steps (1 gpu)
|
||||||
@@ -409,7 +410,7 @@ jobs:
|
|||||||
# python3 -c "import torch; print(torch.__version__)"
|
# python3 -c "import torch; print(torch.__version__)"
|
||||||
# LD_PRELOAD="/opt/rocm/lib/libhsa-runtime64.so" HSA=1 BIG=2 TORCHCUDA=1 python3 test/speed/external_test_speed_v_torch.py | tee torch_speed.txt
|
# LD_PRELOAD="/opt/rocm/lib/libhsa-runtime64.so" HSA=1 BIG=2 TORCHCUDA=1 python3 test/speed/external_test_speed_v_torch.py | tee torch_speed.txt
|
||||||
- name: Test speed vs theoretical
|
- name: Test speed vs theoretical
|
||||||
run: AMD=1 IGNORE_BEAM_CACHE=1 DISABLE_COMPILER_CACHE=1 BEAM_DEBUG=1 DEBUG=1 python -m pytest -rA test/external/speed_v_theoretical.py --durations=20
|
run: AMD=1 IGNORE_BEAM_CACHE=1 CCACHE=0 BEAM_DEBUG=1 DEBUG=1 python -m pytest -rA test/external/speed_v_theoretical.py --durations=20
|
||||||
- name: Test tensor cores
|
- name: Test tensor cores
|
||||||
run: |
|
run: |
|
||||||
AMD=1 AMD_LLVM=0 python3 test/opt/test_tensor_cores.py
|
AMD=1 AMD_LLVM=0 python3 test/opt/test_tensor_cores.py
|
||||||
@@ -524,17 +525,18 @@ jobs:
|
|||||||
run: test/external/process_replay/reset.py
|
run: test/external/process_replay/reset.py
|
||||||
- name: Train MNIST
|
- name: Train MNIST
|
||||||
run: time PYTHONPATH=. AMD=1 TARGET_EVAL_ACC_PCT=96.0 python3 examples/beautiful_mnist.py | tee beautiful_mnist.txt
|
run: time PYTHONPATH=. AMD=1 TARGET_EVAL_ACC_PCT=96.0 python3 examples/beautiful_mnist.py | tee beautiful_mnist.txt
|
||||||
|
# TODO: too slow
|
||||||
- name: Run 10 CIFAR training steps
|
- name: Run 10 CIFAR training steps
|
||||||
run: BENCHMARK_LOG=cifar_10steps ASSERT_MIN_STEP_TIME=330 AMD=1 STEPS=10 python3 examples/hlb_cifar10.py | tee train_cifar.txt
|
run: BENCHMARK_LOG=cifar_10steps ASSERT_MIN_STEP_TIME=2000 AMD=1 STEPS=10 python3 examples/hlb_cifar10.py | tee train_cifar.txt
|
||||||
- name: Run 10 CIFAR training steps w HALF
|
# - name: Run 10 CIFAR training steps w HALF
|
||||||
run: BENCHMARK_LOG=cifar_10steps_half ASSERT_MIN_STEP_TIME=390 AMD=1 STEPS=10 DEFAULT_FLOAT=HALF python3 examples/hlb_cifar10.py | tee train_cifar_half.txt
|
# run: BENCHMARK_LOG=cifar_10steps_half ASSERT_MIN_STEP_TIME=390 AMD=1 STEPS=10 DEFAULT_FLOAT=HALF python3 examples/hlb_cifar10.py | tee train_cifar_half.txt
|
||||||
# - name: Run 10 CIFAR training steps w BF16
|
# - name: Run 10 CIFAR training steps w BF16
|
||||||
# run: BENCHMARK_LOG=cifar_10steps_bf16 ASSERT_MIN_STEP_TIME=288 AMD=1 STEPS=10 DEFAULT_FLOAT=BFLOAT16 python3 examples/hlb_cifar10.py | tee train_cifar_bf16.txt
|
# run: BENCHMARK_LOG=cifar_10steps_bf16 ASSERT_MIN_STEP_TIME=288 AMD=1 STEPS=10 DEFAULT_FLOAT=BFLOAT16 python3 examples/hlb_cifar10.py | tee train_cifar_bf16.txt
|
||||||
# TODO: too slow
|
# TODO: too slow
|
||||||
# - name: Run 10 CIFAR training steps w winograd
|
# - name: Run 10 CIFAR training steps w winograd
|
||||||
# run: BENCHMARK_LOG=cifar_10steps_half_wino ASSERT_MIN_STEP_TIME=66 AMD=1 WINO=1 STEPS=10 DEFAULT_FLOAT=HALF python3 examples/hlb_cifar10.py | tee train_cifar_wino.txt
|
# run: BENCHMARK_LOG=cifar_10steps_half_wino ASSERT_MIN_STEP_TIME=66 AMD=1 WINO=1 STEPS=10 DEFAULT_FLOAT=HALF python3 examples/hlb_cifar10.py | tee train_cifar_wino.txt
|
||||||
- name: Run full CIFAR training w 1 GPU
|
# - name: Run full CIFAR training w 1 GPU
|
||||||
run: time BENCHMARK_LOG=cifar AMD=1 DEFAULT_FLOAT=HALF STEPS=1000 TARGET_EVAL_ACC_PCT=93.0 python3 examples/hlb_cifar10.py | tee train_cifar_one_gpu.txt
|
# run: time BENCHMARK_LOG=cifar AMD=1 DEFAULT_FLOAT=HALF STEPS=1000 TARGET_EVAL_ACC_PCT=93.0 python3 examples/hlb_cifar10.py | tee train_cifar_one_gpu.txt
|
||||||
#- name: Run full CIFAR training steps w 6 GPUS
|
#- name: Run full CIFAR training steps w 6 GPUS
|
||||||
# run: time BENCHMARK_LOG=cifar_6gpu AMD=1 DEFAULT_FLOAT=HALF STEPS=350 BS=1536 GPUS=6 TARGET_EVAL_ACC_PCT=93.0 python3 examples/hlb_cifar10.py | tee train_cifar_six_gpu.txt
|
# run: time BENCHMARK_LOG=cifar_6gpu AMD=1 DEFAULT_FLOAT=HALF STEPS=350 BS=1536 GPUS=6 TARGET_EVAL_ACC_PCT=93.0 python3 examples/hlb_cifar10.py | tee train_cifar_six_gpu.txt
|
||||||
#- name: Run full CIFAR training steps w 6 GPUS (REMOTE)
|
#- name: Run full CIFAR training steps w 6 GPUS (REMOTE)
|
||||||
@@ -623,32 +625,32 @@ jobs:
|
|||||||
rm -f /tmp/staging.db /tmp/staging.db-shm /tmp/staging.db-wal
|
rm -f /tmp/staging.db /tmp/staging.db-shm /tmp/staging.db-wal
|
||||||
- name: reset process replay
|
- name: reset process replay
|
||||||
run: test/external/process_replay/reset.py
|
run: test/external/process_replay/reset.py
|
||||||
- name: openpilot compile3 0.9.9 driving_vision
|
# - name: openpilot compile3 0.9.9 driving_vision
|
||||||
run: BENCHMARK_LOG=openpilot_0_9_9_vision PYTHONPATH=. NOLOCALS=1 FLOAT16=1 IMAGE=2 QCOM=1 taskset -c 4-7 python3 examples/openpilot/compile3.py https://github.com/commaai/openpilot/raw/v0.9.9/selfdrive/modeld/models/driving_vision.onnx
|
# run: BENCHMARK_LOG=openpilot_0_9_9_vision PYTHONPATH=. NOLOCALS=1 FLOAT16=1 IMAGE=2 QCOM=1 taskset -c 4-7 python3 examples/openpilot/compile3.py https://github.com/commaai/openpilot/raw/v0.9.9/selfdrive/modeld/models/driving_vision.onnx
|
||||||
- name: openpilot compile3 0.9.9 driving_policy
|
# - name: openpilot compile3 0.9.9 driving_policy
|
||||||
run: BENCHMARK_LOG=openpilot_0_9_9_policy PYTHONPATH=. NOLOCALS=1 FLOAT16=1 IMAGE=2 QCOM=1 taskset -c 4-7 python3 examples/openpilot/compile3.py https://github.com/commaai/openpilot/raw/v0.9.9/selfdrive/modeld/models/driving_policy.onnx
|
# run: BENCHMARK_LOG=openpilot_0_9_9_policy PYTHONPATH=. NOLOCALS=1 FLOAT16=1 IMAGE=2 QCOM=1 taskset -c 4-7 python3 examples/openpilot/compile3.py https://github.com/commaai/openpilot/raw/v0.9.9/selfdrive/modeld/models/driving_policy.onnx
|
||||||
- name: openpilot compile3 0.9.9 dmonitoring
|
# - name: openpilot compile3 0.9.9 dmonitoring
|
||||||
run: BENCHMARK_LOG=openpilot_0_9_9_dmonitoring PYTHONPATH=. NOLOCALS=1 FLOAT16=1 IMAGE=2 QCOM=1 taskset -c 4-7 python3 examples/openpilot/compile3.py https://github.com/commaai/openpilot/raw/v0.9.9/selfdrive/modeld/models/dmonitoring_model.onnx
|
# run: BENCHMARK_LOG=openpilot_0_9_9_dmonitoring PYTHONPATH=. NOLOCALS=1 FLOAT16=1 IMAGE=2 QCOM=1 taskset -c 4-7 python3 examples/openpilot/compile3.py https://github.com/commaai/openpilot/raw/v0.9.9/selfdrive/modeld/models/dmonitoring_model.onnx
|
||||||
- name: openpilot compile3 0.10.0 driving_policy
|
- name: openpilot compile3 0.10.0 driving_policy
|
||||||
run: BENCHMARK_LOG=openpilot_0_10_0_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/v0.10.0/selfdrive/modeld/models/driving_policy.onnx
|
run: BENCHMARK_LOG=openpilot_0_10_0_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/v0.10.0/selfdrive/modeld/models/driving_policy.onnx
|
||||||
- name: openpilot compile3 0.10.0 dmonitoring
|
- name: openpilot compile3 0.10.0 dmonitoring
|
||||||
run: BENCHMARK_LOG=openpilot_0_10_0_dmonitoring PYTHONPATH="." ASSERT_MIN_STEP_TIME=12 DEV=QCOM FLOAT16=1 IMAGE=2 NOLOCALS=1 taskset -c 4-7 python3 examples/openpilot/compile3.py https://github.com/commaai/openpilot/raw/v0.10.0/selfdrive/modeld/models/dmonitoring_model.onnx
|
run: BENCHMARK_LOG=openpilot_0_10_0_dmonitoring PYTHONPATH="." ASSERT_MIN_STEP_TIME=11 DEV=QCOM FLOAT16=1 IMAGE=2 NOLOCALS=1 taskset -c 4-7 python3 examples/openpilot/compile3.py https://github.com/commaai/openpilot/raw/v0.10.0/selfdrive/modeld/models/dmonitoring_model.onnx
|
||||||
|
- name: DEBUG=2 openpilot compile3 0.10.1 driving_vision
|
||||||
|
run: PYTHONPATH="." DEBUG=2 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_vision.onnx
|
||||||
- name: openpilot compile3 0.10.1 driving_vision
|
- name: openpilot compile3 0.10.1 driving_vision
|
||||||
# TODO: ASSERT_MIN_STEP_TIME=17
|
run: BENCHMARK_LOG=openpilot_0_10_1_vision PYTHONPATH="." ASSERT_MIN_STEP_TIME=17 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_vision.onnx
|
||||||
run: BENCHMARK_LOG=openpilot_0_10_1_vision PYTHONPATH="." ASSERT_MIN_STEP_TIME=21 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_vision.onnx
|
|
||||||
- name: openpilot compile3 0.10.1 driving_policy
|
- name: openpilot compile3 0.10.1 driving_policy
|
||||||
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
|
||||||
# TODO: ASSERT_MIN_STEP_TIME=10
|
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=12 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
|
||||||
|
|
||||||
@@ -704,8 +706,9 @@ jobs:
|
|||||||
run: |
|
run: |
|
||||||
AMD=1 GRAPH_ONE_KERNEL=1 PYTHONPATH=. NSZ=8192 python3 test/speed/external_test_copy_speed.py TestCopySpeed.testCopyDefaulttoCPUJit
|
AMD=1 GRAPH_ONE_KERNEL=1 PYTHONPATH=. NSZ=8192 python3 test/speed/external_test_copy_speed.py TestCopySpeed.testCopyDefaulttoCPUJit
|
||||||
AMD=1 GRAPH_ONE_KERNEL=1 PYTHONPATH=. NSZ=8192 python3 test/speed/external_test_copy_speed.py TestCopySpeed.testCopyCPUtoDefaultJit
|
AMD=1 GRAPH_ONE_KERNEL=1 PYTHONPATH=. NSZ=8192 python3 test/speed/external_test_copy_speed.py TestCopySpeed.testCopyCPUtoDefaultJit
|
||||||
- name: Run full CIFAR training w 1 GPU
|
# TODO: too slow
|
||||||
run: time BENCHMARK_LOG=cifar AMD=1 DEFAULT_FLOAT=HALF STEPS=1000 TARGET_EVAL_ACC_PCT=93.0 python3 examples/hlb_cifar10.py | tee am_train_cifar_one_gpu.txt
|
# - name: Run full CIFAR training w 1 GPU
|
||||||
|
# run: time BENCHMARK_LOG=cifar AMD=1 DEFAULT_FLOAT=HALF STEPS=1000 TARGET_EVAL_ACC_PCT=93.0 python3 examples/hlb_cifar10.py | tee am_train_cifar_one_gpu.txt
|
||||||
# TODO: enable
|
# TODO: enable
|
||||||
# - name: Run 10 MLPerf ResNet50 training steps (1 gpu)
|
# - name: Run 10 MLPerf ResNet50 training steps (1 gpu)
|
||||||
# run: BENCHMARK_LOG=resnet_10steps AMD=1 MNISTMOCK=1 DEFAULT_FLOAT=HALF BENCHMARK=10 BS=256 GPUS=1 MODEL=resnet python3 examples/mlperf/model_train.py | tee am_train_resnet_one_gpu.txt
|
# run: BENCHMARK_LOG=resnet_10steps AMD=1 MNISTMOCK=1 DEFAULT_FLOAT=HALF BENCHMARK=10 BS=256 GPUS=1 MODEL=resnet python3 examples/mlperf/model_train.py | tee am_train_resnet_one_gpu.txt
|
||||||
@@ -767,8 +770,9 @@ jobs:
|
|||||||
NV=1 GRAPH_ONE_KERNEL=1 PYTHONPATH=. NSZ=8192 python3 test/speed/external_test_copy_speed.py TestCopySpeed.testCopyCPUtoDefaultJit
|
NV=1 GRAPH_ONE_KERNEL=1 PYTHONPATH=. NSZ=8192 python3 test/speed/external_test_copy_speed.py TestCopySpeed.testCopyCPUtoDefaultJit
|
||||||
- name: Test LLAMA-3
|
- name: Test LLAMA-3
|
||||||
run: BENCHMARK_LOG=llama3_beam NV=1 JITBEAM=2 IGNORE_BEAM_CACHE=1 python3 examples/llama3.py --size 8B --benchmark --temperature 0 | tee nv_llama3_beam.txt
|
run: BENCHMARK_LOG=llama3_beam NV=1 JITBEAM=2 IGNORE_BEAM_CACHE=1 python3 examples/llama3.py --size 8B --benchmark --temperature 0 | tee nv_llama3_beam.txt
|
||||||
- name: Run full CIFAR training w 1 GPU
|
# TODO: too slow
|
||||||
run: time BENCHMARK_LOG=cifar NV=1 DEFAULT_FLOAT=HALF STEPS=1000 TARGET_EVAL_ACC_PCT=93.0 python3 examples/hlb_cifar10.py | tee nv_train_cifar_one_gpu.txt
|
# - name: Run full CIFAR training w 1 GPU
|
||||||
|
# run: time BENCHMARK_LOG=cifar NV=1 DEFAULT_FLOAT=HALF STEPS=1000 TARGET_EVAL_ACC_PCT=93.0 python3 examples/hlb_cifar10.py | tee nv_train_cifar_one_gpu.txt
|
||||||
#- name: Run 10 MLPerf ResNet50 training steps (1 gpu)
|
#- name: Run 10 MLPerf ResNet50 training steps (1 gpu)
|
||||||
# run: BENCHMARK_LOG=resnet_10steps NV=1 MNISTMOCK=1 DEFAULT_FLOAT=HALF BENCHMARK=10 BS=256 GPUS=1 MODEL=resnet python3 examples/mlperf/model_train.py | tee nv_train_resnet_one_gpu.txt
|
# run: BENCHMARK_LOG=resnet_10steps NV=1 MNISTMOCK=1 DEFAULT_FLOAT=HALF BENCHMARK=10 BS=256 GPUS=1 MODEL=resnet python3 examples/mlperf/model_train.py | tee nv_train_resnet_one_gpu.txt
|
||||||
- name: Run 10 MLPerf Bert training steps (1 gpu)
|
- name: Run 10 MLPerf Bert training steps (1 gpu)
|
||||||
|
|||||||
@@ -22,13 +22,13 @@ jobs:
|
|||||||
- name: Run SDXL with new search
|
- name: Run SDXL with new search
|
||||||
# TODO: GCVM_L2_PROTECTION_FAULT_STATUS with llvm19
|
# TODO: GCVM_L2_PROTECTION_FAULT_STATUS with llvm19
|
||||||
run: |
|
run: |
|
||||||
BENCHMARK_LOG=search_sdxl PYTHONPATH=. AMD=1 JITBEAM=2 IGNORE_BEAM_CACHE=1 DISABLE_COMPILER_CACHE=1 python examples/sdxl.py --noshow --timing --seed 0
|
BENCHMARK_LOG=search_sdxl PYTHONPATH=. AMD=1 JITBEAM=2 IGNORE_BEAM_CACHE=1 CCACHE=0 python examples/sdxl.py --noshow --timing --seed 0
|
||||||
- name: Run SDXL with cached search
|
- name: Run SDXL with cached search
|
||||||
run: |
|
run: |
|
||||||
BENCHMARK_LOG=search_sdxl_cached PYTHONPATH=. AMD=1 JITBEAM=2 python examples/sdxl.py --noshow --timing --seed 0
|
BENCHMARK_LOG=search_sdxl_cached PYTHONPATH=. AMD=1 JITBEAM=2 python examples/sdxl.py --noshow --timing --seed 0
|
||||||
- name: Run winograd cifar with new search
|
- name: Run winograd cifar with new search
|
||||||
run: |
|
run: |
|
||||||
BENCHMARK_LOG=search_wino_cifar WINO=1 DEFAULT_FLOAT=HALF JITBEAM=4 IGNORE_BEAM_CACHE=1 DISABLE_COMPILER_CACHE=1 BS=1024 STEPS=500 python examples/hlb_cifar10.py
|
BENCHMARK_LOG=search_wino_cifar WINO=1 DEFAULT_FLOAT=HALF JITBEAM=4 IGNORE_BEAM_CACHE=1 CCACHE=0 BS=1024 STEPS=500 python examples/hlb_cifar10.py
|
||||||
- name: Run winograd cifar with cached search
|
- name: Run winograd cifar with cached search
|
||||||
run: |
|
run: |
|
||||||
BENCHMARK_LOG=search_wino_cifar_cached WINO=1 DEFAULT_FLOAT=HALF JITBEAM=4 BS=1024 STEPS=500 python examples/hlb_cifar10.py
|
BENCHMARK_LOG=search_wino_cifar_cached WINO=1 DEFAULT_FLOAT=HALF JITBEAM=4 BS=1024 STEPS=500 python examples/hlb_cifar10.py
|
||||||
|
|||||||
@@ -20,11 +20,11 @@ jobs:
|
|||||||
- name: Install dependencies
|
- name: Install dependencies
|
||||||
run: |
|
run: |
|
||||||
python -m pip install --upgrade pip
|
python -m pip install --upgrade pip
|
||||||
pip install setuptools wheel twine
|
pip install setuptools wheel build twine
|
||||||
- name: Build and publish
|
- name: Build and publish
|
||||||
env:
|
env:
|
||||||
TWINE_USERNAME: ${{ secrets.PYPI_USERNAME }}
|
TWINE_USERNAME: ${{ secrets.PYPI_USERNAME }}
|
||||||
TWINE_PASSWORD: ${{ secrets.PYPI_PASSWORD }}
|
TWINE_PASSWORD: ${{ secrets.PYPI_PASSWORD }}
|
||||||
run: |
|
run: |
|
||||||
python setup.py sdist bdist_wheel
|
python -m build
|
||||||
twine upload dist/*
|
twine upload dist/*
|
||||||
|
|||||||
@@ -1,10 +1,7 @@
|
|||||||
name: Unit Tests
|
name: Unit Tests
|
||||||
env:
|
env:
|
||||||
# increment this when downloads substantially change to avoid the internet
|
# increment this when downloads substantially change to avoid the internet
|
||||||
DOWNLOAD_CACHE_VERSION: '12'
|
CACHE_VERSION: '13'
|
||||||
PYTHON_CACHE_VERSION: '4'
|
|
||||||
APT_CACHE_VERSION: '1'
|
|
||||||
BUILD_CACHE_VERSION: '1'
|
|
||||||
CAPTURE_PROCESS_REPLAY: 1
|
CAPTURE_PROCESS_REPLAY: 1
|
||||||
GH_TOKEN: ${{ secrets.GITHUB_TOKEN }}
|
GH_TOKEN: ${{ secrets.GITHUB_TOKEN }}
|
||||||
PYTHONPATH: ${{ github.workspace }}
|
PYTHONPATH: ${{ github.workspace }}
|
||||||
@@ -290,8 +287,8 @@ jobs:
|
|||||||
python extra/optimization/extract_dataset.py
|
python extra/optimization/extract_dataset.py
|
||||||
gzip -c /tmp/sops > extra/datasets/sops.gz
|
gzip -c /tmp/sops > extra/datasets/sops.gz
|
||||||
#DEBUG=1 MIN_ASTS=1 python extra/optimization/get_action_space.py
|
#DEBUG=1 MIN_ASTS=1 python extra/optimization/get_action_space.py
|
||||||
- name: Repo line count < 18500 lines
|
- name: Repo line count < 19000 lines
|
||||||
run: MAX_LINE_COUNT=18500 python sz.py
|
run: MAX_LINE_COUNT=19000 python sz.py
|
||||||
|
|
||||||
spec:
|
spec:
|
||||||
strategy:
|
strategy:
|
||||||
@@ -309,6 +306,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 }}
|
||||||
|
|
||||||
@@ -344,10 +342,11 @@ jobs:
|
|||||||
key: gpu-image
|
key: gpu-image
|
||||||
deps: testing_minimal
|
deps: testing_minimal
|
||||||
opencl: 'true'
|
opencl: 'true'
|
||||||
- name: Test CL IMAGE=2 ops + training
|
- name: Test CL IMAGE=2 ops
|
||||||
run: |
|
run: |
|
||||||
CL=1 IMAGE=2 python -m pytest -n=auto test/test_ops.py --durations=20
|
CL=1 IMAGE=2 python -m pytest -n=auto test/test_ops.py --durations=20
|
||||||
CL=1 IMAGE=2 python test/models/test_end2end.py TestEnd2End.test_linear_mnist
|
# TODO: training is broken
|
||||||
|
# CL=1 IMAGE=2 python test/models/test_end2end.py TestEnd2End.test_linear_mnist
|
||||||
- name: Run process replay tests
|
- name: Run process replay tests
|
||||||
uses: ./.github/actions/process-replay
|
uses: ./.github/actions/process-replay
|
||||||
|
|
||||||
@@ -392,7 +391,7 @@ jobs:
|
|||||||
llvm: 'true'
|
llvm: 'true'
|
||||||
- name: Test openpilot model kernel count and gate usage
|
- name: Test openpilot model kernel count and gate usage
|
||||||
run: |
|
run: |
|
||||||
ALLOWED_KERNEL_COUNT=123 ALLOWED_READ_IMAGE=1452 ALLOWED_GATED_READ_IMAGE=122 FLOAT16=1 CL=1 IMAGE=2 python examples/openpilot/compile3.py https://gitlab.com/commaai/openpilot-lfs.git/gitlab-lfs/objects/cf6376aa9a090f0da26c280ef69eabf9bbdd51d1faac9ed392919c3db69be916
|
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
|
||||||
- name: Test openpilot CL compile fp16
|
- name: Test openpilot CL compile fp16
|
||||||
run: FLOAT16=1 DEBUGCL=1 CL=1 IMAGE=2 python examples/openpilot/compile3.py https://gitlab.com/commaai/openpilot-lfs.git/gitlab-lfs/objects/cf6376aa9a090f0da26c280ef69eabf9bbdd51d1faac9ed392919c3db69be916
|
run: FLOAT16=1 DEBUGCL=1 CL=1 IMAGE=2 python examples/openpilot/compile3.py https://gitlab.com/commaai/openpilot-lfs.git/gitlab-lfs/objects/cf6376aa9a090f0da26c280ef69eabf9bbdd51d1faac9ed392919c3db69be916
|
||||||
- name: Test openpilot CL compile fp32 (test correctness)
|
- name: Test openpilot CL compile fp32 (test correctness)
|
||||||
|
|||||||
@@ -21,17 +21,38 @@ tinygrad: For something between [PyTorch](https://github.com/pytorch/pytorch) an
|
|||||||
|
|
||||||
---
|
---
|
||||||
|
|
||||||
Despite tinygrad's size, it is a fully featured deep learning framework.
|
tinygrad is an end-to-end deep learning stack:
|
||||||
|
|
||||||
Due to its extreme simplicity, it is the easiest framework to add new accelerators to, with support for both inference and training. If XLA is CISC, tinygrad is RISC.
|
- **Tensor library** with autograd
|
||||||
|
- **IR and compiler** that fuse and lower kernels
|
||||||
|
- **JIT + graph execution**
|
||||||
|
- **nn / optim / datasets** for real training
|
||||||
|
|
||||||
tinygrad is now beta software, we [raised some money](https://geohot.github.io/blog/jekyll/update/2023/05/24/the-tiny-corp-raised-5M.html) to make it good. Someday, we will tape out chips.
|
It’s inspired by PyTorch (ergonomics), JAX (functional transforms and IR-based AD), and TVM (scheduling and codegen), but stays intentionally tiny and hackable.
|
||||||
|
|
||||||
## Features
|
---
|
||||||
|
|
||||||
### LLaMA and Stable Diffusion
|
## How tinygrad compares
|
||||||
|
|
||||||
tinygrad can run [LLaMA](/docs/showcase.md#llama) and [Stable Diffusion](/docs/showcase.md#stable-diffusion)!
|
**PyTorch**
|
||||||
|
|
||||||
|
- ✅ Similar: eager `Tensor` API, autograd, `optim`, basic datasets and layers.
|
||||||
|
- ✅ You can write familiar training loops.
|
||||||
|
- 🔁 Unlike PyTorch, the entire compiler and IR are visible and hackable.
|
||||||
|
|
||||||
|
**JAX**
|
||||||
|
|
||||||
|
- ✅ IR-based autodiff over primitives (like JAXPR + XLA).
|
||||||
|
- ✅ Function-level JIT (`TinyJit`) that captures and replays kernels.
|
||||||
|
- 🔁 Fewer functional transforms (no full `vmap`/`pmap` yet), but far easier to read.
|
||||||
|
|
||||||
|
**TVM**
|
||||||
|
|
||||||
|
- ✅ Multiple lowering passes, scheduling, and BEAM search over kernels.
|
||||||
|
- ✅ Device “graphs” for batched execution.
|
||||||
|
- 🔁 tinygrad also ships the **front-end framework** (tensors, nn, optim), not just the compiler.
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
### Laziness
|
### Laziness
|
||||||
|
|
||||||
|
|||||||
@@ -1,568 +0,0 @@
|
|||||||
#!/bin/bash -e
|
|
||||||
|
|
||||||
# setup instructions for clang2py
|
|
||||||
if [[ ! $(clang2py -V) ]]; then
|
|
||||||
pushd .
|
|
||||||
cd /tmp
|
|
||||||
sudo apt-get install -y --no-install-recommends clang
|
|
||||||
pip install --upgrade pip setuptools
|
|
||||||
pip install clang==14.0.6
|
|
||||||
git clone https://github.com/nimlgen/ctypeslib.git
|
|
||||||
cd ctypeslib
|
|
||||||
pip install .
|
|
||||||
clang2py -V
|
|
||||||
popd
|
|
||||||
fi
|
|
||||||
|
|
||||||
BASE=tinygrad/runtime/autogen/
|
|
||||||
|
|
||||||
fixup() {
|
|
||||||
sed -i '1s/^/# mypy: ignore-errors\n/' $1
|
|
||||||
sed -i 's/ *$//' $1
|
|
||||||
grep FIXME_STUB $1 || true
|
|
||||||
}
|
|
||||||
|
|
||||||
patch_dlopen() {
|
|
||||||
path=$1; shift
|
|
||||||
name=$1; shift
|
|
||||||
cat <<EOF | sed -i "/import ctypes.*/r /dev/stdin" $path
|
|
||||||
PATHS_TO_TRY = [
|
|
||||||
$(for p in "$@"; do echo " $p,"; done)
|
|
||||||
]
|
|
||||||
def _try_dlopen_$name():
|
|
||||||
library = ctypes.util.find_library("$name")
|
|
||||||
if library:
|
|
||||||
try: return ctypes.CDLL(library)
|
|
||||||
except OSError: pass
|
|
||||||
for candidate in PATHS_TO_TRY:
|
|
||||||
try: return ctypes.CDLL(candidate)
|
|
||||||
except OSError: pass
|
|
||||||
return None
|
|
||||||
EOF
|
|
||||||
}
|
|
||||||
|
|
||||||
generate_opencl() {
|
|
||||||
clang2py /usr/include/CL/cl.h -o $BASE/opencl.py -l /usr/lib/x86_64-linux-gnu/libOpenCL.so.1 -k cdefstum
|
|
||||||
fixup $BASE/opencl.py
|
|
||||||
# hot patches
|
|
||||||
sed -i "s\import ctypes\import ctypes, ctypes.util\g" $BASE/opencl.py
|
|
||||||
sed -i "s\ctypes.CDLL('/usr/lib/x86_64-linux-gnu/libOpenCL.so.1')\ctypes.CDLL(ctypes.util.find_library('OpenCL'))\g" $BASE/opencl.py
|
|
||||||
python3 -c "import tinygrad.runtime.autogen.opencl"
|
|
||||||
}
|
|
||||||
|
|
||||||
generate_hip() {
|
|
||||||
clang2py /opt/rocm/include/hip/hip_ext.h /opt/rocm/include/hip/hiprtc.h \
|
|
||||||
/opt/rocm/include/hip/hip_runtime_api.h /opt/rocm/include/hip/driver_types.h \
|
|
||||||
--clang-args="-D__HIP_PLATFORM_AMD__ -I/opt/rocm/include -x c++" -o $BASE/hip.py -l /opt/rocm/lib/libamdhip64.so
|
|
||||||
echo "hipDeviceProp_t = hipDeviceProp_tR0600" >> $BASE/hip.py
|
|
||||||
echo "hipGetDeviceProperties = hipGetDevicePropertiesR0600" >> $BASE/hip.py
|
|
||||||
fixup $BASE/hip.py
|
|
||||||
# we can trust HIP is always at /opt/rocm/lib
|
|
||||||
#sed -i "s\import ctypes\import ctypes, ctypes.util\g" $BASE/hip.py
|
|
||||||
#sed -i "s\ctypes.CDLL('/opt/rocm/lib/libhiprtc.so')\ctypes.CDLL(ctypes.util.find_library('hiprtc'))\g" $BASE/hip.py
|
|
||||||
#sed -i "s\ctypes.CDLL('/opt/rocm/lib/libamdhip64.so')\ctypes.CDLL(ctypes.util.find_library('amdhip64'))\g" $BASE/hip.py
|
|
||||||
sed -i "s\import ctypes\import ctypes, os\g" $BASE/hip.py
|
|
||||||
sed -i "s\'/opt/rocm/\os.getenv('ROCM_PATH', '/opt/rocm/')+'/\g" $BASE/hip.py
|
|
||||||
python3 -c "import tinygrad.runtime.autogen.hip"
|
|
||||||
}
|
|
||||||
|
|
||||||
generate_comgr() {
|
|
||||||
clang2py /opt/rocm/include/amd_comgr/amd_comgr.h \
|
|
||||||
--clang-args="-D__HIP_PLATFORM_AMD__ -I/opt/rocm/include -x c++" -o $BASE/comgr.py -l /opt/rocm/lib/libamd_comgr.so
|
|
||||||
fixup $BASE/comgr.py
|
|
||||||
sed -i "s\import ctypes\import ctypes, ctypes.util, os\g" $BASE/comgr.py
|
|
||||||
patch_dlopen $BASE/comgr.py amd_comgr "'/opt/rocm/lib/libamd_comgr.so'" "os.getenv('ROCM_PATH', '')+'/lib/libamd_comgr.so'" "'/usr/local/lib/libamd_comgr.dylib'" "'/opt/homebrew/lib/libamd_comgr.dylib'"
|
|
||||||
sed -i "s\ctypes.CDLL('/opt/rocm/lib/libamd_comgr.so')\_try_dlopen_amd_comgr()\g" $BASE/comgr.py
|
|
||||||
python3 -c "import tinygrad.runtime.autogen.comgr"
|
|
||||||
}
|
|
||||||
|
|
||||||
generate_kfd() {
|
|
||||||
clang2py /usr/include/linux/kfd_ioctl.h -o $BASE/kfd.py -k cdefstum
|
|
||||||
|
|
||||||
fixup $BASE/kfd.py
|
|
||||||
sed -i "s/import ctypes/import ctypes, os/g" $BASE/kfd.py
|
|
||||||
sed -i "s/import fcntl, functools/import functools/g" $BASE/kfd.py
|
|
||||||
sed -i "/import functools/a from tinygrad.runtime.support.hcq import FileIOInterface" $BASE/kfd.py
|
|
||||||
sed -i "s/def _do_ioctl(__idir, __base, __nr, __user_struct, __fd, \*\*kwargs):/def _do_ioctl(__idir, __base, __nr, __user_struct, __fd:FileIOInterface, \*\*kwargs):/g" $BASE/kfd.py
|
|
||||||
sed -i "s/fcntl.ioctl(__fd, (__idir<<30)/__fd.ioctl((__idir<<30)/g" $BASE/kfd.py
|
|
||||||
sed -i "s/!!/not not /g" $BASE/kfd.py
|
|
||||||
python3 -c "import tinygrad.runtime.autogen.kfd"
|
|
||||||
}
|
|
||||||
|
|
||||||
generate_cuda() {
|
|
||||||
clang2py /usr/include/cuda.h --clang-args="-D__CUDA_API_VERSION_INTERNAL" -o $BASE/cuda.py -l /usr/lib/x86_64-linux-gnu/libcuda.so
|
|
||||||
sed -i "s\import ctypes\import ctypes, ctypes.util\g" $BASE/cuda.py
|
|
||||||
sed -i "s\ctypes.CDLL('/usr/lib/x86_64-linux-gnu/libcuda.so')\ctypes.CDLL(ctypes.util.find_library('cuda'))\g" $BASE/cuda.py
|
|
||||||
fixup $BASE/cuda.py
|
|
||||||
python3 -c "import tinygrad.runtime.autogen.cuda"
|
|
||||||
}
|
|
||||||
|
|
||||||
generate_nvrtc() {
|
|
||||||
clang2py /usr/local/cuda/include/nvrtc.h /usr/local/cuda/include/nvJitLink.h -o $BASE/nvrtc.py -l /usr/local/cuda/lib64/libnvrtc.so -l /usr/local/cuda/lib64/libnvJitLink.so
|
|
||||||
sed -i "s\import ctypes\import ctypes, ctypes.util\g" $BASE/nvrtc.py
|
|
||||||
sed -i "s\ctypes.CDLL('/usr/local/cuda/lib64/libnvrtc.so')\ctypes.CDLL(ctypes.util.find_library('nvrtc'))\g" $BASE/nvrtc.py
|
|
||||||
sed -i "s\ctypes.CDLL('/usr/local/cuda/lib64/libnvJitLink.so')\ctypes.CDLL(ctypes.util.find_library('nvJitLink'))\g" $BASE/nvrtc.py
|
|
||||||
fixup $BASE/nvrtc.py
|
|
||||||
python3 -c "import tinygrad.runtime.autogen.nvrtc"
|
|
||||||
}
|
|
||||||
|
|
||||||
generate_nv() {
|
|
||||||
NVKERN_COMMIT_HASH=81fe4fb417c8ac3b9bdcc1d56827d116743892a5
|
|
||||||
NVKERN_SRC=/tmp/open-gpu-kernel-modules-$NVKERN_COMMIT_HASH
|
|
||||||
if [ ! -d "$NVKERN_SRC" ]; then
|
|
||||||
git clone https://github.com/NVIDIA/open-gpu-kernel-modules $NVKERN_SRC
|
|
||||||
pushd .
|
|
||||||
cd $NVKERN_SRC
|
|
||||||
git reset --hard $NVKERN_COMMIT_HASH
|
|
||||||
popd
|
|
||||||
fi
|
|
||||||
|
|
||||||
clang2py -k cdefstum \
|
|
||||||
extra/nv_gpu_driver/clc6c0qmd.h \
|
|
||||||
extra/nv_gpu_driver/clcec0qmd.h \
|
|
||||||
$NVKERN_SRC/src/common/sdk/nvidia/inc/class/cl0000.h \
|
|
||||||
$NVKERN_SRC/src/common/sdk/nvidia/inc/class/cl0080.h \
|
|
||||||
$NVKERN_SRC/src/common/sdk/nvidia/inc/class/cl2080.h \
|
|
||||||
$NVKERN_SRC/src/common/sdk/nvidia/inc/class/cl2080_notification.h \
|
|
||||||
$NVKERN_SRC/src/common/sdk/nvidia/inc/class/clc56f.h \
|
|
||||||
$NVKERN_SRC/src/common/sdk/nvidia/inc/class/clc86f.h \
|
|
||||||
$NVKERN_SRC/src/common/sdk/nvidia/inc/class/clc96f.h \
|
|
||||||
$NVKERN_SRC/src/common/sdk/nvidia/inc/class/clc761.h \
|
|
||||||
$NVKERN_SRC/src/common/sdk/nvidia/inc/class/cl83de.h \
|
|
||||||
$NVKERN_SRC/src/nvidia/generated/g_allclasses.h \
|
|
||||||
$NVKERN_SRC/src/common/sdk/nvidia/inc/class/clc6c0.h \
|
|
||||||
$NVKERN_SRC/src/common/sdk/nvidia/inc/class/clcdc0.h \
|
|
||||||
$NVKERN_SRC/kernel-open/nvidia-uvm/clc6b5.h \
|
|
||||||
$NVKERN_SRC/kernel-open/nvidia-uvm/clc9b5.h \
|
|
||||||
$NVKERN_SRC/kernel-open/nvidia-uvm/uvm_ioctl.h \
|
|
||||||
$NVKERN_SRC/kernel-open/nvidia-uvm/uvm_linux_ioctl.h \
|
|
||||||
$NVKERN_SRC/kernel-open/nvidia-uvm/hwref/ampere/ga100/dev_fault.h \
|
|
||||||
$NVKERN_SRC/src/nvidia/arch/nvalloc/unix/include/nv_escape.h \
|
|
||||||
$NVKERN_SRC/src/nvidia/arch/nvalloc/unix/include/nv-ioctl.h \
|
|
||||||
$NVKERN_SRC/src/nvidia/arch/nvalloc/unix/include/nv-ioctl-numbers.h \
|
|
||||||
$NVKERN_SRC/src/nvidia/arch/nvalloc/unix/include/nv-ioctl-numa.h \
|
|
||||||
$NVKERN_SRC/src/nvidia/arch/nvalloc/unix/include/nv-unix-nvos-params-wrappers.h \
|
|
||||||
$NVKERN_SRC/src/common/sdk/nvidia/inc/alloc/alloc_channel.h \
|
|
||||||
$NVKERN_SRC/src/common/sdk/nvidia/inc/nvos.h \
|
|
||||||
$NVKERN_SRC/src/common/sdk/nvidia/inc/ctrl/ctrl0000/*.h \
|
|
||||||
$NVKERN_SRC/src/common/sdk/nvidia/inc/ctrl/ctrl0080/*.h \
|
|
||||||
$NVKERN_SRC/src/common/sdk/nvidia/inc/ctrl/ctrl2080/*.h \
|
|
||||||
$NVKERN_SRC/src/common/sdk/nvidia/inc/ctrl/ctrl83de/*.h \
|
|
||||||
$NVKERN_SRC/src/common/sdk/nvidia/inc/ctrl/ctrlc36f.h \
|
|
||||||
$NVKERN_SRC/src/common/sdk/nvidia/inc/ctrl/ctrlcb33.h \
|
|
||||||
$NVKERN_SRC/src/common/sdk/nvidia/inc/ctrl/ctrla06c.h \
|
|
||||||
$NVKERN_SRC/src/common/sdk/nvidia/inc/ctrl/ctrl90f1.h \
|
|
||||||
--clang-args="-include $NVKERN_SRC/src/common/sdk/nvidia/inc/nvtypes.h -I$NVKERN_SRC/src/common/inc -I$NVKERN_SRC/kernel-open/nvidia-uvm -I$NVKERN_SRC/kernel-open/common/inc -I$NVKERN_SRC/src/common/sdk/nvidia/inc -I$NVKERN_SRC/src/nvidia/arch/nvalloc/unix/include -I$NVKERN_SRC/src/common/sdk/nvidia/inc/ctrl" \
|
|
||||||
-o $BASE/nv_gpu.py
|
|
||||||
fixup $BASE/nv_gpu.py
|
|
||||||
sed -i "s\(0000000001)\1\g" $BASE/nv_gpu.py
|
|
||||||
sed -i "s\import ctypes\import ctypes, os\g" $BASE/nv_gpu.py
|
|
||||||
sed -i 's/#\?\s\([A-Za-z0-9_]\+\) = MW ( \([0-9]\+\) : \([0-9]\+\) )/\1 = (\2 , \3)/' $BASE/nv_gpu.py # NVC6C0_QMDV03_00 processing
|
|
||||||
sed -i 's/#\sdef NVC6C0_QMD\([A-Za-z0-9_()]\+\):/def NVC6C0_QMD\1:/' $BASE/nv_gpu.py
|
|
||||||
sed -i 's/#\sdef NVCEC0_QMD\([A-Za-z0-9_()]\+\):/def NVCEC0_QMD\1:/' $BASE/nv_gpu.py
|
|
||||||
sed -E -i -n '/^def (NVCEC0_QMDV05_00_RELEASE)(_ENABLE)\(i\):/{p;s//\1'"0"'\2=\1\2(0)\n\1'"1"'\2=\1\2(1)/;H;b};p;${x;s/^\n//;p}' "$BASE/nv_gpu.py"
|
|
||||||
sed -i 's/#\s*return MW(\([0-9i()*+]\+\):\([0-9i()*+]\+\))/ return (\1 , \2)/' $BASE/nv_gpu.py
|
|
||||||
sed -i 's/#\?\s*\(.*\)\s*=\s*\(NV\)\?BIT\(32\)\?\s*(\s*\([0-9]\+\)\s*)/\1 = (1 << \4)/' $BASE/nv_gpu.py # name = BIT(x) -> name = (1 << x)
|
|
||||||
sed -i "s/UVM_\([A-Za-z0-9_]\+\) = \['i', '(', '\([0-9]\+\)', ')'\]/UVM_\1 = \2/" $BASE/nv_gpu.py # UVM_name = ['i', '(', '<num>', ')'] -> UVM_name = <num>
|
|
||||||
|
|
||||||
# Parse status codes
|
|
||||||
sed -n '1i\
|
|
||||||
nv_status_codes = {}
|
|
||||||
/^NV_STATUS_CODE/ { s/^NV_STATUS_CODE(\([^,]*\), *\([^,]*\), *"\([^"]*\)") *.*$/\1 = \2\nnv_status_codes[\1] = "\3"/; p }' $NVKERN_SRC/src/common/sdk/nvidia/inc/nvstatuscodes.h >> $BASE/nv_gpu.py
|
|
||||||
python3 -c "import tinygrad.runtime.autogen.nv_gpu"
|
|
||||||
|
|
||||||
clang2py -k cdefstum \
|
|
||||||
$NVKERN_SRC/src/nvidia/inc/kernel/gpu/fsp/kern_fsp_cot_payload.h \
|
|
||||||
$NVKERN_SRC/src/nvidia/arch/nvalloc/common/inc/gsp/gspifpub.h \
|
|
||||||
$NVKERN_SRC/src/nvidia/arch/nvalloc/common/inc/gsp/gsp_fw_wpr_meta.h \
|
|
||||||
$NVKERN_SRC/src/nvidia/arch/nvalloc/common/inc/gsp/gsp_fw_sr_meta.h \
|
|
||||||
$NVKERN_SRC/src/nvidia/inc/kernel/gpu/gsp/gsp_init_args.h \
|
|
||||||
$NVKERN_SRC/src/nvidia/inc/kernel/gpu/gsp/gsp_init_args.h \
|
|
||||||
$NVKERN_SRC/src/common/uproc/os/common/include/libos_init_args.h \
|
|
||||||
$NVKERN_SRC/src/nvidia/arch/nvalloc/common/inc/rmRiscvUcode.h \
|
|
||||||
$NVKERN_SRC/src/common/shared/msgq/inc/msgq/msgq_priv.h \
|
|
||||||
$NVKERN_SRC/src/nvidia/inc/kernel/vgpu/rpc_headers.h \
|
|
||||||
$NVKERN_SRC/src/nvidia/inc/kernel/vgpu/rpc_global_enums.h \
|
|
||||||
$NVKERN_SRC/src/nvidia/generated/g_rpc-structures.h \
|
|
||||||
$NVKERN_SRC/src/nvidia/arch/nvalloc/common/inc/fsp/fsp_nvdm_format.h \
|
|
||||||
extra/nv_gpu_driver/g_rpc-message-header.h \
|
|
||||||
extra/nv_gpu_driver/gsp_static_config.h \
|
|
||||||
extra/nv_gpu_driver/vbios.h \
|
|
||||||
extra/nv_gpu_driver/pci_exp_table.h \
|
|
||||||
--clang-args="-DRPC_MESSAGE_STRUCTURES -DRPC_STRUCTURES -include $NVKERN_SRC/src/common/sdk/nvidia/inc/nvtypes.h -I$NVKERN_SRC/src/nvidia/generated -I$NVKERN_SRC/src/common/inc -I$NVKERN_SRC/src/nvidia/inc -I$NVKERN_SRC/src/nvidia/interface/ -I$NVKERN_SRC/src/nvidia/inc/kernel -I$NVKERN_SRC/src/nvidia/inc/libraries -I$NVKERN_SRC/src/nvidia/arch/nvalloc/common/inc -I$NVKERN_SRC/kernel-open/nvidia-uvm -I$NVKERN_SRC/kernel-open/common/inc -I$NVKERN_SRC/src/common/sdk/nvidia/inc -I$NVKERN_SRC/src/nvidia/arch/nvalloc/unix/include -I$NVKERN_SRC/src/common/sdk/nvidia/inc/ctrl" \
|
|
||||||
-o $BASE/nv/nv.py
|
|
||||||
|
|
||||||
fixup $BASE/nv/nv.py
|
|
||||||
python3 -c "import tinygrad.runtime.autogen.nv.nv"
|
|
||||||
}
|
|
||||||
|
|
||||||
generate_amd() {
|
|
||||||
# clang2py broken when pass -x c++ to prev headers
|
|
||||||
clang2py -k cdefstum \
|
|
||||||
extra/hip_gpu_driver/sdma_registers.h \
|
|
||||||
extra/hip_gpu_driver/nvd.h \
|
|
||||||
extra/hip_gpu_driver/gc_11_0_0_offset.h \
|
|
||||||
extra/hip_gpu_driver/sienna_cichlid_ip_offset.h \
|
|
||||||
--clang-args="-I/opt/rocm/include -x c++" \
|
|
||||||
-o $BASE/amd_gpu.py
|
|
||||||
|
|
||||||
fixup $BASE/amd_gpu.py
|
|
||||||
sed -i "s\import ctypes\import ctypes, os\g" $BASE/amd_gpu.py
|
|
||||||
python3 -c "import tinygrad.runtime.autogen.amd_gpu"
|
|
||||||
}
|
|
||||||
|
|
||||||
generate_hsa() {
|
|
||||||
clang2py \
|
|
||||||
/opt/rocm/include/hsa/hsa.h \
|
|
||||||
/opt/rocm/include/hsa/hsa_ext_amd.h \
|
|
||||||
/opt/rocm/include/hsa/amd_hsa_signal.h \
|
|
||||||
/opt/rocm/include/hsa/amd_hsa_queue.h \
|
|
||||||
/opt/rocm/include/hsa/amd_hsa_kernel_code.h \
|
|
||||||
/opt/rocm/include/hsa/hsa_ext_finalize.h /opt/rocm/include/hsa/hsa_ext_image.h \
|
|
||||||
/opt/rocm/include/hsa/hsa_ven_amd_aqlprofile.h \
|
|
||||||
--clang-args="-I/opt/rocm/include" \
|
|
||||||
-o $BASE/hsa.py -l /opt/rocm/lib/libhsa-runtime64.so
|
|
||||||
|
|
||||||
fixup $BASE/hsa.py
|
|
||||||
sed -i "s\import ctypes\import ctypes, ctypes.util, os\g" $BASE/hsa.py
|
|
||||||
sed -i "s\ctypes.CDLL('/opt/rocm/lib/libhsa-runtime64.so')\ctypes.CDLL(os.getenv('ROCM_PATH')+'/lib/libhsa-runtime64.so' if os.getenv('ROCM_PATH') else ctypes.util.find_library('hsa-runtime64'))\g" $BASE/hsa.py
|
|
||||||
python3 -c "import tinygrad.runtime.autogen.hsa"
|
|
||||||
}
|
|
||||||
|
|
||||||
generate_io_uring() {
|
|
||||||
clang2py -k cdefstum \
|
|
||||||
/usr/include/liburing.h \
|
|
||||||
/usr/include/linux/io_uring.h \
|
|
||||||
-o $BASE/io_uring.py
|
|
||||||
|
|
||||||
sed -r '/^#define __NR_io_uring/ s/^#define __(NR_io_uring[^ ]+) (.*)$/\1 = \2/; t; d' /usr/include/asm-generic/unistd.h >> $BASE/io_uring.py # io_uring syscalls numbers
|
|
||||||
fixup $BASE/io_uring.py
|
|
||||||
}
|
|
||||||
|
|
||||||
generate_ib() {
|
|
||||||
clang2py -k cdefstum \
|
|
||||||
/usr/include/infiniband/verbs.h \
|
|
||||||
/usr/include/infiniband/verbs_api.h \
|
|
||||||
/usr/include/infiniband/ib_user_ioctl_verbs.h \
|
|
||||||
/usr/include/rdma/ib_user_verbs.h \
|
|
||||||
-o $BASE/ib.py
|
|
||||||
|
|
||||||
sed -i "s\import ctypes\import ctypes, ctypes.util\g" "$BASE/ib.py"
|
|
||||||
sed -i "s\FIXME_STUB\libibverbs\g" "$BASE/ib.py"
|
|
||||||
sed -i "s\FunctionFactoryStub()\ctypes.CDLL(ctypes.util.find_library('ibverbs'), use_errno=True)\g" "$BASE/ib.py"
|
|
||||||
|
|
||||||
fixup $BASE/ib.py
|
|
||||||
}
|
|
||||||
|
|
||||||
generate_libc() {
|
|
||||||
clang2py -k cdefstum \
|
|
||||||
$(dpkg -L libc6-dev | grep sys/mman.h) \
|
|
||||||
$(dpkg -L libc6-dev | grep sys/syscall.h) \
|
|
||||||
/usr/include/string.h \
|
|
||||||
/usr/include/elf.h \
|
|
||||||
/usr/include/unistd.h \
|
|
||||||
/usr/include/asm-generic/mman-common.h \
|
|
||||||
-o $BASE/libc.py
|
|
||||||
|
|
||||||
sed -i "s\import ctypes\import ctypes, ctypes.util, os\g" $BASE/libc.py
|
|
||||||
sed -i "s\FIXME_STUB\libc\g" $BASE/libc.py
|
|
||||||
sed -i "s\FunctionFactoryStub()\None if (libc_path := ctypes.util.find_library('c')) is None else ctypes.CDLL(libc_path, use_errno=True)\g" $BASE/libc.py
|
|
||||||
|
|
||||||
fixup $BASE/libc.py
|
|
||||||
}
|
|
||||||
|
|
||||||
generate_llvm() {
|
|
||||||
INC="$(llvm-config-14 --includedir)"
|
|
||||||
clang2py -k cdefstum \
|
|
||||||
$(find "$INC/llvm-c/" -type f -name '*.h' | sort) \
|
|
||||||
"$INC/llvm/Config/Targets.def" \
|
|
||||||
"$INC/llvm/Config/AsmPrinters.def" \
|
|
||||||
"$INC/llvm/Config/AsmParsers.def" \
|
|
||||||
"$INC/llvm/Config/Disassemblers.def" \
|
|
||||||
--clang-args="$(llvm-config-14 --cflags)" \
|
|
||||||
-o "$BASE/llvm.py"
|
|
||||||
|
|
||||||
sed -i "s\import ctypes\import ctypes, tinygrad.runtime.support.llvm as llvm_support\g" "$BASE/llvm.py"
|
|
||||||
sed -i "s\FIXME_STUB\llvm\g" "$BASE/llvm.py"
|
|
||||||
sed -i "s\FunctionFactoryStub()\ctypes.CDLL(llvm_support.LLVM_PATH)\g" "$BASE/llvm.py"
|
|
||||||
|
|
||||||
fixup "$BASE/llvm.py"
|
|
||||||
}
|
|
||||||
|
|
||||||
generate_kgsl() {
|
|
||||||
clang2py extra/qcom_gpu_driver/msm_kgsl.h -o $BASE/kgsl.py -k cdefstum
|
|
||||||
fixup $BASE/kgsl.py
|
|
||||||
sed -i "s\import ctypes\import ctypes, os\g" $BASE/kgsl.py
|
|
||||||
sed -nE 's/#define ([A-Za-z0-9_]+)_SHIFT\s*[^\S\r\n]*[0-9]*$/def \1(val): return (val << \1_SHIFT) \& \1_MASK/p' extra/qcom_gpu_driver/msm_kgsl.h >> $BASE/kgsl.py
|
|
||||||
sed -i "s\fcntl.ioctl(__fd, (__idir<<30)\__fd.ioctl((__idir<<30)\g" $BASE/kgsl.py
|
|
||||||
python3 -c "import tinygrad.runtime.autogen.kgsl"
|
|
||||||
}
|
|
||||||
|
|
||||||
generate_adreno() {
|
|
||||||
clang2py extra/qcom_gpu_driver/a6xx.xml.h -o $BASE/adreno.py -k cestum
|
|
||||||
sed -nE 's/#define ([A-Za-z0-9_]+)__SHIFT\s*[^\S\r\n]*[0-9]*$/def \1(val): return (val << \1__SHIFT) \& \1__MASK/p' extra/qcom_gpu_driver/a6xx.xml.h >> $BASE/adreno.py
|
|
||||||
fixup $BASE/adreno.py
|
|
||||||
sed -i "s\import ctypes\import ctypes, os\g" $BASE/adreno.py
|
|
||||||
python3 -c "import tinygrad.runtime.autogen.adreno"
|
|
||||||
}
|
|
||||||
|
|
||||||
generate_qcom() {
|
|
||||||
clang2py -k cdefstum \
|
|
||||||
extra/dsp/include/ion.h \
|
|
||||||
extra/dsp/include/msm_ion.h \
|
|
||||||
extra/dsp/include/adsprpc_shared.h \
|
|
||||||
extra/dsp/include/remote_default.h \
|
|
||||||
extra/dsp/include/apps_std.h \
|
|
||||||
-o $BASE/qcom_dsp.py
|
|
||||||
|
|
||||||
fixup $BASE/qcom_dsp.py
|
|
||||||
python3 -c "import tinygrad.runtime.autogen.qcom_dsp"
|
|
||||||
}
|
|
||||||
|
|
||||||
generate_pci() {
|
|
||||||
clang2py -k cdefstum \
|
|
||||||
/usr/include/linux/pci_regs.h \
|
|
||||||
-o $BASE/pci.py
|
|
||||||
fixup $BASE/pci.py
|
|
||||||
}
|
|
||||||
|
|
||||||
generate_vfio() {
|
|
||||||
clang2py -k cdefstum \
|
|
||||||
/usr/include/linux/vfio.h \
|
|
||||||
-o $BASE/vfio.py
|
|
||||||
fixup $BASE/vfio.py
|
|
||||||
sed -i "s\import ctypes\import ctypes, os\g" $BASE/vfio.py
|
|
||||||
sed -i "s\import fcntl, functools\import functools" $BASE/vfio.py
|
|
||||||
sed -i "s\import ctypes,os\a from tinygrad.runtime.support import FileIOInterface\g" $BASE/vfio.py
|
|
||||||
sed -i "s\fcntl.ioctl(__fd, (__idir<<30)\return __fd.ioctl((__idir<<30)\g" $BASE/vfio.py
|
|
||||||
}
|
|
||||||
|
|
||||||
generate_am() {
|
|
||||||
AMKERN_COMMIT_HASH=ceb12c04e2b5b53ec0779362831f5ee40c4921e4
|
|
||||||
AMKERN_SRC=/tmp/ROCK-Kernel-Driver-$AMKERN_COMMIT_HASH
|
|
||||||
if [ ! -d "$AMKERN_SRC" ]; then
|
|
||||||
git clone https://github.com/ROCm/ROCK-Kernel-Driver $AMKERN_SRC --depth 1
|
|
||||||
fi
|
|
||||||
AMKERN_AMD=$AMKERN_SRC/drivers/gpu/drm/amd/
|
|
||||||
AMKERN_INC=$AMKERN_AMD/include/
|
|
||||||
|
|
||||||
clang2py -k cdefstum \
|
|
||||||
extra/amdpci/headers/v11_structs.h \
|
|
||||||
extra/amdpci/headers/v12_structs.h \
|
|
||||||
extra/amdpci/headers/amdgpu_vm.h \
|
|
||||||
extra/amdpci/headers/discovery.h \
|
|
||||||
extra/amdpci/headers/amdgpu_ucode.h \
|
|
||||||
extra/amdpci/headers/psp_gfx_if.h \
|
|
||||||
extra/amdpci/headers/amdgpu_psp.h \
|
|
||||||
extra/amdpci/headers/amdgpu_irq.h \
|
|
||||||
extra/amdpci/headers/amdgpu_doorbell.h \
|
|
||||||
$AMKERN_INC/soc15_ih_clientid.h \
|
|
||||||
--clang-args="-include stdint.h" \
|
|
||||||
-o $BASE/am/am.py
|
|
||||||
fixup $BASE/am/am.py
|
|
||||||
sed -i "s\(int64_t)\ \g" $BASE/am/am.py
|
|
||||||
sed -i "s\AMDGPU_PTE_MTYPE_VG10(2)\AMDGPU_PTE_MTYPE_VG10(0, 2)\g" $BASE/am/am.py # incorrect parsing (TODO: remove when clang2py is gone).
|
|
||||||
|
|
||||||
clang2py -k cdefstum \
|
|
||||||
$AMKERN_AMD/amdkfd/kfd_pm4_headers_ai.h \
|
|
||||||
$AMKERN_AMD/amdgpu/soc15d.h \
|
|
||||||
-o $BASE/am/pm4_soc15.py
|
|
||||||
fixup $BASE/am/pm4_soc15.py
|
|
||||||
|
|
||||||
clang2py -k cdefstum \
|
|
||||||
$AMKERN_AMD/amdkfd/kfd_pm4_headers_ai.h \
|
|
||||||
$AMKERN_AMD/amdgpu/nvd.h \
|
|
||||||
-o $BASE/am/pm4_nv.py
|
|
||||||
fixup $BASE/am/pm4_nv.py
|
|
||||||
|
|
||||||
clang2py -k cdefstum \
|
|
||||||
extra/hip_gpu_driver/sdma_registers.h \
|
|
||||||
$AMKERN_AMD/amdgpu/vega10_sdma_pkt_open.h \
|
|
||||||
--clang-args="-I/opt/rocm/include -x c++" \
|
|
||||||
-o $BASE/am/sdma_4_0_0.py
|
|
||||||
fixup $BASE/am/sdma_4_0_0.py
|
|
||||||
|
|
||||||
clang2py -k cdefstum \
|
|
||||||
extra/hip_gpu_driver/sdma_registers.h \
|
|
||||||
$AMKERN_AMD/amdgpu/navi10_sdma_pkt_open.h \
|
|
||||||
--clang-args="-I/opt/rocm/include -x c++" \
|
|
||||||
-o $BASE/am/sdma_5_0_0.py
|
|
||||||
fixup $BASE/am/sdma_5_0_0.py
|
|
||||||
|
|
||||||
clang2py -k cdefstum \
|
|
||||||
extra/hip_gpu_driver/sdma_registers.h \
|
|
||||||
$AMKERN_AMD/amdgpu/sdma_v6_0_0_pkt_open.h \
|
|
||||||
--clang-args="-I/opt/rocm/include -x c++" \
|
|
||||||
-o $BASE/am/sdma_6_0_0.py
|
|
||||||
fixup $BASE/am/sdma_6_0_0.py
|
|
||||||
|
|
||||||
clang2py -k cdefstum \
|
|
||||||
$AMKERN_AMD/pm/swsmu/inc/pmfw_if/smu_v13_0_0_ppsmc.h \
|
|
||||||
$AMKERN_AMD/pm/swsmu/inc/pmfw_if/smu13_driver_if_v13_0_0.h \
|
|
||||||
extra/amdpci/headers/amdgpu_smu.h \
|
|
||||||
-o $BASE/am/smu_v13_0_0.py
|
|
||||||
fixup $BASE/am/smu_v13_0_0.py
|
|
||||||
|
|
||||||
clang2py -k cdefstum \
|
|
||||||
$AMKERN_AMD/pm/swsmu/inc/pmfw_if/smu_v14_0_0_pmfw.h \
|
|
||||||
$AMKERN_AMD/pm/swsmu/inc/pmfw_if/smu_v14_0_2_ppsmc.h \
|
|
||||||
$AMKERN_AMD/pm/swsmu/inc/pmfw_if/smu14_driver_if_v14_0.h \
|
|
||||||
extra/amdpci/headers/amdgpu_smu.h \
|
|
||||||
--clang-args="-include stdint.h" \
|
|
||||||
-o $BASE/am/smu_v14_0_2.py
|
|
||||||
fixup $BASE/am/smu_v14_0_2.py
|
|
||||||
}
|
|
||||||
|
|
||||||
generate_sqtt() {
|
|
||||||
clang2py -k cdefstum \
|
|
||||||
extra/sqtt/sqtt.h \
|
|
||||||
-o $BASE/sqtt.py
|
|
||||||
fixup $BASE/sqtt.py
|
|
||||||
sed -i "s\import ctypes\import ctypes, os\g" $BASE/sqtt.py
|
|
||||||
python3 -c "import tinygrad.runtime.autogen.sqtt"
|
|
||||||
|
|
||||||
ROCPROF_COMMIT_HASH=dd0485100971522cc4cd8ae136bdda431061a04d
|
|
||||||
ROCPROF_SRC=/tmp/rocprof-trace-decoder-$ROCPROF_COMMIT_HASH
|
|
||||||
if [ ! -d "$ROCPROF_SRC" ]; then
|
|
||||||
git clone https://github.com/ROCm/rocprof-trace-decoder $ROCPROF_SRC
|
|
||||||
pushd .
|
|
||||||
cd $ROCPROF_SRC
|
|
||||||
git reset --hard $ROCPROF_COMMIT_HASH
|
|
||||||
popd
|
|
||||||
fi
|
|
||||||
|
|
||||||
clang2py -k cdefstum \
|
|
||||||
$ROCPROF_SRC/include/rocprof_trace_decoder.h \
|
|
||||||
$ROCPROF_SRC/include/trace_decoder_instrument.h \
|
|
||||||
$ROCPROF_SRC/include/trace_decoder_types.h \
|
|
||||||
-o $BASE/rocprof.py
|
|
||||||
fixup $BASE/rocprof.py
|
|
||||||
sed -i '1s/^/# pylint: skip-file\n/' $BASE/rocprof.py
|
|
||||||
sed -i "s/import ctypes/import ctypes, ctypes.util/g" $BASE/rocprof.py
|
|
||||||
patch_dlopen $BASE/rocprof.py rocprof-trace-decoder "'/usr/local/lib/librocprof-trace-decoder.so'" "'/usr/local/lib/librocprof-trace-decoder.dylib'"
|
|
||||||
sed -i "s/def _try_dlopen_rocprof-trace-decoder():/def _try_dlopen_rocprof_trace_decoder():/g" $BASE/rocprof.py
|
|
||||||
sed -i "s|FunctionFactoryStub()|_try_dlopen_rocprof_trace_decoder()|g" $BASE/rocprof.py
|
|
||||||
}
|
|
||||||
|
|
||||||
generate_webgpu() {
|
|
||||||
clang2py extra/webgpu/webgpu.h -o $BASE/webgpu.py
|
|
||||||
fixup $BASE/webgpu.py
|
|
||||||
sed -i "s/FIXME_STUB/webgpu/g" "$BASE/webgpu.py"
|
|
||||||
sed -i "s/FunctionFactoryStub()/ctypes.CDLL(webgpu_support.WEBGPU_PATH)/g" "$BASE/webgpu.py"
|
|
||||||
sed -i "s/import ctypes/import ctypes, tinygrad.runtime.support.webgpu as webgpu_support/g" "$BASE/webgpu.py"
|
|
||||||
python3 -c "import tinygrad.runtime.autogen.webgpu"
|
|
||||||
}
|
|
||||||
|
|
||||||
generate_libusb() {
|
|
||||||
clang2py -k cdefstum \
|
|
||||||
/usr/include/libusb-1.0/libusb.h \
|
|
||||||
-o $BASE/libusb.py
|
|
||||||
|
|
||||||
fixup $BASE/libusb.py
|
|
||||||
sed -i "s\import ctypes\import ctypes, ctypes.util, os\g" $BASE/libusb.py
|
|
||||||
sed -i "s/FIXME_STUB/libusb/g" "$BASE/libusb.py"
|
|
||||||
sed -i "s/libusb_le16_to_cpu = libusb_cpu_to_le16//g" "$BASE/libusb.py"
|
|
||||||
sed -i "s/FunctionFactoryStub()/None if (lib_path:=os.getenv('LIBUSB_PATH', ctypes.util.find_library('usb-1.0'))) is None else ctypes.CDLL(lib_path)/g" "$BASE/libusb.py"
|
|
||||||
python3 -c "import tinygrad.runtime.autogen.libusb"
|
|
||||||
}
|
|
||||||
|
|
||||||
generate_mesa() {
|
|
||||||
MESA_TAG="mesa-25.2.4"
|
|
||||||
MESA_SRC=/tmp/mesa-$MESA_TAG
|
|
||||||
TINYMESA_TAG=tinymesa-32dc66c
|
|
||||||
TINYMESA_DIR=/tmp/tinymesa-$MESA_TAG-$TINYMESA_TAG/
|
|
||||||
TINYMESA_SO=$TINYMESA_DIR/libtinymesa_cpu.so
|
|
||||||
if [ ! -d "$MESA_SRC" ]; then
|
|
||||||
git clone --depth 1 --branch $MESA_TAG https://gitlab.freedesktop.org/mesa/mesa.git $MESA_SRC
|
|
||||||
pushd .
|
|
||||||
cd $MESA_SRC
|
|
||||||
git reset --hard $MESA_COMMIT_HASH
|
|
||||||
# clang 14 doesn't support packed enums
|
|
||||||
sed -i "s/enum \w\+ \(\w\+\);$/uint8_t \1;/" $MESA_SRC/src/nouveau/headers/nv_device_info.h
|
|
||||||
sed -i "s/enum \w\+ \(\w\+\);$/uint8_t \1;/" $MESA_SRC/src/nouveau/compiler/nak.h
|
|
||||||
sed -i "s/nir_instr_type \(\w\+\);/uint8_t \1;/" $MESA_SRC/src/compiler/nir/nir.h
|
|
||||||
mkdir -p gen/util/format
|
|
||||||
python3 src/util/format/u_format_table.py src/util/format/u_format.yaml --enums > gen/util/format/u_format_gen.h
|
|
||||||
python3 src/compiler/nir/nir_opcodes_h.py > gen/nir_opcodes.h
|
|
||||||
python3 src/compiler/nir/nir_intrinsics_h.py --outdir gen
|
|
||||||
python3 src/compiler/nir/nir_intrinsics_indices_h.py --outdir gen
|
|
||||||
python3 src/compiler/nir/nir_builder_opcodes_h.py > gen/nir_builder_opcodes.h
|
|
||||||
python3 src/compiler/nir/nir_intrinsics_h.py --outdir gen
|
|
||||||
python3 src/compiler/builtin_types_h.py gen/builtin_types.h
|
|
||||||
popd
|
|
||||||
fi
|
|
||||||
|
|
||||||
if [ ! -d "$TINYMESA_DIR" ]; then
|
|
||||||
mkdir $TINYMESA_DIR
|
|
||||||
curl -L https://github.com/sirhcm/tinymesa/releases/download/$TINYMESA_TAG/libtinymesa_cpu-$MESA_TAG-linux-amd64.so -o $TINYMESA_SO
|
|
||||||
fi
|
|
||||||
|
|
||||||
clang2py -k cdefstu \
|
|
||||||
$MESA_SRC/src/compiler/nir/nir.h \
|
|
||||||
$MESA_SRC/src/compiler/nir/nir_builder.h \
|
|
||||||
$MESA_SRC/src/compiler/nir/nir_shader_compiler_options.h \
|
|
||||||
$MESA_SRC/src/compiler/nir/nir_serialize.h \
|
|
||||||
$MESA_SRC/gen/nir_intrinsics.h \
|
|
||||||
$MESA_SRC/src/nouveau/headers/nv_device_info.h \
|
|
||||||
$MESA_SRC/src/nouveau/compiler/nak.h \
|
|
||||||
$MESA_SRC/src/gallium/auxiliary/gallivm/lp_bld.h \
|
|
||||||
$MESA_SRC/src/gallium/auxiliary/gallivm/lp_bld_passmgr.h \
|
|
||||||
$MESA_SRC/src/gallium/auxiliary/gallivm/lp_bld_misc.h \
|
|
||||||
$MESA_SRC/src/gallium/auxiliary/gallivm/lp_bld_type.h \
|
|
||||||
$MESA_SRC/src/gallium/auxiliary/gallivm/lp_bld_init.h \
|
|
||||||
$MESA_SRC/src/gallium/auxiliary/gallivm/lp_bld_nir.h \
|
|
||||||
$MESA_SRC/src/gallium/auxiliary/gallivm/lp_bld_struct.h \
|
|
||||||
$MESA_SRC/src/gallium/auxiliary/gallivm/lp_bld_jit_types.h \
|
|
||||||
$MESA_SRC/src/gallium/auxiliary/gallivm/lp_bld_flow.h \
|
|
||||||
$MESA_SRC/src/gallium/auxiliary/gallivm/lp_bld_const.h \
|
|
||||||
$MESA_SRC/src/compiler/glsl_types.h \
|
|
||||||
$MESA_SRC/src/util/blob.h \
|
|
||||||
$MESA_SRC/src/util/ralloc.h \
|
|
||||||
--clang-args="-DHAVE_ENDIAN_H -DHAVE_STRUCT_TIMESPEC -DHAVE_PTHREAD -I$MESA_SRC/src -I$MESA_SRC/include -I$MESA_SRC/gen -I$MESA_SRC/src/compiler/nir -I$MESA_SRC/src/gallium/auxiliary -I$MESA_SRC/src/gallium/include -I$(llvm-config-20 --includedir)" \
|
|
||||||
-l $TINYMESA_SO \
|
|
||||||
-o $BASE/mesa.py
|
|
||||||
|
|
||||||
LVP_NIR_OPTIONS=$(./extra/mesa/lvp_nir_options.sh $MESA_SRC)
|
|
||||||
|
|
||||||
fixup $BASE/mesa.py
|
|
||||||
patch_dlopen $BASE/mesa.py tinymesa_cpu "(BASE:=os.getenv('MESA_PATH', f\"/usr{'/local/' if helpers.OSX else '/'}lib\"))+'/libtinymesa_cpu'+(EXT:='.dylib' if helpers.OSX else '.so')" "f'{BASE}/libtinymesa{EXT}'" "'/opt/homebrew/lib/libtinymesa_cpu.dylib'" "'/opt/homebrew/lib/libtinymesa.dylib'"
|
|
||||||
echo "lvp_nir_options = gzip.decompress(base64.b64decode('$LVP_NIR_OPTIONS'))" >> $BASE/mesa.py
|
|
||||||
sed -i "/in_dll/s/.*/try: &\nexcept (AttributeError, ValueError): pass/" $BASE/mesa.py
|
|
||||||
sed -i "s/import ctypes/import ctypes, ctypes.util, os, gzip, base64, subprocess, tinygrad.helpers as helpers/" $BASE/mesa.py
|
|
||||||
sed -i "s/ctypes.CDLL('.\+')/(dll := _try_dlopen_tinymesa_cpu())/" $BASE/mesa.py
|
|
||||||
echo "def __getattr__(nm): raise AttributeError('LLVMpipe requires tinymesa_cpu' if 'tinymesa_cpu' not in dll._name else f'attribute {nm} not found') if dll else FileNotFoundError(f'libtinymesa not found (MESA_PATH={BASE}). See https://github.com/sirhcm/tinymesa ($TINYMESA_TAG, $MESA_TAG)')" >> $BASE/mesa.py
|
|
||||||
sed -i "s/ctypes.glsl_base_type/glsl_base_type/" $BASE/mesa.py
|
|
||||||
# bitfield bug in clang2py
|
|
||||||
sed -i "s/('fp_fast_math', ctypes.c_bool, 9)/('fp_fast_math', ctypes.c_uint32, 9)/" $BASE/mesa.py
|
|
||||||
sed -i "s/('\(\w\+\)', pipe_shader_type, 8)/('\1', ctypes.c_ubyte)/" $BASE/mesa.py
|
|
||||||
sed -i "s/\([0-9]\+\)()/\1/" $BASE/mesa.py
|
|
||||||
sed -i '/struct_nir_builder._pack_ = 1 # source:False/d' "$BASE/mesa.py"
|
|
||||||
python3 -c "import tinygrad.runtime.autogen.mesa"
|
|
||||||
}
|
|
||||||
|
|
||||||
if [ "$1" == "opencl" ]; then generate_opencl
|
|
||||||
elif [ "$1" == "hip" ]; then generate_hip
|
|
||||||
elif [ "$1" == "comgr" ]; then generate_comgr
|
|
||||||
elif [ "$1" == "cuda" ]; then generate_cuda
|
|
||||||
elif [ "$1" == "nvrtc" ]; then generate_nvrtc
|
|
||||||
elif [ "$1" == "hsa" ]; then generate_hsa
|
|
||||||
elif [ "$1" == "kfd" ]; then generate_kfd
|
|
||||||
elif [ "$1" == "nv" ]; then generate_nv
|
|
||||||
elif [ "$1" == "amd" ]; then generate_amd
|
|
||||||
elif [ "$1" == "am" ]; then generate_am
|
|
||||||
elif [ "$1" == "sqtt" ]; then generate_sqtt
|
|
||||||
elif [ "$1" == "qcom" ]; then generate_qcom
|
|
||||||
elif [ "$1" == "io_uring" ]; then generate_io_uring
|
|
||||||
elif [ "$1" == "ib" ]; then generate_ib
|
|
||||||
elif [ "$1" == "libc" ]; then generate_libc
|
|
||||||
elif [ "$1" == "llvm" ]; then generate_llvm
|
|
||||||
elif [ "$1" == "kgsl" ]; then generate_kgsl
|
|
||||||
elif [ "$1" == "adreno" ]; then generate_adreno
|
|
||||||
elif [ "$1" == "pci" ]; then generate_pci
|
|
||||||
elif [ "$1" == "vfio" ]; then generate_vfio
|
|
||||||
elif [ "$1" == "webgpu" ]; then generate_webgpu
|
|
||||||
elif [ "$1" == "libusb" ]; then generate_libusb
|
|
||||||
elif [ "$1" == "mesa" ]; then generate_mesa
|
|
||||||
elif [ "$1" == "all" ]; then generate_opencl; generate_hip; generate_comgr; generate_cuda; generate_nvrtc; generate_hsa; generate_kfd; generate_nv; generate_amd; generate_io_uring; generate_libc; generate_am; generate_webgpu; generate_mesa
|
|
||||||
else echo "usage: $0 <type>"
|
|
||||||
fi
|
|
||||||
+2
-2
@@ -1,8 +1,6 @@
|
|||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from typing import List
|
from typing import List
|
||||||
import json, argparse, random, time, os
|
import json, argparse, random, time, os
|
||||||
import tiktoken
|
|
||||||
from tiktoken.load import load_tiktoken_bpe
|
|
||||||
from extra.models.llama import Transformer, convert_from_huggingface, convert_from_gguf, fix_bf16
|
from extra.models.llama import Transformer, convert_from_huggingface, convert_from_gguf, fix_bf16
|
||||||
from tinygrad.nn.state import safe_load, torch_load, load_state_dict, get_parameters, gguf_load
|
from tinygrad.nn.state import safe_load, torch_load, load_state_dict, get_parameters, gguf_load
|
||||||
from tinygrad import Tensor, dtypes, nn, Context, Device, GlobalCounters
|
from tinygrad import Tensor, dtypes, nn, Context, Device, GlobalCounters
|
||||||
@@ -12,6 +10,8 @@ from extra.bench_log import BenchEvent, WallTimeEvent
|
|||||||
class Tokenizer:
|
class Tokenizer:
|
||||||
pat_str = r"(?i:'s|'t|'re|'ve|'m|'ll|'d)|[^\r\n\p{L}\p{N}]?\p{L}+|\p{N}{1,3}| ?[^\s\p{L}\p{N}]+[\r\n]*|\s*[\r\n]+|\s+(?!\S)|\s+"
|
pat_str = r"(?i:'s|'t|'re|'ve|'m|'ll|'d)|[^\r\n\p{L}\p{N}]?\p{L}+|\p{N}{1,3}| ?[^\s\p{L}\p{N}]+[\r\n]*|\s*[\r\n]+|\s+(?!\S)|\s+"
|
||||||
def __init__(self, model_path: str):
|
def __init__(self, model_path: str):
|
||||||
|
import tiktoken
|
||||||
|
from tiktoken.load import load_tiktoken_bpe
|
||||||
mergeable_ranks = load_tiktoken_bpe(model_path)
|
mergeable_ranks = load_tiktoken_bpe(model_path)
|
||||||
self.num_base_tokens = len(mergeable_ranks)
|
self.num_base_tokens = len(mergeable_ranks)
|
||||||
special_tokens = [
|
special_tokens = [
|
||||||
|
|||||||
@@ -1,5 +1,6 @@
|
|||||||
import os, sys, pickle, time, re
|
import os, sys, pickle, time, re
|
||||||
import numpy as np
|
import numpy as np
|
||||||
|
if "JIT_BATCH_SIZE" not in os.environ: os.environ["JIT_BATCH_SIZE"] = "0"
|
||||||
|
|
||||||
from tinygrad import fetch, Tensor, TinyJit, Context, GlobalCounters, Device, dtypes
|
from tinygrad import fetch, Tensor, TinyJit, Context, GlobalCounters, Device, dtypes
|
||||||
from tinygrad.helpers import DEBUG, getenv
|
from tinygrad.helpers import DEBUG, getenv
|
||||||
|
|||||||
@@ -0,0 +1,34 @@
|
|||||||
|
#!/usr/bin/env python3
|
||||||
|
from tinygrad import Tensor, Device, GlobalCounters, Context, dtypes
|
||||||
|
from tinygrad.helpers import getenv, colored
|
||||||
|
|
||||||
|
SZ = 8_000_000_000
|
||||||
|
GPUS = getenv("GPUS", 4) # TODO: expose a way in tinygrad to access this
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
# create tensors
|
||||||
|
tens = [Tensor.ones(SZ, dtype=dtypes.uint8, device=f"{Device.DEFAULT}:{i}").contiguous() for i in range(GPUS)]
|
||||||
|
Tensor.realize(*tens)
|
||||||
|
|
||||||
|
bw = [[0.0]*GPUS for _ in range(GPUS)]
|
||||||
|
for i in range(GPUS):
|
||||||
|
for j in range(GPUS):
|
||||||
|
GlobalCounters.reset()
|
||||||
|
with Context(DEBUG=2):
|
||||||
|
if i == j:
|
||||||
|
# this copy would be optimized out, just add 1
|
||||||
|
(tens[i]+1).realize()
|
||||||
|
else:
|
||||||
|
tens[i].to(f"{Device.DEFAULT}:{j}").realize()
|
||||||
|
t = max(GlobalCounters.time_sum_s, 1e-9)
|
||||||
|
bw[i][j] = SZ / t / 1e9 # GB/s
|
||||||
|
|
||||||
|
def fmt(x):
|
||||||
|
c = "green" if x > 50 else "yellow" if x > 20 else "red"
|
||||||
|
return colored(f"{x:6.1f}", c)
|
||||||
|
|
||||||
|
# header
|
||||||
|
print(" " * 8 + " ".join(f"{'d'+str(j):>6}" for j in range(GPUS)))
|
||||||
|
# rows
|
||||||
|
for i in range(GPUS):
|
||||||
|
print(f"{'s'+str(i):>6} -> " + " ".join(fmt(x) for x in bw[i]))
|
||||||
@@ -4,9 +4,9 @@ from tinygrad.engine.realize import ExecItem, get_runner
|
|||||||
from tinygrad.dtype import AddrSpace
|
from tinygrad.dtype import AddrSpace
|
||||||
from tinygrad.helpers import getenv
|
from tinygrad.helpers import getenv
|
||||||
|
|
||||||
N = 4096
|
N = getenv("N", 4096)
|
||||||
M = K = N
|
M = K = N
|
||||||
run_count = 5
|
run_count = getenv("CNT", 5)
|
||||||
|
|
||||||
# ---------------------------
|
# ---------------------------
|
||||||
# launch/config constants
|
# launch/config constants
|
||||||
@@ -155,14 +155,15 @@ def test_matmul(sink:UOp, N=N):
|
|||||||
ets.append(ei.run(wait=True))
|
ets.append(ei.run(wait=True))
|
||||||
print(f"REAL TFLOPS {N * N * N * 2 / min(ets) * 1e-12:.2f}")
|
print(f"REAL TFLOPS {N * N * N * 2 / min(ets) * 1e-12:.2f}")
|
||||||
|
|
||||||
GlobalCounters.reset()
|
if getenv("VERIFY", 1):
|
||||||
with Context(DEBUG=2):
|
GlobalCounters.reset()
|
||||||
tc = (a @ b).realize()
|
with Context(DEBUG=2):
|
||||||
with Context(DEBUG=0):
|
tc = (a @ b).realize()
|
||||||
err = (hc - tc).square().mean().item()
|
with Context(DEBUG=0):
|
||||||
print(f"mean squared error {err}")
|
err = (hc - tc).square().mean().item()
|
||||||
if err > 1e-06:
|
print(f"mean squared error {err}")
|
||||||
raise RuntimeError("matmul is wrong!")
|
if err > 1e-06:
|
||||||
|
raise RuntimeError("matmul is wrong!")
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
test_matmul(hand_spec_kernel3(), N=N)
|
test_matmul(hand_spec_kernel3(), N=N)
|
||||||
|
|||||||
@@ -0,0 +1,141 @@
|
|||||||
|
import os
|
||||||
|
import numpy as np
|
||||||
|
np.set_printoptions(linewidth=1000000)
|
||||||
|
os.environ["AMD_LLVM"] = "0"
|
||||||
|
|
||||||
|
from tinygrad import Tensor, Context, dtypes, UOp, GlobalCounters
|
||||||
|
from tinygrad.helpers import DEBUG, getenv
|
||||||
|
from tinygrad.dtype import AddrSpace
|
||||||
|
from tinygrad.uop.ops import sint, AxisType, KernelInfo, Ops
|
||||||
|
|
||||||
|
WARP_SIZE = 64
|
||||||
|
|
||||||
|
# Reg tile sizes (tensor cores)
|
||||||
|
TC_M = 16
|
||||||
|
TC_N = 16
|
||||||
|
TC_K = 32
|
||||||
|
|
||||||
|
N,M,K = 4096,4096,4096
|
||||||
|
|
||||||
|
# Threadblock tile sizes (block-level tile of C that a block computes)
|
||||||
|
BLOCK_M = 64
|
||||||
|
BLOCK_N = 64
|
||||||
|
BLOCK_K = 64
|
||||||
|
|
||||||
|
WARPGROUP_SIZE = 1
|
||||||
|
BLOCK_M = BLOCK_M * WARPGROUP_SIZE
|
||||||
|
|
||||||
|
TID_SIZE = WARPGROUP_SIZE*WARP_SIZE
|
||||||
|
|
||||||
|
def copy(dest:UOp, src:UOp, rng:int, set=False, upcast=()):
|
||||||
|
assert dest.shape == src.shape
|
||||||
|
rngs = [UOp.range(s, rng+i, AxisType.UPCAST if i in upcast else AxisType.LOOP) for i,s in enumerate(src.shape)]
|
||||||
|
copy = dest[*rngs].store(src[*rngs]).end(*rngs)
|
||||||
|
return dest.after(copy) if set else copy
|
||||||
|
|
||||||
|
def compute_on_locals(acc:UOp, Asl:UOp, Bsl:UOp, rng:int, afters:tuple[UOp, ...], warpgroup, warp) -> UOp:
|
||||||
|
K_inner_loop = UOp.range(BLOCK_K//TC_K, rng, AxisType.REDUCE)
|
||||||
|
|
||||||
|
# load from locals into registers
|
||||||
|
Ar = UOp.placeholder((BLOCK_M//TC_M//WARPGROUP_SIZE,), dtypes.half.vec(8), slot=1, addrspace=AddrSpace.REG)
|
||||||
|
Br = UOp.placeholder((BLOCK_N//TC_N,), dtypes.half.vec(8), slot=2, addrspace=AddrSpace.REG)
|
||||||
|
|
||||||
|
M_load_loop = UOp.range(BLOCK_M//TC_M//WARPGROUP_SIZE, rng+10)
|
||||||
|
Asl = Asl.reshape(BLOCK_K//TC_K, TC_K, BLOCK_M//TC_M//WARPGROUP_SIZE, WARPGROUP_SIZE, TC_M)
|
||||||
|
load_rng = UOp.range(8, rng+11, axis_type=AxisType.UPCAST)
|
||||||
|
A_in = Asl[K_inner_loop, (warp//16)*8+load_rng, M_load_loop, warpgroup, warp%16].contract(load_rng)
|
||||||
|
Ar = Ar[M_load_loop].set(A_in, end=M_load_loop)
|
||||||
|
|
||||||
|
N_load_loop = UOp.range(BLOCK_N//TC_N, rng+20)
|
||||||
|
Bsl = Bsl.reshape(BLOCK_K//TC_K, TC_K, BLOCK_N//TC_N, TC_N)
|
||||||
|
load_rng = UOp.range(8, rng+21, axis_type=AxisType.UPCAST)
|
||||||
|
B_in = Bsl[K_inner_loop, (warp//16)*8+load_rng, N_load_loop, warp%16].contract(load_rng)
|
||||||
|
Br = Br[N_load_loop].set(B_in, end=N_load_loop)
|
||||||
|
|
||||||
|
M_inner_loop = UOp.range(BLOCK_M//TC_M//WARPGROUP_SIZE, rng+30)
|
||||||
|
N_inner_loop = UOp.range(BLOCK_N//TC_N, rng+31)
|
||||||
|
|
||||||
|
# load values
|
||||||
|
acc_after = acc.after(*afters, M_inner_loop, N_inner_loop, K_inner_loop)
|
||||||
|
acc_load = acc_after[N_inner_loop, M_inner_loop]
|
||||||
|
|
||||||
|
# do WMMA
|
||||||
|
wmma_arg = ('WMMA_16_16_32_half_float', (16, 16, 32), dtypes.half, dtypes.float, 'AMD', 64, ((), (), ((3, 2), (2, 2))), ())
|
||||||
|
out = UOp(Ops.WMMA, dtypes.float.vec(4), (Ar[M_inner_loop], Br[N_inner_loop], acc_load), arg=wmma_arg)
|
||||||
|
|
||||||
|
# store back the acc
|
||||||
|
acc_store = acc[N_inner_loop, M_inner_loop].store(out)
|
||||||
|
return acc_store.end(M_inner_loop, N_inner_loop, K_inner_loop)
|
||||||
|
|
||||||
|
def custom_gemm(C:UOp, A:UOp, B:UOp) -> UOp:
|
||||||
|
gx, gy = UOp.special(M//BLOCK_M, "gidx0"), UOp.special(N//BLOCK_N, "gidx1")
|
||||||
|
K_outer_loop = UOp.range(K//BLOCK_K, 0, AxisType.REDUCE)
|
||||||
|
|
||||||
|
# split out the globals into blocks
|
||||||
|
C = C.src[0].cast(dtypes.float.vec(4).ptr(C.ptrdtype.size)).reshape((M//BLOCK_M, BLOCK_M, N//BLOCK_N, BLOCK_N))
|
||||||
|
A = A.reshape((M//BLOCK_M, BLOCK_M, K//BLOCK_K, BLOCK_K))[gx, :, K_outer_loop, :]
|
||||||
|
B = B.reshape((K//BLOCK_K, BLOCK_K, N//BLOCK_N, BLOCK_N))[K_outer_loop, :, gy, :]
|
||||||
|
|
||||||
|
# ---------------------------
|
||||||
|
# GLOBAL -> LOCAL (As, Bs)
|
||||||
|
# ---------------------------
|
||||||
|
tid = UOp.special(TID_SIZE, "lidx0")
|
||||||
|
warpgroup, warp = tid//WARP_SIZE, tid%WARP_SIZE
|
||||||
|
|
||||||
|
A_view = A.reshape(-1, TID_SIZE, 8)
|
||||||
|
B_view = B.reshape(-1, TID_SIZE, 8)
|
||||||
|
|
||||||
|
# A: read BM x BK tiles (permute on store into locals)
|
||||||
|
As = UOp.placeholder((BLOCK_K, BLOCK_M), dtypes.half, slot=0, addrspace=AddrSpace.LOCAL).shrink_to(BLOCK_K, BLOCK_M)
|
||||||
|
As_view = As.reshape(-1, TID_SIZE, 8)
|
||||||
|
|
||||||
|
Bs = UOp.placeholder((BLOCK_K, BLOCK_N+4), dtypes.half, slot=1, addrspace=AddrSpace.LOCAL).shrink_to(BLOCK_K, BLOCK_N)
|
||||||
|
Bs_view = Bs.reshape(-1, TID_SIZE, 8)
|
||||||
|
|
||||||
|
outer_copy = UOp.range(A_view.shape[0], 100, AxisType.UPCAST)
|
||||||
|
inner_copy = UOp.range(A_view.shape[2], 101, AxisType.UPCAST)
|
||||||
|
As_store = As_view[outer_copy, tid, inner_copy].store(A_view[outer_copy, tid, inner_copy])
|
||||||
|
Bs_store = Bs_view[outer_copy, tid, inner_copy].store(B_view[outer_copy, tid, inner_copy])
|
||||||
|
|
||||||
|
if getenv("NOLOAD"):
|
||||||
|
As_store = As[0,0].store(0)
|
||||||
|
Bs_store = Bs[0,0].store(0)
|
||||||
|
|
||||||
|
# TODO: can we automate barrier?
|
||||||
|
barrier = UOp.barrier(UOp.group(As_store, Bs_store).end(outer_copy, inner_copy))
|
||||||
|
|
||||||
|
if getenv("COMPUTE"):
|
||||||
|
As, Bs = As.after(barrier), Bs.after(barrier)
|
||||||
|
|
||||||
|
acc = UOp.placeholder((BLOCK_N//TC_N, BLOCK_M//TC_M//WARPGROUP_SIZE), dtypes.float.vec(4), 0, AddrSpace.REG)
|
||||||
|
|
||||||
|
sink = compute_on_locals(acc, As, Bs, 200, afters=(barrier,), warpgroup=warpgroup, warp=warp)
|
||||||
|
sink = sink.end(K_outer_loop)
|
||||||
|
|
||||||
|
C_view = C[gx, :, gy, :].reshape(BLOCK_M//TC_M//WARPGROUP_SIZE, WARPGROUP_SIZE, TC_M, BLOCK_N//TC_N, TC_N)[:, warpgroup, warp%16, :, (warp//16)*4]
|
||||||
|
sink = copy(C_view, acc.after(sink), rng=300)
|
||||||
|
else:
|
||||||
|
sink = C.after(barrier.end(K_outer_loop))[0,0,0,0].store(As[0,0]+Bs[0,0])
|
||||||
|
|
||||||
|
return sink.sink(arg=KernelInfo(name="custom_gemm", opts_to_apply=())).simplify()
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
a = Tensor.randn(M, K, dtype=dtypes.half)
|
||||||
|
b = Tensor.randn(K, N, dtype=dtypes.half)
|
||||||
|
c = Tensor.empty(M, N, dtype=dtypes.float)
|
||||||
|
with Context(DEBUG=0): Tensor.realize(a,b)
|
||||||
|
|
||||||
|
|
||||||
|
GlobalCounters.reset()
|
||||||
|
with Context(DEBUG=max(2, DEBUG.value), DEVECTORIZE=2):
|
||||||
|
tst = Tensor.custom_kernel(c, a, b, fxn=custom_gemm)[0]
|
||||||
|
tst.realize()
|
||||||
|
print(f"{(N*M*K*2 / GlobalCounters.time_sum_s)*1e-12:.2f} REAL TFLOPS")
|
||||||
|
|
||||||
|
|
||||||
|
with Context(DEBUG=0):
|
||||||
|
ref = a.dot(b, dtype=dtypes.float)
|
||||||
|
ref.realize()
|
||||||
|
#print(ref.numpy())
|
||||||
|
#print(tst.numpy())
|
||||||
|
assert Tensor.isclose(ref, tst, atol=1e-2).all().item(), "matrix not close"
|
||||||
@@ -12,7 +12,7 @@ MPS = getenv("MPS", 0)
|
|||||||
if getenv("FP16_ACC"): torch.backends.cuda.matmul.allow_fp16_accumulation = True
|
if getenv("FP16_ACC"): torch.backends.cuda.matmul.allow_fp16_accumulation = True
|
||||||
|
|
||||||
for dtype in [torch.float32, torch.float16, torch.bfloat16]:
|
for dtype in [torch.float32, torch.float16, torch.bfloat16]:
|
||||||
for N in [256, 512, 1024, 2048, 4096]:
|
for N in [256, 512, 1024, 2048, 4096] + ([6144, 8192] if getenv("BIG") else []):
|
||||||
FLOPS = N*N*N*2
|
FLOPS = N*N*N*2
|
||||||
|
|
||||||
b = torch.rand((N,N), dtype=dtype)
|
b = torch.rand((N,N), dtype=dtype)
|
||||||
|
|||||||
@@ -0,0 +1,16 @@
|
|||||||
|
from tinygrad import Tensor, Device, TinyJit, dtypes
|
||||||
|
from tinygrad.helpers import getenv
|
||||||
|
|
||||||
|
GPUS = getenv("GPUS", 4) # TODO: expose a way in tinygrad to access this
|
||||||
|
N = 6144
|
||||||
|
|
||||||
|
@TinyJit
|
||||||
|
def many_matmul(A, B):
|
||||||
|
out = A
|
||||||
|
for _ in range(8): out = out@B
|
||||||
|
return out
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
A = Tensor.ones(GPUS, N, N, dtype=dtypes.half).shard(devices=tuple([f"{Device.DEFAULT}:{i}" for i in range(GPUS)]), axis=0).contiguous()
|
||||||
|
B = Tensor.ones(GPUS, N, N, dtype=dtypes.half).shard(devices=tuple([f"{Device.DEFAULT}:{i}" for i in range(GPUS)]), axis=0).contiguous()
|
||||||
|
while 1: many_matmul(A, B)
|
||||||
@@ -51,11 +51,15 @@ def create_report(dev, test, result, stdout, stderr):
|
|||||||
dmesg_output = subprocess.check_output(["sudo", "dmesg", "--ctime", "--color=never"], text=True)
|
dmesg_output = subprocess.check_output(["sudo", "dmesg", "--ctime", "--color=never"], text=True)
|
||||||
with open(dmesg_path, "w") as f: f.write(dmesg_output)
|
with open(dmesg_path, "w") as f: f.write(dmesg_output)
|
||||||
|
|
||||||
|
env_vars = " ".join(f"{k}={v}" for k, v in test.env.items())
|
||||||
|
reproduce_cmd = f"{env_vars} {test.cmd}"
|
||||||
|
|
||||||
summary_path = os.path.join(report_path, "summary.txt")
|
summary_path = os.path.join(report_path, "summary.txt")
|
||||||
with open(summary_path, "w") as f:
|
with open(summary_path, "w") as f:
|
||||||
f.write(f"Test: {test.name()}\n")
|
f.write(f"Test: {test.name()}\n")
|
||||||
f.write(f"Dev params: {vars(dev)}\n")
|
f.write(f"Dev params: {vars(dev)}\n")
|
||||||
f.write(f"Test params: {vars(test)}\n")
|
f.write(f"Test params: {vars(test)}\n")
|
||||||
|
f.write(f"Reproduce cmd: {reproduce_cmd}\n")
|
||||||
f.write(f"Exit Code: {result}\n")
|
f.write(f"Exit Code: {result}\n")
|
||||||
|
|
||||||
print(f"Crash report saved to {report_path}")
|
print(f"Crash report saved to {report_path}")
|
||||||
|
|||||||
@@ -19,5 +19,6 @@ trap 'rm -f "$TMP"' EXIT
|
|||||||
EOF
|
EOF
|
||||||
sed -n '/struct nir_shader_compiler_options/,/^}/{p;/^}/q}' $1/src/gallium/drivers/llvmpipe/lp_screen.c
|
sed -n '/struct nir_shader_compiler_options/,/^}/{p;/^}/q}' $1/src/gallium/drivers/llvmpipe/lp_screen.c
|
||||||
echo "int main(void) { write(1, &gallivm_nir_options, sizeof(gallivm_nir_options)); }"
|
echo "int main(void) { write(1, &gallivm_nir_options, sizeof(gallivm_nir_options)); }"
|
||||||
) | cc -x c -o $TMP - -I$1/src/compiler/nir -I$1/src -I$1/include && $TMP | gzip | base64 -w0
|
) | cc -x c -o $TMP - -I$1/src/compiler/nir -I$1/src -I$1/include || exit 1
|
||||||
|
|
||||||
|
printf 'lvp_nir_options = gzip.decompress(base64.b64decode("%s"))' $("$TMP" | gzip | base64 -w0)
|
||||||
|
|||||||
@@ -1,8 +1,10 @@
|
|||||||
import pathlib
|
import os, pathlib
|
||||||
|
|
||||||
|
# TODO: there is a timing bug without this
|
||||||
|
os.environ["AMD_AQL"] = "1"
|
||||||
|
|
||||||
from tinygrad.device import Device
|
from tinygrad.device import Device
|
||||||
from tinygrad.runtime.ops_amd import AMDProgram, HIPCompiler
|
from tinygrad.runtime.ops_amd import AMDProgram, HIPCompiler
|
||||||
import time
|
|
||||||
import os
|
|
||||||
|
|
||||||
NUM_WORKGROUPS = 96
|
NUM_WORKGROUPS = 96
|
||||||
WAVE_SIZE = 32
|
WAVE_SIZE = 32
|
||||||
@@ -32,7 +34,7 @@ def launchBenchmark(instruction, vgprIndices, dense=True, accum=False, extra="")
|
|||||||
src = src.replace("DIRECTIVE", DIRECTIVE)
|
src = src.replace("DIRECTIVE", DIRECTIVE)
|
||||||
lib = COMPILER.compile(src)
|
lib = COMPILER.compile(src)
|
||||||
fxn = AMDProgram(DEV, "matmul", lib)
|
fxn = AMDProgram(DEV, "matmul", lib)
|
||||||
elapsed = fxn(global_size=(NUM_WORKGROUPS,1,1), local_size=(WAVE_SIZE*NUM_WAVES,1,1), wait=True)
|
elapsed = min([fxn(global_size=(NUM_WORKGROUPS,1,1), local_size=(WAVE_SIZE*NUM_WAVES,1,1), wait=True) for _ in range(2)])
|
||||||
FLOPs = FLOPS_PER_MATMUL * NUM_WAVES * NUM_WORKGROUPS * INTERNAL_LOOP * INSTRUCTIONS_PER_LOOP
|
FLOPs = FLOPS_PER_MATMUL * NUM_WAVES * NUM_WORKGROUPS * INTERNAL_LOOP * INSTRUCTIONS_PER_LOOP
|
||||||
print(f"{instruction:<29} : {FLOPs/elapsed/10**12:.2f} T(FL)OPS")
|
print(f"{instruction:<29} : {FLOPs/elapsed/10**12:.2f} T(FL)OPS")
|
||||||
|
|
||||||
@@ -44,9 +46,9 @@ if __name__=="__main__":
|
|||||||
raise RuntimeError("Error while initiating AMD device")
|
raise RuntimeError("Error while initiating AMD device")
|
||||||
|
|
||||||
COMPILER = HIPCompiler(DEV.arch)
|
COMPILER = HIPCompiler(DEV.arch)
|
||||||
if DEV.arch in {'gfx1100', 'gfx1103'}:
|
if DEV.arch in {'gfx1100', 'gfx1103', 'gfx1151'}:
|
||||||
if DEV.arch == 'gfx1103':
|
if DEV.arch == 'gfx1103': NUM_WORKGROUPS = 8
|
||||||
NUM_WORKGROUPS = 8
|
if DEV.arch == 'gfx1151': NUM_WORKGROUPS = 40
|
||||||
launchBenchmark("v_wmma_bf16_16x16x16_bf16", (7,8,15))
|
launchBenchmark("v_wmma_bf16_16x16x16_bf16", (7,8,15))
|
||||||
launchBenchmark("v_wmma_f16_16x16x16_f16", (7,8,15))
|
launchBenchmark("v_wmma_f16_16x16x16_f16", (7,8,15))
|
||||||
launchBenchmark("v_wmma_f32_16x16x16_bf16", (7,8,15))
|
launchBenchmark("v_wmma_f32_16x16x16_bf16", (7,8,15))
|
||||||
|
|||||||
@@ -3,14 +3,14 @@
|
|||||||
.p2align 8
|
.p2align 8
|
||||||
.type matmul,@function
|
.type matmul,@function
|
||||||
matmul:
|
matmul:
|
||||||
s_mov_b32 s1, INTERNAL_LOOP
|
s_mov_b32 s1, INTERNAL_LOOP
|
||||||
s_mov_b32 s2, 0
|
s_mov_b32 s2, 0
|
||||||
inner_loop:
|
inner_loop:
|
||||||
INSTRUCTION
|
INSTRUCTION
|
||||||
s_sub_u32 s1, s1, 1
|
s_sub_u32 s1, s1, 1
|
||||||
s_cmp_lg_i32 s1, s2
|
s_cmp_lg_i32 s1, s2
|
||||||
s_cbranch_scc1 inner_loop
|
s_cbranch_scc1 inner_loop
|
||||||
s_endpgm
|
s_endpgm
|
||||||
|
|
||||||
.rodata
|
.rodata
|
||||||
.p2align 6
|
.p2align 6
|
||||||
|
|||||||
@@ -0,0 +1,53 @@
|
|||||||
|
/*
|
||||||
|
* NVIDIA_COPYRIGHT_BEGIN
|
||||||
|
*
|
||||||
|
* Copyright (c) 2022, NVIDIA CORPORATION. All rights reserved.
|
||||||
|
*
|
||||||
|
* NVIDIA CORPORATION and its licensors retain all intellectual property
|
||||||
|
* and proprietary rights in and to this software, related documentation
|
||||||
|
* and any modifications thereto. Any use, reproduction, disclosure or
|
||||||
|
* distribution of this software and related documentation without an express
|
||||||
|
* license agreement from NVIDIA CORPORATION is strictly prohibited.
|
||||||
|
*
|
||||||
|
* NVIDIA_COPYRIGHT_END
|
||||||
|
*/
|
||||||
|
|
||||||
|
#include <stdint.h>
|
||||||
|
#include <stdlib.h>
|
||||||
|
|
||||||
|
typedef enum {
|
||||||
|
NVJITLINK_SUCCESS = 0,
|
||||||
|
NVJITLINK_ERROR_UNRECOGNIZED_OPTION,
|
||||||
|
NVJITLINK_ERROR_MISSING_ARCH,
|
||||||
|
NVJITLINK_ERROR_INVALID_INPUT,
|
||||||
|
NVJITLINK_ERROR_PTX_COMPILE,
|
||||||
|
NVJITLINK_ERROR_NVVM_COMPILE,
|
||||||
|
NVJITLINK_ERROR_INTERNAL
|
||||||
|
} nvJitLinkResult;
|
||||||
|
|
||||||
|
typedef enum {
|
||||||
|
NVJITLINK_INPUT_NONE = 0,
|
||||||
|
NVJITLINK_INPUT_CUBIN = 1,
|
||||||
|
NVJITLINK_INPUT_PTX,
|
||||||
|
NVJITLINK_INPUT_LTOIR,
|
||||||
|
NVJITLINK_INPUT_FATBIN,
|
||||||
|
NVJITLINK_INPUT_OBJECT,
|
||||||
|
NVJITLINK_INPUT_LIBRARY
|
||||||
|
} nvJitLinkInputType;
|
||||||
|
|
||||||
|
typedef struct nvJitLink* nvJitLinkHandle;
|
||||||
|
|
||||||
|
nvJitLinkResult nvJitLinkCreate(nvJitLinkHandle *handle, uint32_t numOptions, const char **options);
|
||||||
|
nvJitLinkResult nvJitLinkDestroy(nvJitLinkHandle *handle);
|
||||||
|
nvJitLinkResult nvJitLinkAddData(nvJitLinkHandle handle, nvJitLinkInputType inputType, const void *data, size_t size, const char *name);
|
||||||
|
nvJitLinkResult nvJitLinkAddFile(nvJitLinkHandle handle, nvJitLinkInputType inputType, const char *fileName);
|
||||||
|
nvJitLinkResult nvJitLinkComplete(nvJitLinkHandle handle);
|
||||||
|
nvJitLinkResult nvJitLinkGetLinkedCubinSize(nvJitLinkHandle handle, size_t *size);
|
||||||
|
nvJitLinkResult nvJitLinkGetLinkedCubin(nvJitLinkHandle handle, void *cubin);
|
||||||
|
nvJitLinkResult nvJitLinkGetLinkedPtxSize(nvJitLinkHandle handle, size_t *size);
|
||||||
|
nvJitLinkResult nvJitLinkGetLinkedPtx(nvJitLinkHandle handle, char *ptx);
|
||||||
|
nvJitLinkResult nvJitLinkGetErrorLogSize(nvJitLinkHandle handle, size_t *size);
|
||||||
|
nvJitLinkResult nvJitLinkGetErrorLog(nvJitLinkHandle handle, char *log);
|
||||||
|
nvJitLinkResult nvJitLinkGetInfoLogSize(nvJitLinkHandle handle, size_t *size);
|
||||||
|
nvJitLinkResult nvJitLinkGetInfoLog(nvJitLinkHandle handle, char *log);
|
||||||
|
nvJitLinkResult nvJitLinkVersion(unsigned int *major, unsigned int *minor);
|
||||||
@@ -65,6 +65,8 @@
|
|||||||
#define NVCEC0_QMDV05_00_GRID_HEIGHT_RESUME MW(271:256)
|
#define NVCEC0_QMDV05_00_GRID_HEIGHT_RESUME MW(271:256)
|
||||||
#define NVCEC0_QMDV05_00_GRID_DEPTH_RESUME MW(287:272)
|
#define NVCEC0_QMDV05_00_GRID_DEPTH_RESUME MW(287:272)
|
||||||
#define NVCEC0_QMDV05_00_RELEASE_ENABLE(i) MW((288+(i)*16):(288+(i)*16))
|
#define NVCEC0_QMDV05_00_RELEASE_ENABLE(i) MW((288+(i)*16):(288+(i)*16))
|
||||||
|
#define NVCEC0_QMDV05_00_RELEASE0_ENABLE NVCEC0_QMDV05_00_RELEASE_ENABLE(0)
|
||||||
|
#define NVCEC0_QMDV05_00_RELEASE1_ENABLE NVCEC0_QMDV05_00_RELEASE_ENABLE(1)
|
||||||
#define NVCEC0_QMDV05_00_RELEASE_ENABLE_FALSE 0x00000000
|
#define NVCEC0_QMDV05_00_RELEASE_ENABLE_FALSE 0x00000000
|
||||||
#define NVCEC0_QMDV05_00_RELEASE_ENABLE_TRUE 0x00000001
|
#define NVCEC0_QMDV05_00_RELEASE_ENABLE_TRUE 0x00000001
|
||||||
#define NVCEC0_QMDV05_00_RELEASE_STRUCTURE_SIZE(i) MW((290+(i)*16):(289+(i)*16))
|
#define NVCEC0_QMDV05_00_RELEASE_STRUCTURE_SIZE(i) MW((290+(i)*16):(289+(i)*16))
|
||||||
|
|||||||
@@ -58,20 +58,23 @@ def install_hook(c_function, python_function):
|
|||||||
return orig_func
|
return orig_func
|
||||||
|
|
||||||
# *** ioctl lib end ***
|
# *** ioctl lib end ***
|
||||||
import tinygrad.runtime.autogen.nv_gpu as nv_gpu
|
from tinygrad.runtime.autogen import nv_570 as nv_gpu
|
||||||
nvescs = {getattr(nv_gpu, x):x for x in dir(nv_gpu) if x.startswith("NV_ESC")}
|
nvescs = {getattr(nv_gpu, x):x for x in dir(nv_gpu) if x.startswith("NV_ESC")}
|
||||||
nvcmds = {getattr(nv_gpu, x):(x, getattr(nv_gpu, "struct_"+x+"_PARAMS", getattr(nv_gpu, "struct_"+x.replace("_CMD_", "_")+"_PARAMS", None))) for x in dir(nv_gpu) if \
|
nvcmds = {getattr(nv_gpu, x):(x, getattr(nv_gpu, "struct_"+x+"_PARAMS", getattr(nv_gpu, "struct_"+x.replace("_CMD_", "_")+"_PARAMS", None))) for x in dir(nv_gpu) if \
|
||||||
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_gpu.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 = []
|
||||||
@@ -272,4 +275,4 @@ def compare_launch_state(states, good_states):
|
|||||||
|
|
||||||
return True, "PASS"
|
return True, "PASS"
|
||||||
|
|
||||||
# IOCTL=1 CUDA=1 CUDA_PTX=1 python3 test/test_ops.py TestOps.test_tiny_add
|
# IOCTL=1 CUDA=1 CUDA_PTX=1 python3 test/test_ops.py TestOps.test_tiny_add
|
||||||
|
|||||||
@@ -2,7 +2,8 @@ import os, pathlib, argparse
|
|||||||
from examples.llama3 import Tokenizer
|
from examples.llama3 import Tokenizer
|
||||||
from tabulate import tabulate
|
from tabulate import tabulate
|
||||||
from tinygrad import fetch
|
from tinygrad import fetch
|
||||||
from tinygrad.helpers import flatten
|
from tinygrad.helpers import flatten, getenv
|
||||||
|
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())
|
||||||
@@ -10,19 +11,15 @@ tokenizer = Tokenizer(fetch("https://huggingface.co/bofenghuang/Meta-Llama-3-8B/
|
|||||||
def read_code(base_path):
|
def read_code(base_path):
|
||||||
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
|
||||||
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
|
||||||
fullpath = os.path.join(path, name)
|
fullpath = os.path.join(path, name)
|
||||||
code = pathlib.Path(fullpath).read_text()
|
code = pathlib.Path(fullpath).read_text()
|
||||||
ret.append(("### " + fullpath.split("tinygrad/", 1)[1], code))
|
ret.append((fullpath.split("tinygrad/", 1)[1], code))
|
||||||
return ret
|
return ret
|
||||||
|
|
||||||
def write_code_to_file(filename, code_list):
|
|
||||||
"""Writes the combined code to a specified file."""
|
|
||||||
with open(filename, 'w') as f:
|
|
||||||
f.write('\n'.join(flatten(code_list)))
|
|
||||||
|
|
||||||
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.")
|
||||||
@@ -32,10 +29,11 @@ if __name__ == "__main__":
|
|||||||
|
|
||||||
table = []
|
table = []
|
||||||
for name,code in ret:
|
for name,code in ret:
|
||||||
table.append([name, len(tokenizer.encode(name+"\x00"+code))])
|
table.append([name, len(tokenizer.encode(code))])
|
||||||
print(tabulate([["name", "llm tokens"]]+sorted(table, key=lambda x: -x[1]), headers="firstrow"))
|
print(tabulate([["name", "llm tokens"]]+sorted(table, key=lambda x: -x[1]), headers="firstrow"))
|
||||||
|
|
||||||
code_str = '\x00'.join(flatten(ret))
|
banner = "#"*40
|
||||||
|
code_str = ''.join([f"{banner}\n# {name}\n{banner}\n\n{code}\n" for name,code in ret])
|
||||||
print(f"code has {len(code_str)} chars")
|
print(f"code has {len(code_str)} chars")
|
||||||
newline_count = code_str.count('\n')
|
newline_count = code_str.count('\n')
|
||||||
print(f"code has {newline_count} newlines")
|
print(f"code has {newline_count} newlines")
|
||||||
@@ -44,5 +42,5 @@ if __name__ == "__main__":
|
|||||||
print(f"code has {len(encoded)} tokens")
|
print(f"code has {len(encoded)} tokens")
|
||||||
|
|
||||||
if args.output:
|
if args.output:
|
||||||
write_code_to_file(args.output, ret)
|
with open(args.output, 'w') as f: f.write(code_str)
|
||||||
print(f"Combined code written to {args.output}")
|
print(f"Combined code written to {args.output}")
|
||||||
@@ -0,0 +1,131 @@
|
|||||||
|
import os
|
||||||
|
os.environ["PYTHONPATH"] = "."
|
||||||
|
os.environ["SQTT"] = "1"
|
||||||
|
if "DEV" not in os.environ: os.environ["DEV"] = "AMD"
|
||||||
|
os.environ["PROFILE"] = "1"
|
||||||
|
os.environ["AMD_LLVM"] = "0"
|
||||||
|
|
||||||
|
from dataclasses import replace
|
||||||
|
import atexit, contextlib
|
||||||
|
from tinygrad import Tensor
|
||||||
|
from tinygrad.helpers import system, OSX
|
||||||
|
from tinygrad.runtime.ops_amd import AMDProgram
|
||||||
|
from extra.sqtt.roc import decode, WaveExec, ProfileSQTTEvent
|
||||||
|
from tinygrad.device import Device, ProfileDeviceEvent
|
||||||
|
|
||||||
|
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"]
|
||||||
|
|
||||||
|
@contextlib.contextmanager
|
||||||
|
def save_sqtt():
|
||||||
|
# clear the old traces
|
||||||
|
dev.profile_events.clear()
|
||||||
|
sqtt:dict[str, list[WaveExec]] = {}
|
||||||
|
yield sqtt
|
||||||
|
events = dev.profile_events+[ProfileDeviceEvent("AMD", props=dev.device_props())]
|
||||||
|
|
||||||
|
rctx = decode(events)
|
||||||
|
assert len(rctx.inst_execs) > 0, "empty sqtt output"
|
||||||
|
sqtt.update(rctx.inst_execs)
|
||||||
|
|
||||||
|
for e in events:
|
||||||
|
if isinstance(e, ProfileSQTTEvent):
|
||||||
|
print(replace(e, blob=b''))
|
||||||
|
if e.se == 0:
|
||||||
|
parse_sqtt_print_packets(e.blob)
|
||||||
|
|
||||||
|
template = """.text
|
||||||
|
.globl matmul
|
||||||
|
.p2align 8
|
||||||
|
.type matmul,@function
|
||||||
|
matmul:
|
||||||
|
INSTRUCTION
|
||||||
|
s_endpgm
|
||||||
|
|
||||||
|
.rodata
|
||||||
|
.p2align 6
|
||||||
|
.amdhsa_kernel matmul
|
||||||
|
.amdhsa_user_sgpr_kernarg_segment_ptr 1
|
||||||
|
.amdhsa_next_free_vgpr .amdgcn.next_free_vgpr
|
||||||
|
.amdhsa_next_free_sgpr .amdgcn.next_free_sgpr
|
||||||
|
.amdhsa_wavefront_size32 1
|
||||||
|
.end_amdhsa_kernel
|
||||||
|
|
||||||
|
.amdgpu_metadata
|
||||||
|
---
|
||||||
|
amdhsa.version:
|
||||||
|
- 1
|
||||||
|
- 0
|
||||||
|
amdhsa.kernels:
|
||||||
|
- .name: matmul
|
||||||
|
.symbol: matmul.kd
|
||||||
|
.group_segment_fixed_size: 0
|
||||||
|
.private_segment_fixed_size: 0
|
||||||
|
.wavefront_size: 32
|
||||||
|
.sgpr_count: 8
|
||||||
|
.vgpr_count: 32
|
||||||
|
.max_flat_workgroup_size: 1024
|
||||||
|
.kernarg_segment_align: 8
|
||||||
|
.kernarg_segment_size: 8
|
||||||
|
.args:
|
||||||
|
- .address_space: global
|
||||||
|
.name: a
|
||||||
|
.offset: 0
|
||||||
|
.size: 8
|
||||||
|
.type_name: 'float*'
|
||||||
|
.value_kind: global_buffer
|
||||||
|
...
|
||||||
|
.end_amdgpu_metadata
|
||||||
|
"""
|
||||||
|
|
||||||
|
def run_asm(src):
|
||||||
|
NUM_WORKGROUPS = 1
|
||||||
|
WAVE_SIZE = 32
|
||||||
|
NUM_WAVES = 1
|
||||||
|
t = Tensor.empty(0x1000).realize()
|
||||||
|
buf = t.uop.buffer.ensure_allocated()
|
||||||
|
lib = dev.compiler.compile(template.replace("INSTRUCTION", '\n'.join(src)))
|
||||||
|
dev.compiler.disassemble(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)
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
with save_sqtt() as sqtt:
|
||||||
|
#(Tensor.empty(16,16) @ Tensor.empty(16,16)).elu().realize()
|
||||||
|
Tensor.empty(1).elu().realize()
|
||||||
|
exit(0)
|
||||||
|
|
||||||
|
with save_sqtt() as sqtt:
|
||||||
|
# what's in v0?
|
||||||
|
run_asm([
|
||||||
|
"v_mov_b32_e32 v0, 0",
|
||||||
|
"v_mov_b32_e32 v1, 0",
|
||||||
|
"s_clause 0x1",
|
||||||
|
"s_load_b64 s[0:1], s[0:1], null",
|
||||||
|
"s_waitcnt lgkmcnt(0)",
|
||||||
|
]+[
|
||||||
|
"global_load_b32 v1, v0, s[0:1]",
|
||||||
|
]*10+[
|
||||||
|
"global_load_b32 v10, v1, s[0:1]",
|
||||||
|
"s_waitcnt vmcnt(0)",
|
||||||
|
|
||||||
|
#"v_rcp_f32 v1, v0"
|
||||||
|
#"v_add_f32_e32 v1 v0 v0",
|
||||||
|
#"v_add_f32_e32 v5 v4 v4",
|
||||||
|
#"v_add_f32_e32 v7 v6 v6",
|
||||||
|
#"v_add_f32_e32 v1 v0 v0",
|
||||||
|
#"v_add_f32_e32 v2 v1 v1",
|
||||||
|
#"s_nop 1"
|
||||||
|
]*5+[
|
||||||
|
"v_add_f32_e32 v3 v2 v2",
|
||||||
|
]*5+[
|
||||||
|
"v_mul_f32_e32 v3 v2 v2",
|
||||||
|
]*7)
|
||||||
@@ -0,0 +1,543 @@
|
|||||||
|
import pickle
|
||||||
|
from tinygrad.helpers import getenv
|
||||||
|
from extra.sqtt.roc import decode, ProfileSQTTEvent
|
||||||
|
|
||||||
|
# Instruction packets (one per ISA op)
|
||||||
|
# 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
|
||||||
|
|
||||||
|
OPCODE_NAMES = {
|
||||||
|
# gated by SQ_TT_TOKEN_EXCLUDE_VMEMEXEC_SHIFT
|
||||||
|
0x02: "VMEMEXEC",
|
||||||
|
# gated by SQ_TT_TOKEN_EXCLUDE_ALUEXEC_SHIFT
|
||||||
|
0x03: "ALUEXEC",
|
||||||
|
# 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_WAVERDY_SHIFT
|
||||||
|
0x06: "WAVERDY",
|
||||||
|
# gated by SQ_TT_TOKEN_EXCLUDE_WAVESTARTEND_SHIFT
|
||||||
|
0x08: "WAVEEND",
|
||||||
|
0x09: "WAVESTART",
|
||||||
|
# gated by SQ_TT_TOKEN_EXCLUDE_IMMEDIATE_SHIFT
|
||||||
|
0x04: "IMMEDIATE_4",
|
||||||
|
0x05: "IMMEDIATE_5",
|
||||||
|
# some gated by SQ_TT_TOKEN_EXCLUDE_REG_SHIFT, some always there
|
||||||
|
0x14: "REG",
|
||||||
|
# gated by SQ_TT_TOKEN_EXCLUDE_EVENT_SHIFT
|
||||||
|
0x12: "EVENT",
|
||||||
|
# gated by SQ_TT_TOKEN_EXCLUDE_INST_SHIFT
|
||||||
|
0x18: "INST",
|
||||||
|
# gated by SQ_TT_TOKEN_EXCLUDE_UTILCTR_SHIFT
|
||||||
|
0x19: "UTILCTR",
|
||||||
|
|
||||||
|
# ------------------------------------------------------------------------
|
||||||
|
# 0x07–0x0F: pure timestamp-ish deltas
|
||||||
|
# ------------------------------------------------------------------------
|
||||||
|
0x07: "TS_DELTA_S8_W3", # shift=8, width=3 (small delta)
|
||||||
|
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)
|
||||||
|
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
|
||||||
|
|
||||||
|
# ------------------------------------------------------------------------
|
||||||
|
# 0x10–0x19: 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
|
||||||
|
}
|
||||||
|
|
||||||
|
# these tables are from rocprof trace decoder
|
||||||
|
# rocprof_trace_decoder_parse_data-0x11c6a0
|
||||||
|
# parse_sqtt_180 = b *rocprof_trace_decoder_parse_data-0x11c6a0+0x110040
|
||||||
|
|
||||||
|
# ---------- 1. local_138: 256-byte state->token table ----------
|
||||||
|
|
||||||
|
STATE_TO_TOKEN: bytes = bytes([
|
||||||
|
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, 0x07, 0x18, 0x01, 0x05, 0x0b, 0x0c, 0x00, 0x0f, 0x14, 0x18, 0x01, 0x09, 0x04, 0x03, 0x02,
|
||||||
|
0x10, 0x19, 0x18, 0x01, 0x06, 0x08, 0x0d, 0x00, 0x0f, 0x14, 0x18, 0x01, 0x0a, 0x04, 0x03, 0x02,
|
||||||
|
0x10, 0x00, 0x18, 0x01, 0x05, 0x0b, 0x0c, 0x00, 0x0f, 0x14, 0x18, 0x01, 0x09, 0x04, 0x03, 0x02,
|
||||||
|
0x10, 0x11, 0x18, 0x01, 0x06, 0x08, 0x0d, 0x00, 0x0f, 0x14, 0x18, 0x01, 0x0a, 0x04, 0x03, 0x02,
|
||||||
|
0x10, 0x12, 0x18, 0x01, 0x05, 0x0b, 0x0c, 0x00, 0x0f, 0x14, 0x18, 0x01, 0x09, 0x04, 0x03, 0x02,
|
||||||
|
0x10, 0x15, 0x18, 0x01, 0x06, 0x08, 0x0d, 0x00, 0x0f, 0x14, 0x18, 0x01, 0x0a, 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, 0x07, 0x18, 0x01, 0x05, 0x0b, 0x0c, 0x00, 0x0f, 0x14, 0x18, 0x01, 0x09, 0x04, 0x03, 0x02,
|
||||||
|
0x10, 0x19, 0x18, 0x01, 0x06, 0x08, 0x0d, 0x00, 0x0f, 0x14, 0x18, 0x01, 0x0a, 0x04, 0x03, 0x02,
|
||||||
|
0x10, 0x00, 0x18, 0x01, 0x05, 0x0b, 0x0c, 0x00, 0x0f, 0x14, 0x18, 0x01, 0x09, 0x04, 0x03, 0x02,
|
||||||
|
0x10, 0x11, 0x18, 0x01, 0x06, 0x08, 0x0d, 0x00, 0x0f, 0x14, 0x18, 0x01, 0x0a, 0x04, 0x03, 0x02,
|
||||||
|
0x10, 0x13, 0x18, 0x01, 0x05, 0x0b, 0x0c, 0x00, 0x0f, 0x14, 0x18, 0x01, 0x09, 0x04, 0x03, 0x02,
|
||||||
|
0x10, 0x15, 0x18, 0x01, 0x06, 0x08, 0x0d, 0x00, 0x0f, 0x14, 0x18, 0x01, 0x0a, 0x04, 0x03, 0x02,
|
||||||
|
])
|
||||||
|
|
||||||
|
|
||||||
|
# ---------- 2. DAT_0012e280: nibble budget per opcode&0x1F ----------
|
||||||
|
|
||||||
|
NIBBLE_BUDGET = [
|
||||||
|
0x08, 0x0C, 0x08, 0x08, 0x0C, 0x18, 0x18, 0x40,
|
||||||
|
0x14, 0x20, 0x30, 0x14, 0x34, 0x1C, 0x30, 0x08,
|
||||||
|
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 ----------
|
||||||
|
|
||||||
|
# opcode -> (shift, width)
|
||||||
|
DELTA_MAP_DEFAULT = {
|
||||||
|
0x01: (3, 3), # shift=3, end=6
|
||||||
|
0x02: (4, 2), # shift=4, end=6
|
||||||
|
0x03: (4, 2), # shift=4, end=6
|
||||||
|
0x04: (4, 3), # shift=4, end=7
|
||||||
|
0x05: (5, 3), # shift=5, end=8
|
||||||
|
0x06: (5, 3), # shift=5, end=8
|
||||||
|
0x07: (8, 3), # shift=8, end=11
|
||||||
|
0x08: (5, 3), # shift=5, end=8
|
||||||
|
0x09: (5, 2), # shift=5, end=7
|
||||||
|
0x0A: (5, 2), # shift=5, end=7
|
||||||
|
0x0B: (5, 3), # shift=5, end=8
|
||||||
|
0x0C: (5, 3), # shift=5, end=8
|
||||||
|
0x0D: (5, 3), # shift=5, end=8
|
||||||
|
0x0E: (7, 2), # shift=7, end=9
|
||||||
|
0x0F: (4, 4), # shift=4, end=8
|
||||||
|
0x10: (0, 0), # shift=0, end=0 (no delta)
|
||||||
|
0x11: (7, 9), # shift=7, end=16
|
||||||
|
0x12: (8, 3), # shift=8, end=11
|
||||||
|
0x13: (8, 3), # shift=8, end=11
|
||||||
|
0x14: (4, 3), # shift=4, end=7
|
||||||
|
0x15: (7, 3), # shift=7, end=10
|
||||||
|
0x16: (12, 36), # shift=12, end=48 (36-bit field, matches the 0x16 special-case)
|
||||||
|
0x17: (0, 0), # shift=0, end=0 (no delta)
|
||||||
|
0x18: (4, 3), # shift=4, end=7
|
||||||
|
0x19: (7, 2), # shift=7, end=9
|
||||||
|
}
|
||||||
|
|
||||||
|
# ---------- 4. One-line-per-packet parser ----------
|
||||||
|
|
||||||
|
def decode_packet_fields(opcode: int, reg: int, delta: int) -> str:
|
||||||
|
"""
|
||||||
|
Decode packet payloads conservatively, using:
|
||||||
|
- NIBBLE_BUDGET[opcode & 0x1F] to mask reg down to true width.
|
||||||
|
- DELTA_MAP_DEFAULT[opcode] to expose the "primary" field (often delta).
|
||||||
|
- Per-opcode layouts derived from rocprof's decompiled consumers.
|
||||||
|
"""
|
||||||
|
# --- 0. Restrict to real packet bits ---------------------------------
|
||||||
|
nb_bits = NIBBLE_BUDGET[opcode & 0x1F]
|
||||||
|
if nb_bits <= 0 or nb_bits >= 64:
|
||||||
|
pkt = reg & ((1 << 64) - 1)
|
||||||
|
else:
|
||||||
|
pkt = reg & ((1 << nb_bits) - 1)
|
||||||
|
|
||||||
|
fields: list[str] = []
|
||||||
|
|
||||||
|
shift, width = DELTA_MAP_DEFAULT.get(opcode, (0, 0))
|
||||||
|
if width:
|
||||||
|
field_mask = (1 << width) - 1
|
||||||
|
shaped_field = (pkt >> shift) & field_mask
|
||||||
|
else:
|
||||||
|
field_mask = 0
|
||||||
|
shaped_field = 0
|
||||||
|
|
||||||
|
# =====================================================================
|
||||||
|
# 1. Timestamp-centric opcodes (actually drive 'time')
|
||||||
|
# =====================================================================
|
||||||
|
|
||||||
|
if opcode == 0x0F: # TS_DELTA_SHORT_PLUS4
|
||||||
|
# In the caller, delta already has +4 applied.
|
||||||
|
raw_delta = shaped_field
|
||||||
|
fields.append(f"raw_delta={raw_delta}")
|
||||||
|
fields.append(f"ts_short_plus4={delta}")
|
||||||
|
return ", ".join(fields)
|
||||||
|
|
||||||
|
if opcode == 0x11: # TS_WAVE_STATE_SAMPLE
|
||||||
|
# DELTA_MAP_DEFAULT: shift=7, width=9 -> small delta.
|
||||||
|
raw_delta = shaped_field
|
||||||
|
coarse = (pkt >> (shift + width)) & 0xFF # matches byte at +10 in C
|
||||||
|
fields.append(f"raw_delta={raw_delta}")
|
||||||
|
if coarse:
|
||||||
|
fields.append(f"coarse_state=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")
|
||||||
|
return ", ".join(fields)
|
||||||
|
|
||||||
|
if opcode == 0x16: # TS_DELTA36_OR_MARK
|
||||||
|
# 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, 0x0A–0x0E, we know they drive time (via DELTA_MAP_DEFAULT),
|
||||||
|
# but we don't see any other fields used in the decomp.
|
||||||
|
if opcode in (0x07, 0x0A, 0x0B, 0x0C, 0x0D, 0x0E):
|
||||||
|
if width:
|
||||||
|
raw_delta = shaped_field
|
||||||
|
leftover = pkt & ~(field_mask << shift)
|
||||||
|
fields.append(f"raw_delta={raw_delta}")
|
||||||
|
if leftover:
|
||||||
|
fields.append(f"payload=0x{leftover:x}")
|
||||||
|
return ", ".join(fields)
|
||||||
|
|
||||||
|
# =====================================================================
|
||||||
|
# 2. Small "meta + tiny delta" packets (0x01–0x06)
|
||||||
|
# =====================================================================
|
||||||
|
|
||||||
|
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).
|
||||||
|
|
||||||
|
Now prints only the actual nibbles that belong to each packet, instead of
|
||||||
|
the full 64-bit shift register.
|
||||||
|
"""
|
||||||
|
n = len(data)
|
||||||
|
time = 0
|
||||||
|
reg = 0 # shift register
|
||||||
|
offset = 0 # bit offset, in steps of 4 (one nibble)
|
||||||
|
nib_budget = 0x40
|
||||||
|
flags = 0
|
||||||
|
token_index = 0
|
||||||
|
|
||||||
|
while (offset >> 3) < n and token_index < max_tokens:
|
||||||
|
# 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
|
||||||
|
if nib_budget != 0:
|
||||||
|
target = refill_start + 4 + ((nib_budget - 1) & ~3)
|
||||||
|
cur = refill_start
|
||||||
|
while cur != target and (cur >> 3) < n:
|
||||||
|
byte_index = cur >> 3
|
||||||
|
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)
|
||||||
|
cur += 4
|
||||||
|
offset = cur
|
||||||
|
|
||||||
|
# 2) Decode token from low 8 bits
|
||||||
|
state = reg & 0xFF
|
||||||
|
opcode = STATE_TO_TOKEN[state]
|
||||||
|
|
||||||
|
# 3) Handle pseudo-token 0x10: need more bits, don't print. Looks like a NOP.
|
||||||
|
if opcode == 0x10:
|
||||||
|
# "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
|
||||||
|
nb_index = opcode & 0x1F
|
||||||
|
nib_budget = NIBBLE_BUDGET[nb_index]
|
||||||
|
time_before = time
|
||||||
|
note = ""
|
||||||
|
# 5) Special opcode 0x16 (timestamp / marker)
|
||||||
|
if opcode == 0x16:
|
||||||
|
two_bits = (reg >> 8) & 0x3
|
||||||
|
if two_bits == 1:
|
||||||
|
flags |= 0x01
|
||||||
|
|
||||||
|
# Common 36-bit field at bits [12..47]
|
||||||
|
|
||||||
|
if (reg & 0x200) == 0:
|
||||||
|
# delta mode: add 36-bit delta to time
|
||||||
|
delta = (reg >> 12) & ((1 << 36) - 1)
|
||||||
|
time += delta
|
||||||
|
else:
|
||||||
|
# marker / other modes: no time advance
|
||||||
|
if (reg & 0x100) == 0:
|
||||||
|
# real marker: bit9=1, bit8=0, non-zero payload
|
||||||
|
# "other" 0x16 variants, ignored for timing
|
||||||
|
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:
|
||||||
|
time += delta
|
||||||
|
|
||||||
|
# Append extra decoded fields into the note string
|
||||||
|
note = decode_packet_fields(opcode, reg, delta)
|
||||||
|
|
||||||
|
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}"
|
||||||
|
)
|
||||||
|
|
||||||
|
token_index += 1
|
||||||
|
|
||||||
|
# Optional summary at the end
|
||||||
|
print(f"# done: tokens={token_index}, final_time={time}, flags=0x{flags:02x}")
|
||||||
|
|
||||||
|
def parse(fn:str):
|
||||||
|
dat = pickle.load(open(fn, "rb"))
|
||||||
|
ctx = decode(dat)
|
||||||
|
dat_sqtt = [x for x in dat if isinstance(x, ProfileSQTTEvent)]
|
||||||
|
print(f"got {len(dat_sqtt)} SQTT events in {fn}")
|
||||||
|
return dat_sqtt
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
#dat_sqtt = parse("extra/sqtt/examples/profile_empty_run_0.pkl")
|
||||||
|
#dat_sqtt = parse("extra/sqtt/examples/profile_plus_run_0.pkl")
|
||||||
|
dat_sqtt = parse("extra/sqtt/examples/profile_gemm_run_0.pkl")
|
||||||
|
blob_0 = dat_sqtt[0].blob
|
||||||
|
parse_sqtt_print_packets(blob_0[8:])
|
||||||
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
@@ -12,7 +12,9 @@ if __name__ == "__main__":
|
|||||||
lib = fp.parent/"rocprof-trace-decoder-macos-arm64-0.1.4-Darwin"/"lib"/"librocprof-trace-decoder.dylib"
|
lib = fp.parent/"rocprof-trace-decoder-macos-arm64-0.1.4-Darwin"/"lib"/"librocprof-trace-decoder.dylib"
|
||||||
os.chmod(fp, 0o755)
|
os.chmod(fp, 0o755)
|
||||||
os.system(f"sudo {fp} --prefix={fp.parent} --include-subdir")
|
os.system(f"sudo {fp} --prefix={fp.parent} --include-subdir")
|
||||||
|
shutil.copy2(lib, DEST)
|
||||||
else:
|
else:
|
||||||
lib = fetch("https://github.com/ROCm/rocprof-trace-decoder/raw/43bf0fef74a83c3c25badfc5a09c0bd39ed8c6f9/releases/linux_glibc_2_28_x86_64/librocprof-trace-decoder.so", name="librocprof-trace-decoder.so")
|
lib = DEST/"librocprof-trace-decoder.so"
|
||||||
shutil.copy2(lib, DEST)
|
os.system("sudo curl -L https://github.com/ROCm/rocprof-trace-decoder/raw/43bf0fef74a83c3c25badfc5a09c0bd39ed8c6f9/releases/linux_glibc_2_28_x86_64/librocprof-trace-decoder.so -o"+str(lib))
|
||||||
|
os.system("sudo ldconfig")
|
||||||
print(f"Installed {lib.name} to", DEST)
|
print(f"Installed {lib.name} to", DEST)
|
||||||
|
|||||||
+7
-11
@@ -185,9 +185,7 @@ class RGP:
|
|||||||
magic_number=sqtt.SQTT_FILE_MAGIC_NUMBER,
|
magic_number=sqtt.SQTT_FILE_MAGIC_NUMBER,
|
||||||
version_major=sqtt.SQTT_FILE_VERSION_MAJOR,
|
version_major=sqtt.SQTT_FILE_VERSION_MAJOR,
|
||||||
version_minor=sqtt.SQTT_FILE_VERSION_MINOR,
|
version_minor=sqtt.SQTT_FILE_VERSION_MINOR,
|
||||||
flags=sqtt.struct_sqtt_file_header_flags(
|
flags=sqtt.struct_sqtt_file_header_flags(value=1,),
|
||||||
_0=sqtt.union_sqtt_file_header_flags_0(value=1),
|
|
||||||
),
|
|
||||||
chunk_offset=ctypes.sizeof(sqtt.struct_sqtt_file_header),
|
chunk_offset=ctypes.sizeof(sqtt.struct_sqtt_file_header),
|
||||||
)
|
)
|
||||||
chunks = [
|
chunks = [
|
||||||
@@ -265,7 +263,7 @@ class RGP:
|
|||||||
profiling_mode=sqtt.SQTT_PROFILING_MODE_PRESENT,
|
profiling_mode=sqtt.SQTT_PROFILING_MODE_PRESENT,
|
||||||
instruction_trace_mode=sqtt.SQTT_INSTRUCTION_TRACE_FULL_FRAME if sqtt_itrace_enabled else sqtt.SQTT_INSTRUCTION_TRACE_DISABLED,
|
instruction_trace_mode=sqtt.SQTT_INSTRUCTION_TRACE_FULL_FRAME if sqtt_itrace_enabled else sqtt.SQTT_INSTRUCTION_TRACE_DISABLED,
|
||||||
instruction_trace_data=sqtt.union_sqtt_instruction_trace_data(
|
instruction_trace_data=sqtt.union_sqtt_instruction_trace_data(
|
||||||
shader_engine_filter=sqtt.struct_sqtt_instruction_trace_data_shader_engine_filter(mask=sqtt_itrace_se_mask),
|
shader_engine_filter=sqtt.union_sqtt_instruction_trace_data_shader_engine_filter(mask=sqtt_itrace_se_mask),
|
||||||
),
|
),
|
||||||
)),
|
)),
|
||||||
*flatten([(
|
*flatten([(
|
||||||
@@ -276,13 +274,11 @@ class RGP:
|
|||||||
),
|
),
|
||||||
shader_engine_index=sqtt_event.se,
|
shader_engine_index=sqtt_event.se,
|
||||||
sqtt_version={11: sqtt.SQTT_VERSION_3_2, 12: sqtt.SQTT_VERSION_3_3}.get(gfx_ver),
|
sqtt_version={11: sqtt.SQTT_VERSION_3_2, 12: sqtt.SQTT_VERSION_3_3}.get(gfx_ver),
|
||||||
_0=sqtt.union_sqtt_file_chunk_sqtt_desc_0(
|
v1=sqtt.struct_sqtt_file_chunk_sqtt_desc_0_v1(
|
||||||
v1=sqtt.struct_sqtt_file_chunk_sqtt_desc_0_v1(
|
instrumentation_spec_version=1,
|
||||||
instrumentation_spec_version=1,
|
instrumentation_api_version=0,
|
||||||
instrumentation_api_version=0,
|
compute_unit_index=0,
|
||||||
compute_unit_index=0,
|
)
|
||||||
)
|
|
||||||
),
|
|
||||||
)),
|
)),
|
||||||
RGPChunk(sqtt.struct_sqtt_file_chunk_sqtt_data(
|
RGPChunk(sqtt.struct_sqtt_file_chunk_sqtt_data(
|
||||||
header=sqtt.struct_sqtt_file_chunk_header(
|
header=sqtt.struct_sqtt_file_chunk_header(
|
||||||
|
|||||||
+26
-43
@@ -1,4 +1,4 @@
|
|||||||
import ctypes, pathlib, argparse, pickle, re, functools, dataclasses, itertools
|
import ctypes, pathlib, argparse, pickle, re, functools, dataclasses, itertools, threading
|
||||||
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
|
||||||
@@ -28,18 +28,6 @@ def llvm_disasm(arch:str, lib:bytes) -> dict[int, tuple[str, int]]:
|
|||||||
cur_off += instr_sz
|
cur_off += instr_sz
|
||||||
return addr_table
|
return addr_table
|
||||||
|
|
||||||
@dataclasses.dataclass
|
|
||||||
class InstInfo:
|
|
||||||
typ:str=""
|
|
||||||
inst:str=""
|
|
||||||
hit:int=0
|
|
||||||
lat:int=0
|
|
||||||
stall:int=0
|
|
||||||
def __str__(self): return f"{self.inst:>20} hits:{self.typ:>6} hits:{self.hit:>6} latency:{self.lat:>6} stall:{self.stall:>6}"
|
|
||||||
|
|
||||||
def on_ev(self, ev):
|
|
||||||
self.hit, self.lat, self.stall = self.hit + 1, self.lat + ev.duration, self.stall + ev.stall
|
|
||||||
|
|
||||||
@dataclasses.dataclass(frozen=True)
|
@dataclasses.dataclass(frozen=True)
|
||||||
class InstExec:
|
class InstExec:
|
||||||
typ:str
|
typ:str
|
||||||
@@ -48,25 +36,19 @@ class InstExec:
|
|||||||
dur:int
|
dur:int
|
||||||
time:int
|
time:int
|
||||||
|
|
||||||
@dataclasses.dataclass(frozen=True)
|
|
||||||
class PrgExec:
|
|
||||||
name:str
|
|
||||||
wave:int
|
|
||||||
cu:int
|
|
||||||
simd:int
|
|
||||||
def __str__(self): return f"{self.name},{self.wave},{self.cu},{self.simd}"
|
|
||||||
|
|
||||||
@dataclasses.dataclass(frozen=True)
|
@dataclasses.dataclass(frozen=True)
|
||||||
class WaveExec:
|
class WaveExec:
|
||||||
wave_id:int
|
wave_id:int
|
||||||
cu:int
|
cu:int
|
||||||
simd:int
|
simd:int
|
||||||
|
se:int
|
||||||
|
begin_time:int
|
||||||
|
end_time:int
|
||||||
insts:list[InstExec]
|
insts:list[InstExec]
|
||||||
|
|
||||||
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.wave_events:dict[PrgExec, dict[int, InstInfo]] = {}
|
|
||||||
self.disasms:dict[tuple[str, int], tuple[str, int]] = {}
|
self.disasms:dict[tuple[str, int], tuple[str, int]] = {}
|
||||||
self.inst_execs:dict[str, list[WaveExec]] = {}
|
self.inst_execs:dict[str, list[WaveExec]] = {}
|
||||||
|
|
||||||
@@ -79,27 +61,26 @@ class _ROCParseCtx:
|
|||||||
x = next(self.sqtt_evs, None)
|
x = next(self.sqtt_evs, None)
|
||||||
self.active_kern = x.kern if x is not None else None
|
self.active_kern = x.kern if x is not None else None
|
||||||
self.active_se = x.se if x is not None else None
|
self.active_se = x.se if x is not None else None
|
||||||
return x
|
self.active_blob = (ctypes.c_ubyte * len(x.blob)).from_buffer_copy(x.blob) if x is not None else None
|
||||||
|
return self.active_blob
|
||||||
|
|
||||||
def on_occupancy_ev(self, ev):
|
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("OCC", ev.time, self.active_se, ev.cu, ev.simd, ev.wave_id, ev.start)
|
||||||
|
|
||||||
def on_wave_ev(self, ev):
|
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("WAVE", ev.wave_id, self.active_se, ev.cu, ev.simd, ev.contexts, ev.begin_time, ev.end_time)
|
||||||
|
|
||||||
asm:dict[int, InstInfo] = {}
|
|
||||||
inst_execs:list[InstExec] = []
|
inst_execs:list[InstExec] = []
|
||||||
for j in range(ev.instructions_size):
|
for j in range(ev.instructions_size):
|
||||||
inst_ev = ev.instructions_array[j]
|
inst_ev = ev.instructions_array[j]
|
||||||
inst_typ = rocprof.rocprofiler_thread_trace_decoder_inst_category_t__enumvalues[inst_ev.category]
|
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_disasm = self.disasms[(unwrap(self.active_kern), unwrap(inst_ev.pc.address))][0]
|
||||||
asm.setdefault(inst_ev.pc.address, InstInfo(typ=inst_typ, inst=inst_disasm))
|
|
||||||
asm[inst_ev.pc.address].on_ev(inst_ev)
|
|
||||||
inst_execs.append(InstExec(inst_typ, inst_disasm, inst_ev.stall, inst_ev.duration, inst_ev.time))
|
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:
|
if ev.instructions_size > 0:
|
||||||
self.wave_events[key:=PrgExec(unwrap(self.active_kern), ev.wave_id, ev.cu, ev.simd)] = asm
|
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(key.name, []).append(WaveExec(ev.wave_id, ev.cu, ev.simd, inst_execs))
|
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] = {}
|
||||||
@@ -113,25 +94,25 @@ def decode(profile:list[ProfileEvent]) -> _ROCParseCtx:
|
|||||||
ROCParseCtx = _ROCParseCtx(dev_events, sqtt_events, prog_events)
|
ROCParseCtx = _ROCParseCtx(dev_events, sqtt_events, prog_events)
|
||||||
|
|
||||||
@rocprof.rocprof_trace_decoder_se_data_callback_t
|
@rocprof.rocprof_trace_decoder_se_data_callback_t
|
||||||
def copy_cb(buf, buf_size, data_ptr):
|
def copy_cb(buf, buf_size, _):
|
||||||
if (prof:=ROCParseCtx.next_sqtt()) is None: return 0
|
if (prof_info:=ROCParseCtx.next_sqtt()) is None: return 0
|
||||||
buf[0] = ctypes.cast((ctypes.c_ubyte * len(prof.blob)).from_buffer_copy(prof.blob), ctypes.POINTER(ctypes.c_ubyte))
|
buf[0] = ctypes.cast(prof_info, ctypes.POINTER(ctypes.c_ubyte))
|
||||||
buf_size[0] = len(prof.blob)
|
buf_size[0] = len(prof_info)
|
||||||
return len(prof.blob)
|
return len(prof_info)
|
||||||
|
|
||||||
@rocprof.rocprof_trace_decoder_trace_callback_t
|
@rocprof.rocprof_trace_decoder_trace_callback_t
|
||||||
def trace_cb(record_type, events_ptr, n, data_ptr):
|
def trace_cb(record_type, events_ptr, n, _):
|
||||||
match record_type:
|
match record_type:
|
||||||
case rocprof.ROCPROFILER_THREAD_TRACE_DECODER_RECORD_OCCUPANCY:
|
case rocprof.ROCPROFILER_THREAD_TRACE_DECODER_RECORD_OCCUPANCY:
|
||||||
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 _:
|
case _:
|
||||||
if DEBUG >= 5: print(rocprof.rocprofiler_thread_trace_decoder_record_type_t__enumvalues[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, data_ptr):
|
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
|
||||||
@@ -145,9 +126,11 @@ def decode(profile:list[ProfileEvent]) -> _ROCParseCtx:
|
|||||||
|
|
||||||
return rocprof.ROCPROFILER_THREAD_TRACE_DECODER_STATUS_SUCCESS
|
return rocprof.ROCPROFILER_THREAD_TRACE_DECODER_STATUS_SUCCESS
|
||||||
|
|
||||||
try:
|
def worker():
|
||||||
rocprof.rocprof_trace_decoder_parse_data(copy_cb, trace_cb, isa_cb, None)
|
try: rocprof.rocprof_trace_decoder_parse_data(copy_cb, trace_cb, isa_cb, None)
|
||||||
except AttributeError as e: raise RuntimeError("Failed to find rocprof-trace-decoder. Run ./extra/sqtt/install_sqtt_decoder.py to install") from e
|
except AttributeError as e: raise RuntimeError("Failed to find rocprof-trace-decoder. Run sudo ./extra/sqtt/install_sqtt_decoder.py to install") from e
|
||||||
|
(t:=threading.Thread(target=worker, daemon=True)).start()
|
||||||
|
t.join()
|
||||||
return ROCParseCtx
|
return ROCParseCtx
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
@@ -157,7 +140,7 @@ if __name__ == "__main__":
|
|||||||
|
|
||||||
with args.profile.open("rb") as f: profile = pickle.load(f)
|
with args.profile.open("rb") as f: profile = pickle.load(f)
|
||||||
rctx = decode(profile)
|
rctx = decode(profile)
|
||||||
print('SQTT:', rctx.wave_events.keys())
|
print('SQTT:', rctx.inst_execs.keys())
|
||||||
|
|
||||||
for ev in profile:
|
for ev in profile:
|
||||||
if not isinstance(ev, ProfilePMCEvent): continue
|
if not isinstance(ev, ProfilePMCEvent): continue
|
||||||
|
|||||||
+57
-29
@@ -1,22 +1,20 @@
|
|||||||
import os
|
import os
|
||||||
os.environ["PYTHONPATH"] = "."
|
os.environ["PYTHONPATH"] = "."
|
||||||
os.environ["SQTT"] = "1"
|
os.environ["SQTT"] = "1"
|
||||||
os.environ["AMD"] = "1"
|
if "DEV" not in os.environ: os.environ["DEV"] = "AMD"
|
||||||
os.environ["VIZ"] = "1"
|
os.environ["VIZ"] = "1"
|
||||||
os.environ["AMD_LLVM"] = "0"
|
os.environ["AMD_LLVM"] = "0"
|
||||||
|
|
||||||
import unittest
|
import unittest
|
||||||
import sys, contextlib
|
import sys, contextlib
|
||||||
from tinygrad import Tensor
|
from tinygrad import Tensor, dtypes
|
||||||
from tinygrad.dtype import dtypes
|
from tinygrad.helpers import getenv
|
||||||
from tinygrad.renderer import ProgramSpec
|
|
||||||
from tinygrad.uop.ops import UOp, Ops, KernelInfo
|
from tinygrad.uop.ops import UOp, Ops, KernelInfo
|
||||||
from tinygrad.engine.realize import CompiledRunner
|
|
||||||
from tinygrad.device import Device, ProfileDeviceEvent
|
from tinygrad.device import Device, ProfileDeviceEvent
|
||||||
|
|
||||||
from extra.sqtt.roc import decode, InstExec, PrgExec
|
from extra.sqtt.roc import decode, WaveExec
|
||||||
|
|
||||||
dev = Device["AMD"]
|
dev = Device[os.environ["DEV"]]
|
||||||
|
|
||||||
def custom(arg:str, s:UOp|None=None) -> UOp: return UOp(Ops.CUSTOM, src=(s,) if s is not None else (), arg=arg)
|
def custom(arg:str, s:UOp|None=None) -> UOp: return UOp(Ops.CUSTOM, src=(s,) if s is not None else (), arg=arg)
|
||||||
|
|
||||||
@@ -36,9 +34,10 @@ def asm_kernel(instrs:list[str], l:int=1, g:int=1) -> Tensor:
|
|||||||
def save_sqtt():
|
def save_sqtt():
|
||||||
# clear the old traces
|
# clear the old traces
|
||||||
dev.profile_events.clear()
|
dev.profile_events.clear()
|
||||||
sqtt:dict[PrgExec, list[InstExec]] = {}
|
sqtt:dict[str, list[WaveExec]] = {}
|
||||||
yield sqtt
|
yield sqtt
|
||||||
# decode sqtt
|
# decode sqtt
|
||||||
|
if os.environ["DEV"] != "AMD": return
|
||||||
rctx = decode(dev.profile_events+[ProfileDeviceEvent("AMD", props=dev.device_props())])
|
rctx = decode(dev.profile_events+[ProfileDeviceEvent("AMD", props=dev.device_props())])
|
||||||
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)
|
||||||
@@ -62,28 +61,37 @@ class TestTiming(unittest.TestCase):
|
|||||||
assert all(s.stall == 0 for s in wave)
|
assert all(s.stall == 0 for s in wave)
|
||||||
|
|
||||||
def test_multi_cycle_inst(self):
|
def test_multi_cycle_inst(self):
|
||||||
|
def custom_vrcp(A, B):
|
||||||
|
op = custom("float a = 0.0;")
|
||||||
|
op = custom("float b = (*(data1_1+0));", op)
|
||||||
|
#op = custom('asm volatile("v_mul_f32_e32 %2 %2 %1" : "+v"(a) : "v"(b));', op)
|
||||||
|
op = custom('asm volatile("v_rcp_f32_e32 %2 %1" : "+v"(a) : "v"(b));', op)
|
||||||
|
op = custom('asm volatile("v_add_f32_e64 %1 %1 1.0" : "+v"(a));', op)
|
||||||
|
op = custom("*(data0_1+0) = a;", op)
|
||||||
|
return UOp.sink(op, A, B, arg=KernelInfo(name="custom_vrcp"))
|
||||||
|
out = Tensor([0.]).realize()
|
||||||
|
inp = Tensor([-2.0]).realize()
|
||||||
with save_sqtt() as sqtt:
|
with save_sqtt() as sqtt:
|
||||||
asm_kernel([
|
Tensor.custom_kernel(out, inp, fxn=custom_vrcp)[0].realize()
|
||||||
"v_mov_b32_e32 v4 0x3f800000",
|
wave = list(sqtt.values())[0][0]
|
||||||
"v_rcp_f32_e32 v5 v4",
|
for i in range(len(wave.insts)):
|
||||||
"v_mul_f32_e32 v6 v5 v4",
|
if wave.insts[i].inst.startswith("global_store"):
|
||||||
]).realize()
|
print(f"store diff {wave.insts[i].time-(wave.insts[i-1].time)}")
|
||||||
w = list(sqtt.values())[0]
|
self.assertEqual(out.item(), 0.5)
|
||||||
rcp, mul = w[1], w[2]
|
|
||||||
self.assertGreater(rcp.dur, 1) # 4 cycles on gfx11
|
|
||||||
self.assertEqual(mul.dur, 1)
|
|
||||||
# mul depends on v5, how can it run before rcp is done?
|
|
||||||
self.assertGreaterEqual(mul.time, rcp.time+rcp.dur)
|
|
||||||
|
|
||||||
def test_wmma(self):
|
def test_wmma(self):
|
||||||
with save_sqtt() as sqtt:
|
with save_sqtt() as sqtt:
|
||||||
asm_kernel([
|
for tc in dev.renderer.get_tensor_cores(dev.arch):
|
||||||
"v_wmma_f32_16x16x16_f16 v[16:23], v[0:7], v[8:15], v[16:23]",
|
M, K, N = tc.dims
|
||||||
"v_add_f32_e32 v0 v16 v0",
|
s = 32
|
||||||
], l=32*4).realize()
|
a = Tensor.empty(M*s, K*s, dtype=tc.dtype_in)@Tensor.empty(K*s, N*s, dtype=tc.dtype_in)
|
||||||
assert len(sqtt) == 2, f"expected two waves, got {len(sqtt)} {list(sqtt.keys())}"
|
a.realize()
|
||||||
wmma = list(sqtt.values())[0][0]
|
print(a)
|
||||||
self.assertGreater(wmma.dur, 1) # rgp says 32 clocks
|
for p,waves in sqtt.items():
|
||||||
|
for e in waves[0].insts:
|
||||||
|
if (e.inst.startswith("v_wmma")):
|
||||||
|
instruction = e.inst.split(" ")[0]
|
||||||
|
print(f"{instruction:<29} : {e.dur} cycles")
|
||||||
|
|
||||||
def test_sleep(self):
|
def test_sleep(self):
|
||||||
n = 1
|
n = 1
|
||||||
@@ -91,15 +99,35 @@ class TestTiming(unittest.TestCase):
|
|||||||
assert data0.dtype.base == dtypes.ulong
|
assert data0.dtype.base == dtypes.ulong
|
||||||
op = custom("unsigned long long t0 = __builtin_readcyclecounter();")
|
op = custom("unsigned long long t0 = __builtin_readcyclecounter();")
|
||||||
op = custom(f"__builtin_amdgcn_s_sleep({n});", op)
|
op = custom(f"__builtin_amdgcn_s_sleep({n});", op)
|
||||||
op = custom(f"unsigned long long t1 = __builtin_readcyclecounter();", op)
|
op = custom("unsigned long long t1 = __builtin_readcyclecounter();", op)
|
||||||
op = custom(f"data0_{data0.size}[0] = t1 - t0;", op)
|
op = custom(f"data0_{data0.size}[0] = t1 - t0;", op)
|
||||||
return UOp.sink(data0, op, arg=KernelInfo(name=f"sleep_{n}"))
|
return UOp.sink(data0, op, arg=KernelInfo(name=f"sleep_{n}"))
|
||||||
diff_hw_reg = Tensor.empty(1, dtype=dtypes.ulong)
|
diff_hw_reg = Tensor.empty(1, dtype=dtypes.ulong)
|
||||||
diff_hw_reg = Tensor.custom_kernel(diff_hw_reg, fxn=sleep_kernel)[0]
|
diff_hw_reg = Tensor.custom_kernel(diff_hw_reg, fxn=sleep_kernel)[0]
|
||||||
with save_sqtt() as sqtt:
|
with save_sqtt() as sqtt:
|
||||||
diff_hw_reg.realize()
|
diff_hw_reg.realize()
|
||||||
diff_sqtt = list(sqtt.values())[0][2]
|
sleep = next((e for e in sqtt[f"sleep_{n}"][0].insts if e.inst.startswith("s_sleep")))
|
||||||
self.assertEqual(diff_sqtt.dur, diff_hw_reg.item()-1) # 1 cycle for reading the counter register
|
# cycles = sleep dur + overhead of storing hi/lo REG_SHADER_CYCLES
|
||||||
|
self.assertGreaterEqual(diff_hw_reg.item(), sleep.dur)
|
||||||
|
|
||||||
|
def test_nop(self):
|
||||||
|
with save_sqtt() as sqtt:
|
||||||
|
asm_kernel(["s_nop 1"]*10).realize()
|
||||||
|
wave = list(sqtt.values())[0][0]
|
||||||
|
for e in wave.insts:
|
||||||
|
print(f"{e.inst} {e.dur=} {e.stall=}")
|
||||||
|
|
||||||
|
def test_wave_sched(self):
|
||||||
|
num_waves = getenv("NUM_WAVES", 16)
|
||||||
|
num_wgps = getenv("NUM_WGPS", 2)
|
||||||
|
num_vgpr = getenv("NUM_VGPR", 256)
|
||||||
|
with save_sqtt() as sqtt:
|
||||||
|
# 1 cycle decode, no stall
|
||||||
|
asm_kernel([f"v_mov_b32_e32 v{i} {i}" for i in range(num_vgpr)], l=32*num_waves, g=num_wgps).realize()
|
||||||
|
waves = list(sqtt.values())[0]
|
||||||
|
print(len(waves), "waves decoded")
|
||||||
|
for w in waves:
|
||||||
|
print(f"{w.wave_id:<2} {w.simd=} {w.cu=} {w.se=} @ clk {w.begin_time}")
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
unittest.main()
|
unittest.main()
|
||||||
|
|||||||
@@ -1 +1,6 @@
|
|||||||
WARP_THREADS = 32
|
from tinygrad.device import Device
|
||||||
|
|
||||||
|
if Device.DEFAULT == "AMD":
|
||||||
|
WARP_THREADS = 64
|
||||||
|
else:
|
||||||
|
WARP_THREADS = 32
|
||||||
|
|||||||
+137
-119
@@ -7,7 +7,7 @@ from tinygrad.dtype import AddrSpace, PtrDType
|
|||||||
from tinygrad.helpers import getenv, prod
|
from tinygrad.helpers import getenv, prod
|
||||||
|
|
||||||
from extra.thunder.tiny.tk import WARP_THREADS
|
from extra.thunder.tiny.tk import WARP_THREADS
|
||||||
from extra.thunder.tiny.tk.tiles import TILE_ROW_DIM, TILE_COL_DIM, RT_BASE_TILE_NEPT, slots
|
from extra.thunder.tiny.tk.tiles import ALL_TILES, GL, ST, RT, RV
|
||||||
|
|
||||||
class Group:
|
class Group:
|
||||||
def __init__(self, warps:int, ker):
|
def __init__(self, warps:int, ker):
|
||||||
@@ -27,23 +27,26 @@ class Group:
|
|||||||
# ops that only work on a single warp
|
# ops that only work on a single warp
|
||||||
|
|
||||||
clear_rid = 1000
|
clear_rid = 1000
|
||||||
def clear(self, reg:UOp, value:float=0):
|
def clear(self, reg:ALL_TILES, value:float=0):
|
||||||
|
reg = cast(UOp, reg)
|
||||||
assert self.warps == 1
|
assert self.warps == 1
|
||||||
|
|
||||||
i = UOp.range(reg.size, Group.clear_rid)
|
rngs_for_shape = tuple(UOp.range(dim, Group.clear_rid + i) for i, dim in enumerate(reg.shape))
|
||||||
Group.clear_rid += 1
|
Group.clear_rid += len(reg.shape)
|
||||||
return reg.reshape((reg.size,))[i].set(value, end=i).after(reg).reshape(reg.shape)
|
|
||||||
|
|
||||||
def zero(self, reg:UOp): return self.clear(reg, 0)
|
reg_store = reg[*rngs_for_shape].store(value).end(*rngs_for_shape)
|
||||||
def neg_inf(self, reg:UOp): return self.clear(reg, -math.inf)
|
|
||||||
|
self.ker.push_store(reg_store, reg)
|
||||||
|
return reg.after(reg_store).reshape(reg.shape)
|
||||||
|
|
||||||
|
def zero(self, reg:ALL_TILES): return self.clear(reg, 0)
|
||||||
|
def neg_inf(self, reg:ALL_TILES): return self.clear(reg, -math.inf)
|
||||||
|
|
||||||
copy_rid = 300
|
copy_rid = 300
|
||||||
def copy(self, dst:UOp, src:UOp):
|
def copy(self, dst:ALL_TILES, src:ALL_TILES):
|
||||||
|
dst, src = cast(UOp, dst), cast(UOp, src)
|
||||||
assert self.warps == 1
|
assert self.warps == 1
|
||||||
|
|
||||||
assert dst.shape == src.shape
|
assert dst.shape == src.shape
|
||||||
assert cast(PtrDType, dst.dtype).addrspace == AddrSpace.REG
|
|
||||||
assert cast(PtrDType, src.dtype).addrspace == AddrSpace.REG
|
|
||||||
|
|
||||||
rngs_for_shape = tuple(UOp.range(dim, Group.copy_rid + i) for i, dim in enumerate(dst.shape))
|
rngs_for_shape = tuple(UOp.range(dim, Group.copy_rid + i) for i, dim in enumerate(dst.shape))
|
||||||
Group.copy_rid += len(dst.shape)
|
Group.copy_rid += len(dst.shape)
|
||||||
@@ -53,57 +56,55 @@ 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)
|
||||||
|
|
||||||
mma_rid = 600
|
def mma_AB(self, c:UOp|RT, a:UOp|RT, b:UOp|RT):
|
||||||
def mma_AB(self, c:UOp, a:UOp, b:UOp, after=True):
|
c, a, b = cast(UOp, c), cast(UOp, a), cast(UOp, b)
|
||||||
assert self.warps == 1
|
assert self.warps == 1
|
||||||
|
|
||||||
mma_i_height = UOp.range(c.shape[-3], Group.mma_rid)
|
for height in self.ker.range(c.shape[-3], track=False):
|
||||||
mma_i_width = UOp.range(c.shape[-2], Group.mma_rid+1)
|
for width in self.ker.range(c.shape[-2], track=False):
|
||||||
mma_i_inner = UOp.range(a.shape[-2], Group.mma_rid+2, AxisType.REDUCE)
|
for inner in self.ker.range(a.shape[-2], AxisType.REDUCE, track=False):
|
||||||
Group.mma_rid += 3
|
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)])
|
||||||
|
b_in1 = UOp.vectorize(*([b[inner, width, i] for i in range(2)] + [b[inner, width, 4+i] for i in range(2)]))
|
||||||
|
c_out1 = UOp.vectorize(*[c[height, width, i] for i in range(4)])
|
||||||
|
b_in2 = UOp.vectorize(*([b[inner, width, 2+i] for i in range(2)] + [b[inner, width, 6+i] for i in range(2)]))
|
||||||
|
c_out2 = UOp.vectorize(*[c[height, width, 4+i] for i in range(4)])
|
||||||
|
|
||||||
a_in = UOp.vectorize(*[a[mma_i_height, mma_i_inner, i] for i in range(8)])
|
out1 = UOp(Ops.WMMA, dtypes.float32.vec(4), (a_in, b_in1, c_out1), arg=wmma_arg)
|
||||||
b_in1 = UOp.vectorize(*([b[mma_i_inner, mma_i_width, i] for i in range(2)] + [b[mma_i_inner, mma_i_width, 4+i] for i in range(2)]))
|
out2 = UOp(Ops.WMMA, dtypes.float32.vec(4), (a_in, b_in2, c_out2), arg=wmma_arg)
|
||||||
c_out1 = UOp.vectorize(*[c[mma_i_height, mma_i_width, i] for i in range(4)])
|
c_i = [c[height, width, i].store(out1.gep(i)) for i in range(4)] + [c[height, width, 4+i].store(out2.gep(i)) for i in range(4)]
|
||||||
b_in2 = UOp.vectorize(*([b[mma_i_inner, mma_i_width, 2+i] for i in range(2)] + [b[mma_i_inner, mma_i_width, 6+i] for i in range(2)]))
|
c_store = UOp.group(*c_i).end(height, width, inner)
|
||||||
c_out2 = UOp.vectorize(*[c[mma_i_height, mma_i_width, 4+i] for i in range(4)])
|
|
||||||
|
|
||||||
out1 = UOp(Ops.WMMA, dtypes.float32.vec(4), (a_in, b_in1, c_out1), arg=wmma_arg)
|
|
||||||
out2 = UOp(Ops.WMMA, dtypes.float32.vec(4), (a_in, b_in2, c_out2), arg=wmma_arg)
|
|
||||||
c_i = [c[mma_i_height, mma_i_width, i].store(out1.gep(i)) for i in range(4)] + [c[mma_i_height, mma_i_width, 4+i].store(out2.gep(i)) for i in range(4)]
|
|
||||||
c_store = UOp.group(*c_i).end(mma_i_height, mma_i_width, mma_i_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, a:UOp, b:UOp, 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)
|
||||||
assert self.warps == 1
|
assert self.warps == 1
|
||||||
|
|
||||||
mma_i_height = UOp.range(c.shape[-3], Group.mma_rid)
|
for height in self.ker.range(c.shape[-3], track=False):
|
||||||
mma_i_width = UOp.range(c.shape[-2], Group.mma_rid+1)
|
for width in self.ker.range(c.shape[-2], track=False):
|
||||||
mma_i_inner = UOp.range(a.shape[-2], Group.mma_rid+2, AxisType.REDUCE)
|
for inner in self.ker.range(a.shape[-2], AxisType.REDUCE, track=False):
|
||||||
Group.mma_rid += 3
|
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)])
|
||||||
|
b_in1 = UOp.vectorize(*([b[width, inner, i] for i in range(2)] + [b[width, inner, 4+i] for i in range(2)]))
|
||||||
|
c_out1 = UOp.vectorize(*[c[height, width, i] for i in range(4)])
|
||||||
|
b_in2 = UOp.vectorize(*([b[width, inner, 2+i] for i in range(2)] + [b[width, inner, 6+i] for i in range(2)]))
|
||||||
|
c_out2 = UOp.vectorize(*[c[height, width, 4+i] for i in range(4)])
|
||||||
|
|
||||||
a_in = UOp.vectorize(*[a[mma_i_height, mma_i_inner, i] for i in range(8)])
|
out1 = UOp(Ops.WMMA, dtypes.float32.vec(4), (a_in, b_in1, c_out1), arg=wmma_arg)
|
||||||
b_in1 = UOp.vectorize(*([b[mma_i_width, mma_i_inner, i] for i in range(2)] + [b[mma_i_width, mma_i_inner, 4+i] for i in range(2)]))
|
out2 = UOp(Ops.WMMA, dtypes.float32.vec(4), (a_in, b_in2, c_out2), arg=wmma_arg)
|
||||||
c_out1 = UOp.vectorize(*[c[mma_i_height, mma_i_width, i] for i in range(4)])
|
c_i = [c[height, width, i].store(out1.gep(i)) for i in range(4)] + [c[height, width, 4+i].store(out2.gep(i)) for i in range(4)]
|
||||||
b_in2 = UOp.vectorize(*([b[mma_i_width, mma_i_inner, 2+i] for i in range(2)] + [b[mma_i_width, mma_i_inner, 6+i] for i in range(2)]))
|
c_store = UOp.group(*c_i).end(height, width, inner)
|
||||||
c_out2 = UOp.vectorize(*[c[mma_i_height, mma_i_width, 4+i] for i in range(4)])
|
|
||||||
|
|
||||||
out1 = UOp(Ops.WMMA, dtypes.float32.vec(4), (a_in, b_in1, c_out1), arg=wmma_arg)
|
|
||||||
out2 = UOp(Ops.WMMA, dtypes.float32.vec(4), (a_in, b_in2, c_out2), arg=wmma_arg)
|
|
||||||
c_i = [c[mma_i_height, mma_i_width, i].store(out1.gep(i)) for i in range(4)] + [c[mma_i_height, mma_i_width, 4+i].store(out2.gep(i)) for i in range(4)]
|
|
||||||
c_store = UOp.group(*c_i).end(mma_i_height, mma_i_width, mma_i_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:UOp, op:Callable[[UOp], UOp]|Callable[[UOp, tuple], UOp]):
|
def map(self, a:ALL_TILES, op:Callable[[UOp], UOp]|Callable[[UOp, tuple], UOp]):
|
||||||
|
a = cast(UOp, a)
|
||||||
assert self.warps == 1
|
assert self.warps == 1
|
||||||
|
|
||||||
rngs_for_shape = tuple(UOp.range(dim, Group.map_rid + i) for i, dim in enumerate(a.shape))
|
rngs_for_shape = tuple(UOp.range(dim, Group.map_rid + i) for i, dim in enumerate(a.shape))
|
||||||
@@ -119,70 +120,83 @@ class Group:
|
|||||||
self.ker.push_store(a_store, a)
|
self.ker.push_store(a_store, a)
|
||||||
return a.after(a_store).reshape(a.shape)
|
return a.after(a_store).reshape(a.shape)
|
||||||
|
|
||||||
def row_reduce(self, vec:UOp, src:UOp, op:Callable[[UOp, UOp], UOp]):
|
def row_reduce(self, vec:UOp|RV, src:UOp|RT, op:Callable[[UOp, UOp], UOp]):
|
||||||
|
vec, src = cast(UOp, vec), cast(UOp, src)
|
||||||
assert self.warps == 1
|
assert self.warps == 1
|
||||||
|
|
||||||
red_local = UOp.placeholder((self.group_threads, 2), src.dtype.base, addrspace=AddrSpace.LOCAL, slot=slots.shared_slot)
|
red_local = self.ker.alloc((self.group_threads, 2), src.dtype.base, AddrSpace.LOCAL)
|
||||||
slots.shared_slot += 1
|
red_reg = self.ker.alloc((2,), src.dtype.base, AddrSpace.REG)
|
||||||
|
|
||||||
for height in self.ker.range(src.shape[-3], track=False):
|
for height in self.ker.range(src.shape[-3], track=False):
|
||||||
for i_outer in self.ker.range(2, track=False):
|
i = UOp.range(red_reg.size, Group.clear_rid)
|
||||||
|
Group.clear_rid += 1
|
||||||
|
red_reg = red_reg.after(height, *[tkr._rng for tkr in self.ker.range_stack])
|
||||||
|
reg_store = red_reg.flatten()[i].store(0.).end(i)
|
||||||
|
red_reg = red_reg.after(reg_store).reshape(red_reg.shape)
|
||||||
|
|
||||||
|
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], AxisType.REDUCE, track=False):
|
||||||
for i_inner in self.ker.range(4, AxisType.REDUCE, track=False):
|
for inner in self.ker.range(4, AxisType.REDUCE, track=False):
|
||||||
elem_index = i_inner + 2 * (i_inner // 2) + i_outer * 2
|
elem_index = inner + 2 * (inner // 2) + outer * 2
|
||||||
vec_store = vec[height, 0, i_outer].store(op(vec[height, 0, i_outer], src[height, width, elem_index])).end(width, i_inner, i_outer)
|
reg_store = red_reg[outer].store(op(red_reg[outer], src[height, width, elem_index])).end(inner, width, outer)
|
||||||
vec = vec.after(vec_store).reshape(vec.shape)
|
red_reg = red_reg.after(reg_store).reshape(red_reg.shape)
|
||||||
|
|
||||||
# store to shared memory
|
# store to shared memory
|
||||||
for i_outer in self.ker.range(2, track=False):
|
for outer in self.ker.range(2, track=False):
|
||||||
red_local_store = red_local[self.laneid, i_outer].store(vec[height, 0, i_outer]).end(i_outer)
|
red_local_store = red_local[self.laneid, outer].store(red_reg[outer]).end(outer)
|
||||||
red_local = red_local.after(red_local_store).reshape(red_local.shape)
|
red_local = red_local.after(red_local_store.barrier()).reshape(red_local.shape)
|
||||||
|
|
||||||
# reduce from shared memory
|
# reduce from shared memory
|
||||||
for i_outer in self.ker.range(2, track=False):
|
for outer in self.ker.range(2, track=False):
|
||||||
for i_inner in self.ker.range(3, AxisType.REDUCE, track=False):
|
for inner in self.ker.range(3, AxisType.REDUCE, track=False):
|
||||||
offset = (self.laneid // 4) * 4 + ((self.laneid + 1 + i_inner) % 4)
|
offset = (self.laneid // 4) * 4 + ((self.laneid + inner + 1) % 4)
|
||||||
vec_store = vec[height, 0, i_outer].store(op(vec[height, 0, i_outer], red_local[offset, i_outer])).end(i_inner, i_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)
|
||||||
|
|
||||||
|
# reduce with vec
|
||||||
|
for outer in self.ker.range(2, track=False):
|
||||||
|
vec_store = vec[height, 0, outer].store(op(vec[height, 0, outer], red_reg[outer])).end(outer, height)
|
||||||
|
|
||||||
self.ker.push_store(vec_store, vec)
|
self.ker.push_store(vec_store, vec)
|
||||||
return vec.after(vec_store).reshape(vec.shape)
|
return vec.after(vec_store).reshape(vec.shape)
|
||||||
|
|
||||||
# ops that can work across multiple warps
|
# ops that can work across multiple warps
|
||||||
|
|
||||||
LOAD_INNER = 8
|
LOAD_INNER = 4
|
||||||
load_rid = 100
|
def load(self, dst:ALL_TILES, src:ALL_TILES, dst_idxs:tuple[UOp|int,...]=(), idxs:tuple[UOp|int,...]=(), axis:int=0, transpose:bool=False):
|
||||||
def load(self, dst:UOp, src:UOp, dst_idxs:tuple[UOp|int,...]=(), idxs:tuple[UOp|int,...]=(), axis:int=0, transpose:bool=False):
|
dst, src = cast(UOp, dst), cast(UOp, src)
|
||||||
assert isinstance(dst.dtype, PtrDType) and isinstance(src.dtype, PtrDType)
|
assert isinstance(dst.dtype, PtrDType) and isinstance(src.dtype, PtrDType)
|
||||||
dst_dtype, src_dtype = cast(PtrDType, dst.dtype), cast(PtrDType, src.dtype)
|
dst_dtype, src_dtype = cast(PtrDType, dst.dtype), cast(PtrDType, src.dtype)
|
||||||
if dst_dtype.addrspace == AddrSpace.REG and src_dtype.addrspace == AddrSpace.LOCAL:
|
if dst_dtype.addrspace == AddrSpace.REG and src_dtype.addrspace == AddrSpace.LOCAL:
|
||||||
srcf = src.flatten(-2)
|
srcf = src.flatten(-2)
|
||||||
|
|
||||||
load_i_height = UOp.range(dst.shape[-3], Group.load_rid)
|
|
||||||
load_i_width = UOp.range(dst.shape[-2], Group.load_rid+1)
|
|
||||||
load_i_inner = UOp.range(RT_BASE_TILE_NEPT, Group.load_rid+2)
|
|
||||||
Group.load_rid += 3
|
|
||||||
|
|
||||||
if self.warps % 4 == 0: local_warpid = (self.warpid // 4) + (self.warpid % 4) * (self.warps // 4)
|
if self.warps % 4 == 0: local_warpid = (self.warpid // 4) + (self.warpid % 4) * (self.warps // 4)
|
||||||
else: local_warpid = self.warpid
|
else: local_warpid = self.warpid
|
||||||
warp_laneid = self.threadIdx_x % WARP_THREADS
|
warp_laneid = self.threadIdx_x % WARP_THREADS
|
||||||
|
|
||||||
if not transpose:
|
for height in self.ker.range(dst.shape[-3], track=False):
|
||||||
row = (local_warpid * dst.shape[-3] + load_i_height) * TILE_ROW_DIM + (warp_laneid // 4)
|
for width in self.ker.range(dst.shape[-2], track=False):
|
||||||
col = load_i_width * TILE_COL_DIM + 2 * (warp_laneid % 4)
|
for inner in self.ker.range(RT.BASE_TILE_NEPT, track=False):
|
||||||
|
base_row = (local_warpid * dst.shape[-3] + height) * RT.BASE_TILE_ROWS
|
||||||
|
base_col = width * RT.BASE_TILE_COLS
|
||||||
|
|
||||||
row_offset = ((load_i_inner % 4) // 2) * 8
|
if not transpose:
|
||||||
col_offset = (load_i_inner % 2) + (load_i_inner // 4) * 8
|
row = base_row + (warp_laneid // 4)
|
||||||
else:
|
col = base_col + 2 * (warp_laneid % 4)
|
||||||
row = (local_warpid * dst.shape[-3] + load_i_height) * TILE_ROW_DIM + 2 * (warp_laneid % 4)
|
|
||||||
col = load_i_width * TILE_COL_DIM + (warp_laneid // 4)
|
|
||||||
|
|
||||||
row_offset = (load_i_inner % 2) + (load_i_inner // 4) * 8
|
row_offset = ((inner % 4) // 2) * 8
|
||||||
col_offset = ((load_i_inner % 4) // 2) * 8
|
col_offset = (inner % 2) + (inner // 4) * 8
|
||||||
|
else:
|
||||||
|
row = base_row + 2 * (warp_laneid % 4)
|
||||||
|
col = base_col + (warp_laneid // 4)
|
||||||
|
|
||||||
src_i_last = (row + row_offset) * src.shape[-1] + col + col_offset
|
row_offset = (inner % 2) + (inner // 4) * 8
|
||||||
|
col_offset = ((inner % 4) // 2) * 8
|
||||||
|
|
||||||
dst_store = dst[*dst_idxs, load_i_height, load_i_width, load_i_inner].store(srcf[*idxs[:-2], src_i_last])
|
src_i_last = (row + row_offset) * src.shape[-1] + col + col_offset
|
||||||
dst_store = dst_store.end(load_i_height, load_i_width, load_i_inner)
|
|
||||||
|
dst_store = dst[*dst_idxs, height, width, inner].store(srcf[*idxs[:-2], src_i_last])
|
||||||
|
dst_store = dst_store.end(height, width, inner)
|
||||||
elif dst_dtype.addrspace == AddrSpace.LOCAL and src_dtype.addrspace == AddrSpace.GLOBAL:
|
elif dst_dtype.addrspace == AddrSpace.LOCAL and src_dtype.addrspace == AddrSpace.GLOBAL:
|
||||||
dstf = dst.flatten(-2)
|
dstf = dst.flatten(-2)
|
||||||
|
|
||||||
@@ -196,50 +210,56 @@ class Group:
|
|||||||
memcpy_per_row = dst.shape[-1] // Group.LOAD_INNER
|
memcpy_per_row = dst.shape[-1] // Group.LOAD_INNER
|
||||||
total_calls = prod(dst.shape[-2:]) // (self.group_threads * Group.LOAD_INNER)
|
total_calls = prod(dst.shape[-2:]) // (self.group_threads * Group.LOAD_INNER)
|
||||||
|
|
||||||
load_i_outer = UOp.range(total_calls, Group.load_rid)
|
for outer in self.ker.range(total_calls, track=False):
|
||||||
load_i_inner = UOp.range(Group.LOAD_INNER, Group.load_rid+1)
|
for inner in self.ker.range(Group.LOAD_INNER, track=False):
|
||||||
Group.load_rid += 2
|
load_idx = outer * self.group_threads + self.laneid
|
||||||
|
row = load_idx // memcpy_per_row
|
||||||
|
col = (load_idx * Group.LOAD_INNER) % dst.shape[-1]
|
||||||
|
|
||||||
load_idx = load_i_outer * self.group_threads + self.laneid
|
dst_i = row * dst.shape[-1] + col + inner
|
||||||
row = load_idx // memcpy_per_row
|
src_i += row * row_stride + col + inner
|
||||||
col = (load_idx * Group.LOAD_INNER) % dst.shape[-1]
|
|
||||||
|
|
||||||
dst_i = row * dst.shape[-1] + col + load_i_inner
|
dst_store = dstf[*dst_idxs, dst_i].store(srcf[src_i]).end(outer, inner)
|
||||||
src_i += row * row_stride + col + load_i_inner
|
|
||||||
|
|
||||||
dst_store = dstf[*dst_idxs, dst_i].store(srcf[src_i]).end(load_i_outer, load_i_inner)
|
|
||||||
else:
|
else:
|
||||||
raise NotImplementedError(f"load from {src_dtype.addrspace} to {dst_dtype.addrspace} not implemented")
|
raise NotImplementedError(f"load from {src_dtype.addrspace} to {dst_dtype.addrspace} not implemented")
|
||||||
|
|
||||||
return dst.after(dst_store.barrier()).reshape(dst.shape)
|
return dst.after(dst_store.barrier()).reshape(dst.shape)
|
||||||
|
|
||||||
STORE_INNER = 8
|
STORE_INNER = 4
|
||||||
store_rid = 200
|
def store(self, dst:ALL_TILES, src:ALL_TILES, idxs:tuple[UOp|int,...]=(), src_idxs:tuple[UOp|int,...]=(), axis:int=0, transpose:bool=False):
|
||||||
def store(self, dst:UOp, src:UOp, idxs:tuple[UOp|int,...]=(), src_idxs:tuple[UOp|int,...]=(), axis=0, after=True):
|
dst, src = cast(UOp, dst), cast(UOp, src)
|
||||||
assert isinstance(dst.dtype, PtrDType) and isinstance(src.dtype, PtrDType)
|
assert isinstance(dst.dtype, PtrDType) and isinstance(src.dtype, PtrDType)
|
||||||
dst_dtype, src_dtype = cast(PtrDType, dst.dtype), cast(PtrDType, src.dtype)
|
dst_dtype, src_dtype = cast(PtrDType, dst.dtype), cast(PtrDType, src.dtype)
|
||||||
if src_dtype.addrspace == AddrSpace.REG and dst_dtype.addrspace == AddrSpace.LOCAL:
|
if src_dtype.addrspace == AddrSpace.REG and dst_dtype.addrspace == AddrSpace.LOCAL:
|
||||||
dstf = dst.flatten(-2)
|
dstf = dst.flatten(-2)
|
||||||
|
|
||||||
store_i_height = UOp.range(src.shape[-3], Group.store_rid)
|
|
||||||
store_i_width = UOp.range(src.shape[-2], Group.store_rid+1)
|
|
||||||
store_i_inner = UOp.range(RT_BASE_TILE_NEPT, Group.store_rid+2)
|
|
||||||
Group.store_rid += 3
|
|
||||||
|
|
||||||
if self.warps % 4 == 0: local_warpid = (self.warpid // 4) + (self.warpid % 4) * (self.warps // 4)
|
if self.warps % 4 == 0: local_warpid = (self.warpid // 4) + (self.warpid % 4) * (self.warps // 4)
|
||||||
else: local_warpid = self.warpid
|
else: local_warpid = self.warpid
|
||||||
warp_laneid = self.threadIdx_x % WARP_THREADS
|
warp_laneid = self.threadIdx_x % WARP_THREADS
|
||||||
|
|
||||||
row = (local_warpid * src.shape[-3] + store_i_height) * TILE_ROW_DIM + (warp_laneid // 4)
|
for height in self.ker.range(src.shape[-3], track=False):
|
||||||
col = store_i_width * TILE_COL_DIM + 2 * (warp_laneid % 4)
|
for width in self.ker.range(src.shape[-2], track=False):
|
||||||
|
for inner in self.ker.range(RT.BASE_TILE_NEPT, track=False):
|
||||||
|
base_row = (local_warpid * src.shape[-3] + height) * RT.BASE_TILE_ROWS
|
||||||
|
base_col = width * RT.BASE_TILE_COLS
|
||||||
|
|
||||||
row_offset = ((store_i_inner % 4) // 2) * 8
|
if not transpose:
|
||||||
col_offset = (store_i_inner % 2) + (store_i_inner // 4) * 8
|
row = base_row + (warp_laneid // 4)
|
||||||
|
col = base_col + 2 * (warp_laneid % 4)
|
||||||
|
|
||||||
dst_i_last = (row + row_offset) * dst.shape[-1] + col + col_offset
|
row_offset = ((inner % 4) // 2) * 8
|
||||||
|
col_offset = (inner % 2) + (inner // 4) * 8
|
||||||
|
else:
|
||||||
|
row = base_row + 2 * (warp_laneid % 4)
|
||||||
|
col = base_col + (warp_laneid // 4)
|
||||||
|
|
||||||
dst_store = dstf[*idxs[:-2], dst_i_last].store(src[*src_idxs, store_i_height, store_i_width, store_i_inner])
|
row_offset = (inner % 2) + (inner // 4) * 8
|
||||||
dst_store = dst_store.end(store_i_height, store_i_width, store_i_inner)
|
col_offset = ((inner % 4) // 2) * 8
|
||||||
|
|
||||||
|
dst_i_last = (row + row_offset) * dst.shape[-1] + col + col_offset
|
||||||
|
|
||||||
|
dst_store = dstf[*idxs[:-2], dst_i_last].store(src[*src_idxs, height, width, inner])
|
||||||
|
dst_store = dst_store.end(height, width, inner)
|
||||||
elif src_dtype.addrspace == AddrSpace.LOCAL and dst_dtype.addrspace == AddrSpace.GLOBAL:
|
elif src_dtype.addrspace == AddrSpace.LOCAL and dst_dtype.addrspace == AddrSpace.GLOBAL:
|
||||||
dstf = dst.flatten()
|
dstf = dst.flatten()
|
||||||
row_stride = prod(dst.shape[axis+1:])
|
row_stride = prod(dst.shape[axis+1:])
|
||||||
@@ -253,20 +273,18 @@ class Group:
|
|||||||
memcpy_per_row = src.shape[-1] // Group.STORE_INNER
|
memcpy_per_row = src.shape[-1] // Group.STORE_INNER
|
||||||
total_calls = prod(src.shape[-2:]) // (self.group_threads * Group.STORE_INNER)
|
total_calls = prod(src.shape[-2:]) // (self.group_threads * Group.STORE_INNER)
|
||||||
|
|
||||||
store_i_outer = UOp.range(total_calls, Group.store_rid)
|
for outer in self.ker.range(total_calls, track=False):
|
||||||
store_i_inner = UOp.range(Group.STORE_INNER, Group.store_rid+1)
|
for inner in self.ker.range(Group.STORE_INNER, track=False):
|
||||||
Group.store_rid += 2
|
load_idx = outer * self.group_threads + self.laneid
|
||||||
|
row = load_idx // memcpy_per_row
|
||||||
|
col = (load_idx * Group.STORE_INNER) % src.shape[-1]
|
||||||
|
|
||||||
load_idx = store_i_outer * self.group_threads + self.laneid
|
src_i = row * src.shape[-1] + col + inner
|
||||||
row = load_idx // memcpy_per_row
|
dst_i += row * row_stride + col + inner
|
||||||
col = (load_idx * Group.STORE_INNER) % src.shape[-1]
|
|
||||||
|
|
||||||
src_i = row * src.shape[-1] + col + store_i_inner
|
dst_store = dstf[dst_i].store(srcf[*src_idxs, src_i]).end(outer, inner)
|
||||||
dst_i += row * row_stride + col + store_i_inner
|
|
||||||
|
|
||||||
dst_store = dstf[dst_i].store(srcf[*src_idxs, src_i]).end(store_i_outer, store_i_inner)
|
|
||||||
else:
|
else:
|
||||||
raise NotImplementedError(f"store from {src_dtype.addrspace} to {dst_dtype.addrspace} not implemented")
|
raise NotImplementedError(f"store from {src_dtype.addrspace} to {dst_dtype.addrspace} not implemented")
|
||||||
|
|
||||||
self.ker.push_store(dst_store, dst)
|
self.ker.push_store(dst_store, dst)
|
||||||
return dst.after(dst_store.barrier()).reshape(dst.shape) if after else dst_store
|
return dst.after(dst_store.barrier()).reshape(dst.shape)
|
||||||
|
|||||||
@@ -1,7 +1,8 @@
|
|||||||
from contextlib import AbstractContextManager
|
from contextlib import AbstractContextManager
|
||||||
from tinygrad.uop.ops import UOp, KernelInfo, AxisType
|
from tinygrad.uop.ops import UOp, KernelInfo, AxisType, AddrSpace
|
||||||
from extra.thunder.tiny.tk import WARP_THREADS
|
from extra.thunder.tiny.tk import WARP_THREADS
|
||||||
from extra.thunder.tiny.tk.group import Group
|
from extra.thunder.tiny.tk.group import Group
|
||||||
|
from extra.thunder.tiny.tk.tiles import GL, ST, RT, RV
|
||||||
|
|
||||||
class _tk_range:
|
class _tk_range:
|
||||||
user_rid = 0
|
user_rid = 0
|
||||||
@@ -25,6 +26,11 @@ class Kernel(AbstractContextManager):
|
|||||||
self.range_stack = []
|
self.range_stack = []
|
||||||
self.store_stack = []
|
self.store_stack = []
|
||||||
|
|
||||||
|
self.global_slot = 0
|
||||||
|
self.shared_slot = 0
|
||||||
|
self.register_slot = 0
|
||||||
|
self.allocs = {}
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def warpid(self): return self.threadIdx_x // WARP_THREADS
|
def warpid(self): return self.threadIdx_x // WARP_THREADS
|
||||||
|
|
||||||
@@ -42,6 +48,31 @@ class Kernel(AbstractContextManager):
|
|||||||
if track: self.range_stack.append(rng)
|
if track: self.range_stack.append(rng)
|
||||||
return rng
|
return rng
|
||||||
|
|
||||||
|
def alloc(self, shape, dtype, addrspace:AddrSpace, name:str|None=None):
|
||||||
|
match addrspace:
|
||||||
|
case AddrSpace.GLOBAL:
|
||||||
|
slot = self.global_slot
|
||||||
|
self.global_slot += 1
|
||||||
|
case AddrSpace.LOCAL:
|
||||||
|
slot = self.shared_slot
|
||||||
|
self.shared_slot += 1
|
||||||
|
case AddrSpace.REG:
|
||||||
|
slot = self.register_slot
|
||||||
|
self.register_slot += 1
|
||||||
|
|
||||||
|
uop = UOp.placeholder(shape, dtype, slot=slot, addrspace=addrspace)
|
||||||
|
|
||||||
|
if name:
|
||||||
|
if (name, shape) in self.allocs: return self.allocs[(name, shape)]
|
||||||
|
self.allocs[(name, shape)] = uop
|
||||||
|
|
||||||
|
return uop
|
||||||
|
|
||||||
|
def gl(self, shape, dtype): return GL.create(shape, dtype, self)
|
||||||
|
def st(self, shape, dtype): return ST.create(shape, dtype, self)
|
||||||
|
def rt(self, shape, dtype): return RT.create(shape, dtype, self)
|
||||||
|
def rv(self, length, dtype, layout="naive"): return RV.create(length, dtype, layout, self)
|
||||||
|
|
||||||
def push_store(self, store:UOp, uop:UOp): self.store_stack.append((store, uop))
|
def push_store(self, store:UOp, uop:UOp): self.store_stack.append((store, uop))
|
||||||
|
|
||||||
def finish(self):
|
def finish(self):
|
||||||
@@ -49,7 +80,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].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()
|
||||||
|
|||||||
+133
-42
@@ -1,52 +1,143 @@
|
|||||||
import math
|
import functools
|
||||||
from typing import cast, Callable
|
from tinygrad.dtype import AddrSpace
|
||||||
from tinygrad import Tensor, Device, Context, GlobalCounters, dtypes
|
from tinygrad.mixin import MathMixin
|
||||||
from tinygrad.uop.ops import AxisType, UOp, KernelInfo, Ops
|
from tinygrad.uop.ops import UOp, Ops
|
||||||
from tinygrad.engine.realize import ExecItem, get_runner
|
|
||||||
from tinygrad.dtype import AddrSpace, PtrDType
|
|
||||||
from tinygrad.helpers import getenv, prod
|
|
||||||
|
|
||||||
from extra.thunder.tiny.tk import WARP_THREADS
|
from extra.thunder.tiny.tk import WARP_THREADS
|
||||||
|
|
||||||
class _Slots:
|
def unwrap(x):
|
||||||
def __init__(self):
|
if hasattr(x, "_uop"): return x._uop
|
||||||
self.global_slot = 0
|
if isinstance(x, (list, tuple)): return type(x)(unwrap(y) for y in x)
|
||||||
self.shared_slot = 0
|
if isinstance(x, dict): return {k: unwrap(v) for k,v in x.items()}
|
||||||
self.register_slot = 0
|
return x
|
||||||
slots = _Slots()
|
|
||||||
|
|
||||||
def gl(shape, dtype):
|
def wrap(x, ker, cls):
|
||||||
slots.global_slot += 1
|
if isinstance(x, UOp): return cls(x, ker)
|
||||||
return UOp.placeholder(shape, dtype, slot=slots.global_slot-1)
|
if isinstance(x, (list, tuple)): return type(x)(wrap(y, ker, cls) for y in x)
|
||||||
|
return x
|
||||||
|
|
||||||
shared_slot = 0
|
def autowrap(source_cls, blacklist=None):
|
||||||
def st(shape, dtype):
|
if blacklist is None:
|
||||||
slots.shared_slot += 1
|
blacklist = {
|
||||||
return UOp.placeholder(shape, dtype, addrspace=AddrSpace.LOCAL, slot=slots.shared_slot-1)
|
"__init__", "__new__", "__str__", "__del__", "__repr__", "__dict__", "__getattribute__",
|
||||||
|
"__setattr__", "__delattr__", "__weakref__", "__slots__", "__class__",
|
||||||
|
"__reduce__", "__reduce_ex__", "__getstate__", "__setstate__", "__hash__"
|
||||||
|
}
|
||||||
|
|
||||||
TILE_ROW_DIM, TILE_COL_DIM = 16, 16
|
def decorator(cls):
|
||||||
RT_BASE_TILE_NE = TILE_ROW_DIM * TILE_COL_DIM
|
def __getattr__(self, name):
|
||||||
RT_BASE_TILE_NEPT = RT_BASE_TILE_NE // WARP_THREADS
|
uop = object.__getattribute__(self, "_uop")
|
||||||
register_slot = 0
|
val = getattr(uop, name)
|
||||||
def rt(shape, dtype):
|
if callable(val):
|
||||||
assert len(shape) == 2
|
@functools.wraps(val)
|
||||||
|
def proxy(*args, **kwargs):
|
||||||
|
return wrap(val(*unwrap(args), **unwrap(kwargs)), self.ker, cls)
|
||||||
|
return proxy
|
||||||
|
if name in UOp.__slots__: return val
|
||||||
|
return wrap(val, self.ker, cls)
|
||||||
|
cls.__getattr__ = __getattr__
|
||||||
|
|
||||||
height = shape[0] // TILE_ROW_DIM
|
for name in dir(source_cls):
|
||||||
width = shape[1] // TILE_COL_DIM
|
if name in blacklist or not name.startswith("__"): continue
|
||||||
|
|
||||||
slots.register_slot += 1
|
for base in cls.mro():
|
||||||
return UOp.placeholder((height, width, RT_BASE_TILE_NEPT), dtype, addrspace=AddrSpace.REG, slot=slots.register_slot-1)
|
if base is source_cls: break
|
||||||
|
if name in base.__dict__: break
|
||||||
|
else:
|
||||||
|
original = getattr(source_cls, name)
|
||||||
|
if callable(original):
|
||||||
|
def make_proxy(op_name, func):
|
||||||
|
def proxy(self, *args, **kwargs):
|
||||||
|
return wrap(func(self._uop, *unwrap(args), **unwrap(kwargs)), self.ker, cls)
|
||||||
|
return proxy
|
||||||
|
setattr(cls, name, make_proxy(name, original))
|
||||||
|
|
||||||
def rv(length, dtype, layout="naive"):
|
return cls
|
||||||
tiles = length // TILE_ROW_DIM
|
return decorator
|
||||||
match layout:
|
|
||||||
case "naive":
|
|
||||||
inner_dim = 1
|
|
||||||
outer_dim = (tiles + 1) // 2
|
|
||||||
case "ortho":
|
|
||||||
inner_dim = 1
|
|
||||||
outer_dim = tiles
|
|
||||||
case _: raise NotImplementedError(f"rv layout {layout} not implemented")
|
|
||||||
|
|
||||||
slots.register_slot += 1
|
class TileMathMixin(MathMixin):
|
||||||
return UOp.placeholder((outer_dim, inner_dim, 2), dtype, addrspace=AddrSpace.REG, slot=slots.register_slot-1)
|
def alu(self, op, *src, inner_op=lambda x:x):
|
||||||
|
assert isinstance(self, (RT, RV))
|
||||||
|
if len(src) == 0:
|
||||||
|
if self._uop._shape is None: uop = UOp.alu(self._uop, op)
|
||||||
|
else: uop = self.ker.warp.map(self._uop, lambda x: UOp.alu(x, op))
|
||||||
|
elif len(src) == 1:
|
||||||
|
if self._uop._shape is None: uop = UOp.alu(self._uop, op, inner_op(self._uop.ufix(src[0])))
|
||||||
|
elif isinstance(src[0], (int,float,bool)): uop = self.ker.warp.map(self._uop, lambda x: UOp.alu(x, op, inner_op(x.ufix(src[0]))))
|
||||||
|
elif src[0]._shape is None: uop = UOp.alu(self._uop, op, inner_op(self._uop.ufix(src[0])))
|
||||||
|
else:
|
||||||
|
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: raise NotImplementedError
|
||||||
|
return type(self)(uop, self.ker)
|
||||||
|
def const_like(self, b): return b
|
||||||
|
|
||||||
|
# override ops that do compute on the src uop
|
||||||
|
def sub(self, x, reverse=False):
|
||||||
|
return self.ufix(x).alu(Ops.ADD, self, inner_op=lambda y: -y) if reverse else self.alu(Ops.ADD, self.ufix(x), inner_op=lambda y: -y)
|
||||||
|
def div(self, x, reverse=False):
|
||||||
|
return self.ufix(x).alu(Ops.MUL, self, inner_op=lambda y: 1/y) if reverse else self.alu(Ops.MUL, self.ufix(x), inner_op=lambda y: 1/y)
|
||||||
|
|
||||||
|
@autowrap(UOp)
|
||||||
|
class GL:
|
||||||
|
def __init__(self, uop, ker):
|
||||||
|
self._uop, self.ker = uop, ker
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def create(cls, shape, dtype, ker):
|
||||||
|
uop = ker.alloc(shape, dtype, AddrSpace.GLOBAL)
|
||||||
|
return cls(uop, ker)
|
||||||
|
|
||||||
|
@autowrap(UOp)
|
||||||
|
class ST:
|
||||||
|
def __init__(self, uop, ker):
|
||||||
|
self._uop, self.ker = uop, ker
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def create(cls, shape, dtype, ker):
|
||||||
|
uop = ker.alloc(shape, dtype, AddrSpace.LOCAL)
|
||||||
|
return cls(uop, ker)
|
||||||
|
|
||||||
|
@autowrap(UOp)
|
||||||
|
class RT(TileMathMixin):
|
||||||
|
BASE_TILE_ROWS, BASE_TILE_COLS = 16, 16
|
||||||
|
BASE_TILE_NE = BASE_TILE_ROWS * BASE_TILE_COLS
|
||||||
|
BASE_TILE_NEPT = BASE_TILE_NE // WARP_THREADS
|
||||||
|
|
||||||
|
def __init__(self, uop, ker):
|
||||||
|
self._uop, self.ker = uop, ker
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def create(cls, shape, dtype, ker):
|
||||||
|
assert len(shape) == 2
|
||||||
|
assert shape[0] % RT.BASE_TILE_ROWS == 0
|
||||||
|
assert shape[1] % RT.BASE_TILE_COLS == 0
|
||||||
|
|
||||||
|
height = shape[0] // RT.BASE_TILE_ROWS
|
||||||
|
width = shape[1] // RT.BASE_TILE_COLS
|
||||||
|
|
||||||
|
uop = ker.alloc((height, width, RT.BASE_TILE_NEPT), dtype, AddrSpace.REG)
|
||||||
|
return cls(uop, ker)
|
||||||
|
|
||||||
|
@autowrap(UOp)
|
||||||
|
class RV(TileMathMixin):
|
||||||
|
def __init__(self, uop, ker):
|
||||||
|
self._uop, self.ker = uop, ker
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def create(cls, length, dtype, layout, ker):
|
||||||
|
tiles = length // RT.BASE_TILE_ROWS
|
||||||
|
|
||||||
|
match layout:
|
||||||
|
case "naive":
|
||||||
|
inner_dim = 1
|
||||||
|
outer_dim = (tiles + 1) // 2
|
||||||
|
case "ortho":
|
||||||
|
inner_dim = 1
|
||||||
|
outer_dim = tiles
|
||||||
|
case _: raise NotImplementedError(f"rv layout {layout} not implemented")
|
||||||
|
|
||||||
|
uop = ker.alloc((outer_dim, inner_dim, 2), dtype, AddrSpace.REG)
|
||||||
|
return RV(uop, ker)
|
||||||
|
|
||||||
|
ALL_TILES = UOp | GL | ST | RT | RV
|
||||||
|
|||||||
@@ -1,6 +1,5 @@
|
|||||||
#!/usr/bin/env python3
|
#!/usr/bin/env python3
|
||||||
import sys, os, zlib, struct, hashlib
|
import sys, os, zlib, struct, hashlib
|
||||||
from hexdump import hexdump
|
|
||||||
from tinygrad.helpers import DEBUG, getenv, fetch
|
from tinygrad.helpers import DEBUG, getenv, fetch
|
||||||
from tinygrad.runtime.support.usb import USB3
|
from tinygrad.runtime.support.usb import USB3
|
||||||
|
|
||||||
|
|||||||
@@ -1,10 +0,0 @@
|
|||||||
[mypy]
|
|
||||||
warn_unused_configs = True
|
|
||||||
files = tinygrad
|
|
||||||
ignore_missing_imports = True
|
|
||||||
check_untyped_defs = True
|
|
||||||
explicit_package_bases = True
|
|
||||||
warn_unreachable = True
|
|
||||||
warn_redundant_casts = True
|
|
||||||
# NOTE: had to comment this out to make mypy pass on both CI and OSX
|
|
||||||
#warn_unused_ignores = True
|
|
||||||
+232
@@ -0,0 +1,232 @@
|
|||||||
|
[project]
|
||||||
|
name = "tinygrad"
|
||||||
|
version = "0.11.0"
|
||||||
|
description = "You like pytorch? You like micrograd? You love tinygrad! <3"
|
||||||
|
authors = [{ name = "George Hotz" }]
|
||||||
|
|
||||||
|
classifiers = ["Programming Language :: Python :: 3"]
|
||||||
|
|
||||||
|
license = 'MIT'
|
||||||
|
readme = "README.md"
|
||||||
|
requires-python = ">=3.11"
|
||||||
|
dependencies = []
|
||||||
|
|
||||||
|
[build-system]
|
||||||
|
requires = ["setuptools"]
|
||||||
|
build-backend = "setuptools.build_meta"
|
||||||
|
|
||||||
|
[tool.setuptools]
|
||||||
|
include-package-data = true
|
||||||
|
packages = [
|
||||||
|
'tinygrad',
|
||||||
|
'tinygrad.apps',
|
||||||
|
'tinygrad.codegen',
|
||||||
|
'tinygrad.codegen.opt',
|
||||||
|
'tinygrad.codegen.late',
|
||||||
|
'tinygrad.engine',
|
||||||
|
'tinygrad.mixin',
|
||||||
|
'tinygrad.nn',
|
||||||
|
'tinygrad.renderer',
|
||||||
|
'tinygrad.runtime',
|
||||||
|
'tinygrad.runtime.autogen',
|
||||||
|
'tinygrad.runtime.autogen.am',
|
||||||
|
'tinygrad.runtime.graph',
|
||||||
|
'tinygrad.runtime.support',
|
||||||
|
'tinygrad.runtime.support.am',
|
||||||
|
'tinygrad.runtime.support.nv',
|
||||||
|
'tinygrad.schedule',
|
||||||
|
'tinygrad.uop',
|
||||||
|
'tinygrad.viz',
|
||||||
|
]
|
||||||
|
|
||||||
|
[tool.setuptools.package-data]
|
||||||
|
tinygrad = ["py.typed"]
|
||||||
|
"tinygrad.viz" = ["index.html", "assets/**/*", "js/*"]
|
||||||
|
|
||||||
|
|
||||||
|
[project.optional-dependencies]
|
||||||
|
arm = ["unicorn"]
|
||||||
|
triton = ["triton-nightly>=2.1.0.dev20231014192330"]
|
||||||
|
linting = [
|
||||||
|
"pylint",
|
||||||
|
"mypy==1.18.1",
|
||||||
|
"typing-extensions",
|
||||||
|
"pre-commit",
|
||||||
|
"ruff",
|
||||||
|
"numpy",
|
||||||
|
"typeguard",
|
||||||
|
]
|
||||||
|
# mlperf = [
|
||||||
|
# "mlperf-logging @ git+https://github.com/mlperf/[email protected]",
|
||||||
|
# ]
|
||||||
|
testing_minimal = [
|
||||||
|
"numpy",
|
||||||
|
"torch==2.9.0",
|
||||||
|
"pytest",
|
||||||
|
"pytest-xdist",
|
||||||
|
"pytest-timeout",
|
||||||
|
"pytest-split",
|
||||||
|
"hypothesis",
|
||||||
|
"z3-solver",
|
||||||
|
]
|
||||||
|
testing_unit = ["tinygrad[testing_minimal]", "tqdm", "safetensors", "tabulate"]
|
||||||
|
testing = [
|
||||||
|
"tinygrad[testing_minimal]",
|
||||||
|
"pillow",
|
||||||
|
"onnx==1.18.0",
|
||||||
|
"onnx2torch",
|
||||||
|
"onnxruntime",
|
||||||
|
"opencv-python",
|
||||||
|
"tabulate",
|
||||||
|
"tqdm",
|
||||||
|
"safetensors",
|
||||||
|
"transformers",
|
||||||
|
"sentencepiece",
|
||||||
|
"tiktoken",
|
||||||
|
"blobfile",
|
||||||
|
"librosa",
|
||||||
|
# librosa needs numba but uv ignores python upper bounds and some numba versions require <python3.10
|
||||||
|
"numba>=0.55",
|
||||||
|
"networkx",
|
||||||
|
"nibabel",
|
||||||
|
"bottle",
|
||||||
|
"ggml-python",
|
||||||
|
"capstone",
|
||||||
|
"pycocotools",
|
||||||
|
"boto3",
|
||||||
|
"pandas",
|
||||||
|
"influxdb3-python",
|
||||||
|
]
|
||||||
|
docs = [
|
||||||
|
"mkdocs",
|
||||||
|
"mkdocs-material",
|
||||||
|
"mkdocstrings[python]",
|
||||||
|
"markdown-callouts",
|
||||||
|
"markdown-exec[ansi]",
|
||||||
|
"black",
|
||||||
|
"numpy",
|
||||||
|
]
|
||||||
|
|
||||||
|
|
||||||
|
[tool.mutmut]
|
||||||
|
paths_to_mutate = ["tinygrad/"]
|
||||||
|
do_not_mutate = [
|
||||||
|
"tinygrad/apps/*",
|
||||||
|
"tinygrad/codegen/*",
|
||||||
|
"tinygrad/engine/*",
|
||||||
|
"tinygrad/nn/*",
|
||||||
|
"tinygrad/renderer/*",
|
||||||
|
"tinygrad/runtime/*",
|
||||||
|
"tinygrad/schedule/*",
|
||||||
|
"tinygrad/uop/*",
|
||||||
|
"tinygrad/viz/*",
|
||||||
|
"tinygrad/device.py",
|
||||||
|
"tinygrad/dtype.py",
|
||||||
|
"tinygrad/gradient.py",
|
||||||
|
"tinygrad/helpers.py",
|
||||||
|
"tinygrad/tensor.py",
|
||||||
|
]
|
||||||
|
tests_dir = ["test/test_tiny.py", "test/test_ops.py"]
|
||||||
|
debug = true
|
||||||
|
|
||||||
|
|
||||||
|
[tool.mypy]
|
||||||
|
warn_unused_configs = true
|
||||||
|
files = ["tinygrad"]
|
||||||
|
ignore_missing_imports = true
|
||||||
|
check_untyped_defs = true
|
||||||
|
explicit_package_bases = true
|
||||||
|
warn_unreachable = true
|
||||||
|
warn_redundant_casts = true
|
||||||
|
# NOTE: had to comment this out to make mypy pass on both CI and OSX
|
||||||
|
#warn_unused_ignores = true
|
||||||
|
|
||||||
|
[tool.pytest.ini_options]
|
||||||
|
norecursedirs = [
|
||||||
|
"extra",
|
||||||
|
".hypothesis",
|
||||||
|
".git",
|
||||||
|
]
|
||||||
|
timeout = 300
|
||||||
|
timeout_method = "thread"
|
||||||
|
timeout_func_only = true
|
||||||
|
testpaths = ["test"]
|
||||||
|
|
||||||
|
[tool.ruff]
|
||||||
|
preview = true
|
||||||
|
target-version = "py311"
|
||||||
|
line-length = 150
|
||||||
|
indent-width = 2
|
||||||
|
exclude = [
|
||||||
|
".git/",
|
||||||
|
"docs/",
|
||||||
|
"extra/",
|
||||||
|
"test/external/mlperf_resnet",
|
||||||
|
"test/external/mlperf_unet3d",
|
||||||
|
]
|
||||||
|
|
||||||
|
[tool.ruff.lint]
|
||||||
|
select = [
|
||||||
|
"F", # Pyflakes
|
||||||
|
"W6",
|
||||||
|
"E71",
|
||||||
|
"E72",
|
||||||
|
"E112", # no-indented-block
|
||||||
|
"E113", # unexpected-indentation
|
||||||
|
# "E124",
|
||||||
|
"E203", # whitespace-before-punctuation
|
||||||
|
"E272", # multiple-spaces-before-keyword
|
||||||
|
"E275", # missing-whitespace-after-keyword
|
||||||
|
"E303", # too-many-blank-lines
|
||||||
|
"E304", # blank-line-after-decorator
|
||||||
|
"E501", # line-too-long
|
||||||
|
# "E502",
|
||||||
|
"E702", # multiple-statements-on-one-line-semicolon
|
||||||
|
"E703", # useless-semicolon
|
||||||
|
"E731", # lambda-assignment
|
||||||
|
"W191", # tab-indentation
|
||||||
|
"W291", # trailing-whitespace
|
||||||
|
"W293", # blank-line-with-whitespace
|
||||||
|
"UP039", # unnecessary-class-parentheses
|
||||||
|
"C416", # unnecessary-comprehension
|
||||||
|
"RET506", # superfluous-else-raise
|
||||||
|
"RET507", # superfluous-else-continue
|
||||||
|
"A", # builtin-variable-shadowing, builtin-argument-shadowing, builtin-attribute-shadowing
|
||||||
|
"FURB110",# if-exp-instead-of-or-operator
|
||||||
|
"RUF018", # assignment-in-assert
|
||||||
|
]
|
||||||
|
|
||||||
|
# detect unused imports in examples
|
||||||
|
[tool.ruff.lint.per-file-ignores]
|
||||||
|
"examples/**/*.py" = [
|
||||||
|
"W6",
|
||||||
|
"E71",
|
||||||
|
"E72",
|
||||||
|
"E112",
|
||||||
|
"E113",
|
||||||
|
"E203",
|
||||||
|
"E272",
|
||||||
|
"E275",
|
||||||
|
"E303",
|
||||||
|
"E304",
|
||||||
|
"E501",
|
||||||
|
"E702",
|
||||||
|
"E703",
|
||||||
|
"E731",
|
||||||
|
"W191",
|
||||||
|
"W291",
|
||||||
|
"W293",
|
||||||
|
"UP039",
|
||||||
|
"C416",
|
||||||
|
"RET506",
|
||||||
|
"RET507",
|
||||||
|
"A",
|
||||||
|
"FURB110",
|
||||||
|
"RUF018",
|
||||||
|
"F541",
|
||||||
|
"F841",
|
||||||
|
]
|
||||||
|
"tinygrad/runtime/autogen/**/*.py" = ["E501", "F401", "E722", "E731", "F821", "A006"]
|
||||||
|
|
||||||
|
[tool.ruff.format]
|
||||||
|
exclude = ["*"]
|
||||||
@@ -1,9 +0,0 @@
|
|||||||
[pytest]
|
|
||||||
norecursedirs =
|
|
||||||
extra
|
|
||||||
.hypothesis
|
|
||||||
.git
|
|
||||||
timeout = 300
|
|
||||||
timeout_method = thread
|
|
||||||
timeout_func_only = true
|
|
||||||
testpaths = test
|
|
||||||
@@ -1,56 +0,0 @@
|
|||||||
indent-width = 2
|
|
||||||
preview = true
|
|
||||||
target-version = "py311"
|
|
||||||
|
|
||||||
lint.select = [
|
|
||||||
"F", # Pyflakes
|
|
||||||
"W6",
|
|
||||||
"E71",
|
|
||||||
"E72",
|
|
||||||
"E112", # no-indented-block
|
|
||||||
"E113", # unexpected-indentation
|
|
||||||
# "E124",
|
|
||||||
"E203", # whitespace-before-punctuation
|
|
||||||
"E272", # multiple-spaces-before-keyword
|
|
||||||
"E275", # missing-whitespace-after-keyword
|
|
||||||
"E303", # too-many-blank-lines
|
|
||||||
"E304", # blank-line-after-decorator
|
|
||||||
"E501", # line-too-long
|
|
||||||
# "E502",
|
|
||||||
"E702", # multiple-statements-on-one-line-semicolon
|
|
||||||
"E703", # useless-semicolon
|
|
||||||
"E731", # lambda-assignment
|
|
||||||
"W191", # tab-indentation
|
|
||||||
"W291", # trailing-whitespace
|
|
||||||
"W293", # blank-line-with-whitespace
|
|
||||||
"UP039", # unnecessary-class-parentheses
|
|
||||||
"C416", # unnecessary-comprehension
|
|
||||||
"RET506", # superfluous-else-raise
|
|
||||||
"RET507", # superfluous-else-continue
|
|
||||||
"A", # builtin-variable-shadowing, builtin-argument-shadowing, builtin-attribute-shadowing
|
|
||||||
"FURB110",# if-exp-instead-of-or-operator
|
|
||||||
"RUF018", # assignment-in-assert
|
|
||||||
]
|
|
||||||
|
|
||||||
line-length = 150
|
|
||||||
|
|
||||||
exclude = [
|
|
||||||
".git/",
|
|
||||||
"docs/",
|
|
||||||
"extra/",
|
|
||||||
"tinygrad/runtime/autogen",
|
|
||||||
"test/external/mlperf_resnet",
|
|
||||||
"test/external/mlperf_unet3d",
|
|
||||||
]
|
|
||||||
|
|
||||||
# detect unused imports in examples
|
|
||||||
[lint.per-file-ignores]
|
|
||||||
"examples/**/*.py" = [
|
|
||||||
"W6", "E71", "E72", "E112", "E113", "E203", "E272", "E275",
|
|
||||||
"E303", "E304", "E501", "E702", "E703", "E731", "W191",
|
|
||||||
"W291", "W293", "UP039", "C416", "RET506", "RET507", "A",
|
|
||||||
"FURB110", "RUF018", "F541", "F841"
|
|
||||||
]
|
|
||||||
|
|
||||||
[format]
|
|
||||||
exclude = ["*"]
|
|
||||||
@@ -1,21 +0,0 @@
|
|||||||
[mutmut]
|
|
||||||
paths_to_mutate=tinygrad
|
|
||||||
do_not_mutate=
|
|
||||||
tinygrad/apps/*
|
|
||||||
tinygrad/codegen/*
|
|
||||||
tinygrad/engine/*
|
|
||||||
tinygrad/nn/*
|
|
||||||
tinygrad/renderer/*
|
|
||||||
tinygrad/runtime/*
|
|
||||||
tinygrad/schedule/*
|
|
||||||
tinygrad/uop/*
|
|
||||||
tinygrad/viz/*
|
|
||||||
tinygrad/device.py
|
|
||||||
tinygrad/dtype.py
|
|
||||||
tinygrad/gradient.py
|
|
||||||
tinygrad/helpers.py
|
|
||||||
tinygrad/tensor.py
|
|
||||||
tests_dir=
|
|
||||||
test/test_tiny.py
|
|
||||||
test/test_ops.py
|
|
||||||
debug=true
|
|
||||||
@@ -1,111 +0,0 @@
|
|||||||
#!/usr/bin/env python3
|
|
||||||
|
|
||||||
from pathlib import Path
|
|
||||||
from setuptools import setup
|
|
||||||
|
|
||||||
directory = Path(__file__).resolve().parent
|
|
||||||
with open(directory / 'README.md', encoding='utf-8') as f:
|
|
||||||
long_description = f.read()
|
|
||||||
|
|
||||||
testing_minimal = [
|
|
||||||
"numpy",
|
|
||||||
"torch==2.9.0",
|
|
||||||
"pytest",
|
|
||||||
"pytest-xdist",
|
|
||||||
"pytest-timeout",
|
|
||||||
"pytest-split",
|
|
||||||
"hypothesis",
|
|
||||||
"z3-solver",
|
|
||||||
]
|
|
||||||
|
|
||||||
setup(name='tinygrad',
|
|
||||||
version='0.11.0',
|
|
||||||
description='You like pytorch? You like micrograd? You love tinygrad! <3',
|
|
||||||
author='George Hotz',
|
|
||||||
license='MIT',
|
|
||||||
long_description=long_description,
|
|
||||||
long_description_content_type='text/markdown',
|
|
||||||
packages = [
|
|
||||||
'tinygrad',
|
|
||||||
'tinygrad.apps',
|
|
||||||
'tinygrad.codegen',
|
|
||||||
'tinygrad.codegen.opt',
|
|
||||||
'tinygrad.codegen.late',
|
|
||||||
'tinygrad.engine',
|
|
||||||
'tinygrad.mixin',
|
|
||||||
'tinygrad.nn',
|
|
||||||
'tinygrad.renderer',
|
|
||||||
'tinygrad.runtime',
|
|
||||||
'tinygrad.runtime.autogen',
|
|
||||||
'tinygrad.runtime.autogen.am',
|
|
||||||
'tinygrad.runtime.autogen.nv',
|
|
||||||
'tinygrad.runtime.graph',
|
|
||||||
'tinygrad.runtime.support',
|
|
||||||
'tinygrad.runtime.support.am',
|
|
||||||
'tinygrad.runtime.support.nv',
|
|
||||||
'tinygrad.schedule',
|
|
||||||
'tinygrad.uop',
|
|
||||||
'tinygrad.viz',
|
|
||||||
],
|
|
||||||
package_data = {'tinygrad': ['py.typed'], 'tinygrad.viz': ['index.html', 'assets/**/*', 'js/*']},
|
|
||||||
classifiers=[
|
|
||||||
"Programming Language :: Python :: 3",
|
|
||||||
"License :: OSI Approved :: MIT License"
|
|
||||||
],
|
|
||||||
install_requires=[],
|
|
||||||
python_requires='>=3.11',
|
|
||||||
extras_require={
|
|
||||||
'arm': ["unicorn"],
|
|
||||||
'triton': ["triton-nightly>=2.1.0.dev20231014192330"],
|
|
||||||
'linting': [
|
|
||||||
"pylint",
|
|
||||||
"mypy==1.18.1",
|
|
||||||
"typing-extensions",
|
|
||||||
"pre-commit",
|
|
||||||
"ruff",
|
|
||||||
"numpy",
|
|
||||||
"typeguard",
|
|
||||||
],
|
|
||||||
#'mlperf': ["mlperf-logging @ git+https://github.com/mlperf/[email protected]"],
|
|
||||||
'testing_minimal': testing_minimal,
|
|
||||||
'testing_unit': testing_minimal + [
|
|
||||||
"tqdm",
|
|
||||||
"safetensors",
|
|
||||||
"tabulate", # for sz.py
|
|
||||||
],
|
|
||||||
'testing': testing_minimal + [
|
|
||||||
"pillow",
|
|
||||||
"onnx==1.18.0",
|
|
||||||
"onnx2torch",
|
|
||||||
"onnxruntime",
|
|
||||||
"opencv-python",
|
|
||||||
"tabulate",
|
|
||||||
"tqdm",
|
|
||||||
"safetensors",
|
|
||||||
"transformers",
|
|
||||||
"sentencepiece",
|
|
||||||
"tiktoken",
|
|
||||||
"blobfile",
|
|
||||||
"librosa",
|
|
||||||
"numba>=0.55", # librosa needs numba but uv ignores python upper bounds and some numba versions require <python3.10
|
|
||||||
"networkx",
|
|
||||||
"nibabel",
|
|
||||||
"bottle",
|
|
||||||
"ggml-python",
|
|
||||||
"capstone",
|
|
||||||
"pycocotools",
|
|
||||||
"boto3",
|
|
||||||
"pandas",
|
|
||||||
"influxdb3-python"
|
|
||||||
],
|
|
||||||
'docs': [
|
|
||||||
"mkdocs",
|
|
||||||
"mkdocs-material",
|
|
||||||
"mkdocstrings[python]",
|
|
||||||
"markdown-callouts",
|
|
||||||
"markdown-exec[ansi]",
|
|
||||||
"black",
|
|
||||||
"numpy",
|
|
||||||
],
|
|
||||||
},
|
|
||||||
include_package_data=True)
|
|
||||||
-163
@@ -1,163 +0,0 @@
|
|||||||
# ruff: noqa: E501 E712
|
|
||||||
from tinygrad import dtypes, Device
|
|
||||||
from tinygrad.uop.ops import UOp, AxisType, Ops, KernelInfo
|
|
||||||
from tinygrad.codegen import full_rewrite
|
|
||||||
from tinygrad.renderer import ProgramSpec
|
|
||||||
from tinygrad.engine.realize import CompiledRunner
|
|
||||||
from tinygrad.helpers import dedup
|
|
||||||
from tinygrad.device import Buffer
|
|
||||||
from tinygrad.dtype import ImageDType, Invalid
|
|
||||||
|
|
||||||
c0 = UOp(Ops.DEFINE_GLOBAL, dtypes.float.ptr(1576), (), 0)
|
|
||||||
c2 = UOp.range(1576, 20, AxisType.LOOP)
|
|
||||||
c5 = c2<55
|
|
||||||
c6 = UOp(Ops.DEFINE_GLOBAL, dtypes.imageh((1, 16, 4)), (), 1)
|
|
||||||
c8 = UOp.range(16, 0, AxisType.REDUCE)
|
|
||||||
c11 = UOp.range(4, 1, AxisType.REDUCE)
|
|
||||||
c14 = UOp(Ops.DEFINE_GLOBAL, dtypes.imageh((14, 64, 4)), (), 2)
|
|
||||||
c25 = c5.where((c2%4*4+c11+c8*16+c2//4*256), UOp.const(dtypes.index, Invalid))
|
|
||||||
c27 = c6.index((c8*4+c11))*c14.index(c25)
|
|
||||||
c29 = UOp(Ops.DEFINE_GLOBAL, dtypes.float.ptr(55), (), 3)
|
|
||||||
c30 = c5.where(c2, UOp.const(dtypes.index, Invalid))
|
|
||||||
c34 = c5.where((c27.reduce(c8, c11, arg=Ops.ADD)+c29.index(c30)), UOp.const(dtypes.float, 0.0))
|
|
||||||
c38 = c2<87
|
|
||||||
c39 = (c5!=True)&c38
|
|
||||||
c40 = UOp(Ops.DEFINE_GLOBAL, dtypes.imageh((1, 8, 4)), (), 4)
|
|
||||||
c42 = UOp.range(8, 2, AxisType.REDUCE)
|
|
||||||
c44 = UOp.range(4, 3, AxisType.REDUCE)
|
|
||||||
c47 = UOp(Ops.DEFINE_GLOBAL, dtypes.imageh((8, 32, 4)), (), 5)
|
|
||||||
c49 = c2+1
|
|
||||||
c51 = c49%4*4
|
|
||||||
c57 = c49//4*128
|
|
||||||
c61 = c39.where((c51+c44+c42*16+c57+-1792), UOp.const(dtypes.index, Invalid))
|
|
||||||
c63 = c40.index((c42*4+c44))*c47.index(c61)
|
|
||||||
c65 = UOp(Ops.DEFINE_GLOBAL, dtypes.float.ptr(32), (), 6)
|
|
||||||
c68 = c39.where((c2+-55), UOp.const(dtypes.index, Invalid))
|
|
||||||
c71 = c39.where((c63.reduce(c42, c44, arg=Ops.ADD)+c65.index(c68)), UOp.const(dtypes.float, 0.0))
|
|
||||||
c75 = c2<99
|
|
||||||
c76 = (c38!=True)&c75
|
|
||||||
c77 = UOp(Ops.DEFINE_GLOBAL, dtypes.imageh((1, 8, 4)), (), 7)
|
|
||||||
c78 = UOp.range(8, 4, AxisType.REDUCE)
|
|
||||||
c80 = UOp.range(4, 5, AxisType.REDUCE)
|
|
||||||
c83 = UOp(Ops.DEFINE_GLOBAL, dtypes.imageh((3, 32, 4)), (), 8)
|
|
||||||
c90 = c76.where((c51+c80+c78*16+c57+-2816), UOp.const(dtypes.index, Invalid))
|
|
||||||
c92 = c77.index((c78*4+c80))*c83.index(c90)
|
|
||||||
c94 = UOp(Ops.DEFINE_GLOBAL, dtypes.float.ptr(12), (), 9)
|
|
||||||
c97 = c76.where((c2+-87), UOp.const(dtypes.index, Invalid))
|
|
||||||
c100 = c76.where((c92.reduce(c78, c80, arg=Ops.ADD)+c94.index(c97)), UOp.const(dtypes.float, 0.0))
|
|
||||||
c104 = c2<105
|
|
||||||
c105 = (c75!=True)&c104
|
|
||||||
c106 = UOp(Ops.DEFINE_GLOBAL, dtypes.imageh((1, 8, 4)), (), 10)
|
|
||||||
c107 = UOp.range(8, 6, AxisType.REDUCE)
|
|
||||||
c109 = UOp.range(4, 7, AxisType.REDUCE)
|
|
||||||
c112 = UOp(Ops.DEFINE_GLOBAL, dtypes.imageh((2, 32, 4)), (), 11)
|
|
||||||
c119 = c105.where((c51+c109+c107*16+c57+-3200), UOp.const(dtypes.index, Invalid))
|
|
||||||
c121 = c106.index((c107*4+c109))*c112.index(c119)
|
|
||||||
c123 = UOp(Ops.DEFINE_GLOBAL, dtypes.float.ptr(6), (), 12)
|
|
||||||
c126 = c105.where((c2+-99), UOp.const(dtypes.index, Invalid))
|
|
||||||
c129 = c105.where((c121.reduce(c107, c109, arg=Ops.ADD)+c123.index(c126)), UOp.const(dtypes.float, 0.0))
|
|
||||||
c133 = c2<117
|
|
||||||
c134 = (c104!=True)&c133
|
|
||||||
c135 = UOp(Ops.DEFINE_GLOBAL, dtypes.imageh((1, 8, 4)), (), 13)
|
|
||||||
c136 = UOp.range(8, 8, AxisType.REDUCE)
|
|
||||||
c138 = UOp.range(4, 9, AxisType.REDUCE)
|
|
||||||
c141 = UOp(Ops.DEFINE_GLOBAL, dtypes.imageh((3, 32, 4)), (), 14)
|
|
||||||
c143 = c2+3
|
|
||||||
c145 = c143%4*4
|
|
||||||
c149 = c143//4
|
|
||||||
c150 = c149*128
|
|
||||||
c154 = c134.where((c145+c138+c136*16+c150+-3456), UOp.const(dtypes.index, Invalid))
|
|
||||||
c156 = c135.index((c136*4+c138))*c141.index(c154)
|
|
||||||
c158 = UOp(Ops.DEFINE_GLOBAL, dtypes.float.ptr(12), (), 15)
|
|
||||||
c161 = c134.where((c2+-105), UOp.const(dtypes.index, Invalid))
|
|
||||||
c164 = c134.where((c156.reduce(c136, c138, arg=Ops.ADD)+c158.index(c161)), UOp.const(dtypes.float, 0.0))
|
|
||||||
c168 = c2<645
|
|
||||||
c169 = (c133!=True)&c168
|
|
||||||
c170 = UOp(Ops.DEFINE_GLOBAL, dtypes.imageh((1, 16, 4)), (), 16)
|
|
||||||
c171 = UOp.range(16, 10, AxisType.REDUCE)
|
|
||||||
c173 = UOp.range(4, 11, AxisType.REDUCE)
|
|
||||||
c176 = UOp(Ops.DEFINE_GLOBAL, dtypes.imageh((132, 64, 4)), (), 17)
|
|
||||||
c180 = c149*256
|
|
||||||
c184 = c169.where((c145+c173+c171*16+c180+-7680), UOp.const(dtypes.index, Invalid))
|
|
||||||
c186 = c170.index((c171*4+c173))*c176.index(c184)
|
|
||||||
c188 = UOp(Ops.DEFINE_GLOBAL, dtypes.float.ptr(528), (), 18)
|
|
||||||
c191 = c169.where((c2+-117), UOp.const(dtypes.index, Invalid))
|
|
||||||
c194 = c169.where((c186.reduce(c171, c173, arg=Ops.ADD)+c188.index(c191)), UOp.const(dtypes.float, 0.0))
|
|
||||||
c198 = c2<653
|
|
||||||
c199 = (c168!=True)&c198
|
|
||||||
c200 = UOp(Ops.DEFINE_GLOBAL, dtypes.imageh((1, 4, 4)), (), 19)
|
|
||||||
c201 = UOp.range(4, 12, AxisType.REDUCE)
|
|
||||||
c203 = UOp.range(4, 13, AxisType.REDUCE)
|
|
||||||
c206 = UOp(Ops.DEFINE_GLOBAL, dtypes.imageh((2, 16, 4)), (), 20)
|
|
||||||
c215 = c199.where((c145+c203+c201*16+c149*64+-10368), UOp.const(dtypes.index, Invalid))
|
|
||||||
c217 = c200.index((c201*4+c203))*c206.index(c215)
|
|
||||||
c219 = UOp(Ops.DEFINE_GLOBAL, dtypes.float.ptr(8), (), 21)
|
|
||||||
c222 = c199.where((c2+-645), UOp.const(dtypes.index, Invalid))
|
|
||||||
c225 = c199.where((c217.reduce(c201, c203, arg=Ops.ADD)+c219.index(c222)), UOp.const(dtypes.float, 0.0))
|
|
||||||
c229 = c2<917
|
|
||||||
c230 = (c198!=True)&c229
|
|
||||||
c231 = UOp(Ops.DEFINE_GLOBAL, dtypes.imageh((1, 8, 4)), (), 22)
|
|
||||||
c232 = UOp.range(8, 14, AxisType.REDUCE)
|
|
||||||
c234 = UOp.range(4, 15, AxisType.REDUCE)
|
|
||||||
c237 = UOp(Ops.DEFINE_GLOBAL, dtypes.imageh((66, 32, 4)), (), 23)
|
|
||||||
c244 = c230.where((c145+c234+c232*16+c150+-20992), UOp.const(dtypes.index, Invalid))
|
|
||||||
c246 = c231.index((c232*4+c234))*c237.index(c244)
|
|
||||||
c248 = UOp(Ops.DEFINE_GLOBAL, dtypes.float.ptr(264), (), 24)
|
|
||||||
c251 = c230.where((c2+-653), UOp.const(dtypes.index, Invalid))
|
|
||||||
c254 = c230.where((c246.reduce(c232, c234, arg=Ops.ADD)+c248.index(c251)), UOp.const(dtypes.float, 0.0))
|
|
||||||
c258 = c2<1061
|
|
||||||
c259 = (c229!=True)&c258
|
|
||||||
c260 = UOp(Ops.DEFINE_GLOBAL, dtypes.imageh((1, 16, 4)), (), 25)
|
|
||||||
c261 = UOp.range(16, 16, AxisType.REDUCE)
|
|
||||||
c263 = UOp.range(4, 17, AxisType.REDUCE)
|
|
||||||
c266 = UOp(Ops.DEFINE_GLOBAL, dtypes.imageh((36, 64, 4)), (), 26)
|
|
||||||
c273 = c259.where((c145+c263+c261*16+c180+-58880), UOp.const(dtypes.index, Invalid))
|
|
||||||
c275 = c260.index((c261*4+c263))*c266.index(c273)
|
|
||||||
c277 = UOp(Ops.DEFINE_GLOBAL, dtypes.float.ptr(144), (), 27)
|
|
||||||
c280 = c259.where((c2+-917), UOp.const(dtypes.index, Invalid))
|
|
||||||
c283 = UOp(Ops.DEFINE_GLOBAL, dtypes.float.ptr(144), (), 28)
|
|
||||||
c286 = c259.where(((c275.reduce(c261, c263, arg=Ops.ADD)+c277.index(c280))*c283.index(c280)), UOp.const(dtypes.float, 0.0))
|
|
||||||
c290 = c2<1064
|
|
||||||
c291 = (c258!=True)&c290
|
|
||||||
c292 = UOp(Ops.DEFINE_GLOBAL, dtypes.imageh((1, 4, 4)), (), 29)
|
|
||||||
c293 = UOp.range(4, 18, AxisType.REDUCE)
|
|
||||||
c295 = UOp.range(4, 19, AxisType.REDUCE)
|
|
||||||
c298 = UOp(Ops.DEFINE_GLOBAL, dtypes.imageh((1, 16, 4)), (), 30)
|
|
||||||
c305 = c291.where((c2*4+c295+c293*16+-4244), UOp.const(dtypes.index, Invalid))
|
|
||||||
c307 = c292.index((c293*4+c295))*c298.index(c305)
|
|
||||||
c309 = UOp(Ops.DEFINE_GLOBAL, dtypes.float.ptr(3), (), 31)
|
|
||||||
c312 = c291.where((c2+-1061), UOp.const(dtypes.index, Invalid))
|
|
||||||
c315 = c291.where((c307.reduce(c293, c295, arg=Ops.ADD)+c309.index(c312)), UOp.const(dtypes.float, 0.0))
|
|
||||||
c317 = UOp(Ops.DEFINE_GLOBAL, dtypes.imageh((1, 128, 4)), (), 32)
|
|
||||||
c321 = (c290!=True).where((c2+-1064), UOp.const(dtypes.index, Invalid))
|
|
||||||
c323 = UOp(Ops.DEFINE_GLOBAL, dtypes.float.ptr(1), (), 33)
|
|
||||||
c328 = c290.where(UOp.const(dtypes.float, 0.0), (c317.index(c321)*c323.index(UOp.const(dtypes.index, 0)).reciprocal()))
|
|
||||||
c329 = c34+c71+c100+c129+c164+c194+c225+c254+c286+c315+c328
|
|
||||||
c331 = c0.index(c2, ptr=True).store(c329).end(c2)
|
|
||||||
ast = c331.sink(arg=KernelInfo(name="cat", opts_to_apply=None))
|
|
||||||
|
|
||||||
compiler = Device.default.compiler
|
|
||||||
renderer = Device.default.renderer
|
|
||||||
allocator = Device.default.allocator
|
|
||||||
|
|
||||||
uops = full_rewrite(ast, renderer)
|
|
||||||
src = renderer.render(uops)
|
|
||||||
|
|
||||||
# NOLOCALS=1 IMAGE=2 DEV=CL
|
|
||||||
lib = compiler.compile(src)
|
|
||||||
|
|
||||||
ps = ProgramSpec("cat", src, Device.DEFAULT, ast, uops)
|
|
||||||
# print(ps.src)
|
|
||||||
# print(ps.applied_opts)
|
|
||||||
# NOTE: this is faster with no GROUP and with NOLOCALS
|
|
||||||
# (Opt(op=OptOps.UPCAST, axis=1, arg=4), Opt(op=OptOps.UNROLL, axis=19, arg=4), Opt(op=OptOps.UNROLL, axis=17, arg=4), Opt(op=OptOps.UNROLL, axis=15, arg=4), Opt(op=OptOps.UNROLL, axis=13, arg=4), Opt(op=OptOps.UNROLL, axis=11, arg=4), Opt(op=OptOps.UNROLL, axis=9, arg=4), Opt(op=OptOps.UNROLL, axis=7, arg=4), Opt(op=OptOps.UNROLL, axis=5, arg=4), Opt(op=OptOps.UNROLL, axis=3, arg=4), Opt(op=OptOps.UNROLL, axis=1, arg=4), Opt(op=OptOps.NOLOCALS, axis=None, arg=None))
|
|
||||||
cr = CompiledRunner(ps, precompiled=lib)
|
|
||||||
|
|
||||||
gs = sorted(dedup([u for u in ast.toposort() if u.op is Ops.DEFINE_GLOBAL]), key=lambda u: u.arg)
|
|
||||||
print(len(gs))
|
|
||||||
print([g.dtype for g in gs])
|
|
||||||
|
|
||||||
bufs = [Buffer(ps.device, g.size, g.dtype if isinstance(g.dtype, ImageDType) else g.dtype._base).ensure_allocated() for g in gs]
|
|
||||||
|
|
||||||
t = cr(bufs, wait=True)
|
|
||||||
print(f"{t*1e6:.2f} us")
|
|
||||||
+28
-3
@@ -1,8 +1,8 @@
|
|||||||
# ruff: noqa: E501 E712
|
# ruff: noqa: E501 E712 F401
|
||||||
from tinygrad import dtypes, Device
|
from tinygrad import dtypes, Device
|
||||||
from tinygrad.uop.ops import UOp, AxisType, Ops, KernelInfo
|
from tinygrad.uop.ops import UOp, AxisType, Ops, KernelInfo
|
||||||
from tinygrad.codegen import full_rewrite
|
from tinygrad.codegen import full_rewrite
|
||||||
# from tinygrad.codegen.opt import Opt, OptOps
|
from tinygrad.codegen.opt import Opt, OptOps # pylint: disable=unused-import
|
||||||
from tinygrad.renderer import ProgramSpec
|
from tinygrad.renderer import ProgramSpec
|
||||||
from tinygrad.engine.realize import CompiledRunner
|
from tinygrad.engine.realize import CompiledRunner
|
||||||
from tinygrad.helpers import dedup, getenv
|
from tinygrad.helpers import dedup, getenv
|
||||||
@@ -33,6 +33,8 @@ def vision_conv_143():
|
|||||||
c67 = c0.index((c2*128+c5+c8*4096), ptr=True).store(c65).end(c8, c2, c5)
|
c67 = c0.index((c2*128+c5+c8*4096), ptr=True).store(c65).end(c8, c2, c5)
|
||||||
|
|
||||||
opts = None
|
opts = None
|
||||||
|
# JITBEAM=2
|
||||||
|
# (Opt(op=OptOps.UPCAST, axis=2, arg=4), Opt(op=OptOps.NOLOCALS, axis=None, arg=None), Opt(op=OptOps.UPCAST, axis=2, arg=2), Opt(op=OptOps.UPCAST, axis=1, arg=4), Opt(op=OptOps.SWAP, axis=1, arg=2))
|
||||||
return c67.sink(arg=KernelInfo(name="conv", opts_to_apply=opts))
|
return c67.sink(arg=KernelInfo(name="conv", opts_to_apply=opts))
|
||||||
|
|
||||||
def vision_conv_153():
|
def vision_conv_153():
|
||||||
@@ -57,9 +59,32 @@ def vision_conv_153():
|
|||||||
c67 = c0.index((c2*256+c5+c8*4096), ptr=True).store(c65).end(c8, c2, c5)
|
c67 = c0.index((c2*256+c5+c8*4096), ptr=True).store(c65).end(c8, c2, c5)
|
||||||
|
|
||||||
opts = None
|
opts = None
|
||||||
|
# JITBEAM=2
|
||||||
|
# (Opt(op=OptOps.UPCAST, axis=2, arg=4), Opt(op=OptOps.NOLOCALS, axis=None, arg=None), Opt(op=OptOps.UPCAST, axis=2, arg=2), Opt(op=OptOps.SWAP, axis=1, arg=2))
|
||||||
return c67.sink(arg=KernelInfo(name="conv", opts_to_apply=opts))
|
return c67.sink(arg=KernelInfo(name="conv", opts_to_apply=opts))
|
||||||
|
|
||||||
ast = vision_conv_143() if getenv("NUM", 143) == 143 else vision_conv_153()
|
def dm_conv_172():
|
||||||
|
c0 = UOp(Ops.DEFINE_GLOBAL, dtypes.imageh((1, 240, 4)), (), 0)
|
||||||
|
c2 = UOp.range(960, 4, AxisType.LOOP)
|
||||||
|
c5 = UOp(Ops.DEFINE_GLOBAL, dtypes.imageh((8, 384, 4)), (), 1)
|
||||||
|
c7 = UOp.range(32, 0, AxisType.REDUCE)
|
||||||
|
c10 = UOp.range(4, 1, AxisType.REDUCE)
|
||||||
|
c13 = UOp.range(12, 3, AxisType.REDUCE)
|
||||||
|
c18 = UOp.range(8, 2, AxisType.REDUCE)
|
||||||
|
c23 = UOp(Ops.DEFINE_GLOBAL, dtypes.imageh((240, 128, 4)), (), 2)
|
||||||
|
c35 = c5.index((c7*4+c10+c13*128+c18*1536))*c23.index((c10*4+c2%4+c7*16+c2//4*512))
|
||||||
|
c37 = UOp(Ops.DEFINE_GLOBAL, dtypes.float.ptr(960), (), 3)
|
||||||
|
c39 = c35.reduce(c7, c10, arg=Ops.ADD)+c37.index(c2)
|
||||||
|
c50 = (1.0+((c39+0.044708251953125*(c39*(c39*c39)))*-2.3021129851685216).exp2()).reciprocal()*c39
|
||||||
|
c53 = c50.reduce(c18, c13, arg=Ops.ADD)*0.010416666666666666
|
||||||
|
c55 = c0.index(c2, ptr=True).store(c53).end(c2)
|
||||||
|
|
||||||
|
opts = None
|
||||||
|
# JITBEAM=2
|
||||||
|
# (Opt(op=OptOps.UPCAST, axis=0, arg=4), Opt(op=OptOps.GROUPTOP, axis=1, arg=32), Opt(op=OptOps.UNROLL, axis=1, arg=4), Opt(op=OptOps.LOCAL, axis=0, arg=8), Opt(op=OptOps.UNROLL, axis=0, arg=4), Opt(op=OptOps.GROUP, axis=1, arg=0))
|
||||||
|
return c55.sink(arg=KernelInfo(name="conv", opts_to_apply=opts))
|
||||||
|
|
||||||
|
ast = {143: vision_conv_143, 153: vision_conv_153, 172: dm_conv_172}[getenv("NUM", 143)]()
|
||||||
|
|
||||||
compiler = Device.default.compiler
|
compiler = Device.default.compiler
|
||||||
renderer = Device.default.renderer
|
renderer = Device.default.renderer
|
||||||
|
|||||||
+55
@@ -0,0 +1,55 @@
|
|||||||
|
import time
|
||||||
|
from tinygrad.tensor import Tensor, Device
|
||||||
|
|
||||||
|
MODEL_WIDTH = 512
|
||||||
|
MODEL_HEIGHT = 256
|
||||||
|
MODEL_FRAME_SIZE = MODEL_WIDTH * MODEL_HEIGHT * 3 // 2
|
||||||
|
IMG_INPUT_SHAPE = (1, 12, 128, 256)
|
||||||
|
|
||||||
|
def tensor_arange(end): return Tensor([float(i) for i in range(end)])
|
||||||
|
def tensor_round(tensor:Tensor): return (tensor + 0.5).floor()
|
||||||
|
|
||||||
|
h_src, w_src = 1208, 1928
|
||||||
|
h_dst, w_dst = MODEL_HEIGHT, MODEL_WIDTH
|
||||||
|
x = tensor_arange(w_dst).reshape(1, w_dst).expand(h_dst, w_dst)
|
||||||
|
y = tensor_arange(h_dst).reshape(h_dst, 1).expand(h_dst, w_dst)
|
||||||
|
ones = Tensor.ones_like(x)
|
||||||
|
dst_coords = x.reshape((1,-1)).cat(y.reshape((1,-1))).cat(ones.reshape((1,-1)))
|
||||||
|
|
||||||
|
def warp_perspective_tinygrad(src:Tensor, M_inv:Tensor) -> Tensor:
|
||||||
|
src_coords = M_inv @ dst_coords
|
||||||
|
src_coords = src_coords / src_coords[2:3, :]
|
||||||
|
|
||||||
|
x_src = src_coords[0].reshape(h_dst, w_dst)
|
||||||
|
y_src = src_coords[1].reshape(h_dst, w_dst)
|
||||||
|
|
||||||
|
x_nearest = tensor_round(x_src).clip(0, w_src - 1).cast('int')
|
||||||
|
y_nearest = tensor_round(y_src).clip(0, h_src - 1).cast('int')
|
||||||
|
|
||||||
|
# TODO: make 2d indexing fast
|
||||||
|
idx = y_nearest*src.shape[1] + x_nearest
|
||||||
|
dst = src.flatten()[idx]
|
||||||
|
return dst.reshape(h_dst, w_dst)
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
from tinygrad.engine.jit import TinyJit
|
||||||
|
update_img_jit = TinyJit(warp_perspective_tinygrad, prune=True)
|
||||||
|
|
||||||
|
step_times = []
|
||||||
|
for _ in range(10):
|
||||||
|
# regenerate inputs
|
||||||
|
inputs = [Tensor.randn(1928,1208), Tensor.randn(3,3)]
|
||||||
|
Tensor.realize(*inputs)
|
||||||
|
Device.default.synchronize()
|
||||||
|
|
||||||
|
# do the warp
|
||||||
|
st = time.perf_counter()
|
||||||
|
out = update_img_jit(*inputs)
|
||||||
|
mt = time.perf_counter()
|
||||||
|
val = out.contiguous().realize()
|
||||||
|
Device.default.synchronize()
|
||||||
|
et = time.perf_counter()
|
||||||
|
|
||||||
|
# measure the time
|
||||||
|
step_times.append((et-st)*1e3)
|
||||||
|
print(f"enqueue {(mt-st)*1e3:6.2f} ms -- total run {step_times[-1]:6.2f} ms")
|
||||||
+3
-2
@@ -1,13 +1,14 @@
|
|||||||
import unittest
|
import unittest
|
||||||
from tinygrad import Device
|
from tinygrad import Device
|
||||||
from tinygrad.tensor import Tensor
|
from tinygrad.tensor import Tensor
|
||||||
from tinygrad.helpers import getenv, CI
|
from tinygrad.helpers import getenv, CI, OSX
|
||||||
|
|
||||||
def multidevice_test(fxn):
|
def multidevice_test(fxn):
|
||||||
exclude_devices = getenv("EXCLUDE_DEVICES", "").split(",")
|
exclude_devices = getenv("EXCLUDE_DEVICES", "").split(",")
|
||||||
def ret(self):
|
def ret(self):
|
||||||
for device in Device._devices:
|
for device in Device._devices:
|
||||||
if device in ["REMOTE", "DISK", "NPY", "FAKE", "DSP", "NULL"]: continue
|
# broken on OSX USB AMD, why?
|
||||||
|
if device in ["REMOTE", "DISK", "NPY", "FAKE", "DSP", "NULL"] or (OSX and device in ["AMD"]): continue
|
||||||
if not CI: print(device)
|
if not CI: print(device)
|
||||||
if device in exclude_devices:
|
if device in exclude_devices:
|
||||||
if not CI: print(f"WARNING: {device} test is excluded")
|
if not CI: print(f"WARNING: {device} test is excluded")
|
||||||
|
|||||||
-1
@@ -112,7 +112,6 @@ backend_test.exclude('test_dequantizelinear_e5m2_cpu')
|
|||||||
backend_test.exclude('test_dequantizelinear_float4e2m1_cpu')
|
backend_test.exclude('test_dequantizelinear_float4e2m1_cpu')
|
||||||
|
|
||||||
# we don't support indexes
|
# we don't support indexes
|
||||||
backend_test.exclude('test_nonzero_*')
|
|
||||||
|
|
||||||
# no support for int pow
|
# no support for int pow
|
||||||
backend_test.exclude('test_pow_types_int32_int32_cpu')
|
backend_test.exclude('test_pow_types_int32_int32_cpu')
|
||||||
|
|||||||
+3
-1
@@ -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
|
||||||
if REF == "master": SKIP_PROCESS_REPLAY = True
|
if REF == "master": SKIP_PROCESS_REPLAY = True
|
||||||
class ProcessReplayWarning(Warning): pass
|
class ProcessReplayWarning(Warning): pass
|
||||||
|
|
||||||
|
|||||||
@@ -33,7 +33,7 @@ remu = _try_dlopen_remu()
|
|||||||
def create_sdma_packets():
|
def create_sdma_packets():
|
||||||
# TODO: clean up this, if we want to keep it
|
# TODO: clean up this, if we want to keep it
|
||||||
structs = {}
|
structs = {}
|
||||||
for name,pkt in [(name,s) for name,s in amd_gpu.__dict__.items() if name.startswith("struct_SDMA_PKT_") and name.endswith("_TAG")]:
|
for name,pkt in [(name,s) for name,s in amd_gpu.__dict__.items() if name.startswith("rocr_AMD_SDMA_PKT_") and name.endswith("_TAG")]:
|
||||||
names = set()
|
names = set()
|
||||||
fields = []
|
fields = []
|
||||||
for pkt_fields in pkt._fields_:
|
for pkt_fields in pkt._fields_:
|
||||||
@@ -47,7 +47,7 @@ def create_sdma_packets():
|
|||||||
# merge together 64-bit fields, otherwise just append them
|
# merge together 64-bit fields, otherwise just append them
|
||||||
if fname.endswith("_63_32") and fields[-1][0].endswith("_31_0"): fields[-1] = tuple([fname[:-6], ctypes.c_ulong, 64])
|
if fname.endswith("_63_32") and fields[-1][0].endswith("_31_0"): fields[-1] = tuple([fname[:-6], ctypes.c_ulong, 64])
|
||||||
else: fields.append(tuple([fname, *union_fields[1:]]))
|
else: fields.append(tuple([fname, *union_fields[1:]]))
|
||||||
new_name = name[16:-4].lower()
|
new_name = name[18:-4].lower()
|
||||||
structs[new_name] = init_c_struct_t(tuple(fields))
|
structs[new_name] = init_c_struct_t(tuple(fields))
|
||||||
assert ctypes.sizeof(structs[new_name]) == ctypes.sizeof(pkt), f"{ctypes.sizeof(structs[new_name])} != {ctypes.sizeof(pkt)}"
|
assert ctypes.sizeof(structs[new_name]) == ctypes.sizeof(pkt), f"{ctypes.sizeof(structs[new_name])} != {ctypes.sizeof(pkt)}"
|
||||||
return type("SDMA_PKTS", (object, ), structs)
|
return type("SDMA_PKTS", (object, ), structs)
|
||||||
@@ -124,6 +124,7 @@ class PM4Executor(AMDQueue):
|
|||||||
elif mem_data_sel == 3:
|
elif mem_data_sel == 3:
|
||||||
if mem_event_type == CACHE_FLUSH_AND_INV_TS_EVENT: ptr.cast('Q')[0] = int(time.perf_counter() * 1e8)
|
if mem_event_type == CACHE_FLUSH_AND_INV_TS_EVENT: ptr.cast('Q')[0] = int(time.perf_counter() * 1e8)
|
||||||
else: raise RuntimeError(f"Unknown {mem_data_sel=} {mem_event_type=}")
|
else: raise RuntimeError(f"Unknown {mem_data_sel=} {mem_event_type=}")
|
||||||
|
elif mem_data_sel == 0: pass # no write
|
||||||
else: raise RuntimeError(f"Unknown {mem_data_sel=}")
|
else: raise RuntimeError(f"Unknown {mem_data_sel=}")
|
||||||
|
|
||||||
def _exec_copy_data(self, n):
|
def _exec_copy_data(self, n):
|
||||||
|
|||||||
@@ -164,7 +164,7 @@ def cuStreamWaitEvent(stream: Any, event, flags: int) -> int: return orig_cuda.C
|
|||||||
def cuCtxSynchronize() -> int: return orig_cuda.CUDA_SUCCESS
|
def cuCtxSynchronize() -> int: return orig_cuda.CUDA_SUCCESS
|
||||||
|
|
||||||
def cuGetErrorString(error: int, pStr) -> int:
|
def cuGetErrorString(error: int, pStr) -> int:
|
||||||
error_str = orig_cuda.cudaError_enum__enumvalues.get(error, "Unknown CUDA error").encode()
|
error_str = orig_cuda.enum_cudaError_enum.get(error, "Unknown CUDA error").encode()
|
||||||
buf = ctypes.create_string_buffer(error_str)
|
buf = ctypes.create_string_buffer(error_str)
|
||||||
# Set the pointer to point to our error string buffer
|
# Set the pointer to point to our error string buffer
|
||||||
pStr._obj.value = ctypes.cast(buf, ctypes.POINTER(ctypes.c_char))
|
pStr._obj.value = ctypes.cast(buf, ctypes.POINTER(ctypes.c_char))
|
||||||
|
|||||||
@@ -1,5 +1,5 @@
|
|||||||
import ctypes, mmap, collections, functools, os
|
import ctypes, mmap, collections, functools, os
|
||||||
import tinygrad.runtime.autogen.nv_gpu as nv_gpu
|
from tinygrad.runtime.autogen import nv_570 as nv_gpu
|
||||||
from typing import Any
|
from typing import Any
|
||||||
from tinygrad.helpers import to_mv
|
from tinygrad.helpers import to_mv
|
||||||
from test.mockgpu.driver import VirtDriver, VirtFileDesc, VirtFile
|
from test.mockgpu.driver import VirtDriver, VirtFileDesc, VirtFile
|
||||||
@@ -153,8 +153,10 @@ class NVDriver(VirtDriver):
|
|||||||
51059, 51069, 51071, 51632, 51639, 51639, 51706, 52019, 222, 50287, 50273, 50031, 50017] # from ada102
|
51059, 51069, 51071, 51632, 51639, 51639, 51706, 52019, 222, 50287, 50273, 50031, 50017] # from ada102
|
||||||
params.numClasses = len(classes)
|
params.numClasses = len(classes)
|
||||||
if struct.cmd == nv_gpu.NV0080_CTRL_CMD_GPU_GET_CLASSLIST:
|
if struct.cmd == nv_gpu.NV0080_CTRL_CMD_GPU_GET_CLASSLIST:
|
||||||
clslist = to_mv(params.classList, params.numClasses * 4).cast('I')
|
if params.classList and params.numClasses > 0:
|
||||||
for i,c in enumerate(classes): clslist[i] = c
|
clslist = to_mv(params.classList, params.numClasses * 4).cast('I')
|
||||||
|
for i,c in enumerate(classes): clslist[i] = c
|
||||||
|
else: params.numClasses = len(classes)
|
||||||
else:
|
else:
|
||||||
for i,c in enumerate(classes): params.classList[i] = c
|
for i,c in enumerate(classes): params.classList[i] = c
|
||||||
elif struct.cmd == nv_gpu.NV2080_CTRL_CMD_GR_GET_INFO:
|
elif struct.cmd == nv_gpu.NV2080_CTRL_CMD_GR_GET_INFO:
|
||||||
@@ -192,6 +194,9 @@ class NVDriver(VirtDriver):
|
|||||||
params.mmuFaultInfoList[0].faultAddress = int(os.environ['MOCKGPU_EMU_FAULTADDR'], base=16)
|
params.mmuFaultInfoList[0].faultAddress = int(os.environ['MOCKGPU_EMU_FAULTADDR'], base=16)
|
||||||
params.mmuFaultInfoList[0].faultType = 1
|
params.mmuFaultInfoList[0].faultType = 1
|
||||||
params.mmuFaultInfoList[0].accessType = 1
|
params.mmuFaultInfoList[0].accessType = 1
|
||||||
|
elif struct.cmd == nv_gpu.NV0000_CTRL_CMD_SYSTEM_GET_BUILD_VERSION_V2:
|
||||||
|
params = nv_gpu.NV0000_CTRL_SYSTEM_GET_BUILD_VERSION_V2_PARAMS.from_address(params_ptr)
|
||||||
|
params.driverVersionBuffer = b"570.00.00\0"
|
||||||
else: raise RuntimeError(f"Unknown {struct.cmd} to rm_control")
|
else: raise RuntimeError(f"Unknown {struct.cmd} to rm_control")
|
||||||
return 0
|
return 0
|
||||||
|
|
||||||
@@ -254,4 +259,4 @@ class NVDriver(VirtDriver):
|
|||||||
for gpu in self.gpus.values():
|
for gpu in self.gpus.values():
|
||||||
for q in gpu.queues:
|
for q in gpu.queues:
|
||||||
if q.ctrl.GPGet != q.ctrl.GPPut:
|
if q.ctrl.GPGet != q.ctrl.GPPut:
|
||||||
any_progress |= q.execute()
|
any_progress |= q.execute()
|
||||||
|
|||||||
@@ -1,5 +1,5 @@
|
|||||||
import ctypes, time
|
import ctypes, time
|
||||||
import tinygrad.runtime.autogen.nv_gpu as nv_gpu
|
from tinygrad.runtime.autogen import nv_570 as nv_gpu
|
||||||
from enum import Enum, auto
|
from enum import Enum, auto
|
||||||
from test.mockgpu.gpu import VirtGPU
|
from test.mockgpu.gpu import VirtGPU
|
||||||
from test.mockgpu.helpers import _try_dlopen_gpuocelot
|
from test.mockgpu.helpers import _try_dlopen_gpuocelot
|
||||||
|
|||||||
+22
-8
@@ -5,26 +5,40 @@ from tinygrad.helpers import CI, Context, getenv
|
|||||||
from tinygrad.engine.realize import run_schedule
|
from tinygrad.engine.realize import run_schedule
|
||||||
from tinygrad.engine.realize import CompiledRunner, ExecItem, get_program
|
from tinygrad.engine.realize import CompiledRunner, ExecItem, get_program
|
||||||
from tinygrad.uop.ops import Ops
|
from tinygrad.uop.ops import Ops
|
||||||
|
from tinygrad.renderer import Estimates
|
||||||
|
from tinygrad.renderer.ptx import PTXRenderer
|
||||||
|
|
||||||
class TestArange(unittest.TestCase):
|
class TestArange(unittest.TestCase):
|
||||||
def _get_flops(self, N):
|
def _get_flops(self, tensor, desired):
|
||||||
GlobalCounters.reset()
|
GlobalCounters.reset()
|
||||||
tt = Tensor.arange(N)
|
sched = tensor.schedule()
|
||||||
sched = tt.schedule()
|
|
||||||
self.assertEqual(len(sched), 1)
|
self.assertEqual(len(sched), 1)
|
||||||
p = get_program(sched[-1].ast)
|
p = get_program(sched[-1].ast)
|
||||||
ExecItem(CompiledRunner(p), [tt.uop.buffer]).run()
|
ExecItem(CompiledRunner(p), [tensor.uop.buffer]).run()
|
||||||
np.testing.assert_equal(tt.numpy(), np.arange(N))
|
np.testing.assert_equal(tensor.numpy(), desired)
|
||||||
return p.estimates.ops
|
return p.estimates.ops
|
||||||
|
|
||||||
def test_complexity(self):
|
def test_arange_complexity(self):
|
||||||
self.assertEqual(self._get_flops(256), 0)
|
self.assertEqual(self._get_flops(Tensor.arange(256), np.arange(256)), 0)
|
||||||
self.assertEqual(self._get_flops(2560), 0)
|
self.assertEqual(self._get_flops(Tensor.arange(2560), np.arange(2560)), 0)
|
||||||
|
|
||||||
def test_arange_cat(self):
|
def test_arange_cat(self):
|
||||||
t = Tensor.arange(2, dtype=dtypes.int)+Tensor([3])
|
t = Tensor.arange(2, dtype=dtypes.int)+Tensor([3])
|
||||||
self.assertEqual(t.cat(t).tolist(), [3, 4, 3, 4])
|
self.assertEqual(t.cat(t).tolist(), [3, 4, 3, 4])
|
||||||
|
|
||||||
|
def test_eye_complexity(self):
|
||||||
|
with Context(NOOPT=1):
|
||||||
|
# NOTE: not every backend supports CMPEQ
|
||||||
|
self.assertLessEqual(self._get_flops(Tensor.eye(2560).contiguous(), np.eye(2560)), 2*2560*2560)
|
||||||
|
|
||||||
|
@unittest.skipIf(isinstance(Device[Device.DEFAULT].renderer, PTXRenderer), "PTX indexing is weird")
|
||||||
|
def test_tri_complexity(self):
|
||||||
|
with Context(NOOPT=1):
|
||||||
|
t = Tensor.ones(256, 256).contiguous().realize()
|
||||||
|
sched = t.triu().schedule()
|
||||||
|
p = get_program(sched[-1].ast)
|
||||||
|
self.assertLessEqual(Estimates.from_uops(p.uops).ops, 4 * 256 * 256)
|
||||||
|
|
||||||
DSET, DDIM = 2048, 32
|
DSET, DDIM = 2048, 32
|
||||||
|
|
||||||
class TestIndexing(unittest.TestCase):
|
class TestIndexing(unittest.TestCase):
|
||||||
|
|||||||
@@ -102,6 +102,11 @@ def backward_gemm_custom(gradient:UOp, kernel:UOp) -> tuple[UOp, UOp]:
|
|||||||
# **** tests ****
|
# **** tests ****
|
||||||
|
|
||||||
class TestCustomKernel(unittest.TestCase):
|
class TestCustomKernel(unittest.TestCase):
|
||||||
|
def test_empty(self):
|
||||||
|
a = Tensor.empty(1)
|
||||||
|
a = Tensor.custom_kernel(a, fxn=lambda _: UOp.sink())[0]
|
||||||
|
a.realize()
|
||||||
|
|
||||||
def test_simple(self):
|
def test_simple(self):
|
||||||
a = Tensor.ones(16, 16).contiguous()
|
a = Tensor.ones(16, 16).contiguous()
|
||||||
b = Tensor.ones(16, 16).contiguous()
|
b = Tensor.ones(16, 16).contiguous()
|
||||||
|
|||||||
+20
-1
@@ -14,6 +14,8 @@ from tinygrad.renderer.ptx import PTXRenderer
|
|||||||
from tinygrad.renderer.cstyle import CUDARenderer
|
from tinygrad.renderer.cstyle import CUDARenderer
|
||||||
MOCKGPU = getenv("MOCKGPU")
|
MOCKGPU = getenv("MOCKGPU")
|
||||||
|
|
||||||
|
from tinygrad.uop.ops import print_uops # noqa: F401 # pylint: disable=unused-import
|
||||||
|
|
||||||
class TestLinearizer(unittest.TestCase):
|
class TestLinearizer(unittest.TestCase):
|
||||||
def test_arg_dedup(self):
|
def test_arg_dedup(self):
|
||||||
# NOTE: this realize exists because Tensor.numpy calls .contiguous() internally
|
# NOTE: this realize exists because Tensor.numpy calls .contiguous() internally
|
||||||
@@ -38,6 +40,22 @@ class TestLinearizer(unittest.TestCase):
|
|||||||
np.testing.assert_equal(a.numpy(), ta)
|
np.testing.assert_equal(a.numpy(), ta)
|
||||||
np.testing.assert_equal(b.numpy(), tb)
|
np.testing.assert_equal(b.numpy(), tb)
|
||||||
|
|
||||||
|
@unittest.skip("TODO: some backends insert more casts")
|
||||||
|
def test_cast_there_and_back(self):
|
||||||
|
tst = Tensor.ones(16, dtype=dtypes.int).contiguous().realize()
|
||||||
|
out = tst.neg().cast(dtypes.char).cast(dtypes.int).cast(dtypes.char) * 2
|
||||||
|
ast = helper_linearizer_opt(out)
|
||||||
|
uops = get_program(ast, opts=[]).uops
|
||||||
|
self.assertEqual(len([x for x in uops if x.op is Ops.CAST]), 1)
|
||||||
|
|
||||||
|
@unittest.expectedFailure
|
||||||
|
def test_cast_back_and_there(self):
|
||||||
|
tst = Tensor.ones(16, dtype=dtypes.int).contiguous().realize()
|
||||||
|
out = tst.neg().cast(dtypes.char).cast(dtypes.int) * 2
|
||||||
|
ast = helper_linearizer_opt(out)
|
||||||
|
uops = get_program(ast, opts=[]).uops
|
||||||
|
self.assertEqual(len([x for x in uops if x.op is Ops.CAST]), 0)
|
||||||
|
|
||||||
@unittest.skipIf(isinstance(Device[Device.DEFAULT].renderer, PTXRenderer), "broken on ptx")
|
@unittest.skipIf(isinstance(Device[Device.DEFAULT].renderer, PTXRenderer), "broken on ptx")
|
||||||
def test_late_bias_load(self):
|
def test_late_bias_load(self):
|
||||||
img = Tensor.empty(1, 3, 16, 16)
|
img = Tensor.empty(1, 3, 16, 16)
|
||||||
@@ -78,6 +96,7 @@ class TestLinearizer(unittest.TestCase):
|
|||||||
ranges = [i for i,u in enumerate(uops) if u.op is Ops.RANGE]
|
ranges = [i for i,u in enumerate(uops) if u.op is Ops.RANGE]
|
||||||
assert len(ranges) == 1 # NOTE: it collapses now
|
assert len(ranges) == 1 # NOTE: it collapses now
|
||||||
|
|
||||||
|
@unittest.expectedFailure # TODO: investigate
|
||||||
def test_two_nested_range_alt_indexing(self):
|
def test_two_nested_range_alt_indexing(self):
|
||||||
a = Tensor([2, 2]).realize()
|
a = Tensor([2, 2]).realize()
|
||||||
out = a.reshape(2, 1).pad(((1, 1), (1, 1)), value=2).sum()
|
out = a.reshape(2, 1).pad(((1, 1), (1, 1)), value=2).sum()
|
||||||
@@ -490,7 +509,7 @@ def copyout_outputs(outbufs:list[Buffer]) -> list[np.ndarray]:
|
|||||||
return [np.frombuffer(x.as_buffer(), _to_np_dtype(x.dtype)) for x in outbufs]
|
return [np.frombuffer(x.as_buffer(), _to_np_dtype(x.dtype)) for x in outbufs]
|
||||||
|
|
||||||
def reset_bufs(bufs:list[Buffer]):
|
def reset_bufs(bufs:list[Buffer]):
|
||||||
for buf in bufs: buf.copyin(np.zeros((buf.size, ), dtype=_to_np_dtype(buf.dtype)).data) # Zero to check that all values are filled
|
for buf in bufs: buf.copyin(np.zeros((buf.size*buf.dtype.itemsize,), dtype=np.uint8).data)
|
||||||
|
|
||||||
def _helper_linearizer_opt_ast(realized_ast:UOp, real_bufs:list[Buffer], opts=[],
|
def _helper_linearizer_opt_ast(realized_ast:UOp, real_bufs:list[Buffer], opts=[],
|
||||||
apply_tc=False, atol=1e-4, rtol=1e-4, color_sizes=[], wanna_output=[]):
|
apply_tc=False, atol=1e-4, rtol=1e-4, color_sizes=[], wanna_output=[]):
|
||||||
|
|||||||
+157
-1
@@ -1,5 +1,6 @@
|
|||||||
import unittest
|
import unittest
|
||||||
from tinygrad import Tensor, UOp
|
import numpy as np
|
||||||
|
from tinygrad import Tensor, UOp, nn
|
||||||
from tinygrad.uop.ops import AxisType, Ops
|
from tinygrad.uop.ops import AxisType, Ops
|
||||||
|
|
||||||
class TestOuterworldReduce(unittest.TestCase):
|
class TestOuterworldReduce(unittest.TestCase):
|
||||||
@@ -11,6 +12,81 @@ class TestOuterworldReduce(unittest.TestCase):
|
|||||||
t = Tensor(UOp(Ops.REDUCE, dtype=out.uop.dtype, src=(out.uop, a), arg=Ops.ADD))
|
t = Tensor(UOp(Ops.REDUCE, dtype=out.uop.dtype, src=(out.uop, a), arg=Ops.ADD))
|
||||||
self.assertListEqual(t.tolist(), [5.,5.,5.,5.,5.])
|
self.assertListEqual(t.tolist(), [5.,5.,5.,5.,5.])
|
||||||
|
|
||||||
|
# TODO: delete test_outerworld_range?
|
||||||
|
class TestOuterRange(unittest.TestCase):
|
||||||
|
def test_simple_range(self):
|
||||||
|
a = Tensor.ones(10).contiguous()
|
||||||
|
acc = Tensor.zeros().contiguous()
|
||||||
|
Tensor.realize(a, acc)
|
||||||
|
|
||||||
|
# this is fold
|
||||||
|
i = UOp.range(10, -100, AxisType.OUTER)
|
||||||
|
acc_i = acc.uop.after(i)
|
||||||
|
vi = UOp.variable("i", i.vmin, i.vmax).bind(i)
|
||||||
|
out = Tensor(acc.uop.after(acc_i.store(acc_i + a[vi].uop).end(i)))
|
||||||
|
out.realize()
|
||||||
|
assert out.item() == 10.0
|
||||||
|
|
||||||
|
def test_inner_range(self):
|
||||||
|
a = Tensor.ones(10, 10).contiguous()
|
||||||
|
acc = Tensor.zeros(10).contiguous()
|
||||||
|
Tensor.realize(a, acc)
|
||||||
|
|
||||||
|
# this is fold
|
||||||
|
i = UOp.range(10, -100, AxisType.OUTER)
|
||||||
|
acc_i = acc.uop.after(i)
|
||||||
|
vi = UOp.variable("i", i.vmin, i.vmax).bind(i)
|
||||||
|
out = Tensor(acc.uop.after(acc_i.store(acc_i + a[:, vi].uop).end(i)))
|
||||||
|
out.realize()
|
||||||
|
assert all(x == 10.0 for x in out.tolist())
|
||||||
|
|
||||||
|
def test_range_matmul(self):
|
||||||
|
vec = Tensor.randn(1, 10).realize()
|
||||||
|
mats = Tensor.randn(3, 10, 10).realize()
|
||||||
|
|
||||||
|
# 3 matmuls in "scan"
|
||||||
|
ref = ((vec @ mats[0]) @ mats[1]) @ mats[2]
|
||||||
|
ref.realize()
|
||||||
|
|
||||||
|
# 3 matmuls with outer world range
|
||||||
|
i = UOp.range(3, -100, AxisType.OUTER)
|
||||||
|
vec_i = Tensor(vec.uop.after(i))
|
||||||
|
comp = vec_i.contiguous() @ mats[i]
|
||||||
|
store = vec_i.uop.store(comp.uop).end(i)
|
||||||
|
out = Tensor(vec.uop.after(store))
|
||||||
|
out.realize()
|
||||||
|
|
||||||
|
# TODO: testing allclose
|
||||||
|
assert Tensor.allclose(ref, out, atol=1e-6), f"{ref.numpy()=}, {out.numpy()=}"
|
||||||
|
|
||||||
|
class TestOuterScan(unittest.TestCase):
|
||||||
|
def _test_scan(self):
|
||||||
|
vec = Tensor.randn(1, 10).realize()
|
||||||
|
mats = Tensor.randn(3, 10, 10).realize()
|
||||||
|
|
||||||
|
# 3 matmuls in "scan"
|
||||||
|
vec1 = vec @ mats[0]
|
||||||
|
vec2 = vec1 @ mats[1]
|
||||||
|
vec3 = vec2 @ mats[2]
|
||||||
|
ref = Tensor.stack(vec1, vec2, vec3)
|
||||||
|
ref.realize()
|
||||||
|
return vec, mats, ref
|
||||||
|
|
||||||
|
def test_uop_scan_matmul(self):
|
||||||
|
vec, mats, ref = self._test_scan()
|
||||||
|
|
||||||
|
# 3 matmuls with SCAN
|
||||||
|
i = UOp.range(3, -100, AxisType.OUTER)
|
||||||
|
out = Tensor.empty(3, 1, 10)
|
||||||
|
phi = Tensor(i.eq(0).where(vec.uop, out[(i-1).maximum(0)].uop))
|
||||||
|
comp = phi @ mats[i]
|
||||||
|
store = out[i].uop.store(comp.uop).end(i)
|
||||||
|
out = Tensor(out.uop.after(store))
|
||||||
|
out.realize()
|
||||||
|
|
||||||
|
# TODO: testing allclose
|
||||||
|
assert Tensor.allclose(ref, out, atol=1e-6), f"{ref.numpy()=}, {out.numpy()=}"
|
||||||
|
|
||||||
class TestOuterworld(unittest.TestCase):
|
class TestOuterworld(unittest.TestCase):
|
||||||
def test_range_plus_1(self):
|
def test_range_plus_1(self):
|
||||||
t = Tensor.arange(100).reshape(10,10).realize()
|
t = Tensor.arange(100).reshape(10,10).realize()
|
||||||
@@ -70,5 +146,85 @@ class TestOuterworld(unittest.TestCase):
|
|||||||
out = out.reshape(1, 3).expand(a, 3).contiguous().realize()
|
out = out.reshape(1, 3).expand(a, 3).contiguous().realize()
|
||||||
self.assertListEqual([[0,4,8],[4,8,12],[8,12,16]], out.tolist())
|
self.assertListEqual([[0,4,8],[4,8,12],[8,12,16]], out.tolist())
|
||||||
|
|
||||||
|
class TestVmap(unittest.TestCase):
|
||||||
|
def test_vmap_inner(self, axis_type=AxisType.LOOP, fuse=False, grad=False):
|
||||||
|
x = Tensor.ones(1, 10).contiguous().requires_grad_()
|
||||||
|
mats = Tensor.ones(3, 10, 10).contiguous().requires_grad_()
|
||||||
|
|
||||||
|
ref = x @ mats
|
||||||
|
if fuse: ref = ref * 2
|
||||||
|
|
||||||
|
# vmap across axis 0
|
||||||
|
a = UOp.range(3, -1, axis_type)
|
||||||
|
out = x @ mats[a]
|
||||||
|
out = out.reshape(1, 10).pad(((a,(3-a)-1), None))
|
||||||
|
out = Tensor(out.uop.reduce(a, arg=Ops.ADD))
|
||||||
|
if fuse: out = out * 2
|
||||||
|
if grad:
|
||||||
|
out.mean().backward()
|
||||||
|
np.testing.assert_allclose(mats.grad.numpy(), (2./30) if fuse else (1./30))
|
||||||
|
out.realize()
|
||||||
|
|
||||||
|
# TODO: testing allclose
|
||||||
|
assert Tensor.allclose(ref, out, atol=1e-6), f"{ref.numpy()=}, {out.numpy()=}"
|
||||||
|
def test_vmap_inner_fuse(self): self.test_vmap_inner(fuse=True)
|
||||||
|
def test_vmap_outer(self): self.test_vmap_inner(AxisType.OUTER)
|
||||||
|
def test_vmap_outer_fuse(self): self.test_vmap_inner(AxisType.OUTER, fuse=True)
|
||||||
|
|
||||||
|
def test_vmap_inner_grad(self): self.test_vmap_inner(grad=True)
|
||||||
|
def test_vmap_inner_fuse_grad(self): self.test_vmap_inner(fuse=True, grad=True)
|
||||||
|
def test_vmap_outer_grad(self): self.test_vmap_inner(AxisType.OUTER, grad=True)
|
||||||
|
|
||||||
|
def test_vmap_convs(self):
|
||||||
|
layers = [
|
||||||
|
nn.Conv2d(1, 8, 3), Tensor.relu,
|
||||||
|
nn.Conv2d(8, 8, 3), Tensor.relu]
|
||||||
|
img = Tensor.randn(4, 1, 16, 16).realize(*nn.state.get_parameters(layers))
|
||||||
|
a = UOp.range(4, -1, AxisType.OUTER)
|
||||||
|
out = img[a:a+1].sequential(layers)
|
||||||
|
out = out.pad(((a,(4-a)-1), None, None, None))
|
||||||
|
out = Tensor(out.uop.reduce(a, arg=Ops.ADD))
|
||||||
|
out.realize()
|
||||||
|
np.testing.assert_allclose(out.numpy(), img.sequential(layers).numpy(), atol=1e-6)
|
||||||
|
|
||||||
|
def test_vmap_gemm(self):
|
||||||
|
layers = [
|
||||||
|
nn.Linear(16, 16, bias=False), Tensor.relu,
|
||||||
|
nn.Linear(16, 16, bias=False), Tensor.relu]
|
||||||
|
img = Tensor.randn(4, 16).realize(*nn.state.get_parameters(layers))
|
||||||
|
a = UOp.range(4, -1, AxisType.OUTER)
|
||||||
|
out = img[a:a+1].sequential(layers)
|
||||||
|
out = out.pad(((a,(4-a)-1), None))
|
||||||
|
out = Tensor(out.uop.reduce(a, arg=Ops.ADD))
|
||||||
|
out.realize()
|
||||||
|
np.testing.assert_allclose(out.numpy(), img.sequential(layers).numpy(), atol=1e-6)
|
||||||
|
|
||||||
|
@unittest.skip("this is broken, we need to lower the outer reduce in the outer graph")
|
||||||
|
def test_vmap_gemm_grad(self):
|
||||||
|
layers = [
|
||||||
|
nn.Linear(16, 16, bias=False), Tensor.relu,
|
||||||
|
nn.Linear(16, 16, bias=False), Tensor.relu]
|
||||||
|
layer_tensors = nn.state.get_parameters(layers)
|
||||||
|
img = Tensor.randn(4, 16).realize(*layer_tensors)
|
||||||
|
for l in layer_tensors: l.requires_grad_()
|
||||||
|
a = UOp.range(4, -1, AxisType.OUTER)
|
||||||
|
out = img[a:a+1].sequential(layers)
|
||||||
|
out = out.pad(((a,(4-a)-1), None))
|
||||||
|
out = Tensor(out.uop.reduce(a, arg=Ops.ADD))
|
||||||
|
out.mean().backward()
|
||||||
|
grads = [l.grad for l in layer_tensors]
|
||||||
|
out.realize(*grads)
|
||||||
|
out_grads = [x.numpy() for x in grads]
|
||||||
|
|
||||||
|
# compute reference grads
|
||||||
|
for l in layer_tensors: l.grad = None
|
||||||
|
img.sequential(layers).mean().backward()
|
||||||
|
grads = [l.grad for l in layer_tensors]
|
||||||
|
out.realize(*grads)
|
||||||
|
ref_grads = [x.numpy() for x in grads]
|
||||||
|
|
||||||
|
# compare
|
||||||
|
for o,r in zip(out_grads, ref_grads): np.testing.assert_allclose(o, r, atol=1e-6)
|
||||||
|
|
||||||
if __name__ == '__main__':
|
if __name__ == '__main__':
|
||||||
unittest.main()
|
unittest.main()
|
||||||
@@ -1,18 +0,0 @@
|
|||||||
from tinygrad import Tensor, UOp
|
|
||||||
from tinygrad.uop.ops import Ops, AxisType
|
|
||||||
import unittest
|
|
||||||
# this test is only focused on transformers and using range for the layers
|
|
||||||
|
|
||||||
class TestOuterworldTransformer(unittest.TestCase):
|
|
||||||
def test_three_mats(self):
|
|
||||||
w = Tensor.empty(3, 1024, 1024)
|
|
||||||
inp = Tensor.empty(1, 1024)
|
|
||||||
i = UOp.range(3, -1, AxisType.OUTER)
|
|
||||||
inp_after = Tensor(inp.uop.after(i))
|
|
||||||
inp_gemm = inp_after@w[i]
|
|
||||||
inp = inp.uop.after(inp.uop.store(inp_gemm.uop).end(i)).contiguous()
|
|
||||||
inp = Tensor(inp)
|
|
||||||
inp.realize()
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
|
||||||
unittest.main()
|
|
||||||
@@ -17,7 +17,7 @@ def helper_collect_profile(*devs):
|
|||||||
cpu_events.clear()
|
cpu_events.clear()
|
||||||
|
|
||||||
profile_list = []
|
profile_list = []
|
||||||
with Context(VIZ=1):
|
with Context(VIZ=1, PROFILE=1):
|
||||||
yield profile_list
|
yield profile_list
|
||||||
for dev in devs: dev.synchronize()
|
for dev in devs: dev.synchronize()
|
||||||
for dev in devs: dev._at_profile_finalize()
|
for dev in devs: dev._at_profile_finalize()
|
||||||
|
|||||||
+11
-1
@@ -3,7 +3,7 @@ import torch
|
|||||||
import unittest, copy, mmap, random, math, array
|
import unittest, copy, mmap, random, math, array
|
||||||
from tinygrad import Tensor, Device, dtypes
|
from tinygrad import Tensor, Device, dtypes
|
||||||
from tinygrad.tensor import _METADATA
|
from tinygrad.tensor import _METADATA
|
||||||
from tinygrad.helpers import getenv, temp, mv_address
|
from tinygrad.helpers import Context, getenv, temp, mv_address
|
||||||
from extra.gradcheck import numerical_jacobian, jacobian, gradcheck
|
from extra.gradcheck import numerical_jacobian, jacobian, gradcheck
|
||||||
from hypothesis import given, settings, strategies as strat
|
from hypothesis import given, settings, strategies as strat
|
||||||
from tinygrad.device import is_dtype_supported
|
from tinygrad.device import is_dtype_supported
|
||||||
@@ -846,6 +846,16 @@ class TestTensorMetadata(unittest.TestCase):
|
|||||||
#self.assertEqual(len(bw), 1)
|
#self.assertEqual(len(bw), 1)
|
||||||
#self.assertEqual(bw[0].name, "sigmoid")
|
#self.assertEqual(bw[0].name, "sigmoid")
|
||||||
|
|
||||||
|
def test_tracemeta_0(self):
|
||||||
|
with Context(TRACEMETA=0):
|
||||||
|
x = Tensor.rand(3, requires_grad=True)
|
||||||
|
y = Tensor.rand(3, requires_grad=True)
|
||||||
|
out = (x.relu() * y.sigmoid()).sum()
|
||||||
|
self.assertIsNone(out.uop.metadata)
|
||||||
|
self.assertIsNone(out.uop.src[0].metadata)
|
||||||
|
si = out.schedule()[-1]
|
||||||
|
self.assertEqual(si.metadata, ())
|
||||||
|
|
||||||
class TestIdxUpcast(unittest.TestCase):
|
class TestIdxUpcast(unittest.TestCase):
|
||||||
def _find_op(self, ast: UOp, op: Ops):
|
def _find_op(self, ast: UOp, op: Ops):
|
||||||
if ast.op is op: return ast
|
if ast.op is op: return ast
|
||||||
|
|||||||
+4
-6
@@ -32,8 +32,8 @@ class TestTiny(unittest.TestCase):
|
|||||||
self.assertListEqual(out.tolist(), [2]*16)
|
self.assertListEqual(out.tolist(), [2]*16)
|
||||||
|
|
||||||
def test_cat(self):
|
def test_cat(self):
|
||||||
out = Tensor.cat(Tensor.ones(8).contiguous(), Tensor.ones(8).contiguous())
|
out = Tensor.cat(Tensor.ones(8).contiguous(), Tensor.zeros(8).contiguous())
|
||||||
self.assertListEqual(out.tolist(), [1]*16)
|
self.assertListEqual(out.tolist(), [1]*8+[0]*8)
|
||||||
|
|
||||||
def test_sum(self):
|
def test_sum(self):
|
||||||
out = Tensor.ones(256).contiguous().sum()
|
out = Tensor.ones(256).contiguous().sum()
|
||||||
@@ -62,7 +62,7 @@ class TestTiny(unittest.TestCase):
|
|||||||
out = Tensor.rand(10)
|
out = Tensor.rand(10)
|
||||||
for x in out.tolist():
|
for x in out.tolist():
|
||||||
self.assertGreaterEqual(x, 0.0)
|
self.assertGreaterEqual(x, 0.0)
|
||||||
self.assertLessEqual(x, 1.0)
|
self.assertLess(x, 1.0)
|
||||||
|
|
||||||
# *** JIT (for Python speed) ***
|
# *** JIT (for Python speed) ***
|
||||||
|
|
||||||
@@ -138,9 +138,7 @@ class TestTiny(unittest.TestCase):
|
|||||||
nn.Conv2d(8, 8, 5), Tensor.relu]
|
nn.Conv2d(8, 8, 5), Tensor.relu]
|
||||||
|
|
||||||
# replace random weights with ones
|
# replace random weights with ones
|
||||||
# TODO: there's a bug here where it's tying two of the biases together. we need UNIQUE const
|
Tensor.realize(*[p.replace(Tensor.ones_like(p).contiguous()) for p in nn.state.get_parameters(layers)])
|
||||||
#Tensor.realize(*[p.replace(Tensor.ones_like(p).contiguous()) for p in nn.state.get_parameters(layers)])
|
|
||||||
for p in nn.state.get_parameters(layers): p.replace(Tensor.empty(p.shape))
|
|
||||||
|
|
||||||
# realize gradients
|
# realize gradients
|
||||||
for x in nn.state.get_parameters(layers): x.requires_grad_()
|
for x in nn.state.get_parameters(layers): x.requires_grad_()
|
||||||
|
|||||||
+1
-1
@@ -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])
|
||||||
|
|||||||
@@ -1,31 +1,35 @@
|
|||||||
import unittest
|
import unittest, math
|
||||||
|
|
||||||
from tinygrad import Tensor, Device, dtypes, Context
|
from tinygrad import Tensor, Device, dtypes, Context
|
||||||
from tinygrad.engine.realize import ExecItem, get_runner
|
from tinygrad.engine.realize import ExecItem, get_runner
|
||||||
|
from tinygrad.helpers import CI
|
||||||
|
from tinygrad.renderer.ptx import PTXRenderer
|
||||||
|
import numpy as np
|
||||||
|
|
||||||
from extra.thunder.tiny.tk import WARP_THREADS
|
from extra.thunder.tiny.tk import WARP_THREADS
|
||||||
from extra.thunder.tiny.tk.kernel import Kernel
|
from extra.thunder.tiny.tk.kernel import Kernel
|
||||||
from extra.thunder.tiny.tk.tiles import gl, st, rt, rv
|
|
||||||
|
|
||||||
|
@unittest.skipIf(CI and Device.DEFAULT not in ["CUDA", "NV"], "only cuda")
|
||||||
|
@unittest.skipIf(isinstance(Device[Device.DEFAULT].renderer, PTXRenderer), "no ptx")
|
||||||
class TestTK(unittest.TestCase):
|
class TestTK(unittest.TestCase):
|
||||||
@unittest.skip("store from float rt is wrong")
|
@unittest.skipIf(CI, "no wmma in ci")
|
||||||
def test_simple_matmul(self):
|
def test_simple_matmul(self):
|
||||||
N = 32
|
N = 32
|
||||||
BLOCK_SIZE = 16
|
BLOCK_SIZE = 16
|
||||||
with Kernel((N // BLOCK_SIZE, N // BLOCK_SIZE, 1), WARP_THREADS) as ker:
|
with Kernel((N // BLOCK_SIZE, N // BLOCK_SIZE, 1), WARP_THREADS) as ker:
|
||||||
warp = ker.warp
|
warp = ker.warp
|
||||||
|
|
||||||
c = gl((1, 1, N, N), dtypes.float32)
|
c = ker.gl((1, 1, N, N), dtypes.float32)
|
||||||
a = gl((1, 1, N, N), dtypes.bfloat16)
|
a = ker.gl((1, 1, N, N), dtypes.bfloat16)
|
||||||
b = gl((1, 1, N, N), dtypes.bfloat16)
|
b = ker.gl((1, 1, N, N), dtypes.bfloat16)
|
||||||
|
|
||||||
a_smem = st((BLOCK_SIZE, BLOCK_SIZE), dtypes.bfloat16)
|
a_smem = ker.st((BLOCK_SIZE, BLOCK_SIZE), dtypes.bfloat16)
|
||||||
b_smem = st((BLOCK_SIZE, BLOCK_SIZE), dtypes.bfloat16)
|
b_smem = ker.st((BLOCK_SIZE, BLOCK_SIZE), dtypes.bfloat16)
|
||||||
c_smem = st((BLOCK_SIZE, BLOCK_SIZE), dtypes.float32)
|
c_smem = ker.st((BLOCK_SIZE, BLOCK_SIZE), dtypes.float32)
|
||||||
|
|
||||||
a_reg = rt((BLOCK_SIZE, BLOCK_SIZE), dtypes.bfloat16)
|
a_reg = ker.rt((BLOCK_SIZE, BLOCK_SIZE), dtypes.bfloat16)
|
||||||
b_reg = rt((BLOCK_SIZE, BLOCK_SIZE), dtypes.bfloat16)
|
b_reg = ker.rt((BLOCK_SIZE, BLOCK_SIZE), dtypes.bfloat16)
|
||||||
c_reg = rt((BLOCK_SIZE, BLOCK_SIZE), dtypes.float32)
|
c_reg = ker.rt((BLOCK_SIZE, BLOCK_SIZE), dtypes.float32)
|
||||||
|
|
||||||
col, row = ker.blockIdx_x, ker.blockIdx_y
|
col, row = ker.blockIdx_x, ker.blockIdx_y
|
||||||
|
|
||||||
@@ -57,26 +61,26 @@ class TestTK(unittest.TestCase):
|
|||||||
|
|
||||||
ref = a.matmul(b, dtype=dtypes.float32).float()
|
ref = a.matmul(b, dtype=dtypes.float32).float()
|
||||||
|
|
||||||
assert ref.allclose(c)
|
np.testing.assert_allclose(c.numpy(), ref.numpy())
|
||||||
|
|
||||||
@unittest.skip("store from float rt is wrong")
|
@unittest.skipIf(CI, "no wmma in ci")
|
||||||
def test_simple_matmul_transposed(self):
|
def test_simple_matmul_transposed(self):
|
||||||
N = 32
|
N = 32
|
||||||
BLOCK_SIZE = 16
|
BLOCK_SIZE = 16
|
||||||
with Kernel((N // BLOCK_SIZE, N // BLOCK_SIZE, 1), WARP_THREADS) as ker:
|
with Kernel((N // BLOCK_SIZE, N // BLOCK_SIZE, 1), WARP_THREADS) as ker:
|
||||||
warp = ker.warp
|
warp = ker.warp
|
||||||
|
|
||||||
c = gl((1, 1, N, N), dtypes.float32)
|
c = ker.gl((1, 1, N, N), dtypes.float32)
|
||||||
a = gl((1, 1, N, N), dtypes.bfloat16)
|
a = ker.gl((1, 1, N, N), dtypes.bfloat16)
|
||||||
b = gl((1, 1, N, N), dtypes.bfloat16)
|
b = ker.gl((1, 1, N, N), dtypes.bfloat16)
|
||||||
|
|
||||||
a_smem = st((BLOCK_SIZE, BLOCK_SIZE), dtypes.bfloat16)
|
a_smem = ker.st((BLOCK_SIZE, BLOCK_SIZE), dtypes.bfloat16)
|
||||||
b_smem = st((BLOCK_SIZE, BLOCK_SIZE), dtypes.bfloat16)
|
b_smem = ker.st((BLOCK_SIZE, BLOCK_SIZE), dtypes.bfloat16)
|
||||||
c_smem = st((BLOCK_SIZE, BLOCK_SIZE), dtypes.float32)
|
c_smem = ker.st((BLOCK_SIZE, BLOCK_SIZE), dtypes.float32)
|
||||||
|
|
||||||
a_reg = rt((BLOCK_SIZE, BLOCK_SIZE), dtypes.bfloat16)
|
a_reg = ker.rt((BLOCK_SIZE, BLOCK_SIZE), dtypes.bfloat16)
|
||||||
b_reg = rt((BLOCK_SIZE, BLOCK_SIZE), dtypes.bfloat16)
|
b_reg = ker.rt((BLOCK_SIZE, BLOCK_SIZE), dtypes.bfloat16)
|
||||||
c_reg = rt((BLOCK_SIZE, BLOCK_SIZE), dtypes.float32)
|
c_reg = ker.rt((BLOCK_SIZE, BLOCK_SIZE), dtypes.float32)
|
||||||
|
|
||||||
col, row = ker.blockIdx_x, ker.blockIdx_y
|
col, row = ker.blockIdx_x, ker.blockIdx_y
|
||||||
|
|
||||||
@@ -108,7 +112,7 @@ class TestTK(unittest.TestCase):
|
|||||||
|
|
||||||
ref = a.matmul(b.transpose(2, 3), dtype=dtypes.float32).float()
|
ref = a.matmul(b.transpose(2, 3), dtype=dtypes.float32).float()
|
||||||
|
|
||||||
assert ref.allclose(c)
|
np.testing.assert_allclose(c.numpy(), ref.numpy())
|
||||||
|
|
||||||
def test_load_store(self):
|
def test_load_store(self):
|
||||||
N = 32
|
N = 32
|
||||||
@@ -116,14 +120,14 @@ class TestTK(unittest.TestCase):
|
|||||||
with Kernel((N // BLOCK_SIZE, N // BLOCK_SIZE, 1), WARP_THREADS) as ker:
|
with Kernel((N // BLOCK_SIZE, N // BLOCK_SIZE, 1), WARP_THREADS) as ker:
|
||||||
warp = ker.warp
|
warp = ker.warp
|
||||||
|
|
||||||
b = gl((1, 1, N, N), dtypes.float32)
|
b = ker.gl((1, 1, N, N), dtypes.float32)
|
||||||
a = gl((1, 1, N, N), dtypes.float32)
|
a = ker.gl((1, 1, N, N), dtypes.float32)
|
||||||
|
|
||||||
a_smem = st((BLOCK_SIZE, BLOCK_SIZE), dtypes.float32)
|
a_smem = ker.st((BLOCK_SIZE, BLOCK_SIZE), dtypes.float32)
|
||||||
b_smem = st((BLOCK_SIZE, BLOCK_SIZE), dtypes.float32)
|
b_smem = ker.st((BLOCK_SIZE, BLOCK_SIZE), dtypes.float32)
|
||||||
|
|
||||||
a_reg = rt((BLOCK_SIZE, BLOCK_SIZE), dtypes.float32)
|
a_reg = ker.rt((BLOCK_SIZE, BLOCK_SIZE), dtypes.float32)
|
||||||
b_reg = rt((BLOCK_SIZE, BLOCK_SIZE), dtypes.float32)
|
b_reg = ker.rt((BLOCK_SIZE, BLOCK_SIZE), dtypes.float32)
|
||||||
|
|
||||||
col, row = ker.blockIdx_x, ker.blockIdx_y
|
col, row = ker.blockIdx_x, ker.blockIdx_y
|
||||||
|
|
||||||
@@ -146,7 +150,45 @@ class TestTK(unittest.TestCase):
|
|||||||
|
|
||||||
ref = a.float()
|
ref = a.float()
|
||||||
|
|
||||||
assert ref.allclose(b)
|
np.testing.assert_allclose(b.numpy(), ref.numpy())
|
||||||
|
|
||||||
|
def test_add(self):
|
||||||
|
N = 32
|
||||||
|
BLOCK_SIZE = 16
|
||||||
|
with Kernel((1, 1, 1), WARP_THREADS) as ker:
|
||||||
|
warp = ker.warp
|
||||||
|
|
||||||
|
b = ker.gl((1, 1, N, N), dtypes.float32)
|
||||||
|
a = ker.gl((1, 1, N, N), dtypes.float32)
|
||||||
|
|
||||||
|
a_smem = ker.st((BLOCK_SIZE, BLOCK_SIZE), dtypes.float32)
|
||||||
|
|
||||||
|
a_reg = ker.rt((BLOCK_SIZE, BLOCK_SIZE), dtypes.float32)
|
||||||
|
|
||||||
|
for tile_row in ker.range(N // BLOCK_SIZE):
|
||||||
|
for tile_col in ker.range(N // BLOCK_SIZE):
|
||||||
|
a_smem = warp.load(a_smem, a, (), (0, 0, tile_row, tile_col), axis=2)
|
||||||
|
a_reg = warp.load(a_reg, a_smem)
|
||||||
|
|
||||||
|
a_reg += 1
|
||||||
|
|
||||||
|
a_smem = warp.store(a_smem, a_reg)
|
||||||
|
b = warp.store(b, a_smem, (0, 0, tile_row, tile_col), (), axis=2)
|
||||||
|
|
||||||
|
sink = ker.finish()
|
||||||
|
|
||||||
|
with Context(DEBUG=0):
|
||||||
|
a = Tensor.rand(1, 1, N, N, dtype="float32").contiguous()
|
||||||
|
b = Tensor.empty(1, 1, N, N, dtype="float32")
|
||||||
|
Tensor.realize(a, b)
|
||||||
|
|
||||||
|
ei = ExecItem(get_runner(Device.DEFAULT, sink), [t.uop.buffer for t in (b, a)])
|
||||||
|
for _ in range(5): ei.run(wait=True)
|
||||||
|
b = b.float()
|
||||||
|
|
||||||
|
ref = a.float() + 1
|
||||||
|
|
||||||
|
np.testing.assert_allclose(b.numpy(), ref.numpy())
|
||||||
|
|
||||||
def test_max(self):
|
def test_max(self):
|
||||||
N = 16
|
N = 16
|
||||||
@@ -154,28 +196,27 @@ class TestTK(unittest.TestCase):
|
|||||||
with Kernel((1, 1, 1), WARP_THREADS) as ker:
|
with Kernel((1, 1, 1), WARP_THREADS) as ker:
|
||||||
warp = ker.warp
|
warp = ker.warp
|
||||||
|
|
||||||
b = gl((1, 1, N, N), dtypes.float32)
|
b = ker.gl((1, 1, N, N), dtypes.float32)
|
||||||
a = gl((1, 1, N, N), dtypes.float32)
|
a = ker.gl((1, 1, N, N), dtypes.float32)
|
||||||
|
|
||||||
a_smem = st((BLOCK_SIZE, BLOCK_SIZE), dtypes.float32)
|
a_smem = ker.st((BLOCK_SIZE, BLOCK_SIZE), dtypes.float32)
|
||||||
b_smem = st((BLOCK_SIZE, BLOCK_SIZE), dtypes.float32)
|
b_smem = ker.st((BLOCK_SIZE, BLOCK_SIZE), dtypes.float32)
|
||||||
|
|
||||||
a_reg = rt((BLOCK_SIZE, BLOCK_SIZE), dtypes.float32)
|
a_reg = ker.rt((BLOCK_SIZE, BLOCK_SIZE), dtypes.float32)
|
||||||
b_reg = rt((BLOCK_SIZE, BLOCK_SIZE), dtypes.float32)
|
b_reg = ker.rt((BLOCK_SIZE, BLOCK_SIZE), dtypes.float32)
|
||||||
|
|
||||||
max_reg = rv(BLOCK_SIZE, dtypes.float32, "ortho")
|
max_reg = ker.rv(BLOCK_SIZE, dtypes.float32, "ortho")
|
||||||
|
|
||||||
max_reg = warp.neg_inf(max_reg)
|
|
||||||
|
|
||||||
for tile_row in ker.range(N // BLOCK_SIZE):
|
for tile_row in ker.range(N // BLOCK_SIZE):
|
||||||
|
max_reg = warp.neg_inf(max_reg.after(tile_row))
|
||||||
|
|
||||||
for tile_col in ker.range(N // BLOCK_SIZE):
|
for tile_col in ker.range(N // BLOCK_SIZE):
|
||||||
a_smem = warp.load(a_smem, a, (), (0, 0, tile_row, tile_col), axis=2)
|
a_smem = warp.load(a_smem, a, (), (0, 0, tile_row, tile_col), axis=2)
|
||||||
a_reg = warp.load(a_reg, a_smem)
|
a_reg = warp.load(a_reg, a_smem)
|
||||||
max_reg = warp.row_reduce(max_reg, a_reg, lambda a, b: a.maximum(b))
|
max_reg = warp.row_reduce(max_reg, a_reg, lambda a, b: a.maximum(b))
|
||||||
sum_reg = ker.endrange()
|
max_reg = ker.endrange()
|
||||||
|
|
||||||
b_reg = warp.zero(b_reg).after(tile_row)
|
b_reg = warp.map(b_reg, lambda _, idx: max_reg[idx[0], 0, (idx[2]%4)//2])
|
||||||
b_reg = warp.map(b_reg, lambda _, idx: sum_reg[idx[0], 0, (idx[2]%4)//2])
|
|
||||||
b_smem = warp.store(b_smem, b_reg)
|
b_smem = warp.store(b_smem, b_reg)
|
||||||
|
|
||||||
for tile_col in ker.range(N // BLOCK_SIZE):
|
for tile_col in ker.range(N // BLOCK_SIZE):
|
||||||
@@ -194,7 +235,7 @@ class TestTK(unittest.TestCase):
|
|||||||
|
|
||||||
ref = a.float().max(axis=3, keepdim=True).expand(a.shape)
|
ref = a.float().max(axis=3, keepdim=True).expand(a.shape)
|
||||||
|
|
||||||
assert ref.allclose(b)
|
np.testing.assert_allclose(b.numpy(), ref.numpy())
|
||||||
|
|
||||||
def test_max_nonsquare(self):
|
def test_max_nonsquare(self):
|
||||||
N, M = 16, 64
|
N, M = 16, 64
|
||||||
@@ -202,28 +243,27 @@ class TestTK(unittest.TestCase):
|
|||||||
with Kernel((1, 1, 1), WARP_THREADS) as ker:
|
with Kernel((1, 1, 1), WARP_THREADS) as ker:
|
||||||
warp = ker.warp
|
warp = ker.warp
|
||||||
|
|
||||||
b = gl((1, 1, N, M), dtypes.float32)
|
b = ker.gl((1, 1, N, M), dtypes.float32)
|
||||||
a = gl((1, 1, N, M), dtypes.float32)
|
a = ker.gl((1, 1, N, M), dtypes.float32)
|
||||||
|
|
||||||
a_smem = st((BLOCK_N, BLOCK_M), dtypes.float32)
|
a_smem = ker.st((BLOCK_N, BLOCK_M), dtypes.float32)
|
||||||
b_smem = st((BLOCK_N, BLOCK_M), dtypes.float32)
|
b_smem = ker.st((BLOCK_N, BLOCK_M), dtypes.float32)
|
||||||
|
|
||||||
a_reg = rt((BLOCK_N, BLOCK_M), dtypes.float32)
|
a_reg = ker.rt((BLOCK_N, BLOCK_M), dtypes.float32)
|
||||||
b_reg = rt((BLOCK_N, BLOCK_M), dtypes.float32)
|
b_reg = ker.rt((BLOCK_N, BLOCK_M), dtypes.float32)
|
||||||
|
|
||||||
max_reg = rv(BLOCK_N, dtypes.float32, "ortho")
|
max_reg = ker.rv(BLOCK_N, dtypes.float32, "ortho")
|
||||||
|
|
||||||
max_reg = warp.zero(max_reg)
|
|
||||||
|
|
||||||
for tile_row in ker.range(N // BLOCK_N):
|
for tile_row in ker.range(N // BLOCK_N):
|
||||||
|
max_reg = warp.neg_inf(max_reg.after(tile_row))
|
||||||
|
|
||||||
for tile_col in ker.range(M // BLOCK_M):
|
for tile_col in ker.range(M // BLOCK_M):
|
||||||
a_smem = warp.load(a_smem, a, (), (0, 0, tile_row, tile_col), axis=2)
|
a_smem = warp.load(a_smem, a, (), (0, 0, tile_row, tile_col), axis=2)
|
||||||
a_reg = warp.load(a_reg, a_smem)
|
a_reg = warp.load(a_reg, a_smem)
|
||||||
sum_reg = warp.row_reduce(max_reg, a_reg, lambda a, b: a.maximum(b))
|
max_reg = warp.row_reduce(max_reg, a_reg, lambda a, b: a.maximum(b))
|
||||||
sum_reg = ker.endrange()
|
max_reg = ker.endrange()
|
||||||
|
|
||||||
b_reg = warp.zero(b_reg).after(tile_row)
|
b_reg = warp.map(b_reg, lambda _, idx: max_reg[idx[0], 0, (idx[2]%4)//2])
|
||||||
b_reg = warp.map(b_reg, lambda _, idx: sum_reg[idx[0], 0, (idx[2]%4)//2])
|
|
||||||
b_smem = warp.store(b_smem, b_reg)
|
b_smem = warp.store(b_smem, b_reg)
|
||||||
|
|
||||||
for tile_col in ker.range(M // BLOCK_M):
|
for tile_col in ker.range(M // BLOCK_M):
|
||||||
@@ -242,27 +282,27 @@ class TestTK(unittest.TestCase):
|
|||||||
|
|
||||||
ref = a.float().max(axis=3, keepdim=True).expand(a.shape)
|
ref = a.float().max(axis=3, keepdim=True).expand(a.shape)
|
||||||
|
|
||||||
assert ref.allclose(b)
|
np.testing.assert_allclose(b.numpy(), ref.numpy())
|
||||||
|
|
||||||
def test_sum(self):
|
def test_sum(self):
|
||||||
N = 16
|
N = 32
|
||||||
BLOCK_SIZE = 16
|
BLOCK_SIZE = 16
|
||||||
with Kernel((1, 1, 1), WARP_THREADS) as ker:
|
with Kernel((1, 1, 1), WARP_THREADS) as ker:
|
||||||
warp = ker.warp
|
warp = ker.warp
|
||||||
|
|
||||||
b = gl((1, 1, N, N), dtypes.float32)
|
b = ker.gl((1, 1, N, N), dtypes.float32)
|
||||||
a = gl((1, 1, N, N), dtypes.float32)
|
a = ker.gl((1, 1, N, N), dtypes.float32)
|
||||||
|
|
||||||
a_smem = st((BLOCK_SIZE, BLOCK_SIZE), dtypes.float32)
|
a_smem = ker.st((BLOCK_SIZE, BLOCK_SIZE), dtypes.float32)
|
||||||
b_smem = st((BLOCK_SIZE, BLOCK_SIZE), dtypes.float32)
|
b_smem = ker.st((BLOCK_SIZE, BLOCK_SIZE), dtypes.float32)
|
||||||
|
|
||||||
a_reg = rt((BLOCK_SIZE, BLOCK_SIZE), dtypes.float32)
|
a_reg = ker.rt((BLOCK_SIZE, BLOCK_SIZE), dtypes.float32)
|
||||||
b_reg = rt((BLOCK_SIZE, BLOCK_SIZE), dtypes.float32)
|
b_reg = ker.rt((BLOCK_SIZE, BLOCK_SIZE), dtypes.float32)
|
||||||
|
|
||||||
sum_reg = rv(BLOCK_SIZE, dtypes.float32, "ortho")
|
sum_reg = ker.rv(BLOCK_SIZE, dtypes.float32, "ortho")
|
||||||
|
|
||||||
for tile_row in ker.range(N // BLOCK_SIZE):
|
for tile_row in ker.range(N // BLOCK_SIZE):
|
||||||
sum_reg = warp.zero(sum_reg).after(tile_row)
|
sum_reg = warp.zero(sum_reg.after(tile_row))
|
||||||
|
|
||||||
for tile_col in ker.range(N // BLOCK_SIZE):
|
for tile_col in ker.range(N // BLOCK_SIZE):
|
||||||
a_smem = warp.load(a_smem, a, (), (0, 0, tile_row, tile_col), axis=2)
|
a_smem = warp.load(a_smem, a, (), (0, 0, tile_row, tile_col), axis=2)
|
||||||
@@ -270,7 +310,6 @@ class TestTK(unittest.TestCase):
|
|||||||
sum_reg = warp.row_reduce(sum_reg, a_reg, lambda a, b: a + b)
|
sum_reg = warp.row_reduce(sum_reg, a_reg, lambda a, b: a + b)
|
||||||
sum_reg = ker.endrange()
|
sum_reg = ker.endrange()
|
||||||
|
|
||||||
b_reg = warp.zero(b_reg).after(tile_row)
|
|
||||||
b_reg = warp.map(b_reg, lambda _, idx: sum_reg[idx[0], 0, (idx[2]%4)//2])
|
b_reg = warp.map(b_reg, lambda _, idx: sum_reg[idx[0], 0, (idx[2]%4)//2])
|
||||||
b_smem = warp.store(b_smem, b_reg)
|
b_smem = warp.store(b_smem, b_reg)
|
||||||
|
|
||||||
@@ -281,7 +320,6 @@ class TestTK(unittest.TestCase):
|
|||||||
|
|
||||||
with Context(DEBUG=0):
|
with Context(DEBUG=0):
|
||||||
a = Tensor.rand(1, 1, N, N, dtype="float32").contiguous()
|
a = Tensor.rand(1, 1, N, N, dtype="float32").contiguous()
|
||||||
a = Tensor.arange(1 * 1 * N * N).reshape(1, 1, N, N).cast(dtypes.float32).contiguous()
|
|
||||||
b = Tensor.empty(1, 1, N, N, dtype="float32")
|
b = Tensor.empty(1, 1, N, N, dtype="float32")
|
||||||
Tensor.realize(a, b)
|
Tensor.realize(a, b)
|
||||||
|
|
||||||
@@ -291,7 +329,7 @@ class TestTK(unittest.TestCase):
|
|||||||
|
|
||||||
ref = a.float().sum(axis=3, keepdim=True).expand(a.shape)
|
ref = a.float().sum(axis=3, keepdim=True).expand(a.shape)
|
||||||
|
|
||||||
assert ref.allclose(b)
|
np.testing.assert_allclose(b.numpy(), ref.numpy(), atol=1e-5, rtol=1e-5)
|
||||||
|
|
||||||
def test_sum_nonsquare(self):
|
def test_sum_nonsquare(self):
|
||||||
N, M = 16, 64
|
N, M = 16, 64
|
||||||
@@ -299,27 +337,26 @@ class TestTK(unittest.TestCase):
|
|||||||
with Kernel((1, 1, 1), WARP_THREADS) as ker:
|
with Kernel((1, 1, 1), WARP_THREADS) as ker:
|
||||||
warp = ker.warp
|
warp = ker.warp
|
||||||
|
|
||||||
b = gl((1, 1, N, M), dtypes.float32)
|
b = ker.gl((1, 1, N, M), dtypes.float32)
|
||||||
a = gl((1, 1, N, M), dtypes.float32)
|
a = ker.gl((1, 1, N, M), dtypes.float32)
|
||||||
|
|
||||||
a_smem = st((BLOCK_N, BLOCK_M), dtypes.float32)
|
a_smem = ker.st((BLOCK_N, BLOCK_M), dtypes.float32)
|
||||||
b_smem = st((BLOCK_N, BLOCK_M), dtypes.float32)
|
b_smem = ker.st((BLOCK_N, BLOCK_M), dtypes.float32)
|
||||||
|
|
||||||
a_reg = rt((BLOCK_N, BLOCK_M), dtypes.float32)
|
a_reg = ker.rt((BLOCK_N, BLOCK_M), dtypes.float32)
|
||||||
b_reg = rt((BLOCK_N, BLOCK_M), dtypes.float32)
|
b_reg = ker.rt((BLOCK_N, BLOCK_M), dtypes.float32)
|
||||||
|
|
||||||
sum_reg = rv(BLOCK_N, dtypes.float32, "ortho")
|
sum_reg = ker.rv(BLOCK_N, dtypes.float32, "ortho")
|
||||||
|
|
||||||
sum_reg = warp.zero(sum_reg)
|
|
||||||
|
|
||||||
for tile_row in ker.range(N // BLOCK_N):
|
for tile_row in ker.range(N // BLOCK_N):
|
||||||
|
sum_reg = warp.zero(sum_reg.after(tile_row))
|
||||||
|
|
||||||
for tile_col in ker.range(M // BLOCK_M):
|
for tile_col in ker.range(M // BLOCK_M):
|
||||||
a_smem = warp.load(a_smem, a, (), (0, 0, tile_row, tile_col), axis=2)
|
a_smem = warp.load(a_smem, a, (), (0, 0, tile_row, tile_col), axis=2)
|
||||||
a_reg = warp.load(a_reg, a_smem)
|
a_reg = warp.load(a_reg, a_smem)
|
||||||
sum_reg = warp.row_reduce(sum_reg, a_reg, lambda a, b: a + b)
|
sum_reg = warp.row_reduce(sum_reg, a_reg, lambda a, b: a + b)
|
||||||
sum_reg = ker.endrange()
|
sum_reg = ker.endrange()
|
||||||
|
|
||||||
b_reg = warp.zero(b_reg).after(tile_row)
|
|
||||||
b_reg = warp.map(b_reg, lambda _, idx: sum_reg[idx[0], 0, (idx[2]%4)//2])
|
b_reg = warp.map(b_reg, lambda _, idx: sum_reg[idx[0], 0, (idx[2]%4)//2])
|
||||||
b_smem = warp.store(b_smem, b_reg)
|
b_smem = warp.store(b_smem, b_reg)
|
||||||
|
|
||||||
@@ -339,7 +376,68 @@ class TestTK(unittest.TestCase):
|
|||||||
|
|
||||||
ref = a.float().sum(axis=3, keepdim=True).expand(a.shape)
|
ref = a.float().sum(axis=3, keepdim=True).expand(a.shape)
|
||||||
|
|
||||||
assert ref.allclose(b)
|
np.testing.assert_allclose(b.numpy(), ref.numpy(), atol=1e-5, rtol=1e-5)
|
||||||
|
|
||||||
|
@unittest.skip("fake range not ended")
|
||||||
|
def test_softmax(self):
|
||||||
|
N = 32
|
||||||
|
BLOCK_SIZE = 16
|
||||||
|
with Kernel((1, 1, 1), WARP_THREADS) as ker:
|
||||||
|
warp = ker.warp
|
||||||
|
|
||||||
|
b = ker.gl((1, 1, BLOCK_SIZE, N), dtypes.float32)
|
||||||
|
a = ker.gl((1, 1, BLOCK_SIZE, N), dtypes.float32)
|
||||||
|
|
||||||
|
a_smem = ker.st((BLOCK_SIZE, BLOCK_SIZE), dtypes.float32)
|
||||||
|
|
||||||
|
a_reg = ker.rt((BLOCK_SIZE, BLOCK_SIZE), dtypes.float32)
|
||||||
|
|
||||||
|
max_vec_last = ker.rv(BLOCK_SIZE, dtypes.float32, "ortho")
|
||||||
|
max_vec = ker.rv(BLOCK_SIZE, dtypes.float32, "ortho")
|
||||||
|
norm_vec = ker.rv(BLOCK_SIZE, dtypes.float32, "ortho")
|
||||||
|
|
||||||
|
max_vec = warp.neg_inf(max_vec)
|
||||||
|
norm_vec = warp.zero(norm_vec)
|
||||||
|
|
||||||
|
for tile_col in ker.range(N // BLOCK_SIZE):
|
||||||
|
a_smem = warp.load(a_smem, a, (), (0, 0, 0, tile_col), axis=2)
|
||||||
|
a_reg = warp.load(a_reg, a_smem)
|
||||||
|
|
||||||
|
a_reg *= 1.0 / math.log(2)
|
||||||
|
|
||||||
|
max_vec_last = warp.copy(max_vec_last.after(tile_col), max_vec)
|
||||||
|
max_vec = warp.row_reduce(max_vec, a_reg, lambda a, b: a.maximum(b))
|
||||||
|
a_reg = (a_reg - max_vec).exp2()
|
||||||
|
max_vec_last = (max_vec_last - max_vec).exp2()
|
||||||
|
norm_vec *= max_vec_last
|
||||||
|
norm_vec = warp.row_reduce(norm_vec, a_reg, lambda a, b: a + b)
|
||||||
|
norm_vec = ker.endrange()
|
||||||
|
|
||||||
|
for tile_col in ker.range(N // BLOCK_SIZE):
|
||||||
|
a_smem = warp.load(a_smem, a, (), (0, 0, 0, tile_col), axis=2)
|
||||||
|
a_reg = warp.load(a_reg, a_smem)
|
||||||
|
|
||||||
|
a_reg *= 1.0 / math.log(2)
|
||||||
|
a_reg = (a_reg - max_vec).exp2()
|
||||||
|
a_reg /= norm_vec
|
||||||
|
|
||||||
|
a_smem = warp.store(a_smem, a_reg)
|
||||||
|
b = warp.store(b, a_smem, (0, 0, 0, tile_col), (), axis=2)
|
||||||
|
|
||||||
|
sink = ker.finish()
|
||||||
|
|
||||||
|
with Context(DEBUG=0):
|
||||||
|
a = Tensor.rand(1, 1, BLOCK_SIZE, N, dtype="float32")
|
||||||
|
b = Tensor.empty(1, 1, BLOCK_SIZE, N, dtype="float32")
|
||||||
|
Tensor.realize(a, b)
|
||||||
|
|
||||||
|
ei = ExecItem(get_runner(Device.DEFAULT, sink), [t.uop.buffer for t in (b, a)])
|
||||||
|
for _ in range(5): ei.run(wait=True)
|
||||||
|
b = b.float()
|
||||||
|
|
||||||
|
ref = a.float().softmax(axis=3)
|
||||||
|
|
||||||
|
np.testing.assert_allclose(b.numpy(), ref.numpy(), atol=1e-5, rtol=1e-5)
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
unittest.main()
|
unittest.main()
|
||||||
@@ -0,0 +1,85 @@
|
|||||||
|
import ctypes, subprocess, tempfile, unittest
|
||||||
|
from tinygrad.helpers import WIN
|
||||||
|
from tinygrad.runtime.support.c import Struct
|
||||||
|
|
||||||
|
class TestAutogen(unittest.TestCase):
|
||||||
|
def test_packed_struct_sizeof(self):
|
||||||
|
layout = [('a', ctypes.c_char), ('b', ctypes.c_int, 5), ('c', ctypes.c_char)]
|
||||||
|
class Y(ctypes.Structure): _fields_, _pack_, _layout_ = layout, 1, 'ms'
|
||||||
|
class Z(Struct): pass
|
||||||
|
Z._packed_, Z._fields_ = True, layout
|
||||||
|
self.assertEqual(ctypes.sizeof(Y), 6)
|
||||||
|
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)]
|
||||||
|
class Foo(ctypes.Structure): _fields_, _layout_ = layout, 'gcc-sysv'
|
||||||
|
class Bar(ctypes.Structure): _fields_, _pack_, _layout_ = layout, 1, 'ms'
|
||||||
|
class Baz(Struct): pass
|
||||||
|
Baz._packed_, Baz._fields_ = True, layout
|
||||||
|
self.assertEqual(ctypes.sizeof(Foo), 12)
|
||||||
|
self.assertEqual(ctypes.sizeof(Bar), 12)
|
||||||
|
self.assertEqual(ctypes.sizeof(Baz), 8)
|
||||||
|
|
||||||
|
@unittest.skipIf(WIN, "doesn't compile on windows")
|
||||||
|
def test_packed_struct_interop(self):
|
||||||
|
class Baz(Struct): pass
|
||||||
|
Baz._packed_ = True
|
||||||
|
Baz._fields_ = [('a', ctypes.c_int, 30), ('b', ctypes.c_int, 30), ('c', ctypes.c_int, 2), ('d', ctypes.c_int, 2)]
|
||||||
|
src = '''
|
||||||
|
struct __attribute__((packed)) baz {
|
||||||
|
int a:30;
|
||||||
|
int b:30;
|
||||||
|
int c:2;
|
||||||
|
int d:2;
|
||||||
|
};
|
||||||
|
|
||||||
|
int test(struct baz x) {
|
||||||
|
return x.a + x.b + x.c + x.d;
|
||||||
|
}
|
||||||
|
'''
|
||||||
|
args = ('-x', 'c', '-fPIC', '-shared')
|
||||||
|
with tempfile.NamedTemporaryFile(suffix=".so") as f:
|
||||||
|
subprocess.check_output(('clang',) + args + ('-', '-o', f.name), input=src.encode('utf-8'))
|
||||||
|
b = Baz(0xAA000, 0x00BB0, 0, 1)
|
||||||
|
test = ctypes.CDLL(f.name).test
|
||||||
|
test.argtypes = [Baz]
|
||||||
|
self.assertEqual(test(b), b.a + b.b + b.c + b.d)
|
||||||
|
|
||||||
|
@unittest.skipIf(WIN, "doesn't compile on windows")
|
||||||
|
def test_packed_structs(self):
|
||||||
|
NvU32 = ctypes.c_uint32
|
||||||
|
NvU64 = ctypes.c_uint64
|
||||||
|
class FWSECLIC_READ_VBIOS_DESC(Struct): pass
|
||||||
|
FWSECLIC_READ_VBIOS_DESC._packed_ = True
|
||||||
|
FWSECLIC_READ_VBIOS_DESC._fields_ = [
|
||||||
|
('version', NvU32),
|
||||||
|
('size', NvU32),
|
||||||
|
('gfwImageOffset', NvU64),
|
||||||
|
('gfwImageSize', NvU32),
|
||||||
|
('flags', NvU32),
|
||||||
|
]
|
||||||
|
class FWSECLIC_FRTS_REGION_DESC(Struct): pass
|
||||||
|
FWSECLIC_FRTS_REGION_DESC._packed_ = True
|
||||||
|
FWSECLIC_FRTS_REGION_DESC._fields_ = [
|
||||||
|
('version', NvU32),
|
||||||
|
('size', NvU32),
|
||||||
|
('frtsRegionOffset4K', NvU32),
|
||||||
|
('frtsRegionSize', NvU32),
|
||||||
|
('frtsRegionMediaType', NvU32),
|
||||||
|
]
|
||||||
|
class FWSECLIC_FRTS_CMD(Struct): pass
|
||||||
|
FWSECLIC_FRTS_CMD._packed_ = True
|
||||||
|
FWSECLIC_FRTS_CMD._fields_ = [
|
||||||
|
('readVbiosDesc', FWSECLIC_READ_VBIOS_DESC),
|
||||||
|
('frtsRegionDesc', FWSECLIC_FRTS_REGION_DESC),
|
||||||
|
]
|
||||||
|
read_vbios_desc = FWSECLIC_READ_VBIOS_DESC(version=0x1, size=ctypes.sizeof(FWSECLIC_READ_VBIOS_DESC), flags=2)
|
||||||
|
frst_reg_desc = FWSECLIC_FRTS_REGION_DESC(version=0x1, size=ctypes.sizeof(FWSECLIC_FRTS_REGION_DESC),
|
||||||
|
frtsRegionOffset4K=0xdead, frtsRegionSize=0x100, frtsRegionMediaType=2)
|
||||||
|
frts_cmd = FWSECLIC_FRTS_CMD(readVbiosDesc=read_vbios_desc, frtsRegionDesc=frst_reg_desc)
|
||||||
|
assert int.from_bytes(frts_cmd, 'little') == 0x2000001000000dead0000001400000001000000020000000000000000000000000000001800000001
|
||||||
|
assert int.from_bytes(frts_cmd.readVbiosDesc, 'little') == int.from_bytes(read_vbios_desc, 'little')
|
||||||
|
assert int.from_bytes(frts_cmd.frtsRegionDesc, 'little') == int.from_bytes(frst_reg_desc, 'little')
|
||||||
|
assert frts_cmd.readVbiosDesc.__class__ is FWSECLIC_READ_VBIOS_DESC
|
||||||
|
assert frts_cmd.frtsRegionDesc.__class__ is FWSECLIC_FRTS_REGION_DESC
|
||||||
|
|
||||||
|
if __name__ == "__main__": unittest.main()
|
||||||
@@ -62,6 +62,7 @@ class TestConv(unittest.TestCase):
|
|||||||
np.testing.assert_allclose(r1.numpy(), np.maximum(out.numpy(), 0), atol=1e-5)
|
np.testing.assert_allclose(r1.numpy(), np.maximum(out.numpy(), 0), atol=1e-5)
|
||||||
np.testing.assert_allclose(r2.numpy(), np.where(out.numpy() > 0, out.numpy(), (np.exp(out.numpy()) - 1)), atol=1e-5)
|
np.testing.assert_allclose(r2.numpy(), np.where(out.numpy() > 0, out.numpy(), (np.exp(out.numpy()) - 1)), atol=1e-5)
|
||||||
|
|
||||||
|
@unittest.skip("this test is flaky")
|
||||||
def test_two_overlapping_binops_no_rerun_wino(self):
|
def test_two_overlapping_binops_no_rerun_wino(self):
|
||||||
with Context(WINO=1):
|
with Context(WINO=1):
|
||||||
x = Tensor.randn(1,4,16,16)
|
x = Tensor.randn(1,4,16,16)
|
||||||
|
|||||||
@@ -81,20 +81,20 @@ class TestCompiler(unittest.TestCase):
|
|||||||
def test_compile_cached(self):
|
def test_compile_cached(self):
|
||||||
diskcache_put("key", "123", None) # clear cache
|
diskcache_put("key", "123", None) # clear cache
|
||||||
getenv.cache_clear()
|
getenv.cache_clear()
|
||||||
with Context(DISABLE_COMPILER_CACHE=0):
|
with Context(CCACHE=1):
|
||||||
self.assertEqual(MockCompiler("key").compile_cached("123"), str.encode("123"))
|
self.assertEqual(MockCompiler("key").compile_cached("123"), str.encode("123"))
|
||||||
self.assertEqual(diskcache_get("key", "123"), str.encode("123"))
|
self.assertEqual(diskcache_get("key", "123"), str.encode("123"))
|
||||||
|
|
||||||
def test_compile_cached_disabled(self):
|
def test_compile_cached_disabled(self):
|
||||||
diskcache_put("disabled_key", "123", None) # clear cache
|
diskcache_put("disabled_key", "123", None) # clear cache
|
||||||
getenv.cache_clear()
|
getenv.cache_clear()
|
||||||
with Context(DISABLE_COMPILER_CACHE=1):
|
with Context(CCACHE=0):
|
||||||
self.assertEqual(MockCompiler("disabled_key").compile_cached("123"), str.encode("123"))
|
self.assertEqual(MockCompiler("disabled_key").compile_cached("123"), str.encode("123"))
|
||||||
self.assertIsNone(diskcache_get("disabled_key", "123"))
|
self.assertIsNone(diskcache_get("disabled_key", "123"))
|
||||||
|
|
||||||
def test_device_compile(self):
|
def test_device_compile(self):
|
||||||
getenv.cache_clear()
|
getenv.cache_clear()
|
||||||
with Context(DISABLE_COMPILER_CACHE=1):
|
with Context(CCACHE=0):
|
||||||
a = Tensor([0.,1.], device=Device.DEFAULT).realize()
|
a = Tensor([0.,1.], device=Device.DEFAULT).realize()
|
||||||
(a + 1).realize()
|
(a + 1).realize()
|
||||||
|
|
||||||
|
|||||||
@@ -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()
|
||||||
|
|
||||||
|
|||||||
@@ -66,6 +66,7 @@ class TestProgressBar(unittest.TestCase):
|
|||||||
tqdm_output = tqdm.format_meter(n=total, total=total, elapsed=elapsed, ncols=ncols, prefix="Test")
|
tqdm_output = tqdm.format_meter(n=total, total=total, elapsed=elapsed, ncols=ncols, prefix="Test")
|
||||||
self._compare_bars(tinytqdm_output, tqdm_output)
|
self._compare_bars(tinytqdm_output, tqdm_output)
|
||||||
|
|
||||||
|
@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_unit_scale(self, mock_terminal_size, mock_stderr):
|
def test_unit_scale(self, mock_terminal_size, mock_stderr):
|
||||||
|
|||||||
+14
-4
@@ -6,6 +6,7 @@ from tinygrad.uop.ops import UOp, UPat, Ops, PatternMatcher, TrackedPatternMatch
|
|||||||
from tinygrad.uop.symbolic import sym
|
from tinygrad.uop.symbolic import sym
|
||||||
from tinygrad.dtype import dtypes
|
from tinygrad.dtype import dtypes
|
||||||
from tinygrad.helpers import PROFILE, colored, ansistrip, flatten, TracingKey, ProfileRangeEvent, ProfileEvent, Context, cpu_events, profile_marker
|
from tinygrad.helpers import PROFILE, colored, ansistrip, flatten, TracingKey, ProfileRangeEvent, ProfileEvent, Context, cpu_events, profile_marker
|
||||||
|
from tinygrad.helpers import VIZ
|
||||||
from tinygrad.device import Buffer
|
from tinygrad.device import Buffer
|
||||||
|
|
||||||
@track_rewrites(name=True)
|
@track_rewrites(name=True)
|
||||||
@@ -33,11 +34,14 @@ class BaseTestViz(unittest.TestCase):
|
|||||||
cpu_events.clear()
|
cpu_events.clear()
|
||||||
self.tms = TRACK_MATCH_STATS.value
|
self.tms = TRACK_MATCH_STATS.value
|
||||||
self.profile = PROFILE.value
|
self.profile = PROFILE.value
|
||||||
|
self.viz = VIZ.value
|
||||||
TRACK_MATCH_STATS.value = 2
|
TRACK_MATCH_STATS.value = 2
|
||||||
PROFILE.value = 1
|
PROFILE.value = 1
|
||||||
|
VIZ.value = 1
|
||||||
def tearDown(self):
|
def tearDown(self):
|
||||||
TRACK_MATCH_STATS.value = self.tms
|
TRACK_MATCH_STATS.value = self.tms
|
||||||
PROFILE.value = self.profile
|
PROFILE.value = self.profile
|
||||||
|
VIZ.value = self.viz
|
||||||
|
|
||||||
class TestViz(BaseTestViz):
|
class TestViz(BaseTestViz):
|
||||||
def test_simple(self):
|
def test_simple(self):
|
||||||
@@ -366,8 +370,8 @@ def load_profile(lst:list[ProfileEvent]) -> dict:
|
|||||||
else: v["events"].append({"event":"free", "ts":ts, "key":key, "arg": {"users":[u("<IIBB") for _ in range(u("<I")[0])]}})
|
else: v["events"].append({"event":"free", "ts":ts, "key":key, "arg": {"users":[u("<IIBB") for _ in range(u("<I")[0])]}})
|
||||||
return {"dur":total_dur, "peak":global_peak, "layout":layout, "markers":markers}
|
return {"dur":total_dur, "peak":global_peak, "layout":layout, "markers":markers}
|
||||||
|
|
||||||
class TestVizProfiler(unittest.TestCase):
|
class TestVizProfiler(BaseTestViz):
|
||||||
def test_perfetto_node(self):
|
def test_node(self):
|
||||||
prof = [ProfileRangeEvent(device='NV', name='E_2', st=decimal.Decimal(1000), en=decimal.Decimal(1010), is_copy=False),
|
prof = [ProfileRangeEvent(device='NV', name='E_2', st=decimal.Decimal(1000), en=decimal.Decimal(1010), is_copy=False),
|
||||||
ProfileDeviceEvent(device='NV', comp_tdiff=decimal.Decimal(-1000), copy_tdiff=decimal.Decimal(-100))]
|
ProfileDeviceEvent(device='NV', comp_tdiff=decimal.Decimal(-1000), copy_tdiff=decimal.Decimal(-100))]
|
||||||
|
|
||||||
@@ -381,7 +385,7 @@ class TestVizProfiler(unittest.TestCase):
|
|||||||
self.assertEqual(event['dur'], 10)
|
self.assertEqual(event['dur'], 10)
|
||||||
assert event['ref'] is None
|
assert event['ref'] is None
|
||||||
|
|
||||||
def test_perfetto_copy_node(self):
|
def test_copy_node(self):
|
||||||
prof = [ProfileRangeEvent(device='NV', name='COPYxx', st=decimal.Decimal(1000), en=decimal.Decimal(1010), is_copy=True),
|
prof = [ProfileRangeEvent(device='NV', name='COPYxx', st=decimal.Decimal(1000), en=decimal.Decimal(1010), is_copy=True),
|
||||||
ProfileRangeEvent(device='NV:2', name='COPYxx', st=decimal.Decimal(1000), en=decimal.Decimal(1010), is_copy=True),
|
ProfileRangeEvent(device='NV:2', name='COPYxx', st=decimal.Decimal(1000), en=decimal.Decimal(1010), is_copy=True),
|
||||||
ProfileDeviceEvent(device='NV', comp_tdiff=decimal.Decimal(-1000), copy_tdiff=decimal.Decimal(-100)),
|
ProfileDeviceEvent(device='NV', comp_tdiff=decimal.Decimal(-1000), copy_tdiff=decimal.Decimal(-100)),
|
||||||
@@ -399,7 +403,7 @@ class TestVizProfiler(unittest.TestCase):
|
|||||||
|
|
||||||
self.assertEqual(j["dur"], (event2["st"]+event2["dur"])-event["st"])
|
self.assertEqual(j["dur"], (event2["st"]+event2["dur"])-event["st"])
|
||||||
|
|
||||||
def test_perfetto_graph(self):
|
def test_graph(self):
|
||||||
prof = [ProfileDeviceEvent(device='NV', comp_tdiff=decimal.Decimal(-1000), copy_tdiff=decimal.Decimal(-100)),
|
prof = [ProfileDeviceEvent(device='NV', comp_tdiff=decimal.Decimal(-1000), copy_tdiff=decimal.Decimal(-100)),
|
||||||
ProfileDeviceEvent(device='NV:1', comp_tdiff=decimal.Decimal(-500), copy_tdiff=decimal.Decimal(-50)),
|
ProfileDeviceEvent(device='NV:1', comp_tdiff=decimal.Decimal(-500), copy_tdiff=decimal.Decimal(-50)),
|
||||||
ProfileGraphEvent(ents=[ProfileGraphEntry(device='NV', name='E_25_4n2', st_id=0, en_id=1, is_copy=False),
|
ProfileGraphEvent(ents=[ProfileGraphEntry(device='NV', name='E_25_4n2', st_id=0, en_id=1, is_copy=False),
|
||||||
@@ -436,6 +440,12 @@ class TestVizProfiler(unittest.TestCase):
|
|||||||
sz = len(get_profile(prof))
|
sz = len(get_profile(prof))
|
||||||
self.assertLessEqual(sz/n_events, 26)
|
self.assertLessEqual(sz/n_events, 26)
|
||||||
|
|
||||||
|
def test_calltrace(self):
|
||||||
|
def fxn(): return Tensor.empty(10).mul(2).realize()
|
||||||
|
fxn()
|
||||||
|
trace = get_viz_list()[0]["steps"][0]["trace"]
|
||||||
|
assert any(fxn.__code__.co_filename == f and fxn.__code__.co_firstlineno == l for f,l,*_ in trace), str(trace)
|
||||||
|
|
||||||
# can pack up to 1hr 11 min of trace events
|
# can pack up to 1hr 11 min of trace events
|
||||||
def test_trace_duration(self):
|
def test_trace_duration(self):
|
||||||
dur_mins = 72
|
dur_mins = 72
|
||||||
|
|||||||
@@ -81,10 +81,10 @@ def hand_coded_optimizations(k:Scheduler) -> Scheduler:
|
|||||||
return k
|
return k
|
||||||
|
|
||||||
# are we grouping? (requires local shape support)
|
# are we grouping? (requires local shape support)
|
||||||
if resolve(prod(k.output_shape[i] for i in k.upcastable_dims) <= (128 if NOLOCALS else 2048), False):
|
if resolve(prod(k.output_shape[i] for i in k.upcastable_dims) <= (240 if NOLOCALS else 2048), False):
|
||||||
for sz in [16]:
|
for axis, sz in itertools.product((0, 1, 2), (16,)):
|
||||||
try:
|
try:
|
||||||
k.apply_opt(Opt(OptOps.GROUPTOP, 0, sz))
|
k.apply_opt(Opt(OptOps.GROUPTOP, axis, sz))
|
||||||
break
|
break
|
||||||
except KernelOptError: pass
|
except KernelOptError: pass
|
||||||
|
|
||||||
|
|||||||
@@ -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]:
|
||||||
@@ -52,8 +51,10 @@ class Scheduler:
|
|||||||
def get_optimized_ast(self, name_override:str|None=None):
|
def get_optimized_ast(self, name_override:str|None=None):
|
||||||
if name_override is not None: name = name_override
|
if name_override is not None: name = name_override
|
||||||
else:
|
else:
|
||||||
kernel_type = "r" if self.reduceop is not None else "E"
|
k_type = "r" if self.reduceop is not None else "E"
|
||||||
name = kernel_type + colored('_', 'BLACK').join(['']+[colored(x.src[0].render(), color) for x,color in zip(self.rngs, self.colors())])
|
special_uops = sorted([x for x in self.ast.toposort() if x.op is Ops.SPECIAL], key=lambda x: x.arg)
|
||||||
|
special_ops = [colored(str(x.vmax+1), "blue" if x.arg[0] == "g" else "cyan") for x in special_uops]
|
||||||
|
name = k_type + colored('_', 'BLACK').join(['']+special_ops+[colored(x.src[0].render(), color) for x,color in zip(self.rngs, self.colors())])
|
||||||
Scheduler.kernel_cnt[(function_name := to_function_name(name))] += 1
|
Scheduler.kernel_cnt[(function_name := to_function_name(name))] += 1
|
||||||
num = f"n{Scheduler.kernel_cnt[function_name]-1}" if Scheduler.kernel_cnt[function_name] > 1 else ""
|
num = f"n{Scheduler.kernel_cnt[function_name]-1}" if Scheduler.kernel_cnt[function_name] > 1 else ""
|
||||||
name += colored(num, 'BLACK')
|
name += colored(num, 'BLACK')
|
||||||
@@ -93,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()}")
|
||||||
@@ -229,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]}")
|
||||||
|
|||||||
+6
-5
@@ -4,8 +4,8 @@ from collections import defaultdict
|
|||||||
from typing import Any, Generic, TypeVar, Iterator, Sequence, cast, Generator
|
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, DISABLE_COMPILER_CACHE, 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
|
from tinygrad.helpers import unwrap_class_type, suppress_finalizing, AMD_LLVM, 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
|
||||||
|
|
||||||
@@ -266,7 +266,7 @@ class LRUAllocator(Allocator, Generic[DeviceType]):
|
|||||||
class CompileError(Exception): pass
|
class CompileError(Exception): pass
|
||||||
|
|
||||||
class Compiler:
|
class Compiler:
|
||||||
def __init__(self, cachekey:str|None=None): self.cachekey = None if DISABLE_COMPILER_CACHE else cachekey
|
def __init__(self, cachekey:str|None=None): self.cachekey = cachekey if CCACHE else None
|
||||||
def compile(self, src:str) -> bytes: return src.encode() # NOTE: empty compiler is the default
|
def compile(self, src:str) -> bytes: return src.encode() # NOTE: empty compiler is the default
|
||||||
def compile_cached(self, src:str) -> bytes:
|
def compile_cached(self, src:str) -> bytes:
|
||||||
if self.cachekey is None or (lib := diskcache_get(self.cachekey, src)) is None:
|
if self.cachekey is None or (lib := diskcache_get(self.cachekey, src)) is None:
|
||||||
@@ -355,8 +355,9 @@ if PROFILE:
|
|||||||
|
|
||||||
with open(fn:=temp("profile.pkl", append_user=True), "wb") as f: pickle.dump(cpu_events+Compiled.profile_events+Buffer.profile_events, f)
|
with open(fn:=temp("profile.pkl", append_user=True), "wb") as f: pickle.dump(cpu_events+Compiled.profile_events+Buffer.profile_events, f)
|
||||||
|
|
||||||
from tinygrad.uop.ops import launch_viz
|
if VIZ:
|
||||||
launch_viz("PROFILE", fn)
|
from tinygrad.uop.ops import launch_viz
|
||||||
|
launch_viz("PROFILE", fn)
|
||||||
|
|
||||||
def enumerate_devices_str() -> Generator[str, None, None]:
|
def enumerate_devices_str() -> Generator[str, None, None]:
|
||||||
from tinygrad import Tensor, Device
|
from tinygrad import Tensor, Device
|
||||||
|
|||||||
@@ -3,7 +3,7 @@ import time, pprint, random, itertools, math
|
|||||||
from dataclasses import dataclass, replace, field
|
from dataclasses import dataclass, replace, field
|
||||||
from tinygrad.helpers import all_same, colored, DEBUG, GlobalCounters, ansilen, BEAM, NOOPT, all_int, CAPTURING, Metadata, TRACEMETA, TracingKey
|
from tinygrad.helpers import all_same, colored, DEBUG, GlobalCounters, ansilen, BEAM, NOOPT, all_int, CAPTURING, Metadata, TRACEMETA, TracingKey
|
||||||
from tinygrad.helpers import DEVECTORIZE, time_to_str, VALIDATE_WITH_CPU, getenv, cpu_profile, PROFILE, ProfilePointEvent, cpu_events, prod, Context
|
from tinygrad.helpers import DEVECTORIZE, time_to_str, VALIDATE_WITH_CPU, getenv, cpu_profile, PROFILE, ProfilePointEvent, cpu_events, prod, Context
|
||||||
from tinygrad.helpers import unwrap
|
from tinygrad.helpers import unwrap, disable_gc
|
||||||
from tinygrad.uop.ops import Ops, PatternMatcher, UOp, UPat, sym_infer, graph_rewrite, print_uops, track_rewrites, KernelInfo, pyrender
|
from tinygrad.uop.ops import Ops, PatternMatcher, UOp, UPat, sym_infer, graph_rewrite, print_uops, track_rewrites, KernelInfo, pyrender
|
||||||
from tinygrad.device import Device, Buffer
|
from tinygrad.device import Device, Buffer
|
||||||
from tinygrad.renderer import Renderer, ProgramSpec, Estimates
|
from tinygrad.renderer import Renderer, ProgramSpec, Estimates
|
||||||
@@ -13,6 +13,7 @@ from tinygrad.codegen.opt import Opt
|
|||||||
|
|
||||||
# **************** Program Creation ****************
|
# **************** Program Creation ****************
|
||||||
|
|
||||||
|
@disable_gc()
|
||||||
@track_rewrites(name=lambda *args,ret,**kwargs: TracingKey(ret.name, (ret.function_name, ret.ast), ret=ret), replay=True)
|
@track_rewrites(name=lambda *args,ret,**kwargs: TracingKey(ret.name, (ret.function_name, ret.ast), ret=ret), replay=True)
|
||||||
def get_program(ast:UOp, renderer:Renderer|None=None, opts:list[Opt]|None=None) -> ProgramSpec:
|
def get_program(ast:UOp, renderer:Renderer|None=None, opts:list[Opt]|None=None) -> ProgramSpec:
|
||||||
"""
|
"""
|
||||||
|
|||||||
+58
-25
@@ -1,5 +1,5 @@
|
|||||||
from typing import cast
|
from typing import cast
|
||||||
from dataclasses import dataclass, field
|
from dataclasses import dataclass, field, replace
|
||||||
from collections import deque, defaultdict
|
from collections import deque, defaultdict
|
||||||
from tinygrad.uop.ops import UOp, Ops, buffers
|
from tinygrad.uop.ops import UOp, Ops, buffers
|
||||||
from tinygrad.device import Device, Buffer, MultiBuffer
|
from tinygrad.device import Device, Buffer, MultiBuffer
|
||||||
@@ -13,6 +13,7 @@ class ScheduleItem:
|
|||||||
bufs: tuple[Buffer, ...]
|
bufs: tuple[Buffer, ...]
|
||||||
metadata: tuple[Metadata, ...] = ()
|
metadata: tuple[Metadata, ...] = ()
|
||||||
fixedvars: dict[str, int] = field(default_factory=dict)
|
fixedvars: dict[str, int] = field(default_factory=dict)
|
||||||
|
bound_ranges: tuple[UOp, ...] = ()
|
||||||
|
|
||||||
# **** schedule linearizer
|
# **** schedule linearizer
|
||||||
|
|
||||||
@@ -22,10 +23,13 @@ def create_schedule_with_vars(sched_sink:UOp) -> tuple[list[ScheduleItem], dict[
|
|||||||
in_degree: dict[UOp, int] = {}
|
in_degree: dict[UOp, int] = {}
|
||||||
var_vals: dict[str, int] = {}
|
var_vals: dict[str, int] = {}
|
||||||
for u in sched_sink.toposort():
|
for u in sched_sink.toposort():
|
||||||
if u.op is not Ops.AFTER: continue # anything that's not an ASSIGN doesn't write a kernel, so we can skip
|
if u.op is Ops.RANGE:
|
||||||
|
in_degree.setdefault(u, 0)
|
||||||
|
continue
|
||||||
|
if u.op is not Ops.AFTER or u.src[1].op is Ops.RANGE: continue
|
||||||
k = u.src[1]
|
k = u.src[1]
|
||||||
in_degree.setdefault(k, 0)
|
in_degree.setdefault(k, 0)
|
||||||
for s in k.src:
|
for s in k.src[0].src if k.op is Ops.END else k.src:
|
||||||
if s.op is Ops.AFTER:
|
if s.op is Ops.AFTER:
|
||||||
children[s.src[1]].append(k)
|
children[s.src[1]].append(k)
|
||||||
in_degree[k] += 1
|
in_degree[k] += 1
|
||||||
@@ -39,16 +43,19 @@ def create_schedule_with_vars(sched_sink:UOp) -> tuple[list[ScheduleItem], dict[
|
|||||||
elif s.op is Ops.BUFFER:
|
elif s.op is Ops.BUFFER:
|
||||||
pass # a BUFFER is already realized, nothing to do here
|
pass # a BUFFER is already realized, nothing to do here
|
||||||
elif s.op is Ops.BIND:
|
elif s.op is Ops.BIND:
|
||||||
var, val = s.unbind()
|
# for RANGE this is in fixedvars
|
||||||
assert var.expr not in var_vals or var_vals[var.expr] == val, f"bind mismatch on {var}, {var_vals[var.expr]} != {val}"
|
if s.src[1].op is not Ops.RANGE:
|
||||||
var_vals[var.expr] = val
|
var, val = s.unbind()
|
||||||
|
assert var.expr not in var_vals or var_vals[var.expr] == val, f"bind mismatch on {var}, {var_vals[var.expr]} != {val}"
|
||||||
|
var_vals[var.expr] = val
|
||||||
else:
|
else:
|
||||||
raise RuntimeError(f"input to kernel must be AFTER or BUFFER, not {s.op}")
|
raise RuntimeError(f"input to kernel must be AFTER or BUFFER, not {s.op}")
|
||||||
|
|
||||||
# linearize KERNEL UOps into ScheduleItems in BFS order
|
# linearize KERNEL UOps into ScheduleItems in BFS order
|
||||||
|
|
||||||
def _heuristic(k: UOp):
|
def _heuristic(k: UOp):
|
||||||
if k.arg.ast.op is Ops.COPY and not all_same([Device[cast(Buffer, s.buf_uop.buffer).device].group_id for s in k.src]): return 1000
|
if k.op is Ops.KERNEL and k.arg.ast.op is Ops.COPY and not all_same([Device[cast(Buffer, s.buf_uop.buffer).device].group_id for s in k.src]):
|
||||||
|
return 1000
|
||||||
return 0
|
return 0
|
||||||
|
|
||||||
last_heuristic: int = 0
|
last_heuristic: int = 0
|
||||||
@@ -57,27 +64,53 @@ def create_schedule_with_vars(sched_sink:UOp) -> tuple[list[ScheduleItem], dict[
|
|||||||
for k,v in in_degree.items():
|
for k,v in in_degree.items():
|
||||||
if v == 0: queues[_heuristic(k)].append(k)
|
if v == 0: queues[_heuristic(k)].append(k)
|
||||||
|
|
||||||
schedule: list[ScheduleItem] = []
|
schedule: list[ScheduleItem|UOp] = []
|
||||||
while last_queue or any(queues.values()):
|
while last_queue or any(queues.values()):
|
||||||
if not last_queue: last_heuristic, last_queue = min((it for it in queues.items() if it[1]), key=lambda x: abs(x[0]-last_heuristic))
|
if not last_queue: last_heuristic, last_queue = min((it for it in queues.items() if it[1]), key=lambda x: abs(x[0]-last_heuristic))
|
||||||
k = last_queue.popleft()
|
k = rk = last_queue.popleft()
|
||||||
ast = k.arg.ast
|
if k.op is Ops.END: k = k.src[0]
|
||||||
# create subbuffers if needed
|
if k.op is Ops.RANGE: schedule.append(k)
|
||||||
if ast.op is Ops.BUFFER_VIEW:
|
elif k.op is Ops.KERNEL:
|
||||||
base = k.src[1].buf_uop.buffer
|
ast = k.arg.ast
|
||||||
assert isinstance(base, Buffer), "base can't be MultiBuffer"
|
# create subbuffers if needed
|
||||||
buffers[k.src[0]] = base.view(k.size, ast.dtype, ast.arg[1]*base.dtype.itemsize)
|
if ast.op is Ops.BUFFER_VIEW:
|
||||||
ubufs = tuple(s.buf_uop.buffer for s in k.src if s.op is not Ops.BIND)
|
base = k.src[1].buf_uop.buffer
|
||||||
if any(isinstance(x, MultiBuffer) for x in ubufs):
|
assert isinstance(base, Buffer), "base can't be MultiBuffer"
|
||||||
assert all(isinstance(x, MultiBuffer) for x in ubufs), "kernel must all be multibuffer"
|
buffers[k.src[0]] = base.view(k.size, ast.dtype, ast.arg[1]*base.dtype.itemsize)
|
||||||
dnums = [x for x in ast.variables() if x.arg[0] == '_device_num']
|
ubufs = tuple(s.buf_uop.buffer for s in k.src if s.op is not Ops.BIND)
|
||||||
for i,bufs in enumerate(zip(*[x.bufs for x in cast(tuple[MultiBuffer, ...], ubufs)])):
|
bound_ranges = tuple(s for s in k.src if s.op is Ops.BIND and s.src[1].op is Ops.RANGE)
|
||||||
schedule.append(ScheduleItem(ast, bufs, k.arg.metadata, {dnums[0].expr:i} if len(dnums) else {}))
|
if any(isinstance(x, MultiBuffer) for x in ubufs):
|
||||||
|
assert all(isinstance(x, MultiBuffer) for x in ubufs), "kernel must all be multibuffer"
|
||||||
|
dnums = [x for x in ast.variables() if x.arg[0] == '_device_num']
|
||||||
|
for i,bufs in enumerate(zip(*[x.bufs for x in cast(tuple[MultiBuffer, ...], ubufs)])):
|
||||||
|
schedule.append(ScheduleItem(ast, bufs, k.arg.metadata, {dnums[0].expr:i} if len(dnums) else {}, bound_ranges=bound_ranges))
|
||||||
|
else:
|
||||||
|
# ONE -> ONE
|
||||||
|
schedule.append(ScheduleItem(ast, cast(tuple[Buffer, ...], ubufs), k.arg.metadata, bound_ranges=bound_ranges))
|
||||||
|
if rk.op is Ops.END: schedule.append(rk)
|
||||||
else:
|
else:
|
||||||
# ONE -> ONE
|
raise RuntimeError(f"can't schedule {k.op}")
|
||||||
schedule.append(ScheduleItem(ast, cast(tuple[Buffer, ...], ubufs), k.arg.metadata))
|
for x in children[rk]:
|
||||||
for x in children[k]:
|
|
||||||
in_degree[x] -= 1
|
in_degree[x] -= 1
|
||||||
if in_degree[x] == 0: queues[_heuristic(x)].append(x)
|
if in_degree[x] == 0: queues[_heuristic(x)].append(x)
|
||||||
|
|
||||||
return schedule, var_vals
|
# expand the ranges in the schedule
|
||||||
|
real_schedule: list[ScheduleItem] = []
|
||||||
|
sched_ptr = 0
|
||||||
|
in_ranges = {}
|
||||||
|
range_ptrs = {}
|
||||||
|
while sched_ptr < len(schedule):
|
||||||
|
si = schedule[sched_ptr]
|
||||||
|
if isinstance(si, UOp):
|
||||||
|
if si.op is Ops.RANGE:
|
||||||
|
in_ranges[si] = 0
|
||||||
|
range_ptrs[si] = sched_ptr + 1
|
||||||
|
elif si.op is Ops.END:
|
||||||
|
if in_ranges[si.src[1]] < si.src[1].vmax:
|
||||||
|
in_ranges[si.src[1]] += 1
|
||||||
|
sched_ptr = range_ptrs[si.src[1]]
|
||||||
|
continue
|
||||||
|
else:
|
||||||
|
real_schedule.append(replace(si, fixedvars=si.fixedvars | {s.src[0].arg[0]:in_ranges[s.src[1]] for s in si.bound_ranges}, bound_ranges=()))
|
||||||
|
sched_ptr += 1
|
||||||
|
return real_schedule, var_vals
|
||||||
|
|||||||
+11
-5
@@ -3,14 +3,15 @@ import math, dataclasses
|
|||||||
from tinygrad.uop.ops import UOp, PatternMatcher, UPat, Ops, all_metadata
|
from tinygrad.uop.ops import UOp, PatternMatcher, UPat, Ops, all_metadata
|
||||||
from tinygrad.helpers import argsort
|
from tinygrad.helpers import argsort
|
||||||
|
|
||||||
def reduce_gradient(ctx:UOp, ret:UOp):
|
def reduce_gradient(ctx:UOp, ret:UOp, op:Ops):
|
||||||
def broadcast_to_input(x): return x.reshape(x.shape+(1,)*(len(ret.src[0].shape)-len(x.shape))).expand(ret.src[0].shape)
|
def broadcast_to_input(x): return x.reshape(x.shape+(1,)*(len(ret.src[0].shape)-len(x.shape))).expand(ret.src[0].shape)
|
||||||
if ret.arg[0] == Ops.ADD: return (broadcast_to_input(ctx),)
|
if op == Ops.ADD: return (broadcast_to_input(ctx),)
|
||||||
if ret.arg[0] == Ops.MAX:
|
if op == Ops.MAX:
|
||||||
|
assert ret.op is Ops.REDUCE_AXIS, "only works on REDUCE_AXIS"
|
||||||
mask = ret.src[0].eq(broadcast_to_input(ret)).cast(ctx.dtype)
|
mask = ret.src[0].eq(broadcast_to_input(ret)).cast(ctx.dtype)
|
||||||
count = mask.r(Ops.ADD, ret.arg[1])
|
count = mask.r(Ops.ADD, ret.arg[1])
|
||||||
return ((mask/broadcast_to_input(count)) * broadcast_to_input(ctx),)
|
return ((mask/broadcast_to_input(count)) * broadcast_to_input(ctx),)
|
||||||
if ret.arg[0] == Ops.MUL: return (broadcast_to_input(ctx * ret) / ret.src[0],)
|
if op == Ops.MUL: return (broadcast_to_input(ctx * ret) / ret.src[0],)
|
||||||
|
|
||||||
# ctx is grad_output
|
# ctx is grad_output
|
||||||
pm_gradient = PatternMatcher([
|
pm_gradient = PatternMatcher([
|
||||||
@@ -28,7 +29,8 @@ pm_gradient = PatternMatcher([
|
|||||||
((x>y).where(ctx, (x.eq(y)).where(ctx * 0.5, 0)), (x<y).where(ctx, (x.eq(y)).where(ctx * 0.5, 0)))),
|
((x>y).where(ctx, (x.eq(y)).where(ctx * 0.5, 0)), (x<y).where(ctx, (x.eq(y)).where(ctx * 0.5, 0)))),
|
||||||
(UPat(Ops.MUL, name="ret"), lambda ctx, ret: (ret.src[1]*ctx, ret.src[0]*ctx)),
|
(UPat(Ops.MUL, name="ret"), lambda ctx, ret: (ret.src[1]*ctx, ret.src[0]*ctx)),
|
||||||
(UPat(Ops.WHERE, name="ret"), lambda ctx, ret: (None, ret.src[0].where(ctx, ctx.const_like(0)), ret.src[0].where(ctx.const_like(0), ctx))),
|
(UPat(Ops.WHERE, name="ret"), lambda ctx, ret: (None, ret.src[0].where(ctx, ctx.const_like(0)), ret.src[0].where(ctx.const_like(0), ctx))),
|
||||||
(UPat(Ops.REDUCE_AXIS, name="ret"), reduce_gradient),
|
(UPat(Ops.REDUCE_AXIS, name="ret"), lambda ctx, ret: reduce_gradient(ctx, ret, ret.arg[0])),
|
||||||
|
(UPat(Ops.REDUCE, name="ret"), lambda ctx, ret: reduce_gradient(ctx, ret, ret.arg) + (None,)*(len(ret.src)-1)),
|
||||||
(UPat(Ops.CONTIGUOUS), lambda ctx: (ctx,)),
|
(UPat(Ops.CONTIGUOUS), lambda ctx: (ctx,)),
|
||||||
(UPat(Ops.CONTIGUOUS_BACKWARD), lambda ctx: (ctx.contiguous(),)),
|
(UPat(Ops.CONTIGUOUS_BACKWARD), lambda ctx: (ctx.contiguous(),)),
|
||||||
(UPat(Ops.RESHAPE, name="ret"), lambda ctx, ret: (ctx.reshape(ret.src[0].shape), None)),
|
(UPat(Ops.RESHAPE, name="ret"), lambda ctx, ret: (ctx.reshape(ret.src[0].shape), None)),
|
||||||
@@ -68,4 +70,8 @@ def compute_gradient(root:UOp, root_grad:UOp, targets:set[UOp]) -> dict[UOp, UOp
|
|||||||
# we add the backward metadata to everything new in the graph
|
# we add the backward metadata to everything new in the graph
|
||||||
for bw_uop in v.toposort(lambda x: x not in (t0, *t0.src, grads[t0])):
|
for bw_uop in v.toposort(lambda x: x not in (t0, *t0.src, grads[t0])):
|
||||||
all_metadata[bw_uop] = all_metadata.get(bw_uop, ())+backward_metadata
|
all_metadata[bw_uop] = all_metadata.get(bw_uop, ())+backward_metadata
|
||||||
|
# end any ranges on grads with a reduce sum
|
||||||
|
for k,v in grads.items():
|
||||||
|
if len(v.ranges):
|
||||||
|
grads[k] = v.reduce(*v.ranges, arg=Ops.ADD)
|
||||||
return grads
|
return grads
|
||||||
|
|||||||
+36
-5
@@ -1,5 +1,5 @@
|
|||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
import os, functools, platform, time, re, contextlib, operator, hashlib, pickle, sqlite3, tempfile, pathlib, string, ctypes, sys, gzip, getpass
|
import os, functools, platform, time, re, contextlib, operator, hashlib, pickle, sqlite3, tempfile, pathlib, string, ctypes, sys, gzip, getpass, gc
|
||||||
import urllib.request, subprocess, shutil, math, types, copyreg, inspect, importlib, decimal, itertools
|
import urllib.request, subprocess, shutil, math, types, copyreg, inspect, importlib, decimal, itertools
|
||||||
from dataclasses import dataclass, field
|
from dataclasses import dataclass, field
|
||||||
from typing import ClassVar, Iterable, Any, TypeVar, Callable, Sequence, TypeGuard, Iterator, Generic, Generator, cast, overload
|
from typing import ClassVar, Iterable, Any, TypeVar, Callable, Sequence, TypeGuard, Iterator, Generic, Generator, cast, overload
|
||||||
@@ -173,14 +173,15 @@ TRANSCENDENTAL, NOLOCALS = ContextVar("TRANSCENDENTAL", 1), ContextVar("NOLOCALS
|
|||||||
SPLIT_REDUCEOP, NO_MEMORY_PLANNER, RING = ContextVar("SPLIT_REDUCEOP", 1), ContextVar("NO_MEMORY_PLANNER", 0), ContextVar("RING", 1)
|
SPLIT_REDUCEOP, NO_MEMORY_PLANNER, RING = ContextVar("SPLIT_REDUCEOP", 1), ContextVar("NO_MEMORY_PLANNER", 0), ContextVar("RING", 1)
|
||||||
PICKLE_BUFFERS, LRU = ContextVar("PICKLE_BUFFERS", 1), ContextVar("LRU", 1)
|
PICKLE_BUFFERS, LRU = ContextVar("PICKLE_BUFFERS", 1), ContextVar("LRU", 1)
|
||||||
CACHELEVEL, IGNORE_BEAM_CACHE, DEVECTORIZE = ContextVar("CACHELEVEL", 2), ContextVar("IGNORE_BEAM_CACHE", 0), ContextVar("DEVECTORIZE", 1)
|
CACHELEVEL, IGNORE_BEAM_CACHE, DEVECTORIZE = ContextVar("CACHELEVEL", 2), ContextVar("IGNORE_BEAM_CACHE", 0), ContextVar("DEVECTORIZE", 1)
|
||||||
DISABLE_COMPILER_CACHE = ContextVar("DISABLE_COMPILER_CACHE", 0)
|
|
||||||
VALIDATE_WITH_CPU, DISABLE_FAST_IDIV = ContextVar("VALIDATE_WITH_CPU", 0), ContextVar("DISABLE_FAST_IDIV", 0)
|
VALIDATE_WITH_CPU, DISABLE_FAST_IDIV = ContextVar("VALIDATE_WITH_CPU", 0), ContextVar("DISABLE_FAST_IDIV", 0)
|
||||||
CORRECT_DIVMOD_FOLDING, FUSE_OPTIM = ContextVar("CORRECT_DIVMOD_FOLDING", 0), ContextVar("FUSE_OPTIM", 0)
|
CORRECT_DIVMOD_FOLDING, FUSE_OPTIM = ContextVar("CORRECT_DIVMOD_FOLDING", 0), ContextVar("FUSE_OPTIM", 0)
|
||||||
ALLOW_DEVICE_USAGE, MAX_BUFFER_SIZE = ContextVar("ALLOW_DEVICE_USAGE", 1), ContextVar("MAX_BUFFER_SIZE", 0)
|
ALLOW_DEVICE_USAGE, MAX_BUFFER_SIZE = ContextVar("ALLOW_DEVICE_USAGE", 1), ContextVar("MAX_BUFFER_SIZE", 0)
|
||||||
EMULATE = ContextVar("EMULATE", "")
|
EMULATE = ContextVar("EMULATE", "")
|
||||||
CPU_COUNT = ContextVar("CPU_COUNT", max(1, len(os.sched_getaffinity(0)) if hasattr(os, "sched_getaffinity") else (os.cpu_count() or 1)))
|
CPU_COUNT = ContextVar("CPU_COUNT", max(1, len(os.sched_getaffinity(0)) if hasattr(os, "sched_getaffinity") else (os.cpu_count() or 1)))
|
||||||
CPU_LLVM, CPU_LVP, AMD_LLVM = ContextVar("CPU_LLVM", 0), ContextVar("CPU_LVP", 0), ContextVar("AMD_LLVM", 1)
|
CPU_LLVM, CPU_LVP, AMD_LLVM = ContextVar("CPU_LLVM", 0), ContextVar("CPU_LVP", 0), ContextVar("AMD_LLVM", 0)
|
||||||
VIZ = PROFILE = ContextVar("VIZ", 0)
|
# VIZ implies PROFILE, but you can run PROFILE without VIZ
|
||||||
|
VIZ = ContextVar("VIZ", 0)
|
||||||
|
PROFILE = ContextVar("PROFILE", VIZ.value)
|
||||||
SPEC = ContextVar("SPEC", 1)
|
SPEC = ContextVar("SPEC", 1)
|
||||||
# TODO: disable by default due to speed
|
# TODO: disable by default due to speed
|
||||||
IGNORE_OOB = ContextVar("IGNORE_OOB", 1)
|
IGNORE_OOB = ContextVar("IGNORE_OOB", 1)
|
||||||
@@ -188,6 +189,8 @@ PCONTIG = ContextVar("PCONTIG", 0) # partial contiguous in rangeify
|
|||||||
DEBUG_RANGEIFY = ContextVar("DEBUG_RANGEIFY", 0)
|
DEBUG_RANGEIFY = ContextVar("DEBUG_RANGEIFY", 0)
|
||||||
# set to 1, this uses tuplize in the linearizer sort order
|
# set to 1, this uses tuplize in the linearizer sort order
|
||||||
TUPLE_ORDER = ContextVar("TUPLE_ORDER", 1)
|
TUPLE_ORDER = ContextVar("TUPLE_ORDER", 1)
|
||||||
|
# set to 0 to disable the compiler cache
|
||||||
|
CCACHE = ContextVar("CCACHE", 1)
|
||||||
|
|
||||||
@dataclass(frozen=True)
|
@dataclass(frozen=True)
|
||||||
class Metadata:
|
class Metadata:
|
||||||
@@ -240,11 +243,29 @@ class Profiling(contextlib.ContextDecorator):
|
|||||||
|
|
||||||
def perf_counter_us() -> decimal.Decimal: return decimal.Decimal(time.perf_counter_ns())/1000
|
def perf_counter_us() -> decimal.Decimal: return decimal.Decimal(time.perf_counter_ns())/1000
|
||||||
|
|
||||||
|
@functools.cache
|
||||||
|
def lines(fn) -> list[str]:
|
||||||
|
try:
|
||||||
|
with open(fn, encoding="utf-8") as f: return f.readlines()
|
||||||
|
except (FileNotFoundError, OSError): return []
|
||||||
|
|
||||||
|
def printable(loc:tuple[str, int]) -> str:
|
||||||
|
try: return lines(loc[0])[loc[1]-1].strip()
|
||||||
|
except IndexError: return "<missing>"
|
||||||
|
|
||||||
|
def get_stacktrace(frm, max_frames=30) -> tuple[tuple, ...]:
|
||||||
|
ret:list[tuple] = []
|
||||||
|
for i in range(max_frames):
|
||||||
|
if (frm:=frm.f_back) is None: break
|
||||||
|
ret.append(((fc:=frm.f_code).co_filename, frm.f_lineno, fc.co_name, printable((fc.co_filename, frm.f_lineno))))
|
||||||
|
return tuple(ret)
|
||||||
|
|
||||||
@dataclass(frozen=True)
|
@dataclass(frozen=True)
|
||||||
class TracingKey:
|
class TracingKey:
|
||||||
display_name:str # display name of this trace event
|
display_name:str # display name of this trace event
|
||||||
keys:tuple[Any, ...]=() # optional keys to search for related traces
|
keys:tuple[Any, ...]=() # optional keys to search for related traces
|
||||||
ret:Any=None
|
ret:Any=None
|
||||||
|
tb:tuple[tuple, ...]|None=field(default_factory=lambda: get_stacktrace(sys._getframe(1)) if VIZ else None)
|
||||||
|
|
||||||
class ProfileEvent: pass
|
class ProfileEvent: pass
|
||||||
|
|
||||||
@@ -361,10 +382,12 @@ 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 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:
|
||||||
pathlib.Path(f.name).write_bytes(lib)
|
pathlib.Path(f.name).write_bytes(lib)
|
||||||
print(subprocess.check_output([objdump_tool, '-d', f.name]).decode('utf-8'))
|
print(system(f"{objdump_tool} -d {f.name}"))
|
||||||
|
|
||||||
def capstone_flatdump(lib: bytes):
|
def capstone_flatdump(lib: bytes):
|
||||||
try: import capstone
|
try: import capstone
|
||||||
@@ -395,6 +418,7 @@ def to_mv(ptr:int, sz:int) -> memoryview: return memoryview((ctypes.c_uint8 * sz
|
|||||||
def mv_address(mv): return ctypes.addressof(ctypes.c_char.from_buffer(mv))
|
def mv_address(mv): return ctypes.addressof(ctypes.c_char.from_buffer(mv))
|
||||||
def to_char_p_p(options: list[bytes], to_type=ctypes.c_char):
|
def to_char_p_p(options: list[bytes], to_type=ctypes.c_char):
|
||||||
return (ctypes.POINTER(to_type) * len(options))(*[ctypes.cast(ctypes.create_string_buffer(o), ctypes.POINTER(to_type)) for o in options])
|
return (ctypes.POINTER(to_type) * len(options))(*[ctypes.cast(ctypes.create_string_buffer(o), ctypes.POINTER(to_type)) for o in options])
|
||||||
|
def charptr(s:str|bytes): return ctypes.cast(ctypes.c_char_p(s if isinstance(s, bytes) else s.encode()), ctypes.POINTER(ctypes.c_char))
|
||||||
@functools.cache
|
@functools.cache
|
||||||
def init_c_struct_t(fields: tuple[tuple[str, type[ctypes._SimpleCData]], ...]):
|
def init_c_struct_t(fields: tuple[tuple[str, type[ctypes._SimpleCData]], ...]):
|
||||||
class CStruct(ctypes.Structure):
|
class CStruct(ctypes.Structure):
|
||||||
@@ -442,6 +466,13 @@ class tqdm(Generic[T]):
|
|||||||
class trange(tqdm):
|
class trange(tqdm):
|
||||||
def __init__(self, n:int, **kwargs): super().__init__(iterable=range(n), total=n, **kwargs)
|
def __init__(self, n:int, **kwargs): super().__init__(iterable=range(n), total=n, **kwargs)
|
||||||
|
|
||||||
|
class disable_gc(contextlib.ContextDecorator):
|
||||||
|
def __enter__(self):
|
||||||
|
self._was_enabled = gc.isenabled()
|
||||||
|
if self._was_enabled: gc.disable()
|
||||||
|
def __exit__(self, *exc):
|
||||||
|
if self._was_enabled: gc.enable()
|
||||||
|
|
||||||
# *** universal support for code object pickling
|
# *** universal support for code object pickling
|
||||||
|
|
||||||
def _reconstruct_code(*args): return types.CodeType(*args)
|
def _reconstruct_code(*args): return types.CodeType(*args)
|
||||||
|
|||||||
@@ -1,4 +1,6 @@
|
|||||||
from tinygrad.mixin.math import MathMixin
|
from tinygrad.mixin.math import MathMixin
|
||||||
from tinygrad.mixin.movement import MovementMixin
|
from tinygrad.mixin.movement import MovementMixin
|
||||||
|
|
||||||
class OpMixin(MathMixin, MovementMixin): pass
|
|
||||||
|
class OpMixin(MathMixin, MovementMixin):
|
||||||
|
pass
|
||||||
|
|||||||
+173
-66
@@ -2,24 +2,38 @@ from typing import Self
|
|||||||
from tinygrad.uop import Ops
|
from tinygrad.uop import Ops
|
||||||
from tinygrad.dtype import dtypes, ConstType
|
from tinygrad.dtype import dtypes, ConstType
|
||||||
|
|
||||||
|
|
||||||
class MathMixin:
|
class MathMixin:
|
||||||
# required to implement
|
# required to implement
|
||||||
def alu(self, op:Ops, *src:Self) -> Self: raise NotImplementedError
|
def alu(self, op: Ops, *src: Self) -> Self:
|
||||||
def const_like(self, b:ConstType) -> Self: raise NotImplementedError
|
raise NotImplementedError
|
||||||
|
|
||||||
|
def const_like(self, b: ConstType) -> Self:
|
||||||
|
raise NotImplementedError
|
||||||
|
|
||||||
# great functions you get!
|
# great functions you get!
|
||||||
def ufix(self, x:Self|ConstType) -> Self: return self.const_like(x) if not isinstance(x, MathMixin) else x
|
def ufix(self, x: Self | ConstType) -> Self:
|
||||||
def _binop(self, op:Ops, x:Self|ConstType, reverse:bool) -> Self:
|
return self.const_like(x) if not isinstance(x, MathMixin) else x
|
||||||
|
|
||||||
|
def _binop(self, op: Ops, x: Self | ConstType, reverse: bool) -> Self:
|
||||||
return self.ufix(x).alu(op, self) if reverse else self.alu(op, self.ufix(x))
|
return self.ufix(x).alu(op, self) if reverse else self.alu(op, self.ufix(x))
|
||||||
def logical_not(self): return self.ne(True)
|
|
||||||
|
def logical_not(self):
|
||||||
|
return self.ne(True)
|
||||||
|
|
||||||
def neg(self):
|
def neg(self):
|
||||||
if (dtype:=getattr(self, 'dtype')) is None: raise TypeError(f"MathTraits __neg__ requires a dtype, {self=}")
|
if (dtype := getattr(self, "dtype")) is None:
|
||||||
return self.logical_not() if dtype.scalar() == dtypes.bool else self*(-1)
|
raise TypeError(f"MathTraits __neg__ requires a dtype, {self=}")
|
||||||
|
return self.logical_not() if dtype.scalar() == dtypes.bool else self * (-1)
|
||||||
|
|
||||||
def _check_dtype(self):
|
def _check_dtype(self):
|
||||||
if (dtype:=getattr(self, 'dtype')) is not None:
|
if (dtype := getattr(self, "dtype")) is not None:
|
||||||
if isinstance(dtype, tuple): dtype = dtype[0]
|
if isinstance(dtype, tuple):
|
||||||
if not (dtypes.is_bool(dtype) or dtypes.is_int(dtype)): raise RuntimeError(f"{dtype} is not supported")
|
dtype = dtype[0]
|
||||||
def add(self, x:Self|ConstType, reverse:bool=False):
|
if not (dtypes.is_bool(dtype) or dtypes.is_int(dtype)):
|
||||||
|
raise RuntimeError(f"{dtype} is not supported")
|
||||||
|
|
||||||
|
def add(self, x: Self | ConstType, reverse: bool = False):
|
||||||
"""
|
"""
|
||||||
Adds `self` and `x`.
|
Adds `self` and `x`.
|
||||||
Equivalent to `self + x`.
|
Equivalent to `self + x`.
|
||||||
@@ -37,7 +51,8 @@ class MathMixin:
|
|||||||
```
|
```
|
||||||
"""
|
"""
|
||||||
return self._binop(Ops.ADD, x, reverse)
|
return self._binop(Ops.ADD, x, reverse)
|
||||||
def mul(self, x:Self|ConstType, reverse:bool=False):
|
|
||||||
|
def mul(self, x: Self | ConstType, reverse: bool = False):
|
||||||
"""
|
"""
|
||||||
Multiplies `self` and `x`.
|
Multiplies `self` and `x`.
|
||||||
Equivalent to `self * x`.
|
Equivalent to `self * x`.
|
||||||
@@ -56,7 +71,8 @@ class MathMixin:
|
|||||||
```
|
```
|
||||||
"""
|
"""
|
||||||
return self._binop(Ops.MUL, x, reverse)
|
return self._binop(Ops.MUL, x, reverse)
|
||||||
def bitwise_and(self, x:Self|ConstType, reverse:bool=False):
|
|
||||||
|
def bitwise_and(self, x: Self | ConstType, reverse: bool = False):
|
||||||
"""
|
"""
|
||||||
Computes the bitwise AND of `self` and `x`.
|
Computes the bitwise AND of `self` and `x`.
|
||||||
Equivalent to `self & x`.
|
Equivalent to `self & x`.
|
||||||
@@ -70,7 +86,8 @@ class MathMixin:
|
|||||||
"""
|
"""
|
||||||
self._check_dtype()
|
self._check_dtype()
|
||||||
return self._binop(Ops.AND, x, reverse)
|
return self._binop(Ops.AND, x, reverse)
|
||||||
def bitwise_or(self, x:Self|ConstType, reverse:bool=False):
|
|
||||||
|
def bitwise_or(self, x: Self | ConstType, reverse: bool = False):
|
||||||
"""
|
"""
|
||||||
Computes the bitwise OR of `self` and `x`.
|
Computes the bitwise OR of `self` and `x`.
|
||||||
Equivalent to `self | x`.
|
Equivalent to `self | x`.
|
||||||
@@ -84,7 +101,8 @@ class MathMixin:
|
|||||||
"""
|
"""
|
||||||
self._check_dtype()
|
self._check_dtype()
|
||||||
return self._binop(Ops.OR, x, reverse)
|
return self._binop(Ops.OR, x, reverse)
|
||||||
def bitwise_xor(self, x:Self|ConstType, reverse:bool=False):
|
|
||||||
|
def bitwise_xor(self, x: Self | ConstType, reverse: bool = False):
|
||||||
"""
|
"""
|
||||||
Computes bitwise xor of `self` and `x`.
|
Computes bitwise xor of `self` and `x`.
|
||||||
Equivalent to `self ^ x`.
|
Equivalent to `self ^ x`.
|
||||||
@@ -99,7 +117,8 @@ class MathMixin:
|
|||||||
"""
|
"""
|
||||||
self._check_dtype()
|
self._check_dtype()
|
||||||
return self._binop(Ops.XOR, x, reverse)
|
return self._binop(Ops.XOR, x, reverse)
|
||||||
def idiv(self, x:Self|ConstType, reverse:bool=False):
|
|
||||||
|
def idiv(self, x: Self | ConstType, reverse: bool = False):
|
||||||
"""
|
"""
|
||||||
Divides `self` by `x`.
|
Divides `self` by `x`.
|
||||||
Equivalent to `self // x`.
|
Equivalent to `self // x`.
|
||||||
@@ -111,62 +130,150 @@ class MathMixin:
|
|||||||
```
|
```
|
||||||
"""
|
"""
|
||||||
return self._binop(Ops.IDIV, x, reverse)
|
return self._binop(Ops.IDIV, x, reverse)
|
||||||
def mod(self, x:Self|ConstType, reverse:bool=False): return self._binop(Ops.MOD, x, reverse)
|
|
||||||
def sub(self, x:Self|ConstType, reverse:bool=False): return self.ufix(x).alu(Ops.ADD, -self) if reverse else self.alu(Ops.ADD, self.ufix(-x))
|
|
||||||
def div(self, x:Self|ConstType, reverse:bool=False):
|
|
||||||
return (self.ufix(x)*self.alu(Ops.RECIPROCAL)) if reverse else (self*self.ufix(x).alu(Ops.RECIPROCAL))
|
|
||||||
|
|
||||||
def __neg__(self): return self.neg()
|
def mod(self, x: Self | ConstType, reverse: bool = False):
|
||||||
|
return self._binop(Ops.MOD, x, reverse)
|
||||||
|
|
||||||
def __add__(self, x:Self|ConstType): return self.add(x)
|
def sub(self, x: Self | ConstType, reverse: bool = False):
|
||||||
def __sub__(self, x:Self|ConstType): return self.sub(x)
|
return self.ufix(x).alu(Ops.ADD, -self) if reverse else self.alu(Ops.ADD, self.ufix(-x))
|
||||||
def __mul__(self, x:Self|ConstType): return self.mul(x)
|
|
||||||
def __truediv__(self, x:Self|ConstType): return self.div(x)
|
|
||||||
def __floordiv__(self, x:Self|ConstType): return self.idiv(x) # TODO: idiv is trunc div, not floordiv
|
|
||||||
def __mod__(self, x:Self|ConstType): return self.mod(x)
|
|
||||||
def __and__(self, x:Self|ConstType): return self.bitwise_and(x)
|
|
||||||
def __or__(self, x:Self|ConstType): return self.bitwise_or(x)
|
|
||||||
def __xor__(self, x:Self|ConstType): return self.bitwise_xor(x)
|
|
||||||
|
|
||||||
def __radd__(self, x:Self|ConstType): return self.add(x, True)
|
def div(self, x: Self | ConstType, reverse: bool = False):
|
||||||
def __rsub__(self, x:Self|ConstType): return self.sub(x, True)
|
return (self.ufix(x) * self.alu(Ops.RECIPROCAL)) if reverse else (self * self.ufix(x).alu(Ops.RECIPROCAL))
|
||||||
def __rmul__(self, x:Self|ConstType): return self.mul(x, True)
|
|
||||||
def __rtruediv__(self, x:Self|ConstType): return self.div(x, True)
|
|
||||||
def __rfloordiv__(self, x:Self|ConstType): return self.idiv(x, True)
|
|
||||||
def __rand__(self, x:Self|ConstType): return self.bitwise_and(x, True)
|
|
||||||
def __ror__(self, x:Self|ConstType): return self.bitwise_or(x, True)
|
|
||||||
def __rxor__(self, x:Self|ConstType): return self.bitwise_xor(x, True)
|
|
||||||
def __rmod__(self, x:Self|ConstType): return self.mod(x, True)
|
|
||||||
|
|
||||||
def __lt__(self, x:Self|ConstType): return self.alu(Ops.CMPLT, self.ufix(x))
|
def __neg__(self):
|
||||||
def __gt__(self, x:Self|ConstType): return self.ufix(x).alu(Ops.CMPLT, self)
|
return self.neg()
|
||||||
def __ge__(self, x:Self|ConstType): return (self < x).logical_not()
|
|
||||||
def __le__(self, x:Self|ConstType): return (self > x).logical_not()
|
def __add__(self, x: Self | ConstType):
|
||||||
|
return self.add(x)
|
||||||
|
|
||||||
|
def __sub__(self, x: Self | ConstType):
|
||||||
|
return self.sub(x)
|
||||||
|
|
||||||
|
def __mul__(self, x: Self | ConstType):
|
||||||
|
return self.mul(x)
|
||||||
|
|
||||||
|
def __truediv__(self, x: Self | ConstType):
|
||||||
|
return self.div(x)
|
||||||
|
|
||||||
|
def __floordiv__(self, x: Self | ConstType):
|
||||||
|
return self.idiv(x) # TODO: idiv is trunc div, not floordiv
|
||||||
|
|
||||||
|
def __mod__(self, x: Self | ConstType):
|
||||||
|
return self.mod(x)
|
||||||
|
|
||||||
|
def __and__(self, x: Self | ConstType):
|
||||||
|
return self.bitwise_and(x)
|
||||||
|
|
||||||
|
def __or__(self, x: Self | ConstType):
|
||||||
|
return self.bitwise_or(x)
|
||||||
|
|
||||||
|
def __xor__(self, x: Self | ConstType):
|
||||||
|
return self.bitwise_xor(x)
|
||||||
|
|
||||||
|
def __radd__(self, x: Self | ConstType):
|
||||||
|
return self.add(x, True)
|
||||||
|
|
||||||
|
def __rsub__(self, x: Self | ConstType):
|
||||||
|
return self.sub(x, True)
|
||||||
|
|
||||||
|
def __rmul__(self, x: Self | ConstType):
|
||||||
|
return self.mul(x, True)
|
||||||
|
|
||||||
|
def __rtruediv__(self, x: Self | ConstType):
|
||||||
|
return self.div(x, True)
|
||||||
|
|
||||||
|
def __rfloordiv__(self, x: Self | ConstType):
|
||||||
|
return self.idiv(x, True)
|
||||||
|
|
||||||
|
def __rand__(self, x: Self | ConstType):
|
||||||
|
return self.bitwise_and(x, True)
|
||||||
|
|
||||||
|
def __ror__(self, x: Self | ConstType):
|
||||||
|
return self.bitwise_or(x, True)
|
||||||
|
|
||||||
|
def __rxor__(self, x: Self | ConstType):
|
||||||
|
return self.bitwise_xor(x, True)
|
||||||
|
|
||||||
|
def __rmod__(self, x: Self | ConstType):
|
||||||
|
return self.mod(x, True)
|
||||||
|
|
||||||
|
def __lt__(self, x: Self | ConstType):
|
||||||
|
return self.alu(Ops.CMPLT, self.ufix(x))
|
||||||
|
|
||||||
|
def __gt__(self, x: Self | ConstType):
|
||||||
|
return self.ufix(x).alu(Ops.CMPLT, self)
|
||||||
|
|
||||||
|
def __ge__(self, x: Self | ConstType):
|
||||||
|
return (self < x).logical_not()
|
||||||
|
|
||||||
|
def __le__(self, x: Self | ConstType):
|
||||||
|
return (self > x).logical_not()
|
||||||
|
|
||||||
|
def ne(self, x: Self | ConstType):
|
||||||
|
return self.alu(Ops.CMPNE, self.ufix(x))
|
||||||
|
|
||||||
|
def eq(self, x: Self | ConstType):
|
||||||
|
return self.ne(x).logical_not()
|
||||||
|
|
||||||
|
def __ne__(self, x: Self | ConstType): # type: ignore[override]
|
||||||
|
return self.ne(x)
|
||||||
|
|
||||||
def ne(self, x:Self|ConstType): return self.alu(Ops.CMPNE, self.ufix(x))
|
|
||||||
def eq(self, x:Self|ConstType): return self.ne(x).logical_not()
|
|
||||||
def __ne__(self, x:Self|ConstType): return self.ne(x) # type: ignore[override]
|
|
||||||
# NOTE: __eq__ isn't overridden, and means the same thing as is by default
|
# NOTE: __eq__ isn't overridden, and means the same thing as is by default
|
||||||
|
|
||||||
def lshift(self, x:Self|int, reverse:bool=False): return self._binop(Ops.SHL, x, reverse)
|
def lshift(self, x: Self | int, reverse: bool = False):
|
||||||
def rshift(self, x:Self|int, reverse:bool=False): return self._binop(Ops.SHR, x, reverse)
|
return self._binop(Ops.SHL, x, reverse)
|
||||||
def __lshift__(self, x:Self|int): return self.lshift(x)
|
|
||||||
def __rshift__(self, x:Self|int): return self.rshift(x)
|
|
||||||
def __rlshift__(self, x:Self|int): return self.lshift(x, True)
|
|
||||||
def __rrshift__(self, x:Self|int): return self.rshift(x, True)
|
|
||||||
|
|
||||||
def maximum(self, x:Self|ConstType): return self.alu(Ops.MAX, self.ufix(x))
|
def rshift(self, x: Self | int, reverse: bool = False):
|
||||||
def minimum(self, x:Self|ConstType): return -(-self).maximum(-x)
|
return self._binop(Ops.SHR, x, reverse)
|
||||||
def where(self, x:Self|ConstType, y:Self|ConstType):
|
|
||||||
if isinstance(x, type(self)): return self.alu(Ops.WHERE, x, x.ufix(y))
|
def __lshift__(self, x: Self | int):
|
||||||
if isinstance(y, type(self)): return self.alu(Ops.WHERE, y.ufix(x), y)
|
return self.lshift(x)
|
||||||
|
|
||||||
|
def __rshift__(self, x: Self | int):
|
||||||
|
return self.rshift(x)
|
||||||
|
|
||||||
|
def __rlshift__(self, x: Self | int):
|
||||||
|
return self.lshift(x, True)
|
||||||
|
|
||||||
|
def __rrshift__(self, x: Self | int):
|
||||||
|
return self.rshift(x, True)
|
||||||
|
|
||||||
|
def maximum(self, x: Self | ConstType):
|
||||||
|
return self.alu(Ops.MAX, self.ufix(x))
|
||||||
|
|
||||||
|
def minimum(self, x: Self | ConstType):
|
||||||
|
return -(-self).maximum(-x)
|
||||||
|
|
||||||
|
def where(self, x: Self | ConstType, y: Self | ConstType):
|
||||||
|
if isinstance(x, type(self)):
|
||||||
|
return self.alu(Ops.WHERE, x, x.ufix(y))
|
||||||
|
if isinstance(y, type(self)):
|
||||||
|
return self.alu(Ops.WHERE, y.ufix(x), y)
|
||||||
raise RuntimeError("where needs at least one UOp arg")
|
raise RuntimeError("where needs at least one UOp arg")
|
||||||
def threefry(self, seed:Self): return self.alu(Ops.THREEFRY, seed)
|
|
||||||
def reciprocal(self): return self.alu(Ops.RECIPROCAL)
|
def threefry(self, seed: Self):
|
||||||
def trunc(self): return self.alu(Ops.TRUNC)
|
return self.alu(Ops.THREEFRY, seed)
|
||||||
def sqrt(self): return self.alu(Ops.SQRT)
|
|
||||||
def sin(self): return self.alu(Ops.SIN)
|
def reciprocal(self):
|
||||||
def log2(self): return self.alu(Ops.LOG2)
|
return self.alu(Ops.RECIPROCAL)
|
||||||
def exp2(self): return self.alu(Ops.EXP2)
|
|
||||||
def pow(self, x:Self|ConstType): return self.alu(Ops.POW, self.ufix(x))
|
def trunc(self):
|
||||||
def __pow__(self, x:Self|ConstType): return self.pow(x)
|
return self.alu(Ops.TRUNC)
|
||||||
|
|
||||||
|
def sqrt(self):
|
||||||
|
return self.alu(Ops.SQRT)
|
||||||
|
|
||||||
|
def sin(self):
|
||||||
|
return self.alu(Ops.SIN)
|
||||||
|
|
||||||
|
def log2(self):
|
||||||
|
return self.alu(Ops.LOG2)
|
||||||
|
|
||||||
|
def exp2(self):
|
||||||
|
return self.alu(Ops.EXP2)
|
||||||
|
|
||||||
|
def pow(self, x: Self | ConstType):
|
||||||
|
return self.alu(Ops.POW, self.ufix(x))
|
||||||
|
|
||||||
|
def __pow__(self, x: Self | ConstType):
|
||||||
|
return self.pow(x)
|
||||||
|
|||||||
+88
-40
@@ -2,20 +2,28 @@
|
|||||||
import functools
|
import functools
|
||||||
from typing import TypeAlias, TYPE_CHECKING, Self
|
from typing import TypeAlias, TYPE_CHECKING, Self
|
||||||
from tinygrad.uop import Ops
|
from tinygrad.uop import Ops
|
||||||
from tinygrad.helpers import prod, argfix, flatten, dedup
|
from tinygrad.helpers import prod, argfix, flatten, dedup, make_tuple, ceildiv
|
||||||
if TYPE_CHECKING: from tinygrad.uop.ops import UOp
|
from tinygrad.uop.ops import resolve, smax
|
||||||
|
|
||||||
|
if TYPE_CHECKING:
|
||||||
|
from tinygrad.uop.ops import UOp
|
||||||
sint: TypeAlias = "UOp | int"
|
sint: TypeAlias = "UOp | int"
|
||||||
|
|
||||||
def _align_left(*shapes:tuple[sint, ...]) -> tuple[tuple[sint, ...], ...]:
|
|
||||||
|
def _align_left(*shapes: tuple[sint, ...]) -> tuple[tuple[sint, ...], ...]:
|
||||||
# unsqueeze left to make every shape same length
|
# unsqueeze left to make every shape same length
|
||||||
max_dim = max(len(shape) for shape in shapes)
|
max_dim = max(len(shape) for shape in shapes)
|
||||||
return tuple((1,) * (max_dim - len(shape)) + shape for shape in shapes)
|
return tuple((1,) * (max_dim - len(shape)) + shape for shape in shapes)
|
||||||
|
|
||||||
|
|
||||||
class MovementMixin:
|
class MovementMixin:
|
||||||
# required to implement
|
# required to implement
|
||||||
def _mop(self, op:Ops, arg) -> Self: raise NotImplementedError
|
def _mop(self, op: Ops, arg) -> Self:
|
||||||
|
raise NotImplementedError
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def shape(self) -> tuple[sint, ...]: raise NotImplementedError
|
def shape(self) -> tuple[sint, ...]:
|
||||||
|
raise NotImplementedError
|
||||||
|
|
||||||
# great functions you get!
|
# great functions you get!
|
||||||
@property
|
@property
|
||||||
@@ -41,18 +49,21 @@ class MovementMixin:
|
|||||||
"""
|
"""
|
||||||
return prod(self.shape)
|
return prod(self.shape)
|
||||||
|
|
||||||
def _resolve_dim(self, dim:int, *, extra:bool=False) -> int:
|
def _resolve_dim(self, dim: int, *, extra: bool = False) -> int:
|
||||||
total = self.ndim + int(extra)
|
total = self.ndim + int(extra)
|
||||||
if not -max(1, total) <= dim <= max(1, total)-1: raise IndexError(f"{dim=} out of range {[-max(1, total), max(1, total)-1]}")
|
if not -max(1, total) <= dim <= max(1, total) - 1:
|
||||||
|
raise IndexError(f"{dim=} out of range {[-max(1, total), max(1, total) - 1]}")
|
||||||
return dim + total if dim < 0 else dim
|
return dim + total if dim < 0 else dim
|
||||||
|
|
||||||
def _broadcast_to(self, new_shape:tuple[sint, ...]) -> Self:
|
def _broadcast_to(self, new_shape: tuple[sint, ...]) -> Self:
|
||||||
if self.shape == new_shape: return self
|
if self.shape == new_shape:
|
||||||
if self.ndim > len(new_shape): raise ValueError(f"cannot broadcast tensor to fewer dimensions. shape={self.shape} to {new_shape=}")
|
return self
|
||||||
|
if self.ndim > len(new_shape):
|
||||||
|
raise ValueError(f"cannot broadcast tensor to fewer dimensions. shape={self.shape} to {new_shape=}")
|
||||||
# first unsqueeze left with 1s https://data-apis.org/array-api/latest/API_specification/broadcasting.html
|
# first unsqueeze left with 1s https://data-apis.org/array-api/latest/API_specification/broadcasting.html
|
||||||
shape, _ = _align_left(self.shape, new_shape)
|
shape, _ = _align_left(self.shape, new_shape)
|
||||||
# for each dimension, check either dim is 1, or it does not change
|
# for each dimension, check either dim is 1, or it does not change
|
||||||
if not all(s == ns or s == 1 for s,ns in zip(shape, new_shape)):
|
if not all(s == ns or s == 1 for s, ns in zip(shape, new_shape)):
|
||||||
raise ValueError(f"cannot broadcast {self.shape} to {new_shape=}")
|
raise ValueError(f"cannot broadcast {self.shape} to {new_shape=}")
|
||||||
reshaped = self.reshape(shape)
|
reshaped = self.reshape(shape)
|
||||||
ret = reshaped._mop(Ops.EXPAND, arg=new_shape)
|
ret = reshaped._mop(Ops.EXPAND, arg=new_shape)
|
||||||
@@ -84,15 +95,18 @@ class MovementMixin:
|
|||||||
```
|
```
|
||||||
"""
|
"""
|
||||||
# resolve None and args
|
# resolve None and args
|
||||||
new_shape = tuple([s if s is not None else self.shape[i] for i,s in enumerate(argfix(shape, *args))])
|
new_shape = tuple([s if s is not None else self.shape[i] for i, s in enumerate(argfix(shape, *args))])
|
||||||
# resolve -1
|
# resolve -1
|
||||||
if (c := new_shape.count(-1)) > 1: raise RuntimeError(f"only one dimension can be inferred using -1, getting {new_shape}")
|
if (c := new_shape.count(-1)) > 1:
|
||||||
if c: new_shape = tuple([-prod(self.shape) // prod(new_shape) if s == -1 else s for s in new_shape])
|
raise RuntimeError(f"only one dimension can be inferred using -1, getting {new_shape}")
|
||||||
if prod(self.shape) != prod(new_shape): raise ValueError(f"size mismatch, can't reshape ({self.shape}) -> ({new_shape})")
|
if c:
|
||||||
|
new_shape = tuple([-prod(self.shape) // prod(new_shape) if s == -1 else s for s in new_shape])
|
||||||
|
if prod(self.shape) != prod(new_shape):
|
||||||
|
raise ValueError(f"size mismatch, can't reshape ({self.shape}) -> ({new_shape})")
|
||||||
ret = self._mop(Ops.RESHAPE, arg=new_shape)
|
ret = self._mop(Ops.RESHAPE, arg=new_shape)
|
||||||
return self if ret.shape == self.shape else ret
|
return self if ret.shape == self.shape else ret
|
||||||
|
|
||||||
def shrink(self, arg:tuple[tuple[sint, sint]|None, ...]) -> Self:
|
def shrink(self, arg: tuple[tuple[sint, sint] | None, ...]) -> Self:
|
||||||
"""
|
"""
|
||||||
Returns a tensor that shrinks the each axis based on input arg.
|
Returns a tensor that shrinks the each axis based on input arg.
|
||||||
`arg` must have the same length as `self.ndim`.
|
`arg` must have the same length as `self.ndim`.
|
||||||
@@ -109,8 +123,9 @@ class MovementMixin:
|
|||||||
print(t.shrink((((0, 2), (0, 2)))).numpy())
|
print(t.shrink((((0, 2), (0, 2)))).numpy())
|
||||||
```
|
```
|
||||||
"""
|
"""
|
||||||
if self.ndim != len(arg): raise ValueError(f"{self.ndim=} != {len(arg)=}")
|
if self.ndim != len(arg):
|
||||||
ret = self._mop(Ops.SHRINK, arg=[x if x is not None else (0,s) for x,s in zip(arg, self.shape)])
|
raise ValueError(f"{self.ndim=} != {len(arg)=}")
|
||||||
|
ret = self._mop(Ops.SHRINK, arg=[x if x is not None else (0, s) for x, s in zip(arg, self.shape)])
|
||||||
return self if ret.shape == self.shape else ret
|
return self if ret.shape == self.shape else ret
|
||||||
|
|
||||||
def permute(self, order, *args) -> Self:
|
def permute(self, order, *args) -> Self:
|
||||||
@@ -128,7 +143,8 @@ class MovementMixin:
|
|||||||
```
|
```
|
||||||
"""
|
"""
|
||||||
order_arg = tuple(self._resolve_dim(x) for x in argfix(order, *args))
|
order_arg = tuple(self._resolve_dim(x) for x in argfix(order, *args))
|
||||||
if sorted(order_arg) != list(range(self.ndim)): raise RuntimeError(f"order is not a valid permutation, getting {order_arg}")
|
if sorted(order_arg) != list(range(self.ndim)):
|
||||||
|
raise RuntimeError(f"order is not a valid permutation, getting {order_arg}")
|
||||||
return self._mop(Ops.PERMUTE, arg=order_arg) if order_arg != tuple(range(self.ndim)) else self
|
return self._mop(Ops.PERMUTE, arg=order_arg) if order_arg != tuple(range(self.ndim)) else self
|
||||||
|
|
||||||
def flip(self, axis, *args) -> Self:
|
def flip(self, axis, *args) -> Self:
|
||||||
@@ -149,7 +165,8 @@ class MovementMixin:
|
|||||||
"""
|
"""
|
||||||
axis_arg = tuple(self._resolve_dim(x) for x in argfix(axis, *args))
|
axis_arg = tuple(self._resolve_dim(x) for x in argfix(axis, *args))
|
||||||
assert all(not isinstance(x, bool) and x >= 0 and x < self.ndim for x in axis_arg), f"flip args must be axis ints {axis_arg}"
|
assert all(not isinstance(x, bool) and x >= 0 and x < self.ndim for x in axis_arg), f"flip args must be axis ints {axis_arg}"
|
||||||
if len(axis_arg) != len(dedup(axis_arg)): raise RuntimeError(f"dim can appear at most once, getting {axis_arg}")
|
if len(axis_arg) != len(dedup(axis_arg)):
|
||||||
|
raise RuntimeError(f"dim can appear at most once, getting {axis_arg}")
|
||||||
flip_arg = tuple([i in axis_arg for i in range(len(self.shape))])
|
flip_arg = tuple([i in axis_arg for i in range(len(self.shape))])
|
||||||
return self._mop(Ops.FLIP, arg=flip_arg) if any(flip_arg) else self
|
return self._mop(Ops.FLIP, arg=flip_arg) if any(flip_arg) else self
|
||||||
|
|
||||||
@@ -162,7 +179,7 @@ class MovementMixin:
|
|||||||
"""`.view` is an alias for `.reshape`."""
|
"""`.view` is an alias for `.reshape`."""
|
||||||
return self.reshape(shape, *args)
|
return self.reshape(shape, *args)
|
||||||
|
|
||||||
def squeeze(self, dim:int|None=None) -> Self:
|
def squeeze(self, dim: int | None = None) -> Self:
|
||||||
"""
|
"""
|
||||||
Returns a tensor with specified dimensions of input of size 1 removed.
|
Returns a tensor with specified dimensions of input of size 1 removed.
|
||||||
If `dim` is not specified, all dimensions with size 1 are removed.
|
If `dim` is not specified, all dimensions with size 1 are removed.
|
||||||
@@ -178,11 +195,12 @@ class MovementMixin:
|
|||||||
print(t.squeeze(1).shape)
|
print(t.squeeze(1).shape)
|
||||||
```
|
```
|
||||||
"""
|
"""
|
||||||
if dim is None: return self.reshape(tuple(dim for dim in self.shape if dim != 1))
|
if dim is None:
|
||||||
|
return self.reshape(tuple(dim for dim in self.shape if dim != 1))
|
||||||
dim = self._resolve_dim(dim)
|
dim = self._resolve_dim(dim)
|
||||||
return self if not self.ndim or self.shape[dim] != 1 else self.reshape(self.shape[:dim] + self.shape[dim+1:])
|
return self if not self.ndim or self.shape[dim] != 1 else self.reshape(self.shape[:dim] + self.shape[dim + 1 :])
|
||||||
|
|
||||||
def unsqueeze(self, dim:int) -> Self:
|
def unsqueeze(self, dim: int) -> Self:
|
||||||
"""
|
"""
|
||||||
Returns a tensor with a new dimension of size 1 inserted at the specified `dim`.
|
Returns a tensor with a new dimension of size 1 inserted at the specified `dim`.
|
||||||
|
|
||||||
@@ -233,9 +251,9 @@ class MovementMixin:
|
|||||||
```
|
```
|
||||||
"""
|
"""
|
||||||
start_dim, end_dim = self._resolve_dim(start_dim), self._resolve_dim(end_dim)
|
start_dim, end_dim = self._resolve_dim(start_dim), self._resolve_dim(end_dim)
|
||||||
return self.reshape(self.shape[:start_dim] + (prod(self.shape[start_dim:end_dim+1]), ) + self.shape[end_dim+1:])
|
return self.reshape(self.shape[:start_dim] + (prod(self.shape[start_dim : end_dim + 1]),) + self.shape[end_dim + 1 :])
|
||||||
|
|
||||||
def unflatten(self, dim:int, sizes:tuple[int,...]) -> Self:
|
def unflatten(self, dim: int, sizes: tuple[int, ...]) -> Self:
|
||||||
"""
|
"""
|
||||||
Unflattens dimension `dim` of the tensor into multiple dimensions specified by `sizes`. `Tensor.flatten()` is the inverse of this function.
|
Unflattens dimension `dim` of the tensor into multiple dimensions specified by `sizes`. `Tensor.flatten()` is the inverse of this function.
|
||||||
|
|
||||||
@@ -250,9 +268,9 @@ class MovementMixin:
|
|||||||
```
|
```
|
||||||
"""
|
"""
|
||||||
dim = self._resolve_dim(dim)
|
dim = self._resolve_dim(dim)
|
||||||
return self.reshape(self.shape[:dim] + sizes + self.shape[dim+1:])
|
return self.reshape(self.shape[:dim] + sizes + self.shape[dim + 1 :])
|
||||||
|
|
||||||
def rearrange(self, formula:str, **sizes) -> Self:
|
def rearrange(self, formula: str, **sizes) -> Self:
|
||||||
"""
|
"""
|
||||||
Rearranges input according to formula
|
Rearranges input according to formula
|
||||||
|
|
||||||
@@ -263,38 +281,43 @@ class MovementMixin:
|
|||||||
print(Tensor.rearrange(x, "batch channel -> (batch channel)").numpy())
|
print(Tensor.rearrange(x, "batch channel -> (batch channel)").numpy())
|
||||||
```
|
```
|
||||||
"""
|
"""
|
||||||
|
|
||||||
def parse_formula(formula: str):
|
def parse_formula(formula: str):
|
||||||
tokens = f" {formula} ".replace("…", "...").replace("(", " ( ").replace(")", " ) ").replace(" ", " ").replace(" 1 ", " ( ) ").split()
|
tokens = f" {formula} ".replace("…", "...").replace("(", " ( ").replace(")", " ) ").replace(" ", " ").replace(" 1 ", " ( ) ").split()
|
||||||
lparens, rparens = map(lambda x: [i for i, ch in enumerate(tokens) if ch == x], ("(", ")"))
|
lparens, rparens = map(lambda x: [i for i, ch in enumerate(tokens) if ch == x], ("(", ")"))
|
||||||
pairs = list(zip(lparens, rparens))
|
pairs = list(zip(lparens, rparens))
|
||||||
assert len(lparens) == len(rparens) and sorted(flatten(pairs)) == flatten(pairs), "bracket mismatch"
|
assert len(lparens) == len(rparens) and sorted(flatten(pairs)) == flatten(pairs), "bracket mismatch"
|
||||||
return [name for name in tokens if name not in ("(", ")")], [(s - 2*i, e - 1 - 2*i) for i, (s, e) in enumerate(pairs)]
|
return [name for name in tokens if name not in ("(", ")")], [(s - 2 * i, e - 1 - 2 * i) for i, (s, e) in enumerate(pairs)]
|
||||||
|
|
||||||
assert formula.count("->") == 1, 'need exactly one "->" in formula'
|
assert formula.count("->") == 1, 'need exactly one "->" in formula'
|
||||||
|
|
||||||
(lhs, unflatten_dims), (rhs, flatten_dims) = map(parse_formula, formula.split("->"))
|
(lhs, unflatten_dims), (rhs, flatten_dims) = map(parse_formula, formula.split("->"))
|
||||||
|
|
||||||
for name in sizes: assert name in lhs, f"axis {name} is not used in transform"
|
for name in sizes:
|
||||||
|
assert name in lhs, f"axis {name} is not used in transform"
|
||||||
assert sorted(lhs) == sorted(rhs) and len(lhs) == len(set(lhs)), f"name mismatch in {formula}"
|
assert sorted(lhs) == sorted(rhs) and len(lhs) == len(set(lhs)), f"name mismatch in {formula}"
|
||||||
for name in flatten((lhs, rhs)): assert name == "..." or (name.isidentifier() and "_" not in (name[0], name[-1])), f"invalid axis name {name}"
|
for name in flatten((lhs, rhs)):
|
||||||
|
assert name == "..." or (name.isidentifier() and "_" not in (name[0], name[-1])), f"invalid axis name {name}"
|
||||||
assert "..." not in flatten([lhs[s:e] for s, e in unflatten_dims]), f"cannot have collapsed ellipsis (...) in lhs of {formula}"
|
assert "..." not in flatten([lhs[s:e] for s, e in unflatten_dims]), f"cannot have collapsed ellipsis (...) in lhs of {formula}"
|
||||||
assert lhs.count("...") <= 1, f"too many ellipses in {formula}"
|
assert lhs.count("...") <= 1, f"too many ellipses in {formula}"
|
||||||
|
|
||||||
# resolve ellipsis
|
# resolve ellipsis
|
||||||
if "..." in lhs: ell_len = len(self.shape) - len(lhs) + 1 + sum(e - s - 1 for s, e in unflatten_dims)
|
if "..." in lhs:
|
||||||
lhs, rhs = map(lambda l: l[:(i:=l.index("..."))] + [f"...{j}" for j in range(ell_len)] + l[i + 1:] if "..." in l else l, (lhs, rhs))
|
ell_len = len(self.shape) - len(lhs) + 1 + sum(e - s - 1 for s, e in unflatten_dims)
|
||||||
|
lhs, rhs = map(lambda l: l[: (i := l.index("..."))] + [f"...{j}" for j in range(ell_len)] + l[i + 1 :] if "..." in l else l, (lhs, rhs))
|
||||||
unflatten_dims = [(s + (ell_len - 1 if "...0" in lhs[:s] else 0), e + (ell_len - 1 if "...0" in lhs[:e] else 0)) for s, e in unflatten_dims]
|
unflatten_dims = [(s + (ell_len - 1 if "...0" in lhs[:s] else 0), e + (ell_len - 1 if "...0" in lhs[:e] else 0)) for s, e in unflatten_dims]
|
||||||
flatten_dims = [(s + (ell_len - 1 if "...0" in rhs[:s] else 0), e + (ell_len - 1 if "...0" in rhs[:e] else 0)) for s, e in flatten_dims]
|
flatten_dims = [(s + (ell_len - 1 if "...0" in rhs[:s] else 0), e + (ell_len - 1 if "...0" in rhs[:e] else 0)) for s, e in flatten_dims]
|
||||||
|
|
||||||
# apply movement ops in order unflatten -> permute -> flatten/unsqueeze
|
# apply movement ops in order unflatten -> permute -> flatten/unsqueeze
|
||||||
t = functools.reduce(lambda x, dims: x.unflatten(dims[0], tuple(sizes.get(lhs[d], -1) for d in range(*dims))), unflatten_dims, self)
|
t = functools.reduce(lambda x, dims: x.unflatten(dims[0], tuple(sizes.get(lhs[d], -1) for d in range(*dims))), unflatten_dims, self)
|
||||||
for i, name in enumerate(lhs): assert (name not in sizes) or sizes[name] == t.shape[i], f"size provided for dimension {name} incorrect"
|
for i, name in enumerate(lhs):
|
||||||
|
assert (name not in sizes) or sizes[name] == t.shape[i], f"size provided for dimension {name} incorrect"
|
||||||
t = t.permute([lhs.index(name) for name in rhs])
|
t = t.permute([lhs.index(name) for name in rhs])
|
||||||
return functools.reduce(lambda x, dims: x.flatten(dims[0], dims[1] - 1) if dims[0]<dims[1] else x.unsqueeze(dims[0]), reversed(flatten_dims), t)
|
return functools.reduce(lambda x, dims: x.flatten(dims[0], dims[1] - 1) if dims[0] < dims[1] else x.unsqueeze(dims[0]), reversed(flatten_dims), t)
|
||||||
|
|
||||||
# *** movement ops with expand ***
|
# *** movement ops with expand ***
|
||||||
|
|
||||||
def repeat_interleave(self, repeats:int, dim:int|None=None) -> Self:
|
def repeat_interleave(self, repeats: int, dim: int | None = None) -> Self:
|
||||||
"""
|
"""
|
||||||
Repeats elements of a tensor.
|
Repeats elements of a tensor.
|
||||||
|
|
||||||
@@ -305,7 +328,10 @@ class MovementMixin:
|
|||||||
"""
|
"""
|
||||||
x, dim = (self.flatten(), 0) if dim is None else (self, self._resolve_dim(dim))
|
x, dim = (self.flatten(), 0) if dim is None else (self, self._resolve_dim(dim))
|
||||||
shp = x.shape
|
shp = x.shape
|
||||||
return x.reshape(*shp[:dim+1], 1, *shp[dim+1:]).expand(*shp[:dim+1], repeats, *shp[dim+1:]).reshape(*shp[:dim], shp[dim]*repeats, *shp[dim+1:])
|
x = x.reshape(*shp[: dim + 1], 1, *shp[dim + 1 :])
|
||||||
|
x = x.expand(*shp[: dim + 1], repeats, *shp[dim + 1 :])
|
||||||
|
x = x.reshape(*shp[:dim], shp[dim] * repeats, *shp[dim + 1 :])
|
||||||
|
return x
|
||||||
|
|
||||||
def repeat(self, repeats, *args) -> Self:
|
def repeat(self, repeats, *args) -> Self:
|
||||||
"""
|
"""
|
||||||
@@ -322,7 +348,29 @@ class MovementMixin:
|
|||||||
"""
|
"""
|
||||||
repeats = argfix(repeats, *args)
|
repeats = argfix(repeats, *args)
|
||||||
base_shape = _align_left(self.shape, repeats)[0]
|
base_shape = _align_left(self.shape, repeats)[0]
|
||||||
unsqueezed_shape = flatten([[s] if r == 1 else [1, s] for r,s in zip(repeats, base_shape)])
|
unsqueezed_shape = flatten([[s] if r == 1 else [1, s] for r, s in zip(repeats, base_shape)])
|
||||||
expanded_shape = flatten([[s] if r == 1 else [r, s] for r,s in zip(repeats, base_shape)])
|
expanded_shape = flatten([[s] if r == 1 else [r, s] for r, s in zip(repeats, base_shape)])
|
||||||
final_shape = [r*s for r,s in zip(repeats, base_shape)]
|
final_shape = [r * s for r, s in zip(repeats, base_shape)]
|
||||||
return self.reshape(unsqueezed_shape).expand(expanded_shape).reshape(final_shape)
|
return self.reshape(unsqueezed_shape).expand(expanded_shape).reshape(final_shape)
|
||||||
|
|
||||||
|
# **** pool level ****
|
||||||
|
|
||||||
|
def _pool(self, k_: tuple[sint, ...], stride: int | tuple[int, ...] = 1, dilation: int | tuple[int, ...] = 1) -> Self:
|
||||||
|
assert len(self.shape) >= len(k_), f"can't pool {self.shape} with {k_}"
|
||||||
|
s_, d_ = make_tuple(stride, len(k_)), make_tuple(dilation, len(k_))
|
||||||
|
assert len(k_) == len(s_) == len(d_), f"stride/dilation mismatch kernel:{k_} stride:{s_} dilation:{d_}"
|
||||||
|
noop, i_ = [None] * (self.ndim - len(k_)), self.shape[-len(k_) :]
|
||||||
|
assert all(resolve(d * (k - 1) + 1 <= i) for k, d, i in zip(k_, d_, i_)), "kernel size cannot be greater than actual input size"
|
||||||
|
o_ = [ceildiv(i - d * (k - 1), s) for i, d, k, s in zip(i_, d_, k_, s_)]
|
||||||
|
# input size scaling factor to make sure shrink for stride is possible
|
||||||
|
f_ = [smax(1, ceildiv(o * s - d, i)) for o, s, i, d in zip(o_, s_, i_, d_)]
|
||||||
|
# repeats such that we don't need padding
|
||||||
|
x = self.repeat([1] * len(noop) + [ceildiv(k * (i * f + d), i) for k, i, d, f in zip(k_, i_, d_, f_)])
|
||||||
|
# handle dilation
|
||||||
|
x = x.shrink_to(noop + [k * (i * f + d) for k, i, d, f in zip(k_, i_, d_, f_)])
|
||||||
|
x = x.reshape(noop + flatten((k, (i * f + d)) for k, i, d, f in zip(k_, i_, d_, f_)))
|
||||||
|
# handle stride
|
||||||
|
x = x.shrink_to(noop + flatten((k, o * s) for k, o, s in zip(k_, o_, s_))).reshape(noop + flatten((k, o, s) for k, o, s in zip(k_, o_, s_)))
|
||||||
|
x = x.shrink_to(noop + flatten((k, o, 1) for k, o in zip(k_, o_))).reshape(noop + flatten((k, o) for k, o in zip(k_, o_)))
|
||||||
|
# permute to move reduce to the end
|
||||||
|
return x.permute(*range(len(noop)), *[len(noop) + i * 2 + 1 for i in range(len(i_))], *[len(noop) + i * 2 for i in range(len(i_))])
|
||||||
|
|||||||
@@ -1124,6 +1124,16 @@ def get_onnx_ops() -> dict[str, types.FunctionType|dict[OpSetId, types.FunctionT
|
|||||||
return output.flatten(start_dim=2) if len(original_input_shape) == 3 else output.permute(0, 2, 1, 3)
|
return output.flatten(start_dim=2) if len(original_input_shape) == 3 else output.permute(0, 2, 1, 3)
|
||||||
|
|
||||||
# ***** Indexing Ops *****
|
# ***** Indexing Ops *****
|
||||||
|
def NonZero(x:Tensor):
|
||||||
|
mask = (x!=0).flatten()
|
||||||
|
flat_idx = Tensor.arange(mask.numel(), dtype=dtypes.int64, device=x.device).masked_select(mask)
|
||||||
|
if flat_idx.ndim == 0: flat_idx = flat_idx.reshape(1)
|
||||||
|
if x.ndim == 0:
|
||||||
|
return Tensor.zeros((0, flat_idx.shape[0]), dtype=dtypes.int64, device=x.device, requires_grad=False)
|
||||||
|
strides = [prod(int(s) for s in x.shape[i+1:]) if i+1 < x.ndim else 1 for i in range(x.ndim)]
|
||||||
|
coords = [((flat_idx // stride) % int(dim)) for stride, dim in zip(strides, x.shape)]
|
||||||
|
return Tensor.stack(*coords, dim=0)
|
||||||
|
|
||||||
def ArrayFeatureExtractor(x:Tensor, indices:Tensor): return x[..., indices]
|
def ArrayFeatureExtractor(x:Tensor, indices:Tensor): return x[..., indices]
|
||||||
|
|
||||||
def Gather(x:Tensor, indices:Tensor, axis:int=0):
|
def Gather(x:Tensor, indices:Tensor, axis:int=0):
|
||||||
|
|||||||
@@ -194,6 +194,10 @@ def torch_load(t:Tensor) -> dict[str, Tensor]:
|
|||||||
"""
|
"""
|
||||||
offsets: dict[str|int, int] = {}
|
offsets: dict[str|int, int] = {}
|
||||||
lens: dict[str|int, int] = {}
|
lens: dict[str|int, int] = {}
|
||||||
|
|
||||||
|
def _rebuild_tensor(storage, storage_offset, size, stride):
|
||||||
|
return _rebuild_tensor_v2(storage, storage_offset, size, stride)
|
||||||
|
|
||||||
def _rebuild_tensor_v2(storage, storage_offset, size, stride, requires_grad=None, backward_hooks=None, metadata=None):
|
def _rebuild_tensor_v2(storage, storage_offset, size, stride, requires_grad=None, backward_hooks=None, metadata=None):
|
||||||
#print(storage, storage_offset, size, stride, requires_grad, backward_hooks, metadata)
|
#print(storage, storage_offset, size, stride, requires_grad, backward_hooks, metadata)
|
||||||
lens[storage[2]] = storage[4] * storage[1].itemsize
|
lens[storage[2]] = storage[4] * storage[1].itemsize
|
||||||
@@ -220,7 +224,8 @@ def torch_load(t:Tensor) -> dict[str, Tensor]:
|
|||||||
deserialized_objects: dict[str, Any] = {}
|
deserialized_objects: dict[str, Any] = {}
|
||||||
intercept = {"HalfStorage": dtypes.float16, "FloatStorage": dtypes.float32, "BFloat16Storage": dtypes.bfloat16,
|
intercept = {"HalfStorage": dtypes.float16, "FloatStorage": dtypes.float32, "BFloat16Storage": dtypes.bfloat16,
|
||||||
"IntStorage": dtypes.int32, "BoolStorage": dtypes.bool,
|
"IntStorage": dtypes.int32, "BoolStorage": dtypes.bool,
|
||||||
"LongStorage": dtypes.int64, "_rebuild_tensor_v2": _rebuild_tensor_v2, "FloatTensor": None, "Parameter": Parameter}
|
"LongStorage": dtypes.int64, "_rebuild_tensor": _rebuild_tensor, "_rebuild_tensor_v2": _rebuild_tensor_v2,
|
||||||
|
"FloatTensor": None, "Parameter": Parameter}
|
||||||
whitelist = {"torch", "collections", "numpy", "_codecs"} # NOTE: this is not for security, only speed
|
whitelist = {"torch", "collections", "numpy", "_codecs"} # NOTE: this is not for security, only speed
|
||||||
class Dummy: pass
|
class Dummy: pass
|
||||||
class TorchPickle(pickle.Unpickler):
|
class TorchPickle(pickle.Unpickler):
|
||||||
|
|||||||
@@ -450,16 +450,9 @@ class AMDRenderer(CStyleLanguage):
|
|||||||
]) + base_rewrite
|
]) + base_rewrite
|
||||||
def __reduce__(self): return self.__class__, (self.arch,)
|
def __reduce__(self): return self.__class__, (self.arch,)
|
||||||
|
|
||||||
# language options
|
|
||||||
ockl = [(f"__ockl_get_{name}", "unsigned int", "size_t", "const") for name in ["local_id", "group_id", "local_size"]]
|
|
||||||
ocml = [(f"__ocml_{name}_f{n}", f"{dt}, {dt}" if "fmax" == name else dt, dt, atr)
|
|
||||||
for dt, n in [(dtype.name, dtype.itemsize * 8) for dtype in [dtypes.float, dtypes.double, dtypes.half]]
|
|
||||||
for name, atr in [("fmax", "const"), ("exp2", "pure"), ("log2", "pure"), ("sqrt", "const"), ("sin", ""), ("trunc", "")]]
|
|
||||||
|
|
||||||
kernel_typedef = "\n".join(f'extern "C" __attribute__((device{f", {atr}" if atr else ""})) {dto} {meth}({dti});' for meth,dti,dto,atr in ockl+ocml)
|
|
||||||
# https://clang.llvm.org/docs/AttributeReference.html#amdgpu-flat-work-group-size
|
# https://clang.llvm.org/docs/AttributeReference.html#amdgpu-flat-work-group-size
|
||||||
# NOTE: this makes hlb_cifar10 twice as fast, there may be more gains in tweaking these parameters
|
# NOTE: this makes hlb_cifar10 twice as fast, there may be more gains in tweaking these parameters
|
||||||
kernel_typedef += '\nextern "C" __attribute__((global)) void __attribute__((amdgpu_flat_work_group_size(1, {launch_bounds})))'
|
kernel_typedef = 'extern "C" __attribute__((global)) void __attribute__((amdgpu_flat_work_group_size(1, {launch_bounds})))'
|
||||||
code_for_workitem = {"g": lambda x: f"__ockl_get_group_id({x})", "l": lambda x: f"__ockl_get_local_id({x})",
|
code_for_workitem = {"g": lambda x: f"__ockl_get_group_id({x})", "l": lambda x: f"__ockl_get_local_id({x})",
|
||||||
"i": lambda x: f"(__ockl_get_group_id({x})*__ockl_get_local_size({x})+__ockl_get_local_id({x}))"}
|
"i": lambda x: f"(__ockl_get_group_id({x})*__ockl_get_local_size({x})+__ockl_get_local_id({x}))"}
|
||||||
code_for_op = { **CStyleLanguage.code_for_op,
|
code_for_op = { **CStyleLanguage.code_for_op,
|
||||||
@@ -490,15 +483,25 @@ class AMDRenderer(CStyleLanguage):
|
|||||||
f"{vec} make_{vec}({', '.join([f'{scal} {x}' for x in _nms[:dtype.count]])}) {{ return {{ {', '.join(_nms[:dtype.count])} }}; }}"
|
f"{vec} make_{vec}({', '.join([f'{scal} {x}' for x in _nms[:dtype.count]])}) {{ return {{ {', '.join(_nms[:dtype.count])} }}; }}"
|
||||||
|
|
||||||
def render_kernel(self, function_name, kernel, bufs, uops, prefix=None) -> str:
|
def render_kernel(self, function_name, kernel, bufs, uops, prefix=None) -> str:
|
||||||
prefix = ["#define INFINITY (__builtin_inff())","#define NAN (__builtin_nanf(\"\"))","typedef long unsigned int size_t;","#define half _Float16"]
|
prefix, ockl = [], []
|
||||||
type_map = { dtypes.bfloat16: "bf16", dtypes.float: "f32", dtypes.half: "f16", dtypes.fp8e4m3: "_fp8_fp8", dtypes.fp8e5m2: "_bf8_bf8" }
|
type_map = { dtypes.bfloat16: "bf16", dtypes.float: "f32", dtypes.half: "f16", dtypes.fp8e4m3: "_fp8_fp8", dtypes.fp8e5m2: "_bf8_bf8" }
|
||||||
used_dtypes = uops_to_dtypes(uops)
|
used_dtypes = uops_to_dtypes(uops)
|
||||||
|
if any(u.op is Ops.CONST and not math.isfinite(u.arg) for u in uops):
|
||||||
|
prefix += ["#define INFINITY (__builtin_inff())", "#define NAN (__builtin_nanf(\"\"))"]
|
||||||
|
if any(u.op is Ops.SPECIAL for u in uops):
|
||||||
|
prefix.append("typedef long unsigned int size_t;")
|
||||||
|
ockl = [(f"__ockl_get_{name}", "unsigned int", "size_t", "const") for name in ["local_id", "group_id", "local_size"]]
|
||||||
|
ocml_ops = {Ops.EXP2: ("exp2", "pure"), Ops.LOG2: ("log2", "pure"), Ops.SQRT: ("sqrt", "const"), Ops.SIN: ("sin", ""), Ops.TRUNC: ("trunc", "")}
|
||||||
|
ocml = [(f"__ocml_{ocml_ops[op][0]}_f{dt.itemsize * 8}", dt.name, dt.name, ocml_ops[op][1])
|
||||||
|
for op, dt in dedup((u.op, u.dtype.scalar()) for u in uops) if op in ocml_ops and dt in (dtypes.half, dtypes.float, dtypes.double)]
|
||||||
if any(dt.scalar() == dtypes.bfloat16 for dt in used_dtypes): prefix.append("typedef unsigned short hip_bfloat16;")
|
if any(dt.scalar() == dtypes.bfloat16 for dt in used_dtypes): prefix.append("typedef unsigned short hip_bfloat16;")
|
||||||
|
if any(dt.scalar() == dtypes.half for dt in used_dtypes): prefix.append("#define half _Float16")
|
||||||
if any(dt.scalar() in dtypes.fp8s for dt in used_dtypes):
|
if any(dt.scalar() in dtypes.fp8s for dt in used_dtypes):
|
||||||
prefix += ["typedef unsigned char hip_bf8;", "typedef unsigned char hip_fp8;"]
|
prefix += ["typedef unsigned char hip_bf8;", "typedef unsigned char hip_fp8;"]
|
||||||
prefix.append("""static inline __attribute__((device)) unsigned char f32_to_fp8(float v, int is_bf8) {
|
prefix.append("""static inline __attribute__((device)) unsigned char f32_to_fp8(float v, int is_bf8) {
|
||||||
v = (((*(unsigned*)&v)&0x7F800000)!=0x7F800000)?__builtin_amdgcn_fmed3f(v,is_bf8?57344.0f:448.0f,is_bf8?-57344.0f:-448.0f) : v;
|
v = (((*(unsigned*)&v)&0x7F800000)!=0x7F800000)?__builtin_amdgcn_fmed3f(v,is_bf8?57344.0f:448.0f,is_bf8?-57344.0f:-448.0f) : v;
|
||||||
return (unsigned char)(is_bf8?__builtin_amdgcn_cvt_pk_bf8_f32(v,v,0,false):__builtin_amdgcn_cvt_pk_fp8_f32(v,v,0,false));\n}""")
|
return (unsigned char)(is_bf8?__builtin_amdgcn_cvt_pk_bf8_f32(v,v,0,false):__builtin_amdgcn_cvt_pk_fp8_f32(v,v,0,false));\n}""")
|
||||||
|
prefix += [f'extern "C" __attribute__((device{f", {atr}" if atr else ""})) {dto} {meth}({dti});' for meth,dti,dto,atr in ockl+ocml]
|
||||||
prefix += [self.render_vector_prefix(dt) for dt in used_dtypes if dt.count > 1]
|
prefix += [self.render_vector_prefix(dt) for dt in used_dtypes if dt.count > 1]
|
||||||
|
|
||||||
for name, (N, M, K), dtype_in, dtype_out, _, _, _, _ in wmma_args(uops): # TODO: handle TCs f32_bf16 and bf16_bf16 w/ wrapper
|
for name, (N, M, K), dtype_in, dtype_out, _, _, _, _ in wmma_args(uops): # TODO: handle TCs f32_bf16 and bf16_bf16 w/ wrapper
|
||||||
|
|||||||
@@ -1,11 +1,11 @@
|
|||||||
from typing import Callable, cast, Any
|
from typing import Callable, cast, Any
|
||||||
from tinygrad.dtype import AddrSpace, DType, PtrDType, dtypes
|
from tinygrad.dtype import AddrSpace, DType, PtrDType, dtypes
|
||||||
from tinygrad.helpers import DEBUG, OSX, unwrap
|
from tinygrad.helpers import DEBUG, OSX, unwrap, charptr
|
||||||
from tinygrad.renderer import Renderer
|
from tinygrad.renderer import Renderer
|
||||||
from tinygrad.renderer.cstyle import CUDARenderer
|
from tinygrad.renderer.cstyle import CUDARenderer
|
||||||
from tinygrad.uop.ops import GroupOp, Ops, UOp, PatternMatcher, UPat, range_str
|
from tinygrad.uop.ops import GroupOp, Ops, UOp, PatternMatcher, UPat, range_str
|
||||||
import tinygrad.runtime.autogen.mesa as mesa
|
from tinygrad.runtime.autogen import mesa
|
||||||
import base64, ctypes, ctypes.util, struct, functools, inspect
|
import base64, contextlib, ctypes, ctypes.util, struct, functools, inspect
|
||||||
|
|
||||||
def g(s:str): return getattr(mesa, s)
|
def g(s:str): return getattr(mesa, s)
|
||||||
def nsrc(d:mesa.nir_def) -> mesa.nir_src: return mesa.nir_src(ssa=ctypes.pointer(d))
|
def nsrc(d:mesa.nir_def) -> mesa.nir_src: return mesa.nir_src(ssa=ctypes.pointer(d))
|
||||||
@@ -51,7 +51,7 @@ def nir_instr(nc=1, bs=lambda: None, intrins=None, srcs=None, has_def=True, df=N
|
|||||||
instr = f(*args, **kwargs)
|
instr = f(*args, **kwargs)
|
||||||
if has_def: mesa.nir_def_init(instr.contents.instr, getattr(instr.contents, "def"), go(nc), go(bs))
|
if has_def: mesa.nir_def_init(instr.contents.instr, getattr(instr.contents, "def"), go(nc), go(bs))
|
||||||
for k, v in go(intrins or {}).items():
|
for k, v in go(intrins or {}).items():
|
||||||
idx = mesa.nir_intrinsic_infos[instr.contents.intrinsic].index_map[g(f"NIR_INTRINSIC_{k}")]
|
idx = mesa.nir_intrinsic_infos[instr.contents.intrinsic.value].index_map[g(f"NIR_INTRINSIC_{k}")]
|
||||||
assert idx > 0
|
assert idx > 0
|
||||||
instr.contents.const_index[idx - 1] = go(v)
|
instr.contents.const_index[idx - 1] = go(v)
|
||||||
for i, src in enumerate(go(srcs or [])): ctypes.cast(instr.contents.src, ctypes.POINTER(mesa.nir_src))[i] = go(src)
|
for i, src in enumerate(go(srcs or [])): ctypes.cast(instr.contents.src, ctypes.POINTER(mesa.nir_src))[i] = go(src)
|
||||||
@@ -157,8 +157,7 @@ class NIRRenderer(Renderer):
|
|||||||
def __init__(self): mesa.glsl_type_singleton_init_or_ref()
|
def __init__(self): mesa.glsl_type_singleton_init_or_ref()
|
||||||
|
|
||||||
def __del__(self):
|
def __del__(self):
|
||||||
try: mesa.glsl_type_singleton_decref()
|
with contextlib.suppress(AttributeError):mesa.glsl_type_singleton_decref()
|
||||||
except FileNotFoundError: pass
|
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def nir_options(self): raise NotImplementedError("needs nir_options")
|
def nir_options(self): raise NotImplementedError("needs nir_options")
|
||||||
@@ -177,7 +176,7 @@ class NIRRenderer(Renderer):
|
|||||||
elif u.op is Ops.AFTER:
|
elif u.op is Ops.AFTER:
|
||||||
self.r[u] = self.r[u.src[0]]
|
self.r[u] = self.r[u.src[0]]
|
||||||
elif u.op == Ops.SINK:
|
elif u.op == Ops.SINK:
|
||||||
if u.arg is not None: self.b.shader.contents.info.name = mesa.char_pointer_cast(u.arg.function_name)
|
if u.arg is not None: self.b.shader.contents.info.name = charptr(u.arg.function_name.encode())
|
||||||
elif u.op == Ops.DEFINE_LOCAL:
|
elif u.op == Ops.DEFINE_LOCAL:
|
||||||
self.r[u] = nimm(self.b, self.b.shader.contents.info.shared_size, dtypes.long)
|
self.r[u] = nimm(self.b, self.b.shader.contents.info.shared_size, dtypes.long)
|
||||||
self.b.shader.contents.info.shared_size += u.dtype.nbytes()
|
self.b.shader.contents.info.shared_size += u.dtype.nbytes()
|
||||||
|
|||||||
@@ -0,0 +1,132 @@
|
|||||||
|
import glob, importlib, pathlib, subprocess, tarfile
|
||||||
|
from tinygrad.helpers import fetch, flatten, system, getenv
|
||||||
|
|
||||||
|
root = (here:=pathlib.Path(__file__).parent).parents[2]
|
||||||
|
nv_src = {"nv_570": "https://github.com/NVIDIA/open-gpu-kernel-modules/archive/81fe4fb417c8ac3b9bdcc1d56827d116743892a5.tar.gz",
|
||||||
|
"nv_580": "https://github.com/NVIDIA/open-gpu-kernel-modules/archive/2af9f1f0f7de4988432d4ae875b5858ffdb09cc2.tar.gz"}
|
||||||
|
macossdk = "/var/db/xcode_select_link/Platforms/MacOSX.platform/Developer/SDKs/MacOSX.sdk"
|
||||||
|
|
||||||
|
def load(name, dll, files, **kwargs):
|
||||||
|
if not (f:=(root/(path:=kwargs.pop("path", __name__)).replace('.','/')/f"{name}.py")).exists() or getenv('REGEN'):
|
||||||
|
files, kwargs['args'] = files() if callable(files) else files, args() if callable(args:=kwargs.get('args', [])) else args
|
||||||
|
if (tarball:=kwargs.pop('tarball', None)):
|
||||||
|
# dangerous for arbitrary urls!
|
||||||
|
with tarfile.open(fetch(tarball, gunzip=tarball.endswith("gz"))) as tf:
|
||||||
|
tf.extractall("/tmp")
|
||||||
|
base = f"/tmp/{tf.getnames()[0]}"
|
||||||
|
files, kwargs['args'] = [str(f).format(base) for f in files], [a.format(base) for a in kwargs.get('args', [])]
|
||||||
|
kwargs['anon_names'] = {k.format(base):v for k,v in kwargs.get('anon_names', {}).items()}
|
||||||
|
if (preprocess:=kwargs.pop('preprocess', None)): preprocess(base)
|
||||||
|
files = flatten(sorted(glob.glob(p, recursive=True)) if isinstance(p, str) and '*' in p else [p] for p in files)
|
||||||
|
kwargs['epilog'] = (epi(base) if tarball else epi()) if callable(epi:=kwargs.get('epilog', [])) else epi
|
||||||
|
f.write_text(importlib.import_module("tinygrad.runtime.support.autogen").gen(dll, files, **kwargs))
|
||||||
|
return importlib.import_module(f"{path}.{name.replace('/', '.')}")
|
||||||
|
|
||||||
|
def __getattr__(nm):
|
||||||
|
match nm:
|
||||||
|
case "libc": return load("libc", ["find_library('c')"], lambda: (
|
||||||
|
[i for i in system("dpkg -L libc6-dev").split() if 'sys/mman.h' in i or 'sys/syscall.h' in i] +
|
||||||
|
["/usr/include/string.h", "/usr/include/elf.h", "/usr/include/unistd.h", "/usr/include/asm-generic/mman-common.h"]), use_errno=True)
|
||||||
|
case "opencl": return load("opencl", ["find_library('OpenCL')"], ["/usr/include/CL/cl.h"])
|
||||||
|
case "cuda": return load("cuda", ["find_library('cuda')"], ["/usr/include/cuda.h"], args=["-D__CUDA_API_VERSION_INTERNAL"], parse_macros=False)
|
||||||
|
case "nvrtc": return load("nvrtc", ["find_library('nvrtc')"], ["/usr/include/nvrtc.h"])
|
||||||
|
case "nvjitlink": load("nvjitlink", ["find_library('nvJitLink')"], [root/"extra/nvJitLink.h"])
|
||||||
|
case "kfd": return load("kfd", [], ["/usr/include/linux/kfd_ioctl.h"])
|
||||||
|
case "nv_570" | "nv_580":
|
||||||
|
return load(nm, [], [
|
||||||
|
*[root/"extra/nv_gpu_driver"/s for s in ["clc6c0qmd.h","clcec0qmd.h"]], "{}/kernel-open/common/inc/nvmisc.h",
|
||||||
|
*[f"{{}}/src/common/sdk/nvidia/inc/class/cl{s}.h" for s in ["0000", "0080", "2080", "2080_notification", "c56f", "c86f", "c96f", "c761",
|
||||||
|
"83de", "c6c0", "cdc0"]],
|
||||||
|
*[f"{{}}/kernel-open/nvidia-uvm/{s}.h" for s in ["clc6b5", "clc9b5", "uvm_ioctl", "uvm_linux_ioctl", "hwref/ampere/ga100/dev_fault"]],
|
||||||
|
*[f"{{}}/src/nvidia/arch/nvalloc/unix/include/nv{s}.h" for s in ["_escape", "-ioctl", "-ioctl-numbers",
|
||||||
|
"-ioctl-numa", "-unix-nvos-params-wrappers"]],
|
||||||
|
*[f"{{}}/src/common/sdk/nvidia/inc/{s}.h" for s in ["alloc/alloc_channel", "nvos", "ctrl/ctrlc36f", "ctrl/ctrlcb33",
|
||||||
|
"ctrl/ctrla06c", "ctrl/ctrl90f1"]],
|
||||||
|
*[f"{{}}/src/common/sdk/nvidia/inc/ctrl/ctrl{s}/*.h" for s in ["0000", "0080", "2080", "83de"]],
|
||||||
|
"{}/kernel-open/common/inc/nvstatus.h", "{}/src/nvidia/generated/g_allclasses.h"
|
||||||
|
], args=[
|
||||||
|
"-include", "{}/src/common/sdk/nvidia/inc/nvtypes.h", "-I{}/src/common/inc", "-I{}/kernel-open/nvidia-uvm", "-I{}/kernel-open/common/inc",
|
||||||
|
"-I{}/src/common/sdk/nvidia/inc", "-I{}/src/nvidia/arch/nvalloc/unix/include", "-I{}/src/common/sdk/nvidia/inc/ctrl"
|
||||||
|
], rules=[(r'MW\(([^:]+):(.+)\)',r'(\1, \2)')], tarball=nv_src[nm], anon_names={"{}/kernel-open/common/inc/nvstatus.h:37":"nv_status_codes"})
|
||||||
|
case "nv": return load("nv", [], [
|
||||||
|
*[f"{{}}/src/nvidia/inc/kernel/gpu/{s}.h" for s in ["fsp/kern_fsp_cot_payload", "gsp/gsp_init_args"]],
|
||||||
|
*[f"{{}}/src/nvidia/arch/nvalloc/common/inc/{s}.h" for s in ["gsp/gspifpub", "gsp/gsp_fw_wpr_meta", "gsp/gsp_fw_sr_meta", "rmRiscvUcode",
|
||||||
|
"fsp/fsp_nvdm_format"]],
|
||||||
|
*[f"{{}}/src/nvidia/inc/kernel/vgpu/{s}.h" for s in ["rpc_headers", "rpc_global_enums"]],
|
||||||
|
"{}/src/common/uproc/os/common/include/libos_init_args.h", "{}/src/common/shared/msgq/inc/msgq/msgq_priv.h",
|
||||||
|
"{}/src/nvidia/generated/g_rpc-structures.h", root/"extra/nv_gpu_driver/g_rpc-message-header.h", root/"extra/nv_gpu_driver/gsp_static_config.h",
|
||||||
|
root/"extra/nv_gpu_driver/vbios.h", root/"extra/nv_gpu_driver/pci_exp_table.h"
|
||||||
|
], args=[
|
||||||
|
"-DRPC_MESSAGE_STRUCTURES", "-DRPC_STRUCTURES", "-include", "{}/src/common/sdk/nvidia/inc/nvtypes.h", "-I{}/src/nvidia/generated",
|
||||||
|
"-I{}/src/common/inc", "-I{}/src/nvidia/inc", "-I{}/src/nvidia/interface/", "-I{}/src/nvidia/inc/kernel", "-I{}/src/nvidia/inc/libraries",
|
||||||
|
"-I{}/src/nvidia/arch/nvalloc/common/inc", "-I{}/kernel-open/nvidia-uvm", "-I{}/kernel-open/common/inc", "-I{}/src/common/sdk/nvidia/inc",
|
||||||
|
"-I{}/src/nvidia/arch/nvalloc/unix/include", "-I{}/src/common/sdk/nvidia/inc/ctrl"
|
||||||
|
], tarball=nv_src["nv_570"], anon_names={
|
||||||
|
"{}/src/nvidia/inc/kernel/vgpu/rpc_global_enums.h:8": "rpc_fns",
|
||||||
|
"{}/src/nvidia/inc/kernel/vgpu/rpc_global_enums.h:244": "rpc_events"
|
||||||
|
})
|
||||||
|
# this defines all syscall numbers. should probably unify linux autogen?
|
||||||
|
case "io_uring": return load("io_uring", [], ["/usr/include/liburing.h", "/usr/include/linux/io_uring.h", "/usr/include/asm-generic/unistd.h"],
|
||||||
|
rules=[('__NR', 'NR')])
|
||||||
|
case "ib": return load("ib", ["ibverbs"], ["/usr/include/infiniband/verbs.h", "/usr/include/infiniband/verbs_api.h",
|
||||||
|
"/usr/include/infiniband/ib_user_ioctl_verbs.h","/usr/include/rdma/ib_user_verbs.h"], use_errno=True)
|
||||||
|
case "llvm": return load("llvm", ["LLVM_PATH"], lambda: [system("llvm-config-20 --includedir")+"/llvm-c/**/*.h"],
|
||||||
|
args=lambda: system("llvm-config-20 --cflags").split(), recsym=True,
|
||||||
|
prolog=["from tinygrad.runtime.support.llvm import LLVM_PATH"])
|
||||||
|
case "pci": return load("pci", [], ["/usr/include/linux/pci_regs.h"])
|
||||||
|
case "vfio": return load("vfio", [], ["/usr/include/linux/vfio.h"])
|
||||||
|
# could add rule: WGPU_COMMA -> ','
|
||||||
|
case "webgpu":
|
||||||
|
return load("webgpu", ["WEBGPU_PATH"], [root/"extra/webgpu/webgpu.h"], prolog=["from tinygrad.runtime.support.webgpu import WEBGPU_PATH"])
|
||||||
|
case "libusb": return load("libusb", ["os.getenv('LIBUSB_PATH', find_library('usb-1.0'))"], ["/usr/include/libusb-1.0/libusb.h"])
|
||||||
|
case "hip": return load("hip", ["os.getenv('ROCM_PATH', '/opt/rocm')+'/lib/libamdhip64.so'"], ["/opt/rocm/include/hip/hip_ext.h",
|
||||||
|
"/opt/rocm/include/hip/hiprtc.h", "/opt/rocm/include/hip/hip_runtime_api.h", "/opt/rocm/include/hip/driver_types.h"],
|
||||||
|
args=["-D__HIP_PLATFORM_AMD__", "-I/opt/rocm/include", "-x", "c++"])
|
||||||
|
case "comgr" | "comgr_3":
|
||||||
|
return load("comgr_3" if nm == "comgr_3" else "comgr", [
|
||||||
|
"os.getenv('ROCM_PATH', '/opt/rocm')+'/lib/libamd_comgr.so'", "'/usr/local/lib/libamd_comgr.dylib'", "'/opt/homebrew/lib/libamd_comgr.dylib'"
|
||||||
|
], ["/opt/rocm/include/amd_comgr/amd_comgr.h"], args=["-D__HIP_PLATFORM_AMD__", "-I/opt/rocm/include", "-x", "c++"])
|
||||||
|
case "hsa": return load("hsa", ["os.getenv('ROCM_PATH', '/opt/rocm')+'/lib/libhsa-runtime64.so'", "find_library('hsa-runtime64')"], [
|
||||||
|
f"/opt/rocm/include/hsa/{s}.h" for s in ["hsa", "hsa_ext_amd", "amd_hsa_signal", "amd_hsa_queue", "amd_hsa_kernel_code", "hsa_ext_finalize",
|
||||||
|
"hsa_ext_image", "hsa_ven_amd_aqlprofile"] ], args=["-I/opt/rocm/include"])
|
||||||
|
case "amd_gpu": return load("amd_gpu", [], [root/f"extra/hip_gpu_driver/{s}.h" for s in ["sdma_registers", "nvd", "gc_11_0_0_offset",
|
||||||
|
"sienna_cichlid_ip_offset"]],
|
||||||
|
args=["-I/opt/rocm/include", "-x", "c++"])
|
||||||
|
case "kgsl": return load("kgsl", [], [root/"extra/qcom_gpu_driver/msm_kgsl.h"], args=["-D__user="])
|
||||||
|
case "adreno": return load("adreno", [], [root/"extra/qcom_gpu_driver/a6xx.xml.h"])
|
||||||
|
case "qcom_dsp":
|
||||||
|
return load("qcom_dsp", [], [root/f"extra/dsp/include/{s}.h" for s in ["ion", "msm_ion", "adsprpc_shared", "remote_default", "apps_std"]])
|
||||||
|
case "sqtt": return load("sqtt", [], [root/"extra/sqtt/sqtt.h"])
|
||||||
|
case "rocprof":
|
||||||
|
return load("rocprof", ["find_library('rocprof-trace-decoder')", p:="'/usr/local/lib/rocprof-trace-decoder.so'", p.replace('so','dylib')],
|
||||||
|
[f"{{}}/include/{s}.h" for s in ["rocprof_trace_decoder", "trace_decoder_instrument", "trace_decoder_types"]],
|
||||||
|
tarball="https://github.com/ROCm/rocprof-trace-decoder/archive/dd0485100971522cc4cd8ae136bdda431061a04d.tar.gz")
|
||||||
|
case "mesa": return load("mesa", ["find_library('tinymesa_cpu')",
|
||||||
|
"(BASE:=os.getenv('MESA_PATH', f\"/usr{'/local/' if OSX else '/'}lib\"))+'/libtinymesa_cpu'+(EXT:='.dylib' if OSX else '.so')",
|
||||||
|
"f'{BASE}/libtinymesa{EXT}'", "'/opt/homebrew/lib/libtinymesa_cpu.dylib'", "'/opt/homebrew/lib/libtinymesa.dylib'"], [
|
||||||
|
*[f"{{}}/src/compiler/nir/{s}.h" for s in ["nir", "nir_builder", "nir_shader_compiler_options", "nir_serialize"]], "{}/gen/nir_intrinsics.h",
|
||||||
|
*[f"{{}}/src/nouveau/{s}.h" for s in ["headers/nv_device_info", "compiler/nak"]],
|
||||||
|
*[f"{{}}/src/gallium/auxiliary/gallivm/lp_bld{s}.h" for s in ["", "_passmgr", "_misc", "_type", "_init", "_nir", "_struct", "_jit_types",
|
||||||
|
"_flow", "_const"]],
|
||||||
|
"{}/src/compiler/glsl_types.h", "{}/src/util/blob.h", "{}/src/util/ralloc.h"], args=lambda:[
|
||||||
|
"-DHAVE_ENDIAN_H", "-DHAVE_STRUCT_TIMESPEC", "-DHAVE_PTHREAD", "-DHAVE_FUNC_ATTRIBUTE_PACKED", "-I{}/src", "-I{}/include", "-I{}/gen",
|
||||||
|
"-I{}/src/compiler/nir", "-I{}/src/gallium/auxiliary", "-I{}/src/gallium/include", f"-I{system('llvm-config-20 --includedir')}"],
|
||||||
|
preprocess=lambda path: subprocess.run("""mkdir -p gen/util/format
|
||||||
|
python3 src/util/format/u_format_table.py src/util/format/u_format.yaml --enums > gen/util/format/u_format_gen.h
|
||||||
|
python3 src/compiler/nir/nir_opcodes_h.py > gen/nir_opcodes.h
|
||||||
|
python3 src/compiler/nir/nir_intrinsics_h.py --outdir gen
|
||||||
|
python3 src/compiler/nir/nir_intrinsics_indices_h.py --outdir gen
|
||||||
|
python3 src/compiler/nir/nir_builder_opcodes_h.py > gen/nir_builder_opcodes.h
|
||||||
|
python3 src/compiler/nir/nir_intrinsics_h.py --outdir gen
|
||||||
|
python3 src/compiler/builtin_types_h.py gen/builtin_types.h""", cwd=path, shell=True, check=True),
|
||||||
|
tarball="https://gitlab.freedesktop.org/mesa/mesa/-/archive/mesa-25.2.4/mesa-25.2.4.tar.gz",
|
||||||
|
prolog=["import gzip, base64", "from tinygrad.helpers import OSX"], epilog=lambda path: [system(f"{root}/extra/mesa/lvp_nir_options.sh {path}")])
|
||||||
|
case "libclang":
|
||||||
|
return load("libclang", ["os.getenv('LIBCLANG_PATH', find_library('clang-20'))"],
|
||||||
|
lambda: [f"{system('llvm-config-20 --includedir')}/clang-c/{s}.h" for s in ["Index", "CXString", "CXSourceLocation", "CXFile"]],
|
||||||
|
args=lambda: system("llvm-config-20 --cflags").split())
|
||||||
|
case "metal":
|
||||||
|
return load("metal", ["find_library('Metal')"],[f"{macossdk}/System/Library/Frameworks/Metal.framework/Headers/MTL{s}.h" for s in
|
||||||
|
["ComputeCommandEncoder", "ComputePipeline", "CommandQueue", "Device", "IndirectCommandBuffer", "Resource", "CommandEncoder"]],
|
||||||
|
args=["-xobjective-c","-isysroot",macossdk], types={"dispatch_data_t":"objc.id_"})
|
||||||
|
case _: raise AttributeError(f"no such autogen: {nm}")
|
||||||
+7806
-17903
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,23 @@
|
|||||||
|
from tinygrad.runtime.autogen import load, root
|
||||||
|
|
||||||
|
am_src="https://github.com/ROCm/ROCK-Kernel-Driver/archive/ceb12c04e2b5b53ec0779362831f5ee40c4921e4.tar.gz"
|
||||||
|
AMD="{}/drivers/gpu/drm/amd"
|
||||||
|
inc = ["-include", "stdint.h"]
|
||||||
|
|
||||||
|
def __getattr__(nm):
|
||||||
|
match nm:
|
||||||
|
case "am": return load("am/am", [], [root/f"extra/amdpci/headers/{s}.h" for s in ["v11_structs", "v12_structs", "amdgpu_vm", "discovery",
|
||||||
|
"amdgpu_ucode", "psp_gfx_if", "amdgpu_psp", "amdgpu_irq", "amdgpu_doorbell"]]+[f"{AMD}/include/soc15_ih_clientid.h"], args=inc, tarball=am_src)
|
||||||
|
case "pm4_soc15": return load("am/pm4_soc15", [], [f"{AMD}/amdkfd/kfd_pm4_headers_ai.h", f"{AMD}/amdgpu/soc15d.h"], tarball=am_src)
|
||||||
|
case "pm4_nv": return load("am/pm4_nv", [], [f"{AMD}/amdkfd/kfd_pm4_headers_ai.h", f"{AMD}/amdgpu/nvd.h"], tarball=am_src)
|
||||||
|
case "sdma_4_0_0": return load("am/sdma_4_0_0", [], [root/"extra/hip_gpu_driver/sdma_registers.h", f"{AMD}/amdgpu/vega10_sdma_pkt_open.h"],
|
||||||
|
args=["-I/opt/rocm/include", "-x", "c++"], tarball=am_src),
|
||||||
|
case "sdma_5_0_0": return load("am/sdma_5_0_0", [], [root/"extra/hip_gpu_driver/sdma_registers.h", f"{AMD}/amdgpu/navi10_sdma_pkt_open.h"],
|
||||||
|
args=["-I/opt/rocm/include", "-x", "c++"], tarball=am_src),
|
||||||
|
case "sdma_6_0_0": return load("am/sdma_6_0_0", [], [root/"extra/hip_gpu_driver/sdma_registers.h", f"{AMD}//amdgpu/sdma_v6_0_0_pkt_open.h"],
|
||||||
|
args=["-I/opt/rocm/include", "-x", "c++"], tarball=am_src),
|
||||||
|
case "smu_v13_0_0": return load("am/smu_v13_0_0",[],[f"{AMD}/pm/swsmu/inc/pmfw_if/{s}.h" for s in ["smu_v13_0_0_ppsmc","smu13_driver_if_v13_0_0"]]
|
||||||
|
+[root/"extra/amdpci/headers/amdgpu_smu.h"], tarball=am_src),
|
||||||
|
case "smu_v14_0_2": return load("am/smu_v14_0_2", [], [f"{AMD}/pm/swsmu/inc/pmfw_if/{s}.h" for s in ["smu_v14_0_0_pmfw", "smu_v14_0_2_ppsmc",
|
||||||
|
"smu14_driver_if_v14_0"]]+[root/"extra/amdpci/headers/amdgpu_smu.h"], args=inc, tarball=am_src)
|
||||||
|
case _: raise AttributeError(f"no such autogen: {nm}")
|
||||||
+3899
-5626
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
+2637
-5209
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
+3589
-7103
File diff suppressed because it is too large
Load Diff
+4085
-8085
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
+1871
-3494
File diff suppressed because it is too large
Load Diff
+13009
-21999
File diff suppressed because it is too large
Load Diff
+329
-911
File diff suppressed because it is too large
Load Diff
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user